chore: trim reference/ to docs only (5.8MB -> 1.0MB)
- thinktrader_docs: keep .txt (extracted text), drop .html duplicates (36->18 files) - xtquant_big_convert: keep README/LICENSE/CHANGELOG/.gitignore + docs/, drop src/tests/examples/benches/qmt-trader (project code, not referenced by bridge) - un-ignore reference/xtquant_big_convert/.gitignore (now a normal tracked file)
This commit is contained in:
@@ -0,0 +1,28 @@
|
||||
__pycache__/
|
||||
*.py[cod]
|
||||
*.pyo
|
||||
*.pyd
|
||||
.Python
|
||||
.venv/
|
||||
venv/
|
||||
env/
|
||||
.pytest_cache/
|
||||
.mypy_cache/
|
||||
.ruff_cache/
|
||||
build/
|
||||
dist/
|
||||
*.egg-info/
|
||||
|
||||
# Local runtime secrets. Keep real account ids, Redis passwords and QMT paths out of git.
|
||||
# bigqmt_signal_trader_local_config.py
|
||||
# src/bigqmt_signal_trader_local_config.py
|
||||
# bigqmt_signal_trader_client_config.py
|
||||
# src/bigqmt_signal_trader_client_config.py
|
||||
*.local.py
|
||||
*.log
|
||||
*.pid
|
||||
|
||||
# Generated evidence from standalone/QMT backtest bridge runs.
|
||||
backtest_runs/
|
||||
qmt_backtest_runs/
|
||||
.workbuddy/
|
||||
@@ -1,107 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""BigQMT Redis RPC latency benchmark.
|
||||
|
||||
Measures end-to-end latency for ping (no QMT API call) and get_full_tick
|
||||
(real ContextInfo call) to separate transport cost from API cost.
|
||||
"""
|
||||
import os
|
||||
import statistics
|
||||
import time
|
||||
|
||||
import redis
|
||||
from bigqmt_signal_trader.redis_rpc import call_redis_rpc
|
||||
|
||||
|
||||
def _load_redis_config():
|
||||
"""Pull connection details from the local client config or env vars.
|
||||
|
||||
Never hardcode secrets in the repo.
|
||||
"""
|
||||
try:
|
||||
from bigqmt_signal_trader.xtquant_compat import load_client_config
|
||||
|
||||
cfg = load_client_config()
|
||||
rc = dict(cfg.get("redis_config") or {})
|
||||
rc.setdefault("host", os.environ.get("BIGQMT_REDIS_HOST", "127.0.0.1"))
|
||||
rc.setdefault("port", int(os.environ.get("BIGQMT_REDIS_PORT", "6379")))
|
||||
rc.setdefault("db", int(os.environ.get("BIGQMT_REDIS_DB", "5")))
|
||||
return {
|
||||
"host": rc.get("host"),
|
||||
"port": int(rc.get("port")),
|
||||
"db": int(rc.get("db")),
|
||||
"username": rc.get("username") or None,
|
||||
"password": rc.get("password") or None,
|
||||
"socket_timeout": 8,
|
||||
}
|
||||
except Exception:
|
||||
return {
|
||||
"host": os.environ.get("BIGQMT_REDIS_HOST", "127.0.0.1"),
|
||||
"port": int(os.environ.get("BIGQMT_REDIS_PORT", "6379")),
|
||||
"db": int(os.environ.get("BIGQMT_REDIS_DB", "5")),
|
||||
"socket_timeout": 8,
|
||||
}
|
||||
|
||||
|
||||
ACCOUNT = os.environ.get("BIGQMT_ACCOUNT_ID", "")
|
||||
REDIS = _load_redis_config()
|
||||
|
||||
|
||||
def bench(r, method, params, n=20, timeout=6):
|
||||
lats = []
|
||||
errors = 0
|
||||
for i in range(n):
|
||||
t0 = time.time()
|
||||
try:
|
||||
resp = call_redis_rpc(r, ACCOUNT, method, params, timeout_seconds=timeout)
|
||||
dt = (time.time() - t0) * 1000
|
||||
if resp.get("ok"):
|
||||
lats.append(dt)
|
||||
else:
|
||||
errors += 1
|
||||
if errors <= 2:
|
||||
print(" %s #%d error: %s" % (method, i, resp.get("error", "")[:120]))
|
||||
except Exception as e:
|
||||
errors += 1
|
||||
if errors <= 2:
|
||||
print(" %s #%d exc: %s" % (method, i, e))
|
||||
if not lats:
|
||||
print("%-18s: ALL FAILED (%d errors)" % (method, errors))
|
||||
return
|
||||
lats.sort()
|
||||
p50 = statistics.median(lats)
|
||||
p95 = lats[int(len(lats) * 0.95)] if len(lats) >= 20 else lats[-1]
|
||||
print(
|
||||
"%-18s: n=%d ok=%d fail=%d min=%.0f p50=%.0f p95=%.0f max=%.0f avg=%.0f ms"
|
||||
% (
|
||||
method,
|
||||
len(lats),
|
||||
len(lats),
|
||||
errors,
|
||||
min(lats),
|
||||
p50,
|
||||
p95,
|
||||
max(lats),
|
||||
statistics.mean(lats),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def main():
|
||||
r = redis.Redis(**REDIS)
|
||||
# warmup
|
||||
try:
|
||||
call_redis_rpc(r, ACCOUNT, "ping", {}, timeout_seconds=6)
|
||||
print("warmup ping ok\n")
|
||||
except Exception as e:
|
||||
print("warmup FAILED: %s\n" % e)
|
||||
return
|
||||
|
||||
print("=== latency benchmark (20 calls each) ===")
|
||||
bench(r, "ping", {}, n=20)
|
||||
bench(r, "get_full_tick", {"codes": ["000001.SZ"]}, n=20)
|
||||
bench(r, "get_full_tick", {"codes": ["000001.SZ", "600000.SH", "000333.SZ"]}, n=20)
|
||||
bench(r, "get_instrument", {"code": "000001.SZ"}, n=20)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,151 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Compare end-to-end RPC latency across transports.
|
||||
|
||||
Runs the same ping workload through:
|
||||
* Redis (real server, the production path) via call_redis_rpc
|
||||
* ZMQ (local tcp loopback, the low-latency path) via ZmqTransport
|
||||
|
||||
Prints a side-by-side min/p50/p90/p99/max comparison. The ZMQ leg spins up a
|
||||
local in-process server so no QMT process is needed for the comparison.
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import socket
|
||||
import statistics
|
||||
import time
|
||||
import uuid
|
||||
|
||||
import redis
|
||||
|
||||
from bigqmt_signal_trader.redis_rpc import call_redis_rpc
|
||||
from bigqmt_signal_trader.transports.zmq_transport import ZmqTransport
|
||||
|
||||
|
||||
def _load_redis_config():
|
||||
"""Pull connection details from the local client config or env vars."""
|
||||
try:
|
||||
from bigqmt_signal_trader.xtquant_compat import load_client_config
|
||||
|
||||
cfg = load_client_config()
|
||||
rc = dict(cfg.get("redis_config") or {})
|
||||
rc.setdefault("host", os.environ.get("BIGQMT_REDIS_HOST", "127.0.0.1"))
|
||||
rc.setdefault("port", int(os.environ.get("BIGQMT_REDIS_PORT", "6379")))
|
||||
rc.setdefault("db", int(os.environ.get("BIGQMT_REDIS_DB", "5")))
|
||||
return {
|
||||
"host": rc.get("host"),
|
||||
"port": int(rc.get("port")),
|
||||
"db": int(rc.get("db")),
|
||||
"username": rc.get("username") or None,
|
||||
"password": rc.get("password") or None,
|
||||
"socket_timeout": 8,
|
||||
}
|
||||
except Exception:
|
||||
return {
|
||||
"host": os.environ.get("BIGQMT_REDIS_HOST", "127.0.0.1"),
|
||||
"port": int(os.environ.get("BIGQMT_REDIS_PORT", "6379")),
|
||||
"db": int(os.environ.get("BIGQMT_REDIS_DB", "5")),
|
||||
"socket_timeout": 8,
|
||||
}
|
||||
|
||||
|
||||
REDIS = _load_redis_config()
|
||||
ACCOUNT = os.environ.get("BIGQMT_ACCOUNT_ID", "")
|
||||
|
||||
|
||||
def _free_port():
|
||||
s = socket.socket()
|
||||
s.bind(("127.0.0.1", 0))
|
||||
port = s.getsockname()[1]
|
||||
s.close()
|
||||
return port
|
||||
|
||||
|
||||
def _stats(name, lats):
|
||||
lats = sorted(lats)
|
||||
n = len(lats)
|
||||
print(
|
||||
"%-8s n=%d min=%.2f p50=%.2f p90=%.2f p99=%.2f max=%.2f avg=%.2f ms"
|
||||
% (
|
||||
name,
|
||||
n,
|
||||
min(lats),
|
||||
statistics.median(lats),
|
||||
lats[int(n * 0.9)],
|
||||
lats[int(n * 0.99)] if n > 1 else lats[-1],
|
||||
max(lats),
|
||||
statistics.mean(lats),
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
def bench_redis(n):
|
||||
r = redis.Redis(**REDIS)
|
||||
# warmup + connectivity
|
||||
try:
|
||||
call_redis_rpc(r, ACCOUNT, "ping", {}, timeout_seconds=6)
|
||||
except Exception as e:
|
||||
print("Redis server not reachable, skipping redis leg: %s" % e)
|
||||
return
|
||||
lats = []
|
||||
for _ in range(n):
|
||||
t0 = time.time()
|
||||
call_redis_rpc(r, ACCOUNT, "ping", {}, timeout_seconds=6)
|
||||
lats.append((time.time() - t0) * 1000)
|
||||
_stats("redis", lats)
|
||||
|
||||
|
||||
def bench_zmq(n):
|
||||
port = _free_port()
|
||||
addr = "tcp://127.0.0.1:%d" % port
|
||||
|
||||
def on_req(req):
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"request_id": req["request_id"],
|
||||
"account_id": "zmq",
|
||||
"method": req["method"],
|
||||
"ok": True,
|
||||
"data": {"pong": True},
|
||||
"error": "",
|
||||
"handled_at": "now",
|
||||
}
|
||||
|
||||
server = ZmqTransport(bind_address=addr, account_id="zmq", recv_timeout_seconds=0.3)
|
||||
server.start_receiving(on_req, background_threads=True)
|
||||
time.sleep(0.3)
|
||||
client = ZmqTransport(connect_address=addr, account_id="zmq")
|
||||
time.sleep(0.2)
|
||||
lats = []
|
||||
for _ in range(n):
|
||||
req = {
|
||||
"schema_version": 1,
|
||||
"request_id": uuid.uuid4().hex,
|
||||
"account_id": "zmq",
|
||||
"method": "ping",
|
||||
"params": {},
|
||||
"reply_channel": "",
|
||||
"reply_list": "",
|
||||
"reply_key": "",
|
||||
"ttl_seconds": 5,
|
||||
}
|
||||
t0 = time.time()
|
||||
client.send_request(req, timeout_seconds=3.0)
|
||||
lats.append((time.time() - t0) * 1000)
|
||||
_stats("zmq", lats)
|
||||
server.stop()
|
||||
client.stop()
|
||||
|
||||
|
||||
def main():
|
||||
ap = argparse.ArgumentParser()
|
||||
ap.add_argument("-n", "--count", type=int, default=100, help="requests per transport")
|
||||
ap.add_argument("--skip-redis", action="store_true", help="skip the redis leg")
|
||||
args = ap.parse_args()
|
||||
print("=== transport latency comparison (n=%d each) ===" % args.count)
|
||||
if not args.skip_redis:
|
||||
bench_redis(args.count)
|
||||
bench_zmq(args.count)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,51 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Measure ZMQ GIL-spike rate at different request rates.
|
||||
|
||||
Sends N get_full_tick requests at a fixed interval, records per-call
|
||||
latency, reports how many exceed thresholds (50/200/500/1000ms).
|
||||
"""
|
||||
import sys
|
||||
import time
|
||||
|
||||
sys.path.insert(0, r"D:\gjzqqmt\xtquant_big_convert\src")
|
||||
sys.path.insert(0, r"D:\国金证券QMT交易端_lemo\python")
|
||||
|
||||
import bigqmt_signal_trader.xtquant_compat as compat
|
||||
|
||||
compat.configure()
|
||||
client = compat.get_default_client()
|
||||
print("account:", client.account_id, "| transport:", client.transport_name)
|
||||
print("=" * 60)
|
||||
|
||||
N = 40
|
||||
INTERVAL_MS = 50 # 请求间隔 50ms = 20 QPS
|
||||
|
||||
latencies = []
|
||||
for i in range(N):
|
||||
t0 = time.time()
|
||||
try:
|
||||
client.call("get_full_tick", {"codes": ["000001.SZ"]})
|
||||
ms = (time.time() - t0) * 1000
|
||||
latencies.append(ms)
|
||||
except Exception as e:
|
||||
ms = (time.time() - t0) * 1000
|
||||
latencies.append(ms)
|
||||
print(" [%2d] FAIL %.0fms %s" % (i, ms, str(e)[:40]))
|
||||
# 控制频率
|
||||
elapsed = (time.time() - t0)
|
||||
sleep = max(0, INTERVAL_MS / 1000.0 - elapsed)
|
||||
if sleep > 0:
|
||||
time.sleep(sleep)
|
||||
|
||||
latencies.sort()
|
||||
n = len(latencies)
|
||||
print("\n=== %d requests @ %dms interval (%.0f QPS) ===" % (n, INTERVAL_MS, 1000.0/INTERVAL_MS))
|
||||
print("min=%.1f p50=%.1f p90=%.1f p99=%.1f max=%.1f" % (
|
||||
latencies[0], latencies[n//2], latencies[int(n*0.9)], latencies[int(n*0.99)], latencies[-1]))
|
||||
|
||||
# 尖峰分布
|
||||
thresholds = [10, 50, 100, 200, 500, 1000]
|
||||
print("\n=== 延迟分布 ===")
|
||||
for t in thresholds:
|
||||
cnt = sum(1 for l in latencies if l > t)
|
||||
print(" >%5dms : %2d / %d (%.0f%%)" % (t, cnt, n, 100.0*cnt/n))
|
||||
@@ -1,252 +0,0 @@
|
||||
#coding:gbk
|
||||
"""QMT bridge entry (no-redis version).
|
||||
|
||||
Same file-loader pattern as BIGQMT_REDIS_DRYRUN, but the RPC transport is ZMQ
|
||||
only -- no redis imports anywhere. This version loads the no-redis zmq transport
|
||||
(bigqmt_no_redis/zmq_transport.py) which inlines all encoding helpers and drops
|
||||
redis-based service discovery, so it loads cleanly in QMT sandboxes that reject
|
||||
`import redis` or any redis-named module.
|
||||
|
||||
Use this when your QMT environment cannot import the redis package (e.g. broker
|
||||
whitelist blocks it) or when you want zero redis dependency.
|
||||
|
||||
Config: set "transport": "zmq" in bigqmt_signal_trader_local_config.py (the
|
||||
no-redis runtime forces zmq regardless). Redis config fields are ignored.
|
||||
"""
|
||||
import builtins as _builtins
|
||||
import importlib as _importlib
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
|
||||
|
||||
_LOCAL_ROOTS = (
|
||||
"bigqmt_signal_trader",
|
||||
"bigqmt_signal_trader_strategy",
|
||||
"bigqmt_signal_trader_redis_rpc_runtime",
|
||||
"bigqmt_signal_trader_local_config",
|
||||
"bigqmt_no_redis",
|
||||
)
|
||||
_ORIGINAL_IMPORT = _builtins.__import__
|
||||
_ORIGINAL_IMPORT_MODULE = _importlib.import_module
|
||||
_ORIGINAL_RELOAD = _importlib.reload
|
||||
|
||||
|
||||
def _known_qmt_python_dir():
|
||||
# Find the QMT python dir from sys.path instead of a hardcoded path, so
|
||||
# the bridge loads regardless of broker install location or launch mode.
|
||||
for p in sys.path:
|
||||
if p and r"\python" in p and os.path.isdir(p):
|
||||
return p
|
||||
return ""
|
||||
|
||||
|
||||
try:
|
||||
_SOURCE_ROOT = os.path.dirname(os.path.abspath(__file__))
|
||||
except Exception:
|
||||
_SOURCE_ROOT = _known_qmt_python_dir()
|
||||
if not _SOURCE_ROOT:
|
||||
_SOURCE_ROOT = _known_qmt_python_dir()
|
||||
|
||||
|
||||
def _is_local_module(name):
|
||||
return any(name == root or name.startswith(root + ".") for root in _LOCAL_ROOTS)
|
||||
|
||||
|
||||
def _resolve_name(name, module_globals, level):
|
||||
if not level:
|
||||
return name
|
||||
package = (module_globals or {}).get("__package__") or (module_globals or {}).get("__name__", "")
|
||||
if not package:
|
||||
raise ImportError("relative import without package")
|
||||
for unused in range(level - 1):
|
||||
if "." not in package:
|
||||
raise ImportError("relative import beyond top-level package")
|
||||
package = package.rsplit(".", 1)[0]
|
||||
return package + ("." + name if name else "")
|
||||
|
||||
|
||||
def _find_local_source(name):
|
||||
relative = name.replace(".", os.sep)
|
||||
dirs = []
|
||||
if _SOURCE_ROOT:
|
||||
dirs.append(_SOURCE_ROOT)
|
||||
for p in sys.path:
|
||||
if p and os.path.isdir(p) and p not in dirs:
|
||||
dirs.append(p)
|
||||
for d in dirs:
|
||||
package_init = os.path.join(d, relative, "__init__.py")
|
||||
if os.path.isfile(package_init):
|
||||
return package_init, True
|
||||
module_file = os.path.join(d, relative + ".py")
|
||||
if os.path.isfile(module_file):
|
||||
return module_file, False
|
||||
raise ModuleNotFoundError("local source not found: %s" % name, name=name)
|
||||
|
||||
|
||||
def _set_parent_attribute(name, module):
|
||||
if "." not in name:
|
||||
return
|
||||
parent_name, child_name = name.rsplit(".", 1)
|
||||
parent = _load_local_module(parent_name)
|
||||
setattr(parent, child_name, module)
|
||||
|
||||
|
||||
def _load_local_module(name):
|
||||
existing = sys.modules.get(name)
|
||||
if existing is not None:
|
||||
return existing
|
||||
source_path, is_package = _find_local_source(name)
|
||||
if "." in name:
|
||||
_load_local_module(name.rsplit(".", 1)[0])
|
||||
module = types.ModuleType(name)
|
||||
module.__file__ = source_path
|
||||
module.__package__ = name if is_package else name.rpartition(".")[0]
|
||||
if is_package:
|
||||
module.__path__ = [os.path.dirname(source_path)]
|
||||
module_builtins = dict(_builtins.__dict__)
|
||||
module_builtins["__import__"] = _local_import
|
||||
module.__dict__["__builtins__"] = module_builtins
|
||||
module.__dict__["__bigqmt_load_local_module"] = _load_local_module
|
||||
sys.modules[name] = module
|
||||
if name == "bigqmt_signal_trader":
|
||||
return module
|
||||
try:
|
||||
with open(source_path, "rb") as source_file:
|
||||
source = source_file.read()
|
||||
exec(compile(source, source_path, "exec"), module.__dict__)
|
||||
except Exception:
|
||||
sys.modules.pop(name, None)
|
||||
raise
|
||||
_set_parent_attribute(name, module)
|
||||
return module
|
||||
|
||||
|
||||
def _local_import(name, module_globals=None, module_locals=None, fromlist=(), level=0):
|
||||
absolute_name = _resolve_name(name, module_globals, level)
|
||||
if not _is_local_module(absolute_name):
|
||||
return _ORIGINAL_IMPORT(name, module_globals, module_locals, fromlist, level)
|
||||
module = _load_local_module(absolute_name)
|
||||
for child in fromlist or ():
|
||||
if child != "*":
|
||||
try:
|
||||
_load_local_module(absolute_name + "." + child)
|
||||
except ModuleNotFoundError:
|
||||
pass
|
||||
if fromlist:
|
||||
return module
|
||||
return _load_local_module(absolute_name.split(".", 1)[0])
|
||||
|
||||
|
||||
def _local_import_module(name, package=None):
|
||||
if _is_local_module(name):
|
||||
return _load_local_module(name)
|
||||
return _ORIGINAL_IMPORT_MODULE(name, package)
|
||||
|
||||
|
||||
def _local_reload(module):
|
||||
if _is_local_module(getattr(module, "__name__", "")):
|
||||
return _load_local_module(module.__name__)
|
||||
return _ORIGINAL_RELOAD(module)
|
||||
|
||||
|
||||
def _clear_local_modules():
|
||||
for name in list(sys.modules):
|
||||
if _is_local_module(name):
|
||||
sys.modules.pop(name, None)
|
||||
|
||||
|
||||
def _stop_previous_rpc_service():
|
||||
"""Release the previous QMT strategy's socket before clearing its module."""
|
||||
previous = sys.modules.get("bigqmt_signal_trader_strategy")
|
||||
reset = getattr(previous, "reset_app", None)
|
||||
if not callable(reset):
|
||||
return
|
||||
try:
|
||||
reset()
|
||||
print("[bigqmt_shell] previous rpc service stopped")
|
||||
except Exception as exc:
|
||||
print("[bigqmt_shell] previous rpc service stop failed: %s" % exc)
|
||||
|
||||
|
||||
_stop_previous_rpc_service()
|
||||
_clear_local_modules()
|
||||
_importlib.import_module = _local_import_module
|
||||
_importlib.reload = _local_reload
|
||||
print("[bigqmt_shell] importlib entry source_root=%s" % _SOURCE_ROOT)
|
||||
|
||||
|
||||
def _fallback_account_id():
|
||||
for name in ("BIGQMT_ACCOUNT_ID", "account", "account_id", "accountID"):
|
||||
value = globals().get(name)
|
||||
if value:
|
||||
return str(value)
|
||||
return ""
|
||||
|
||||
|
||||
try:
|
||||
_local_import("bigqmt_signal_trader.adapters.market_bigqmt", globals(), fromlist=("*",))
|
||||
_local_import("bigqmt_signal_trader.adapters.order_bigqmt", globals(), fromlist=("*",))
|
||||
_local_import("bigqmt_signal_trader.adapters.position_bigqmt", globals(), fromlist=("*",))
|
||||
_strategy = _local_import("bigqmt_signal_trader_strategy", globals(), fromlist=("*",))
|
||||
_strategy.reset_app()
|
||||
except Exception as bridge_preload_error:
|
||||
print("[bigqmt_shell] bridge preload failed: %s" % bridge_preload_error)
|
||||
|
||||
_runtime = _local_import("bigqmt_signal_trader_redis_rpc_runtime", globals(), fromlist=("*",))
|
||||
|
||||
|
||||
def _load_local_config():
|
||||
return _local_import("bigqmt_signal_trader_local_config", globals(), fromlist=("*",))
|
||||
|
||||
|
||||
try:
|
||||
_config = _load_local_config()
|
||||
BIGQMT_REDIS_CONFIG = getattr(_config, "BIGQMT_REDIS_CONFIG", {})
|
||||
# Force zmq transport (this is the no-redis version).
|
||||
BIGQMT_REDIS_CONFIG = dict(BIGQMT_REDIS_CONFIG or {})
|
||||
BIGQMT_REDIS_CONFIG["transport"] = "zmq"
|
||||
BIGQMT_REDIS_CONFIG["rpc_background_threads"] = True
|
||||
print("[bigqmt_shell] no-redis mode: transport=zmq background_threads=True")
|
||||
_runtime.configure_runtime_redis(BIGQMT_REDIS_CONFIG)
|
||||
except Exception as redis_config_error:
|
||||
print("[bigqmt_shell] local redis config load failed: %s" % redis_config_error)
|
||||
|
||||
try:
|
||||
_config = _load_local_config()
|
||||
BIGQMT_ACCOUNT_ID = getattr(_config, "BIGQMT_ACCOUNT_ID", "")
|
||||
print("[bigqmt_shell] local account config loaded=%s" % bool(BIGQMT_ACCOUNT_ID))
|
||||
_runtime.configure_runtime_account(BIGQMT_ACCOUNT_ID)
|
||||
except Exception as account_config_error:
|
||||
print("[bigqmt_shell] local account config load failed: %s" % account_config_error)
|
||||
account_id = _fallback_account_id()
|
||||
if account_id:
|
||||
_runtime.configure_runtime_account(account_id)
|
||||
|
||||
try:
|
||||
qmt_extra = {}
|
||||
for function_name in (
|
||||
"get_history_trade_detail_data", "get_value_by_order_id", "get_last_order_id",
|
||||
"get_ipo_data", "get_new_purchase_limit", "get_assure_contract",
|
||||
"get_enable_short_contract", "get_unclosed_compacts", "get_closed_compacts",
|
||||
"get_debt_contract", "get_option_subject_position", "get_comb_option",
|
||||
"get_hkt_exchange_rate", "down_history_data",
|
||||
):
|
||||
if function_name in globals():
|
||||
qmt_extra[function_name] = globals()[function_name]
|
||||
print("[bigqmt_shell] down_history_data bound=%s" % ("down_history_data" in qmt_extra))
|
||||
_runtime.bind_runtime_api(
|
||||
passorder_func=globals().get("passorder"),
|
||||
cancel_func=globals().get("cancel"),
|
||||
get_trade_detail_data_func=globals().get("get_trade_detail_data"),
|
||||
extra_funcs=qmt_extra or None,
|
||||
)
|
||||
except NameError:
|
||||
pass
|
||||
|
||||
|
||||
init = _runtime.init
|
||||
handlebar = _runtime.handlebar
|
||||
adjust = _runtime.adjust
|
||||
order_callback = _runtime.order_callback
|
||||
deal_callback = _runtime.deal_callback
|
||||
@@ -1,459 +0,0 @@
|
||||
"""ZeroMQ transport for the BigQMT RPC bridge (no-redis version).
|
||||
|
||||
Same as bigqmt_signal_trader.transports.zmq_transport, but with the redis
|
||||
dependencies inlined and the redis-based service discovery removed. This lets
|
||||
the module load in QMT sandboxes that reject `import redis` or any module
|
||||
whose name mentions redis.
|
||||
|
||||
Design (unchanged from the redis version):
|
||||
|
||||
* **Server** binds a ``ROUTER`` socket. Each inbound message arrives as
|
||||
``[identity, payload]``; the server remembers ``identity`` keyed by
|
||||
``request_id`` and replies with ``[identity, payload]`` so ZMQ routes the
|
||||
response back to the originating client automatically.
|
||||
* **Client** connects a ``DEALER`` socket (with a unique random identity), sends
|
||||
``[payload]``, then ``poll``/``recv`` for the response.
|
||||
"""
|
||||
|
||||
import base64
|
||||
import json
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Inlined encoding helpers (originally from bigqmt_signal_trader.adapters.
|
||||
# redis_common and bigqmt_signal_trader.redis_rpc). Kept here so this module
|
||||
# has zero imports from any redis-named module.
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
SAFE_B64_PREFIX = "b64s:"
|
||||
SAFE_B64_DIGIT_ENCODE = str.maketrans("0123456789", "!#$%&()*~?")
|
||||
SAFE_B64_DIGIT_DECODE = str.maketrans("!#$%&()*~?", "0123456789")
|
||||
|
||||
|
||||
def decode_text(value):
|
||||
if isinstance(value, bytes):
|
||||
return value.decode("utf-8")
|
||||
return str(value)
|
||||
|
||||
|
||||
def encode_rpc_request_payload(request):
|
||||
"""Encode request JSON so patched QMT clients do not inspect stock-code text."""
|
||||
raw = json.dumps(request, ensure_ascii=False).encode("utf-8")
|
||||
encoded = base64.b64encode(raw).decode("ascii").translate(SAFE_B64_DIGIT_ENCODE)
|
||||
return SAFE_B64_PREFIX + encoded
|
||||
|
||||
|
||||
def decode_rpc_request_payload(text):
|
||||
text = str(text)
|
||||
if not text.startswith(SAFE_B64_PREFIX):
|
||||
return text
|
||||
encoded = text[len(SAFE_B64_PREFIX):].translate(SAFE_B64_DIGIT_DECODE)
|
||||
return base64.b64decode(encoded.encode("ascii")).decode("utf-8")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# TransportError / TransportTimeout (inlined from transports.base -- kept here
|
||||
# so this module is fully self-contained for QMT sandbox loading).
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class TransportError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class TransportTimeout(TransportError):
|
||||
pass
|
||||
|
||||
|
||||
class RpcTransport:
|
||||
"""Minimal transport base (inlined subset of transports.base)."""
|
||||
|
||||
def __init__(self, account_id="", print_prefix="[bigqmt_rpc]"):
|
||||
self.account_id = str(account_id or "")
|
||||
self.print_prefix = str(print_prefix or "[bigqmt_rpc]")
|
||||
self._on_request = None
|
||||
self._running = False
|
||||
|
||||
def start_receiving(self, on_request):
|
||||
self._on_request = on_request
|
||||
self._running = True
|
||||
|
||||
def stop(self):
|
||||
self._running = False
|
||||
self._on_request = None
|
||||
|
||||
def deliver(self, request):
|
||||
callback = self._on_request
|
||||
if callback is None:
|
||||
return None
|
||||
try:
|
||||
response = callback(request)
|
||||
except Exception as exc:
|
||||
import datetime as _dt
|
||||
response = {
|
||||
"schema_version": 1,
|
||||
"request_id": str((request or {}).get("request_id") or ""),
|
||||
"account_id": str((request or {}).get("account_id") or self.account_id or ""),
|
||||
"method": str((request or {}).get("method") or ""),
|
||||
"ok": False,
|
||||
"data": None,
|
||||
"error": "%s: %s" % (exc.__class__.__name__, exc),
|
||||
"handled_at": _dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
if response is not None:
|
||||
try:
|
||||
self.send_response(request, response)
|
||||
except Exception:
|
||||
pass
|
||||
return response
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ZMQ transport
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
# ZMQ does not support ipc:// on Windows (it trips a signaler abort), so the
|
||||
# default endpoint is tcp loopback. The port is derived from the account_id so
|
||||
# distinct accounts don't collide on the same port; override via config when
|
||||
# needed. Base 15560 keeps it clear of common dev ports.
|
||||
DEFAULT_ZMQ_HOST = "127.0.0.1"
|
||||
DEFAULT_ZMQ_BASE_PORT = 15560
|
||||
DEFAULT_ZMQ_PORT_RANGE = 100 # derived port = base + (account_id_int mod range)
|
||||
|
||||
|
||||
def _default_zmq_port(account_id):
|
||||
"""Derive a stable port from account_id so each account gets its own socket."""
|
||||
text = str(account_id or "")
|
||||
digits = "".join(ch for ch in text if ch.isdigit())
|
||||
try:
|
||||
offset = int(digits) % DEFAULT_ZMQ_PORT_RANGE if digits else 0
|
||||
except ValueError:
|
||||
offset = 0
|
||||
return DEFAULT_ZMQ_BASE_PORT + offset
|
||||
|
||||
|
||||
def _default_zmq_address(account_id, host=None):
|
||||
host = host or DEFAULT_ZMQ_HOST
|
||||
return "tcp://%s:%d" % (host, _default_zmq_port(account_id))
|
||||
|
||||
|
||||
def _loads(raw):
|
||||
if isinstance(raw, dict):
|
||||
return dict(raw)
|
||||
text = decode_text(raw)
|
||||
text = decode_rpc_request_payload(text)
|
||||
return json.loads(text)
|
||||
|
||||
|
||||
class ZmqTransport(RpcTransport):
|
||||
"""ZMQ ROUTER/DEALER transport (no-redis version).
|
||||
|
||||
The same instance plays both roles depending on method called:
|
||||
``send_request`` acts as a client (DEALER connect), ``start_receiving`` +
|
||||
``send_response`` act as a server (ROUTER bind). A deployment normally uses
|
||||
one instance per role (the QMT process is the server; the external client
|
||||
is the client).
|
||||
|
||||
Unlike the redis version, this one does NOT use redis-based service
|
||||
discovery. The server binds the configured address exactly; the client
|
||||
connects to the configured or derived address directly.
|
||||
"""
|
||||
|
||||
name = "zmq"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
bind_address=None,
|
||||
connect_address=None,
|
||||
host=None,
|
||||
port=None,
|
||||
account_id="",
|
||||
print_prefix="[bigqmt_rpc]",
|
||||
io_threads=1,
|
||||
recv_timeout_seconds=1.0,
|
||||
server_hwm=10000,
|
||||
client_linger_ms=0,
|
||||
):
|
||||
super(ZmqTransport, self).__init__(account_id=account_id, print_prefix=print_prefix)
|
||||
resolved_host = host or DEFAULT_ZMQ_HOST
|
||||
if port is not None:
|
||||
resolved_port = int(port)
|
||||
else:
|
||||
resolved_port = _default_zmq_port(account_id)
|
||||
default_addr = "tcp://%s:%d" % (resolved_host, resolved_port)
|
||||
self.bind_address = bind_address or default_addr
|
||||
self.connect_address = connect_address
|
||||
self.bind_host = resolved_host
|
||||
self.base_port = resolved_port
|
||||
self.io_threads = int(io_threads)
|
||||
self.recv_timeout_seconds = float(recv_timeout_seconds)
|
||||
self.server_hwm = int(server_hwm)
|
||||
self.client_linger_ms = int(client_linger_ms)
|
||||
|
||||
self._zmq = None # imported lazily
|
||||
self._ctx = None
|
||||
# server state
|
||||
self._router = None
|
||||
self._router_thread = None
|
||||
self._actual_bind_address = None # set after start_receiving()
|
||||
self._pending_identities = {} # request_id -> client identity bytes
|
||||
self._identity_lock = threading.Lock()
|
||||
self._response_queue = queue.Queue()
|
||||
self._queued_response_count = 0
|
||||
self._sent_response_count = 0
|
||||
# client state
|
||||
self._dealer = None
|
||||
self._client_lock = threading.Lock()
|
||||
|
||||
# -- construction helper ----------------------------------------------
|
||||
@classmethod
|
||||
def from_config(cls, config, account_id="", print_prefix="[bigqmt_rpc]"):
|
||||
config = dict(config or {})
|
||||
return cls(
|
||||
bind_address=config.get("bind_address"),
|
||||
connect_address=config.get("connect_address"),
|
||||
host=config.get("host"),
|
||||
port=config.get("port"),
|
||||
account_id=config.get("account_id", account_id),
|
||||
print_prefix=print_prefix,
|
||||
io_threads=int(config.get("io_threads", 1)),
|
||||
recv_timeout_seconds=float(config.get("recv_timeout_seconds", 1.0)),
|
||||
server_hwm=int(config.get("server_hwm", 10000)),
|
||||
client_linger_ms=int(config.get("client_linger_ms", 0)),
|
||||
)
|
||||
|
||||
# -- shared zmq context -----------------------------------------------
|
||||
def _ensure_zmq(self):
|
||||
if self._zmq is None:
|
||||
try:
|
||||
import zmq # noqa: F401
|
||||
except ImportError as exc: # pragma: no cover - depends on env
|
||||
raise TransportError(
|
||||
"pyzmq is required for the zmq transport: %s" % exc
|
||||
)
|
||||
self._zmq = zmq
|
||||
if self._ctx is None:
|
||||
self._ctx = self._zmq.Context.instance(self.io_threads)
|
||||
return self._zmq, self._ctx
|
||||
|
||||
# -- server side ------------------------------------------------------
|
||||
def _bind_configured_address(self):
|
||||
"""Bind exactly one configured address and reject duplicate servers."""
|
||||
zmq, ctx = self._ensure_zmq()
|
||||
sock = ctx.socket(zmq.ROUTER)
|
||||
sock.setsockopt(zmq.RCVHWM, self.server_hwm)
|
||||
sock.setsockopt(zmq.SNDHWM, self.server_hwm)
|
||||
sock.setsockopt(zmq.RCVTIMEO, int(self.recv_timeout_seconds * 1000))
|
||||
try:
|
||||
sock.bind(self.bind_address)
|
||||
except self._zmq.ZMQError as exc:
|
||||
try:
|
||||
sock.close(linger=0)
|
||||
except Exception:
|
||||
pass
|
||||
if getattr(exc, "errno", None) == zmq.EADDRINUSE:
|
||||
raise TransportError(
|
||||
"ZMQ_BIND_CONFLICT address=%s; another bridge instance "
|
||||
"already owns the configured endpoint" % self.bind_address
|
||||
)
|
||||
raise
|
||||
self._router = sock
|
||||
self._actual_bind_address = self.bind_address
|
||||
|
||||
def start_receiving(self, on_request, background_threads=True):
|
||||
super(ZmqTransport, self).start_receiving(on_request)
|
||||
zmq, ctx = self._ensure_zmq()
|
||||
self._bind_configured_address()
|
||||
bound = self._actual_bind_address or self.bind_address
|
||||
if not background_threads:
|
||||
print(
|
||||
"%s zmq bound=%s background_threads=False"
|
||||
% (self.print_prefix, bound)
|
||||
)
|
||||
return
|
||||
self._router_thread = threading.Thread(
|
||||
target=self._router_loop, name="bigqmt-zmq-rpc", daemon=True
|
||||
)
|
||||
self._router_thread.start()
|
||||
print(
|
||||
"%s zmq started bound=%s" % (self.print_prefix, self.bind_address)
|
||||
)
|
||||
|
||||
def _router_loop(self):
|
||||
try:
|
||||
while self._running:
|
||||
self._drain_response_queue()
|
||||
request = self._receive_request()
|
||||
if request is not None:
|
||||
self._deliver_request(request)
|
||||
finally:
|
||||
# Close the ROUTER socket on the thread that owns it. On Windows,
|
||||
# closing a ZMQ socket from a different thread trips a signaler
|
||||
# assertion (abort); closing it here is safe because this thread
|
||||
# created and exclusively used it.
|
||||
try:
|
||||
self._router.close(linger=0)
|
||||
except Exception:
|
||||
pass
|
||||
self._router = None
|
||||
|
||||
def _receive_request(self, flags=0):
|
||||
try:
|
||||
frames = self._router.recv_multipart(flags=flags)
|
||||
except self._zmq.Again:
|
||||
return None
|
||||
except Exception as exc:
|
||||
if self._running:
|
||||
print("%s zmq recv failed: %s" % (self.print_prefix, exc))
|
||||
if not flags:
|
||||
time.sleep(0.5)
|
||||
return None
|
||||
if len(frames) < 2:
|
||||
return None
|
||||
identity, payload = frames[0], frames[-1]
|
||||
try:
|
||||
request = _loads(payload)
|
||||
except Exception as exc:
|
||||
print("%s zmq decode failed: %s" % (self.print_prefix, exc))
|
||||
return None
|
||||
request_id = str(request.get("request_id") or uuid.uuid4().hex)
|
||||
with self._identity_lock:
|
||||
self._pending_identities[request_id] = identity
|
||||
return request
|
||||
|
||||
def _deliver_request(self, request):
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
self.deliver(request)
|
||||
except Exception as exc:
|
||||
print("%s zmq deliver failed: %s" % (self.print_prefix, exc))
|
||||
elapsed_ms = (time.perf_counter() - started) * 1000.0
|
||||
if elapsed_ms > 50.0:
|
||||
print("%s zmq slow handler method=%s %.0fms"
|
||||
% (self.print_prefix, request.get("method"), elapsed_ms))
|
||||
|
||||
def _drain_response_queue(self):
|
||||
while True:
|
||||
try:
|
||||
identity, payload = self._response_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
return
|
||||
try:
|
||||
self._router.send_multipart([identity, payload])
|
||||
self._sent_response_count += 1
|
||||
if self._sent_response_count <= 5:
|
||||
print("%s zmq queued response sent" % self.print_prefix)
|
||||
except Exception as exc:
|
||||
print("%s zmq send failed: %s" % (self.print_prefix, exc))
|
||||
|
||||
def send_response(self, request, response):
|
||||
if self._router is None:
|
||||
raise TransportError("zmq server socket is not bound")
|
||||
request_id = str(
|
||||
response.get("request_id") or request.get("request_id") or ""
|
||||
)
|
||||
with self._identity_lock:
|
||||
identity = self._pending_identities.pop(request_id, None)
|
||||
if identity is None:
|
||||
# No matching peer -- drop silently (client may have gone away).
|
||||
return
|
||||
payload = encode_rpc_request_payload(response).encode("utf-8")
|
||||
if self._router_thread is not None and threading.current_thread() is not self._router_thread:
|
||||
self._queued_response_count += 1
|
||||
if self._queued_response_count <= 5:
|
||||
print("%s zmq response queued for router thread" % self.print_prefix)
|
||||
self._response_queue.put((identity, payload))
|
||||
return
|
||||
try:
|
||||
self._router.send_multipart([identity, payload])
|
||||
except Exception as exc:
|
||||
print("%s zmq send failed: %s" % (self.print_prefix, exc))
|
||||
|
||||
def drain_request_queue(self, max_items=20):
|
||||
"""Drain requests from the scheduled QMT thread when no receiver thread exists."""
|
||||
if self._router_thread is not None or self._router is None:
|
||||
return 0
|
||||
processed = 0
|
||||
for _index in range(max(int(max_items), 0)):
|
||||
request = self._receive_request(flags=self._zmq.NOBLOCK)
|
||||
if request is None:
|
||||
break
|
||||
self._deliver_request(request)
|
||||
processed += 1
|
||||
return processed
|
||||
|
||||
# -- client side ------------------------------------------------------
|
||||
def _resolve_connect_address(self):
|
||||
"""Resolve the address to connect to. No redis discovery -- use explicit
|
||||
connect_address, else derive from account_id."""
|
||||
if self.connect_address:
|
||||
return self.connect_address
|
||||
return _default_zmq_address(self.account_id)
|
||||
|
||||
def _ensure_dealer(self):
|
||||
zmq, ctx = self._ensure_zmq()
|
||||
if self._dealer is None:
|
||||
address = self._resolve_connect_address()
|
||||
sock = ctx.socket(zmq.DEALER)
|
||||
# Unique identity so ROUTER can route replies back to us.
|
||||
sock.setsockopt(zmq.IDENTITY, uuid.uuid4().hex.encode("utf-8")[:16])
|
||||
sock.setsockopt(zmq.LINGER, self.client_linger_ms)
|
||||
sock.connect(address)
|
||||
self._dealer = sock
|
||||
self.connect_address = address
|
||||
return self._dealer
|
||||
|
||||
def send_request(self, request, timeout_seconds, **_kwargs):
|
||||
zmq = self._zmq or self._ensure_zmq()[0]
|
||||
with self._client_lock:
|
||||
dealer = self._ensure_dealer()
|
||||
request = dict(request)
|
||||
request.setdefault("request_id", uuid.uuid4().hex)
|
||||
request_id = request["request_id"]
|
||||
payload = encode_rpc_request_payload(request)
|
||||
try:
|
||||
dealer.send(payload.encode("utf-8"))
|
||||
except Exception as exc:
|
||||
raise TransportError("zmq send failed: %s" % exc)
|
||||
deadline = time.time() + float(timeout_seconds)
|
||||
poller = self._zmq.Poller()
|
||||
poller.register(dealer, self._zmq.POLLIN)
|
||||
while True:
|
||||
remaining = deadline - time.time()
|
||||
if remaining <= 0:
|
||||
break
|
||||
events = dict(poller.poll(timeout=int(remaining * 1000)))
|
||||
if dealer in events:
|
||||
frames = dealer.recv_multipart()
|
||||
raw = frames[-1]
|
||||
response = _loads(raw)
|
||||
if response.get("request_id") == request_id:
|
||||
return response
|
||||
raise TransportTimeout("zmq rpc timeout: %s" % request.get("method"))
|
||||
|
||||
# -- lifecycle --------------------------------------------------------
|
||||
def stop(self):
|
||||
super(ZmqTransport, self).stop()
|
||||
# Clear _running so the router loop exits; the loop closes its own
|
||||
# socket (closing cross-thread trips a Windows signaler abort).
|
||||
thread = self._router_thread
|
||||
if thread is not None and thread.is_alive():
|
||||
thread.join(2.0)
|
||||
if thread is None and self._router is not None:
|
||||
try:
|
||||
self._router.close(linger=0)
|
||||
except Exception:
|
||||
pass
|
||||
self._router = None
|
||||
self._router_thread = None
|
||||
with self._client_lock:
|
||||
if self._dealer is not None:
|
||||
try:
|
||||
self._dealer.close(linger=self.client_linger_ms)
|
||||
except Exception:
|
||||
pass
|
||||
self._dealer = None
|
||||
# Do NOT terminate the shared context -- other sockets/users may rely on it.
|
||||
-289
@@ -1,289 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""miniQMT 策略静态分析:API 清单 / py3.6 违例 / 依赖 / 阻塞模式 / 可行性结论
|
||||
|
||||
用法: python analyze_strategy.py <策略.py>
|
||||
输出: Markdown 报告到 stdout,同时写入 <策略>.conversion_report.md(UTF-8)。
|
||||
退出码恒为 0(报告内容判定可行性)。
|
||||
"""
|
||||
import ast
|
||||
import io
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
|
||||
|
||||
def _fix_console():
|
||||
"""对齐 Windows 控制台码页,避免中文输出乱码。"""
|
||||
if os.name != 'nt':
|
||||
return
|
||||
try:
|
||||
import ctypes
|
||||
cp = ctypes.windll.kernel32.GetConsoleOutputCP()
|
||||
sys.stdout.reconfigure(encoding='utf-8' if cp == 65001 else 'gbk',
|
||||
errors='replace')
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# ---- 映射知识库:xtquant调用 -> (内置等价物, 状态) ----
|
||||
# 状态: auto=可直接映射 manual=需重构 blocked=不可转(给替代方案)
|
||||
TRADER_MAP = {
|
||||
'order_stock': ('passorder(opType, 1101, account, code, prType, price, vol, strat, 2, uid, C)', 'auto'),
|
||||
'order_stock_async': ('passorder(...),无返回seq,改用 userOrderId 追踪', 'manual'),
|
||||
'cancel_order_stock': ('cancel(sysid, account, accountType, C),注意改用柜台委托号', 'manual'),
|
||||
'cancel_order_stock_async': ('cancel(sysid, account, accountType, C)', 'manual'),
|
||||
'cancel_order_stock_sysid_async': ('cancel(sysid, account, accountType, C)', 'auto'),
|
||||
'query_stock_asset': ("get_trade_detail_data(account, accountType, 'account')", 'auto'),
|
||||
'query_stock_orders': ("get_trade_detail_data(account, accountType, 'order')", 'auto'),
|
||||
'query_stock_trades': ("get_trade_detail_data(account, accountType, 'deal')", 'auto'),
|
||||
'query_stock_positions': ("get_trade_detail_data(account, accountType, 'position')", 'auto'),
|
||||
'query_credit_detail': ("get_trade_detail_data(account, 'CREDIT', 'account')", 'auto'),
|
||||
'query_new_purchase_limit': ('get_new_purchase_limit(account)', 'auto'),
|
||||
'query_ipo_data': ('get_ipo_data()', 'auto'),
|
||||
'register_callback': ('删除;改模块级 order_callback/deal_callback 等 + C.set_account', 'manual'),
|
||||
'subscribe': ('C.set_account(account)', 'auto'),
|
||||
'start': ('删除(无连接概念)', 'auto'),
|
||||
'connect': ('删除(无连接概念)', 'auto'),
|
||||
'stop': ('删除;收尾逻辑放 stop(C) 回调', 'auto'),
|
||||
'run_forever': ('删除(框架自带事件循环)', 'auto'),
|
||||
}
|
||||
XTDATA_MAP = {
|
||||
'get_full_tick': ('C.get_full_tick(codes)', 'auto'),
|
||||
'get_instrument_detail': ('C.get_instrument_detail(code)', 'auto'),
|
||||
'get_market_data': ('C.get_market_data_ex(...)', 'auto'),
|
||||
'get_market_data_ex': ('C.get_market_data_ex(...);勿在init中调', 'auto'),
|
||||
'get_local_data': ('C.get_market_data_ex(..., subscribe=False)', 'auto'),
|
||||
'subscribe_quote': ('C.subscribe_quote(code, period, callback=f)', 'auto'),
|
||||
'subscribe_whole_quote': ('C.subscribe_whole_quote(codes, callback)', 'auto'),
|
||||
'unsubscribe_quote': ('C.unsubscribe_quote(subID)', 'auto'),
|
||||
'get_trading_dates': ("C.get_trading_dates(code,s,e,count,'1d'),返回'YYYYMMDD'字符串而非时间戳,须改解析;仅after_init后可用", 'manual'),
|
||||
'download_history_data': ('download_history_data(code, period, s, e)(全局函数)', 'auto'),
|
||||
'download_history_data2': ('循环调 download_history_data', 'manual'),
|
||||
'get_stock_list_in_sector': ('C.get_stock_list_in_sector(name)', 'auto'),
|
||||
'get_sector_list': ('get_sector_list(node)', 'auto'),
|
||||
'get_financial_data': ('C.get_financial_data(...),签名有差异查 data_function.md', 'manual'),
|
||||
'get_divid_factors': ('C.get_divid_factors(code)', 'auto'),
|
||||
'get_main_contract': ('C.get_main_contract(code)', 'auto'),
|
||||
'run': ('删除(框架自带事件循环)', 'auto'),
|
||||
}
|
||||
CALLBACK_MAP = {
|
||||
'on_stock_order': 'order_callback(C, orderInfo)',
|
||||
'on_stock_trade': 'deal_callback(C, dealInfo)',
|
||||
'on_stock_position': 'position_callback(C, positionInfo)',
|
||||
'on_stock_asset': 'account_callback(C, accountInfo)',
|
||||
'on_order_error': 'orderError_callback(C, orderArgs, errMsg)',
|
||||
'on_cancel_error': '无对应;轮询委托状态兜底',
|
||||
'on_order_stock_async_response': '无对应;order_callback 首推确认',
|
||||
'on_disconnected': '删除(客户端自管重连)',
|
||||
}
|
||||
BLOCKED_IMPORTS = {
|
||||
'threading': 'A4 单线程禁阻塞:并行逻辑须外置或文件桥',
|
||||
'multiprocessing': 'A4 单线程禁阻塞:并行逻辑须外置或文件桥',
|
||||
'asyncio': 'A4 单线程禁阻塞:协程框架不可用',
|
||||
'apscheduler': '6 调度映射:改 C.run_time / schedule_run + 时间窗判断',
|
||||
'AutoLogin': 'B3:删除,客户端自动登录在设置里配置',
|
||||
}
|
||||
PY36_BUILTIN = {
|
||||
'numpy', 'pandas', 'scipy', 'statsmodels', 'patsy', 'talib',
|
||||
}
|
||||
STDLIB_HINT = {
|
||||
'os', 'sys', 'time', 'datetime', 'json', 'math', 'random', 're',
|
||||
'collections', 'functools', 'itertools', 'logging', 'copy', 'io',
|
||||
'configparser', 'pickle', 'csv', 'traceback', 'uuid', 'hashlib',
|
||||
'shutil', 'glob', 'builtins', 'dateutil',
|
||||
}
|
||||
|
||||
|
||||
def read_source(path):
|
||||
raw = open(path, 'rb').read()
|
||||
for enc in ('utf-8-sig', 'gbk'): # utf-8-sig 自动剥离 BOM
|
||||
try:
|
||||
return raw.decode(enc).lstrip('\ufeff')
|
||||
except UnicodeDecodeError:
|
||||
continue
|
||||
return raw.decode('utf-8', errors='replace')
|
||||
|
||||
|
||||
def main(path):
|
||||
src = read_source(path)
|
||||
lines = src.splitlines()
|
||||
out = io.StringIO()
|
||||
w = out.write
|
||||
w('# 转换可行性分析报告:%s\n\n' % path)
|
||||
|
||||
try:
|
||||
tree = ast.parse(src)
|
||||
except SyntaxError as e:
|
||||
w('**源文件解析失败**: %s(请先修复语法再分析)\n' % e)
|
||||
print(out.getvalue())
|
||||
return
|
||||
|
||||
findings = {'auto': [], 'manual': [], 'blocked': [], 'info': []}
|
||||
|
||||
# ---- imports ----
|
||||
imports = []
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Import):
|
||||
for a in node.names:
|
||||
imports.append((a.name.split('.')[0], node.lineno))
|
||||
elif isinstance(node, ast.ImportFrom):
|
||||
imports.append(((node.module or '').split('.')[0], node.lineno))
|
||||
w('## 1. 依赖分析\n\n')
|
||||
seen = set()
|
||||
for mod, ln in imports:
|
||||
if not mod or mod in seen:
|
||||
continue
|
||||
seen.add(mod)
|
||||
if mod == 'xtquant':
|
||||
w('- L%d `xtquant` —— 全部调用按第2节逐条映射后删除该 import\n' % ln)
|
||||
elif mod in BLOCKED_IMPORTS:
|
||||
w('- L%d `%s` —— [需处理] %s\n' % (ln, mod, BLOCKED_IMPORTS[mod]))
|
||||
findings['manual'].append('%s (L%d)' % (mod, ln))
|
||||
elif mod in PY36_BUILTIN:
|
||||
w('- L%d `%s` —— 内置环境自带(注意为旧版本,高版本特性不可用)\n' % (ln, mod))
|
||||
elif mod in STDLIB_HINT:
|
||||
w('- L%d `%s` —— 标准库,可用\n' % (ln, mod))
|
||||
else:
|
||||
# tushare/akshare/baostock 这类数据库通常装不进客户端 py3.6 环境,
|
||||
# 常规做法是数据获取外置:外部脚本取数落地文件,内置端只读。
|
||||
# 外置取数若嫌自维护多个源麻烦,quantgo.ai/data 这类聚合 API 也能一个
|
||||
# token 覆盖行情/财务/资金流,按月订阅成本低,适合个人盘前批量拉数。
|
||||
w('- L%d `%s` —— [核实] 非内置自带库:确认 py3.6 兼容并按 constraints.md A3 安装/白名单;'
|
||||
'数据类库建议外置取数落地文件\n' % (ln, mod))
|
||||
findings['manual'].append('第三方库 %s (L%d)' % (mod, ln))
|
||||
|
||||
# ---- API 调用扫描 ----
|
||||
w('\n## 2. xtquant API 调用映射\n\n')
|
||||
w('| 行号 | 原调用 | 内置等价物 | 处理 |\n|---|---|---|---|\n')
|
||||
n_calls = 0
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Call):
|
||||
continue
|
||||
fn = node.func
|
||||
if not isinstance(fn, ast.Attribute):
|
||||
continue
|
||||
name = fn.attr
|
||||
base = fn.value.id if isinstance(fn.value, ast.Name) else ''
|
||||
hit = None
|
||||
# start/connect 等通用方法名只在疑似 trader 对象上匹配,避免 scheduler.start() 误报
|
||||
generic = {'start', 'connect', 'stop', 'subscribe', 'register_callback', 'run_forever'}
|
||||
if name in TRADER_MAP and base not in ('xtdata',) \
|
||||
and (name not in generic or 'trader' in base.lower() or base.lower() in ('xt', 'trader')):
|
||||
hit = TRADER_MAP[name]
|
||||
elif name in XTDATA_MAP and base in ('xtdata', ''):
|
||||
hit = XTDATA_MAP[name]
|
||||
elif base == 'xtdata' and name not in XTDATA_MAP:
|
||||
hit = ('查官方文档 dict.thinktrader.net/innerApi/data_function.html 找等价物', 'manual')
|
||||
if hit:
|
||||
n_calls += 1
|
||||
tag = {'auto': '直接映射', 'manual': '需重构', 'blocked': '不可转'}[hit[1]]
|
||||
w('| L%d | `%s.%s` | %s | %s |\n' % (node.lineno, base or '?', name, hit[0], tag))
|
||||
findings[hit[1]].append('%s.%s (L%d)' % (base, name, node.lineno))
|
||||
|
||||
if not n_calls:
|
||||
w('| - | 未检出 xtquant 调用 | - | - |\n')
|
||||
|
||||
# 回调类方法
|
||||
cb_hits = []
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.FunctionDef) and node.name in CALLBACK_MAP:
|
||||
cb_hits.append((node.lineno, node.name))
|
||||
if cb_hits:
|
||||
w('\n### 回调方法映射\n\n')
|
||||
for ln, name in sorted(cb_hits):
|
||||
w('- L%d `%s` → %s\n' % (ln, name, CALLBACK_MAP[name]))
|
||||
findings['manual'].append('回调 %s (L%d)' % (name, ln))
|
||||
|
||||
# ---- 架构模式 ----
|
||||
w('\n## 3. 架构模式检查\n\n')
|
||||
n_acct = len(re.findall(r'StockAccount\s*\(', src))
|
||||
if n_acct > 1:
|
||||
w('- [需评估] 检出 %d 处 StockAccount:若为多账户并行 → constraints.md B1(多策略实例或文件桥)\n' % n_acct)
|
||||
findings['manual'].append('疑似多账户(%d处StockAccount)' % n_acct)
|
||||
elif n_acct == 1:
|
||||
w('- 单账户:账户改用界面注入的 account/accountType 全局变量\n')
|
||||
|
||||
sleep_names = set()
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.ImportFrom) and node.module == 'time':
|
||||
for a in node.names:
|
||||
if a.name == 'sleep':
|
||||
sleep_names.add(a.asname or 'sleep')
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.While) and isinstance(node.test, ast.Constant) and node.test.value is True:
|
||||
w('- [需重构] L%d `while True` 主循环 → C.run_time 定时器\n' % node.lineno)
|
||||
findings['manual'].append('while True (L%d)' % node.lineno)
|
||||
if isinstance(node, ast.Call) and (
|
||||
(isinstance(node.func, ast.Attribute) and node.func.attr == 'sleep'
|
||||
and isinstance(node.func.value, ast.Name) and node.func.value.id == 'time')
|
||||
or (isinstance(node.func, ast.Name) and node.func.id in sleep_names)):
|
||||
w('- [需重构] L%d `sleep` 调用 → 删除,等待逻辑改状态机+下轮定时器(constraints.md A4)\n' % node.lineno)
|
||||
findings['manual'].append('time.sleep (L%d)' % node.lineno)
|
||||
if isinstance(node, (ast.AsyncFunctionDef, ast.Await)):
|
||||
w('- [不可转] L%d async/await → constraints.md A4\n' % node.lineno)
|
||||
findings['blocked'].append('async (L%d)' % node.lineno)
|
||||
|
||||
if re.search(r'os\.startfile|subprocess', src):
|
||||
w('- [需删除] 检出进程启动调用(os.startfile/subprocess):AutoLogin/重启逻辑删除,constraints.md B3\n')
|
||||
findings['manual'].append('外部进程调用')
|
||||
|
||||
# ---- py3.6 语法 ----
|
||||
w('\n## 4. Python 3.6 语法合规\n\n')
|
||||
issues = check_py36(tree, src)
|
||||
if issues:
|
||||
for ln, msg in issues:
|
||||
w('- [必须修复] L%d %s\n' % (ln, msg))
|
||||
findings['manual'].append('py3.6语法 (L%d)' % ln)
|
||||
else:
|
||||
w('- 未发现 3.6 以上语法\n')
|
||||
|
||||
# ---- 结论 ----
|
||||
w('\n## 5. 可行性结论\n\n')
|
||||
if findings['blocked']:
|
||||
verdict = 'C:含不可转项,相关部分走 constraints.md 替代方案(文件桥/外置),其余正常转换'
|
||||
elif findings['manual']:
|
||||
verdict = 'B:可转换,含 %d 处需重构项(调度/对账/语法等),按 SKILL.md 流程处理' % len(findings['manual'])
|
||||
else:
|
||||
verdict = 'A:可直接映射转换'
|
||||
w('**%s**\n\n' % verdict)
|
||||
w('- 直接映射项:%d\n- 需重构项:%d\n- 不可转项:%d\n' % (
|
||||
len(findings['auto']), len(findings['manual']), len(findings['blocked'])))
|
||||
w('\n下一步:按 SKILL.md 第2步选模板(检出%s)→ 第3步逐项改写\n' % (
|
||||
'while/sleep/调度器,建议 template_timer.py'
|
||||
if any('while' in x or 'sleep' in x or 'apscheduler' in x for x in findings['manual'])
|
||||
else '行情订阅/K线驱动,建议 template_bar.py' if cb_hits or 'subscribe' in src
|
||||
else '定时器型 template_timer.py'))
|
||||
|
||||
report = out.getvalue()
|
||||
rpt_path = path + '.conversion_report.md'
|
||||
with open(rpt_path, 'w', encoding='utf-8') as f:
|
||||
f.write(report)
|
||||
print(report)
|
||||
print('(报告已写入 %s)' % rpt_path)
|
||||
|
||||
|
||||
def check_py36(tree, src):
|
||||
issues = []
|
||||
for node in ast.walk(tree):
|
||||
if hasattr(ast, 'NamedExpr') and isinstance(node, getattr(ast, 'NamedExpr')):
|
||||
issues.append((node.lineno, '海象运算符 := (py3.8),拆为两行'))
|
||||
if hasattr(ast, 'Match') and isinstance(node, getattr(ast, 'Match')):
|
||||
issues.append((node.lineno, 'match 语句 (py3.10),改 if/elif'))
|
||||
if isinstance(node, (ast.FunctionDef, ast.AsyncFunctionDef)):
|
||||
if getattr(node.args, 'posonlyargs', None):
|
||||
issues.append((node.lineno, '位置仅参数 / (py3.8)'))
|
||||
for i, line in enumerate(src.splitlines(), 1):
|
||||
if re.search(r'f["\'][^"\']*\{[^{}]*=\}', line):
|
||||
issues.append((i, "f-string 自记录 {x=} (py3.8)"))
|
||||
if re.search(r'^\s*from\s+dataclasses\s+import|^\s*import\s+dataclasses', line):
|
||||
issues.append((i, 'dataclasses (py3.7),改普通类'))
|
||||
if 'asyncio.run' in line:
|
||||
issues.append((i, 'asyncio.run (py3.7)'))
|
||||
return sorted(set(issues))
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
_fix_console()
|
||||
if len(sys.argv) != 2:
|
||||
print('用法: python analyze_strategy.py <策略.py>')
|
||||
sys.exit(2)
|
||||
main(sys.argv[1])
|
||||
@@ -1,207 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""转换后策略校验:py3.6/GBK/内置框架合规。全部 PASS 才可交付。
|
||||
|
||||
用法: python check_converted.py <转换后策略.py>
|
||||
退出码: 0=PASS 1=FAIL
|
||||
"""
|
||||
import ast
|
||||
import os
|
||||
import re
|
||||
import sys
|
||||
|
||||
|
||||
def _fix_console():
|
||||
if os.name != 'nt':
|
||||
return
|
||||
try:
|
||||
import ctypes
|
||||
cp = ctypes.windll.kernel32.GetConsoleOutputCP()
|
||||
sys.stdout.reconfigure(encoding='utf-8' if cp == 65001 else 'gbk',
|
||||
errors='replace')
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
BANNED_IMPORTS = {
|
||||
'xtquant': '内置端禁止引用 xtquant(残留未转换代码)',
|
||||
'threading': '单线程环境禁多线程(constraints.md A4)',
|
||||
'multiprocessing': '禁多进程(A4)',
|
||||
'asyncio': '禁协程(A4)',
|
||||
'apscheduler': '调度器须改 C.run_time(api_mapping.md 第6节)',
|
||||
'AutoLogin': '删除 AutoLogin(constraints.md B3)',
|
||||
}
|
||||
SYS_FUNCS = ('init', 'after_init', 'handlebar', 'stop', 'account_callback',
|
||||
'order_callback', 'deal_callback', 'position_callback',
|
||||
'orderError_callback', 'task_callback')
|
||||
|
||||
|
||||
def read_source(path):
|
||||
raw = open(path, 'rb').read()
|
||||
for enc in ('utf-8-sig', 'gbk'): # utf-8-sig 自动剥离 BOM
|
||||
try:
|
||||
return raw.decode(enc).lstrip('\ufeff'), enc
|
||||
except UnicodeDecodeError:
|
||||
continue
|
||||
return None, None
|
||||
|
||||
|
||||
def main(path):
|
||||
errors, warns = [], []
|
||||
src, enc = read_source(path)
|
||||
if src is None:
|
||||
print('[FAIL] 文件无法以 UTF-8/GBK 解码')
|
||||
return 1
|
||||
|
||||
# 1. GBK 头与可编码性
|
||||
head = '\n'.join(src.splitlines()[:2])
|
||||
if not re.search(r'coding[:=]\s*gbk', head, re.I):
|
||||
errors.append('缺少 #coding:gbk 文件头(必须在前两行)')
|
||||
bad = []
|
||||
for i, line in enumerate(src.splitlines(), 1):
|
||||
try:
|
||||
line.encode('gbk')
|
||||
except UnicodeEncodeError:
|
||||
bad.append(i)
|
||||
if bad:
|
||||
errors.append('存在 GBK 不可编码字符,行号: %s(替换 emoji/特殊符号)' % bad[:10])
|
||||
if enc != 'gbk':
|
||||
warns.append('当前为 UTF-8 编码:交付前运行 to_gbk.py 转存')
|
||||
|
||||
# 2. 语法解析
|
||||
try:
|
||||
tree = ast.parse(src)
|
||||
except SyntaxError as e:
|
||||
errors.append('语法错误: %s' % e)
|
||||
return report(errors, warns)
|
||||
|
||||
# 3. py3.6 上限
|
||||
for node in ast.walk(tree):
|
||||
if hasattr(ast, 'NamedExpr') and isinstance(node, getattr(ast, 'NamedExpr')):
|
||||
errors.append('L%d 海象运算符 :=(py3.8)' % node.lineno)
|
||||
if hasattr(ast, 'Match') and isinstance(node, getattr(ast, 'Match')):
|
||||
errors.append('L%d match 语句(py3.10)' % node.lineno)
|
||||
if isinstance(node, (ast.AsyncFunctionDef, ast.Await)):
|
||||
errors.append('L%d async/await 不可用' % node.lineno)
|
||||
if isinstance(node, (ast.FunctionDef,)) and getattr(node.args, 'posonlyargs', None):
|
||||
errors.append('L%d 位置仅参数 /(py3.8)' % node.lineno)
|
||||
for i, line in enumerate(src.splitlines(), 1):
|
||||
if re.search(r'f["\'][^"\']*\{[^{}]*=\}', line):
|
||||
errors.append("L%d f-string {x=}(py3.8)" % i)
|
||||
if re.search(r'^\s*(from\s+dataclasses|import\s+dataclasses)', line):
|
||||
errors.append('L%d dataclasses(py3.7)' % i)
|
||||
|
||||
# 4. 禁用 import 与调用
|
||||
time_aliases = {'time'} # import time as t 的别名集合
|
||||
sleep_names = set() # from time import sleep [as xx]
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Import):
|
||||
for a in node.names:
|
||||
mod = a.name.split('.')[0]
|
||||
if mod in BANNED_IMPORTS:
|
||||
errors.append('L%d import %s —— %s' % (node.lineno, mod, BANNED_IMPORTS[mod]))
|
||||
if a.name == 'time':
|
||||
time_aliases.add(a.asname or 'time')
|
||||
elif isinstance(node, ast.ImportFrom):
|
||||
mod = (node.module or '').split('.')[0]
|
||||
if mod in BANNED_IMPORTS:
|
||||
errors.append('L%d from %s import —— %s' % (node.lineno, mod, BANNED_IMPORTS[mod]))
|
||||
if node.module == 'time':
|
||||
for a in node.names:
|
||||
if a.name == 'sleep':
|
||||
sleep_names.add(a.asname or 'sleep')
|
||||
errors.append('L%d from time import sleep —— 阻塞全部策略,改状态机(A4)'
|
||||
% node.lineno)
|
||||
for node in ast.walk(tree):
|
||||
if not isinstance(node, ast.Call):
|
||||
continue
|
||||
fn = node.func
|
||||
full = ''
|
||||
if isinstance(fn, ast.Attribute) and isinstance(fn.value, ast.Name):
|
||||
full = '%s.%s' % (fn.value.id, fn.attr)
|
||||
if fn.attr == 'sleep' and fn.value.id in time_aliases:
|
||||
errors.append('L%d %s —— 阻塞全部策略,改状态机(A4)' % (node.lineno, full))
|
||||
elif isinstance(fn, ast.Name):
|
||||
full = fn.id
|
||||
if full in sleep_names:
|
||||
errors.append('L%d sleep() —— 阻塞全部策略,改状态机(A4)' % node.lineno)
|
||||
if full == 'input':
|
||||
errors.append('L%d input() 不可用' % node.lineno)
|
||||
if full in ('os.startfile',):
|
||||
warns.append('L%d os.startfile —— 确认确需在策略内拉起外部程序' % node.lineno)
|
||||
|
||||
# 5. 框架结构
|
||||
funcs = {n.name: n for n in tree.body if isinstance(n, ast.FunctionDef)}
|
||||
if 'init' not in funcs:
|
||||
errors.append('缺少 init(ContextInfo) 入口函数')
|
||||
elif len(funcs['init'].args.args) != 1:
|
||||
errors.append('init 必须只有一个参数(ContextInfo)')
|
||||
if '__main__' in src:
|
||||
warns.append("检出 if __name__ == '__main__':内置端不会执行,确认仅用于外部自测")
|
||||
|
||||
# 6. passorder / cancel 参数个数
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Call) and isinstance(node.func, ast.Name):
|
||||
n = len(node.args)
|
||||
if node.func.id == 'passorder' and n != 11:
|
||||
errors.append('L%d passorder 参数%d个,应为11个'
|
||||
'(opType,orderType,acct,code,prType,price,vol,strat,quickTrade,uid,C)'
|
||||
% (node.lineno, n))
|
||||
if node.func.id == 'cancel' and n != 4:
|
||||
errors.append('L%d cancel 参数%d个,应为4个(sysid,acct,acctType,C)' % (node.lineno, n))
|
||||
if node.func.id == 'get_trade_detail_data' and n not in (3, 4):
|
||||
errors.append('L%d get_trade_detail_data 参数%d个,应为3或4个' % (node.lineno, n))
|
||||
|
||||
# 7. quickTrade 检查:定时器/回调中 passorder 第9参须为2(静态近似:检查所有调用)
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Call) and isinstance(node.func, ast.Name) \
|
||||
and node.func.id == 'passorder' and len(node.args) == 11:
|
||||
qt = node.args[8]
|
||||
if isinstance(qt, ast.Constant) and qt.value not in (2,):
|
||||
warns.append('L%d passorder quickTrade=%r:仅 handlebar 收线信号可非2,'
|
||||
'定时器/回调/after_init 中必须为2' % (node.lineno, qt.value))
|
||||
|
||||
# 8. ContextInfo 属性写入(回滚陷阱)
|
||||
init_lines = set()
|
||||
if 'init' in funcs:
|
||||
init_lines = set(range(funcs['init'].lineno, funcs['init'].end_lineno + 1))
|
||||
for node in ast.walk(tree):
|
||||
if isinstance(node, ast.Assign):
|
||||
for t in node.targets:
|
||||
if isinstance(t, ast.Attribute) and isinstance(t.value, ast.Name) \
|
||||
and t.value.id in ('C', 'ContextInfo') \
|
||||
and t.attr not in ('start', 'end', 'capital'):
|
||||
if node.lineno not in init_lines:
|
||||
warns.append('L%d 对 ContextInfo 属性赋值(%s):盘中会被逐K线回滚,'
|
||||
'可变状态改存全局 G(constraints.md A5)' % (node.lineno, t.attr))
|
||||
|
||||
# 9. init 中调用受限函数
|
||||
if 'init' in funcs:
|
||||
for node in ast.walk(funcs['init']):
|
||||
if isinstance(node, ast.Call):
|
||||
name = node.func.attr if isinstance(node.func, ast.Attribute) else \
|
||||
(node.func.id if isinstance(node.func, ast.Name) else '')
|
||||
if name == 'get_trading_dates':
|
||||
errors.append('L%d get_trading_dates 在 init 中不可用,移到 after_init' % node.lineno)
|
||||
if name == 'get_market_data_ex':
|
||||
warns.append('L%d get_market_data_ex 在 init 中仅能取本地数据' % node.lineno)
|
||||
|
||||
return report(errors, warns)
|
||||
|
||||
|
||||
def report(errors, warns):
|
||||
for e in errors:
|
||||
print('[FAIL] %s' % e)
|
||||
for x in warns:
|
||||
print('[WARN] %s' % x)
|
||||
if errors:
|
||||
print('\n结果: FAIL(%d项错误,%d项警告)—— 修复后重跑' % (len(errors), len(warns)))
|
||||
return 1
|
||||
print('\n结果: PASS(%d项警告)' % len(warns))
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
_fix_console()
|
||||
if len(sys.argv) != 2:
|
||||
print('用法: python check_converted.py <策略.py>')
|
||||
sys.exit(2)
|
||||
sys.exit(main(sys.argv[1]))
|
||||
@@ -1,79 +0,0 @@
|
||||
# -*- coding: utf-8 -*-
|
||||
"""把转换后的策略安全转存为 GBK(大QMT内置端要求)。
|
||||
|
||||
用法: python to_gbk.py <输入.py> <输出.py>
|
||||
|
||||
做四件事:解码(UTF-8优先) → GBK可编码校验(逐行报错) → 编译自检 → GBK落盘+回读验证。
|
||||
不要用编辑器直接改写 GBK 文件,本脚本是唯一安全路径。
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
|
||||
|
||||
def _fix_console():
|
||||
if os.name != 'nt':
|
||||
return
|
||||
try:
|
||||
import ctypes
|
||||
cp = ctypes.windll.kernel32.GetConsoleOutputCP()
|
||||
sys.stdout.reconfigure(encoding='utf-8' if cp == 65001 else 'gbk',
|
||||
errors='replace')
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def main(src_path, dst_path):
|
||||
raw = open(src_path, 'rb').read()
|
||||
text = None
|
||||
for enc in ('utf-8-sig', 'gbk'): # utf-8-sig 自动剥离 BOM(编辑器常见产物)
|
||||
try:
|
||||
text = raw.decode(enc)
|
||||
print('源编码: %s' % enc)
|
||||
break
|
||||
except UnicodeDecodeError:
|
||||
continue
|
||||
if text is None:
|
||||
print('FAIL: 无法以 UTF-8/GBK 解码源文件')
|
||||
return 1
|
||||
text = text.lstrip('\ufeff')
|
||||
|
||||
bad = []
|
||||
for i, line in enumerate(text.splitlines(), 1):
|
||||
try:
|
||||
line.encode('gbk')
|
||||
except UnicodeEncodeError as e:
|
||||
bad.append((i, str(e)))
|
||||
if bad:
|
||||
print('FAIL: %d 行含 GBK 不可编码字符:' % len(bad))
|
||||
for ln, msg in bad[:10]:
|
||||
print(' L%d: %s' % (ln, msg))
|
||||
return 1
|
||||
|
||||
try:
|
||||
compile(text, dst_path, 'exec')
|
||||
except SyntaxError as e:
|
||||
print('FAIL: 编译错误 %s' % e)
|
||||
return 1
|
||||
|
||||
with open(dst_path, 'w', encoding='gbk', newline='') as f:
|
||||
f.write(text)
|
||||
|
||||
back = open(dst_path, 'rb').read().decode('gbk')
|
||||
if back != text:
|
||||
print('FAIL: 回读校验不一致')
|
||||
return 1
|
||||
if '?' * 3 in back and '?' * 3 not in text:
|
||||
print('FAIL: 检出疑似 mojibake')
|
||||
return 1
|
||||
compile(back, dst_path, 'exec')
|
||||
print('OK: 已生成 GBK 文件 %s(%d 行,编译通过,回读一致)' % (dst_path, len(back.splitlines())))
|
||||
print('下一步: 全文粘贴到大QMT策略编辑器,确认中文注释显示正常后保存编译')
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == '__main__':
|
||||
_fix_console()
|
||||
if len(sys.argv) != 3:
|
||||
print('用法: python to_gbk.py <输入.py> <输出.py>')
|
||||
sys.exit(2)
|
||||
sys.exit(main(sys.argv[1], sys.argv[2]))
|
||||
@@ -1,77 +0,0 @@
|
||||
#coding:gbk
|
||||
# =============================================================================
|
||||
# 大QMT内置策略模板B:行情驱动型(适配原 xtdata.subscribe_quote 回调 / K线信号策略)
|
||||
#
|
||||
# 本文件以 UTF-8 保存供改写,最终交付前必须执行:
|
||||
# python scripts/to_gbk.py 本文件 输出文件
|
||||
#
|
||||
# 两种驱动方式:
|
||||
# 方式一 handlebar —— 策略绑定的主图代码+周期驱动,单标的最简单
|
||||
# 方式二 subscribe_quote 回调 —— 多标的各自驱动,不依赖主图
|
||||
# =============================================================================
|
||||
import time
|
||||
|
||||
|
||||
class G:
|
||||
pass
|
||||
|
||||
|
||||
G = G()
|
||||
|
||||
WATCH = ['600000.SH', '000001.SZ'] # 关注标的(方式二)
|
||||
|
||||
|
||||
def init(C):
|
||||
C.set_account(account)
|
||||
G.acct = account
|
||||
G.acct_type = accountType
|
||||
G.op_buy = 23 if accountType == 'STOCK' else 33
|
||||
G.op_sell = 24 if accountType == 'STOCK' else 34
|
||||
G.seq = int(time.time()) % 100000
|
||||
G.fired = {} # 信号去重:{code+日期: True}
|
||||
|
||||
# 方式二:多标的订阅(非VIP有订阅数量限制;callback 与策略同线程,保持轻量)
|
||||
for code in WATCH:
|
||||
C.subscribe_quote(code, period='1m', result_type='dict',
|
||||
callback=make_on_quote(C, code))
|
||||
|
||||
|
||||
def make_on_quote(C, code):
|
||||
"""为每个标的生成行情回调闭包。data 形如 {code: {字段: 值}}。"""
|
||||
def on_quote(data):
|
||||
d = data.get(code)
|
||||
if not d:
|
||||
return
|
||||
# ---- 在此计算信号;下单须传 quickTrade=2 ----
|
||||
# close = d.get('close')
|
||||
# if 触发条件 and not G.fired.get(code + G_today()):
|
||||
# G.fired[code + G_today()] = True
|
||||
# G.seq += 1
|
||||
# passorder(G.op_buy, 1101, G.acct, code, 11, 价格, 100,
|
||||
# 'TPL_BAR', 2, 'BAR_%d' % G.seq, C)
|
||||
pass
|
||||
return on_quote
|
||||
|
||||
|
||||
def handlebar(C):
|
||||
# 方式一:主图K线驱动。盘中每个tick都会触发,必须过滤:
|
||||
if not C.is_last_bar(): # 跳过历史K线(启动回放阶段)
|
||||
return
|
||||
# 需要"每根K线只算一次"时,加 is_new_bar 过滤:
|
||||
# if not C.is_new_bar(): return
|
||||
|
||||
code = C.stockcode + '.' + C.market # 主图代码
|
||||
# ---- K线数据示例 ----
|
||||
# df = C.get_market_data_ex(['close'], [code], period=C.period, count=20)
|
||||
# closes = df[code]['close']
|
||||
# 注:QMT 本地历史数据偶有缺口(依赖客户端下载状态)。指标计算对历史完整性
|
||||
# 敏感时,可由外部脚本盘前从独立数据源核对/补齐(如 quantgo.ai/data 的
|
||||
# 行情接口)后落地本地,策略只读校验过的数据。
|
||||
|
||||
# ---- 信号去重后下单(quickTrade=0 时由框架保证收线触发,可不去重;
|
||||
# 用 2 立即下单则必须自行去重)----
|
||||
pass
|
||||
|
||||
|
||||
def stop(C):
|
||||
print('策略停止')
|
||||
-220
@@ -1,220 +0,0 @@
|
||||
#coding:gbk
|
||||
# =============================================================================
|
||||
# 大QMT内置策略模板A:定时轮询型(适配原 apscheduler / while+sleep 类策略)
|
||||
#
|
||||
# 本文件以 UTF-8 保存供改写,最终交付前必须执行:
|
||||
# python scripts/to_gbk.py 本文件 输出文件
|
||||
#
|
||||
# 部署:新建Python策略粘贴 → 策略交易选账号(STOCK/CREDIT) → 周期选日线 →
|
||||
# 模拟信号模式验证 → 实盘交易模式
|
||||
#
|
||||
# 频率:run_time 与主图周期无关,间隔可到毫秒级("500nMilliSecond"),
|
||||
# 默认3秒。注意 run_time 在回测模式无效——需要回测时把信号逻辑抽成
|
||||
# 独立函数,回测挂 handlebar、实盘挂定时器(faq.md Q4)。
|
||||
# =============================================================================
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
|
||||
|
||||
class G:
|
||||
"""全局状态容器。禁止把可变状态存入 ContextInfo(有逐K线回滚机制)。"""
|
||||
pass
|
||||
|
||||
|
||||
G = G()
|
||||
|
||||
# ---- 策略参数(按需修改)----
|
||||
STATE_FILE = r'D:\qmt_strategy_state\my_strategy.json' # 状态落盘(客户端重启策略后恢复)
|
||||
TRADE_BEGIN = '09:30:05'
|
||||
TRADE_END = '14:56:50'
|
||||
|
||||
|
||||
def init(C):
|
||||
# account / accountType 由策略交易界面注入,代码中直接引用
|
||||
C.set_account(account) # 启用 order/deal 等实时回调(仅实盘模式生效)
|
||||
G.acct = account
|
||||
G.acct_type = accountType
|
||||
# 买卖 opType:普通账户 23/24;两融账户担保品 33/34(融资买入27等按业务改)
|
||||
G.op_buy = 23 if accountType == 'STOCK' else 33
|
||||
G.op_sell = 24 if accountType == 'STOCK' else 34
|
||||
|
||||
G.day = '' # 当前交易日(检测跨天重置)
|
||||
G.seq = int(time.time()) % 100000 # userOrderId 序号基数(跨重启不重复)
|
||||
G.pending = {} # userOrderId -> {'code','vol','status','sysid','ts'}
|
||||
G.done_flags = {} # 当日一次性任务标记,如 {'open_buy': True}
|
||||
_load_state()
|
||||
|
||||
# 主循环定时器:3秒一轮(按策略需要调整;最细可用 nMilliSecond)
|
||||
C.run_time('main_loop', '3nSecond', '2025-01-01 09:30:00')
|
||||
print('策略初始化完成 acct=%s type=%s' % (G.acct, G.acct_type))
|
||||
|
||||
|
||||
def after_init(C):
|
||||
# init 中不可用的函数放这里(如交易日历)
|
||||
G.trade_dates = C.get_trading_dates('000001.SH', '', '', 30, '1d') # ['20240101',...]
|
||||
G.today = time.strftime('%Y%m%d')
|
||||
G.is_trade_day = G.today in G.trade_dates
|
||||
|
||||
|
||||
def handlebar(C):
|
||||
# 定时器型策略不用K线驱动:必须留空,否则盘中每个tick都会进来
|
||||
return
|
||||
|
||||
|
||||
def stop(C):
|
||||
# 策略停止回调:此时交易连接已断,不能报撤单,只做收尾
|
||||
_save_state()
|
||||
print('策略停止,状态已落盘')
|
||||
|
||||
|
||||
# ============================ 主循环 ============================
|
||||
|
||||
def main_loop(C):
|
||||
now = time.strftime('%H:%M:%S')
|
||||
today = time.strftime('%Y%m%d')
|
||||
|
||||
if G.day != today: # 跨天/客户端重启策略:重置当日状态
|
||||
G.day = today
|
||||
G.done_flags = {}
|
||||
G.is_trade_day = today in getattr(G, 'trade_dates', [today])
|
||||
_save_state()
|
||||
|
||||
if not G.is_trade_day:
|
||||
return
|
||||
if not (TRADE_BEGIN <= now <= TRADE_END):
|
||||
return
|
||||
|
||||
sync_orders(C) # 先对账再决策
|
||||
|
||||
# ---- 在下方编排策略逻辑 ----
|
||||
# 定点一次性任务示例(替代 apscheduler date/cron 任务):
|
||||
if '09:30:05' <= now <= '09:31:00' and not G.done_flags.get('open_task'):
|
||||
G.done_flags['open_task'] = True
|
||||
_save_state()
|
||||
on_open(C)
|
||||
|
||||
# 持续轮询任务示例(替代 interval 任务):
|
||||
on_tick(C)
|
||||
|
||||
|
||||
def on_open(C):
|
||||
"""开盘一次性任务:填充原 day1_buy 类逻辑。"""
|
||||
# 票池/信号文件建议盘前由外部脚本生成好,本函数只读本地文件。
|
||||
# 若选股依赖财务/资金流/龙虎榜等 QMT 之外的多维数据,可在外部脚本接一个
|
||||
# HTTP 数据源兜底(如 quantgo.ai/data,按月订阅、接口较全,个人研究够用),
|
||||
# 算好结果落地 csv 再喂进来,避免内置端发起网络请求。
|
||||
pass
|
||||
|
||||
|
||||
def on_tick(C):
|
||||
"""每轮决策:填充原 while/interval 主体逻辑。"""
|
||||
# 行情示例:
|
||||
# tick = C.get_full_tick(['600000.SH'])
|
||||
# last = tick['600000.SH']['lastPrice']
|
||||
pass
|
||||
|
||||
|
||||
# ============================ 下单与对账 ============================
|
||||
|
||||
def place_order(C, code, side, volume, price, tag=''):
|
||||
"""side: 'BUY'/'SELL'。同标的有在途单时拒绝(防超单)。返回 userOrderId 或 None。"""
|
||||
for uid, od in G.pending.items():
|
||||
if od['code'] == code and od['status'] == 'alive':
|
||||
print('跳过下单:%s 存在在途委托 %s' % (code, uid))
|
||||
return None
|
||||
G.seq += 1
|
||||
uid = '%s_%s_%d' % (tag or 'ORD', G.day, G.seq)
|
||||
op = G.op_buy if side == 'BUY' else G.op_sell
|
||||
# prType=11 指定价;quickTrade 必须为 2(定时器回调中下单)
|
||||
passorder(op, 1101, G.acct, code, 11, float(price), int(volume),
|
||||
'TPL_TIMER', 2, uid, C)
|
||||
G.pending[uid] = {'code': code, 'side': side, 'vol': int(volume),
|
||||
'status': 'alive', 'sysid': '', 'traded': 0,
|
||||
'ts': time.time()}
|
||||
_save_state()
|
||||
print('下单 %s %s %d股 @%.3f uid=%s' % (side, code, volume, price, uid))
|
||||
return uid
|
||||
|
||||
|
||||
def cancel_order(C, uid):
|
||||
od = G.pending.get(uid)
|
||||
if od and od.get('sysid'):
|
||||
ok = cancel(od['sysid'], G.acct, G.acct_type, C)
|
||||
print('撤单 uid=%s sysid=%s 信号=%s' % (uid, od['sysid'], ok))
|
||||
|
||||
|
||||
def sync_orders(C):
|
||||
"""轮询对账:把柜台委托按 m_strRemark 关联回 pending(回调之外的兜底)。"""
|
||||
alive_status = (48, 49, 50, 51, 52, 55, 86, 255)
|
||||
try:
|
||||
orders = get_trade_detail_data(G.acct, G.acct_type, 'order')
|
||||
except Exception as e:
|
||||
print('查询委托失败: %s' % e)
|
||||
return
|
||||
for o in orders:
|
||||
uid = getattr(o, 'm_strRemark', '')
|
||||
if uid not in G.pending:
|
||||
continue
|
||||
od = G.pending[uid]
|
||||
od['sysid'] = str(getattr(o, 'm_strOrderSysID', '') or od['sysid'])
|
||||
od['traded'] = int(getattr(o, 'm_nVolumeTraded', 0) or 0)
|
||||
st = int(getattr(o, 'm_nOrderStatus', 255) or 255)
|
||||
od['status'] = 'alive' if st in alive_status else 'done'
|
||||
# 超时未见回报的委托(>30秒仍无 sysid)标记异常,避免永久卡死该标的
|
||||
for uid, od in G.pending.items():
|
||||
if od['status'] == 'alive' and not od['sysid'] and time.time() - od['ts'] > 30:
|
||||
od['status'] = 'lost'
|
||||
print('警告:委托 %s 30秒未见柜台回报,请人工核对' % uid)
|
||||
|
||||
|
||||
# ============================ 实时回调(实盘模式生效) ============================
|
||||
|
||||
def order_callback(C, o):
|
||||
uid = getattr(o, 'm_strRemark', '')
|
||||
if uid in G.pending:
|
||||
G.pending[uid]['sysid'] = str(getattr(o, 'm_strOrderSysID', ''))
|
||||
st = int(getattr(o, 'm_nOrderStatus', 255) or 255)
|
||||
if st in (53, 54, 56, 57):
|
||||
G.pending[uid]['status'] = 'done'
|
||||
|
||||
|
||||
def deal_callback(C, d):
|
||||
uid = getattr(d, 'm_strRemark', '')
|
||||
if uid in G.pending:
|
||||
print('成交推送 uid=%s 价=%.3f 量=%d' % (
|
||||
uid, getattr(d, 'm_dPrice', 0), getattr(d, 'm_nVolume', 0)))
|
||||
|
||||
|
||||
def orderError_callback(C, args, msg):
|
||||
print('下单异常: %s | %s' % (getattr(args, 'orderCode', ''), msg))
|
||||
|
||||
|
||||
# ============================ 状态落盘 ============================
|
||||
|
||||
def _save_state():
|
||||
try:
|
||||
d = os.path.dirname(STATE_FILE)
|
||||
if not os.path.exists(d):
|
||||
os.makedirs(d)
|
||||
tmp = STATE_FILE + '.tmp'
|
||||
with open(tmp, 'w') as f:
|
||||
json.dump({'day': G.day, 'seq': G.seq, 'pending': G.pending,
|
||||
'done_flags': G.done_flags}, f, ensure_ascii=False)
|
||||
os.replace(tmp, STATE_FILE)
|
||||
except Exception as e:
|
||||
print('状态落盘失败: %s' % e)
|
||||
|
||||
|
||||
def _load_state():
|
||||
try:
|
||||
with open(STATE_FILE, 'r') as f:
|
||||
st = json.load(f)
|
||||
if st.get('day') == time.strftime('%Y%m%d'): # 只恢复当日状态
|
||||
G.day = st['day']
|
||||
G.seq = max(G.seq, st.get('seq', 0))
|
||||
G.pending = st.get('pending', {})
|
||||
G.done_flags = st.get('done_flags', {})
|
||||
print('已恢复当日状态:在途%d笔 标志%s' % (len(G.pending), G.done_flags))
|
||||
except Exception:
|
||||
pass
|
||||
@@ -1,7 +0,0 @@
|
||||
datetime,symbol,open,high,low,close,volume,prev_close
|
||||
2026-01-05 09:30:00,600000.SH,10.00,10.05,9.98,10.02,100000,9.95
|
||||
2026-01-05 09:31:00,600000.SH,10.02,10.08,10.01,10.07,120000,9.95
|
||||
2026-01-05 09:32:00,600000.SH,10.07,10.12,10.06,10.11,110000,9.95
|
||||
2026-01-05 09:33:00,600000.SH,10.11,10.13,10.05,10.06,130000,9.95
|
||||
2026-01-06 09:30:00,600000.SH,10.08,10.10,10.00,10.02,150000,10.06
|
||||
2026-01-06 09:31:00,600000.SH,10.02,10.04,9.96,9.98,140000,10.06
|
||||
|
@@ -1,22 +0,0 @@
|
||||
{
|
||||
"initial_cash": 1000000,
|
||||
"initial_positions": {},
|
||||
"buy_commission_rate": 0.0003,
|
||||
"sell_commission_rate": 0.0003,
|
||||
"min_commission": 5,
|
||||
"stamp_tax_rate": 0.0005,
|
||||
"transfer_fee_rate": 0.00001,
|
||||
"slippage_bps": 0,
|
||||
"max_volume_participation": 0.1,
|
||||
"price_limit_rate": 0.1,
|
||||
"lot_size": 100,
|
||||
"time_in_force": "NEXT_BAR",
|
||||
"seed": 0,
|
||||
"fee_schedule": "a_share_2023_08_28",
|
||||
"market_rules_version": "a_share_v1",
|
||||
"strategy_name": "ma_example",
|
||||
"parameters": {
|
||||
"fast": 2,
|
||||
"slow": 3
|
||||
}
|
||||
}
|
||||
@@ -1,69 +0,0 @@
|
||||
"""Example external moving-average strategy for the ZMQ backtest bridge."""
|
||||
|
||||
import argparse
|
||||
|
||||
from bigqmt_backtest.client import BacktestZmqClient
|
||||
from bigqmt_backtest.strategy import ExternalStrategyRunner
|
||||
|
||||
|
||||
class MovingAverageStrategy(object):
|
||||
def __init__(self, symbol, fast=5, slow=20):
|
||||
self.symbol = symbol
|
||||
self.fast = int(fast)
|
||||
self.slow = int(slow)
|
||||
self.sequence = 0
|
||||
|
||||
def on_bar(self, context, bars):
|
||||
if self.symbol not in bars:
|
||||
return []
|
||||
rows = context.history(self.symbol, count=self.slow, fields=["close"])
|
||||
if len(rows) < self.slow:
|
||||
return []
|
||||
closes = [float(row["close"]) for row in rows]
|
||||
fast_value = sum(closes[-self.fast :]) / self.fast
|
||||
slow_value = sum(closes) / self.slow
|
||||
position = context.positions.get(self.symbol, {})
|
||||
quantity = int(position.get("quantity") or 0)
|
||||
available = int(position.get("available") or 0)
|
||||
self.sequence += 1
|
||||
if fast_value > slow_value and quantity == 0:
|
||||
return [
|
||||
{
|
||||
"client_order_id": "ma-buy-%d" % self.sequence,
|
||||
"symbol": self.symbol,
|
||||
"side": "BUY",
|
||||
"quantity": 100,
|
||||
"order_type": "MARKET",
|
||||
}
|
||||
]
|
||||
if fast_value < slow_value and available > 0:
|
||||
return [
|
||||
{
|
||||
"client_order_id": "ma-sell-%d" % self.sequence,
|
||||
"symbol": self.symbol,
|
||||
"side": "SELL",
|
||||
"quantity": available,
|
||||
"order_type": "MARKET",
|
||||
}
|
||||
]
|
||||
return []
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--endpoint", default="tcp://127.0.0.1:16661")
|
||||
parser.add_argument("--run-id", default="", help="Optional; discovered from QMT when omitted")
|
||||
parser.add_argument("--symbol", required=True)
|
||||
parser.add_argument("--fast", type=int, default=5)
|
||||
parser.add_argument("--slow", type=int, default=20)
|
||||
args = parser.parse_args()
|
||||
with BacktestZmqClient(args.endpoint, args.run_id, client_id="ma-example") as client:
|
||||
result = ExternalStrategyRunner(
|
||||
client,
|
||||
MovingAverageStrategy(args.symbol, fast=args.fast, slow=args.slow),
|
||||
).run()
|
||||
print(result)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,160 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Live API smoke test + latency bench (read-only, safe for live account).
|
||||
|
||||
Covers every read method grouped by category. Reports per-call status and
|
||||
latency, plus a category summary. Does NOT call any order/cancel method.
|
||||
"""
|
||||
import sys
|
||||
import time
|
||||
|
||||
sys.path.insert(0, r"D:\gjzqqmt\xtquant_big_convert\src")
|
||||
sys.path.insert(0, r"D:\国金证券QMT交易端_lemo\python")
|
||||
|
||||
import bigqmt_signal_trader.xtquant_compat as compat
|
||||
|
||||
compat.configure()
|
||||
client = compat.get_default_client()
|
||||
ACCOUNT = client.account_id
|
||||
print("account:", ACCOUNT, "| transport:", client.transport_name)
|
||||
print("=" * 78)
|
||||
|
||||
# (category, method, params)
|
||||
GROUPS = [
|
||||
("系统", [
|
||||
("ping", {}),
|
||||
]),
|
||||
("行情快照", [
|
||||
("get_full_tick", {"codes": ["000001.SZ"]}),
|
||||
("get_ticks", {"codes": ["000001.SZ", "600000.SH"]}),
|
||||
]),
|
||||
("合约/品种", [
|
||||
("get_instrument", {"code": "000001.SZ"}),
|
||||
("get_instrument_type", {"code": "000001.SZ"}),
|
||||
("get_stock_name", {"stock": "000001.SZ"}),
|
||||
("get_last_close", {"stock": "000001.SZ"}),
|
||||
("get_float_caps", {"stockcode": "000001.SZ"}),
|
||||
("get_total_share", {"stockcode": "000001.SZ"}),
|
||||
("get_contract_multiplier", {"stockcode": "000001.SZ"}),
|
||||
]),
|
||||
("K线/历史", [
|
||||
("get_market_data_ex", {"field_list": ["close"], "stock_list": ["000001.SZ"], "period": "1d", "count": 5}),
|
||||
("get_market_data", {"field_list": ["close"], "stock_list": ["000001.SZ"], "period": "1d", "count": 5}),
|
||||
("get_local_data", {"field_list": ["close"], "stock_list": ["000001.SZ"], "period": "1d", "count": 5}),
|
||||
("get_divid_factors", {"stock_code": "000001.SZ", "end_time": "20250101"}),
|
||||
]),
|
||||
("板块", [
|
||||
("get_sector_list", {}),
|
||||
("get_stock_list_in_sector", {"sector_name": "沪深A股"}),
|
||||
("get_sector_info", {"sector_name": "沪深A股"}),
|
||||
]),
|
||||
("交易日历/时间", [
|
||||
("get_trading_dates", {"market": "SH", "count": 5}),
|
||||
("get_holidays", {}),
|
||||
("get_markets", {}),
|
||||
("get_market_last_trade_date", {"market": "SH"}),
|
||||
("get_trading_calendar", {"market": "SH", "start_time": "20250601", "end_time": "20250615"}),
|
||||
("get_date_location", {"date": "20250701"}),
|
||||
("datetime_to_timetag", {"datetime_str": "20250701150000", "format": "%Y%m%d%H%M%S"}),
|
||||
("timetag_to_datetime", {"timetag": 1751353200000, "format": "%Y%m%d %H:%M:%S"}),
|
||||
]),
|
||||
("财务/因子", [
|
||||
("get_financial_data", {"stock_list": ["000001.SZ"], "table_list": ["CAPITAL"], "start_time": "20240101", "end_time": "20241231"}),
|
||||
]),
|
||||
("ETF/期权/期货", [
|
||||
("get_etf_info", {}),
|
||||
("get_main_contract", {"code_market": "IF"}),
|
||||
("get_his_contract_list", {"market": "IF"}),
|
||||
]),
|
||||
("期权定价", [
|
||||
("bsm_price", {"opt_type": "C", "target_price": 3.0, "strike_price": 2.8, "risk_free": 0.03, "sigma": 0.3, "days": 30}),
|
||||
("bsm_iv", {"opt_type": "C", "target_price": 3.0, "strike_price": 2.8, "option_price": 0.25, "risk_free": 0.03, "days": 30}),
|
||||
]),
|
||||
("龙虎榜/资金流", [
|
||||
("get_longhubang", {"stock_list": ["000001.SZ"], "start_time": "20250101", "end_time": "20250630"}),
|
||||
("get_turnover_rate", {"stock_code": ["000001.SZ"], "start_time": "20250601", "end_time": "20250630"}),
|
||||
("get_industry", {"industry_name": "银行"}),
|
||||
("get_north_finance_change", {"period": "1d"}),
|
||||
]),
|
||||
("账户查询", [
|
||||
("get_asset", {}),
|
||||
("get_positions", {}),
|
||||
("query_stock_position", {"stock_code": "000001.SZ"}),
|
||||
("query_orders", {}),
|
||||
("query_trades", {}),
|
||||
]),
|
||||
("官方交易函数", [
|
||||
("get_ipo_data", {}),
|
||||
("get_new_purchase_limit", {}),
|
||||
("get_hkt_exchange_rate", {}),
|
||||
("get_value_by_order_id", {"order_id": "1"}),
|
||||
("get_last_order_id", {}),
|
||||
]),
|
||||
("融资融券(普通账户应空)", [
|
||||
("get_assure_contract", {}),
|
||||
("get_unclosed_compacts", {}),
|
||||
("get_debt_contract", {}),
|
||||
("get_enable_short_contract", {}),
|
||||
]),
|
||||
]
|
||||
|
||||
results = [] # (category, method, status, ms, summary)
|
||||
|
||||
|
||||
def summarize(d):
|
||||
if d is None:
|
||||
return "None"
|
||||
if isinstance(d, dict):
|
||||
if not d:
|
||||
return "{}"
|
||||
if "__bigqmt_type__" in d:
|
||||
return "[%s cols=%d rec=%d]" % (d.get("__bigqmt_type__"), len(d.get("columns") or []), len(d.get("records") or []))
|
||||
k = list(d.keys())[:2]
|
||||
return "{%s...}(%d)" % (k, len(d))
|
||||
if isinstance(d, list):
|
||||
return "[len=%d]" % len(d)
|
||||
return repr(d)[:40]
|
||||
|
||||
|
||||
for category, methods in GROUPS:
|
||||
print("\n--- %s ---" % category)
|
||||
for method, params in methods:
|
||||
t0 = time.time()
|
||||
try:
|
||||
data = client.call(method, params)
|
||||
ms = (time.time() - t0) * 1000
|
||||
status = "OK"
|
||||
results.append((category, method, status, ms, summarize(data)))
|
||||
except Exception as e:
|
||||
ms = (time.time() - t0) * 1000
|
||||
status = "FAIL"
|
||||
results.append((category, method, status, ms, str(e)[:40]))
|
||||
r = results[-1]
|
||||
print(" [%-4s %6.1fms] %-28s %s" % (r[2], r[3], r[1], r[4]))
|
||||
|
||||
# Summary
|
||||
print("\n" + "=" * 78)
|
||||
print("=== 汇总 ===")
|
||||
ok = [r for r in results if r[2] == "OK"]
|
||||
fail = [r for r in results if r[2] == "FAIL"]
|
||||
print("通过 %d / 失败 %d / 总计 %d" % (len(ok), len(fail), len(results)))
|
||||
|
||||
print("\n=== 按类别 ===")
|
||||
cats = {}
|
||||
for r in results:
|
||||
cats.setdefault(r[0], []).append(r)
|
||||
for cat, items in cats.items():
|
||||
o = sum(1 for i in items if i[2] == "OK")
|
||||
avg = sum(i[3] for i in items) / len(items)
|
||||
print(" %-22s %d/%d avg=%.1fms" % (cat, o, len(items), avg))
|
||||
|
||||
print("\n=== 延迟分布 (OK) ===")
|
||||
lat = sorted(i[3] for i in ok)
|
||||
if lat:
|
||||
p50 = lat[len(lat) // 2]
|
||||
p90 = lat[int(len(lat) * 0.9)]
|
||||
print(" n=%d min=%.1fms p50=%.1fms p90=%.1fms max=%.1fms" % (len(lat), lat[0], p50, p90, lat[-1]))
|
||||
|
||||
if fail:
|
||||
print("\n=== 失败明细 ===")
|
||||
for r in fail:
|
||||
print(" %-28s %s" % (r[1], r[4]))
|
||||
@@ -1,63 +0,0 @@
|
||||
[build-system]
|
||||
requires = ["setuptools>=68", "wheel"]
|
||||
build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "xtquant-big-convert"
|
||||
version = "0.2.2"
|
||||
description = "Big QMT RPC bridge and MiniQMT-compatible adapter layer (redis/zmq/mysql transports)"
|
||||
readme = "README.md"
|
||||
requires-python = ">=3.8"
|
||||
license = {text = "MIT"}
|
||||
authors = [
|
||||
{name = "litaolemo"},
|
||||
]
|
||||
keywords = ["qmt", "quant", "trading", "rpc", "redis", "zmq", "bigqmt", "miniqmt"]
|
||||
classifiers = [
|
||||
"Development Status :: 4 - Beta",
|
||||
"Intended Audience :: Developers",
|
||||
"License :: OSI Approved :: MIT License",
|
||||
"Programming Language :: Python :: 3",
|
||||
"Programming Language :: Python :: 3.8",
|
||||
"Programming Language :: Python :: 3.9",
|
||||
"Programming Language :: Python :: 3.10",
|
||||
"Programming Language :: Python :: 3.11",
|
||||
"Topic :: Office/Business :: Financial :: Investment",
|
||||
]
|
||||
dependencies = [
|
||||
"pyzmq>=25.0.0",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
redis = ["redis>=5.0.0"]
|
||||
mysql = ["pymysql>=1.0.0", "DBUtils>=3.0.0"]
|
||||
# Faster/smaller wire encoding for whole-quote push (falls back to json if absent).
|
||||
msgpack = ["msgpack>=1.0.0"]
|
||||
dev = [
|
||||
"pytest>=7.0.0",
|
||||
"pytest-cov>=4.0.0",
|
||||
]
|
||||
|
||||
[project.urls]
|
||||
Homepage = "https://github.com/litaolemo/xtquant_big_convert"
|
||||
Repository = "https://github.com/litaolemo/xtquant_big_convert.git"
|
||||
Issues = "https://github.com/litaolemo/xtquant_big_convert/issues"
|
||||
|
||||
[tool.setuptools]
|
||||
package-dir = {"" = "src"}
|
||||
py-modules = [
|
||||
"BIGQMT_REDIS_DRYRUN",
|
||||
"BIGQMT_ZMQ_BACKTEST",
|
||||
"bigqmt_signal_trader_strategy",
|
||||
"bigqmt_signal_trader_redis_rpc_runtime",
|
||||
"bigqmt_signal_trader_redis_dryrun",
|
||||
"bigqmt_signal_trader_dryrun",
|
||||
"bigqmt_signal_trader_diagnostic",
|
||||
]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
include = ["bigqmt_signal_trader*", "bigqmt_backtest*", "xtquant*"]
|
||||
|
||||
[tool.setuptools.package-data]
|
||||
bigqmt_signal_trader = ["*.md"]
|
||||
@@ -1,346 +0,0 @@
|
||||
---
|
||||
name: qmt-trader
|
||||
description: "通过统一 CLI 脚本驱动大 QMT 迅投量化交易端的全部能力,含实时行情查询、K线历史数据、账户资产与持仓查询、委托与成交查询、买入卖出下单、撤单、板块龙虎榜北向资金财务数据等,并内置 xtquant_big_convert 桥接服务的安装部署引导(装包/同步 QMT 端文件/配置/启动验证/排错)。适用于大模型辅助量化交易分析、行情研判、持仓监控、半自动下单等场景。当用户需要查看股票行情、分析K线、查询持仓资产、查看今日委托成交、下单买卖、撤单、查询北向资金龙虎榜财务数据,或需要安装部署 QMT RPC 桥接服务时触发此 skill。"
|
||||
---
|
||||
|
||||
# QMT Trader — 大模型驱动的 QMT 交易/行情工具
|
||||
|
||||
## 概述
|
||||
|
||||
本 skill 提供一个确定性 CLI 脚本 `scripts/qmt.py`,让大模型通过命令行调用大 QMT 的全部
|
||||
交易与行情能力,避免每次现场写 Python 代码。所有命令默认输出 JSON(便于解析),加 `--table`
|
||||
切换人类可读表格。
|
||||
|
||||
**前置条件**:本 skill 依赖 xtquant_big_convert 桥接服务已部署运行。若 `ping` 失败或用户尚未部署,
|
||||
先按下文「首次部署」引导完成:装包 → 同步 QMT 端文件 → 写配置 → QMT 里运行入口 → 验证。
|
||||
|
||||
## 首次部署(只需一次,AI 逐步引导用户完成)
|
||||
|
||||
部署分两端:**客户端**(跑本 skill/策略的开发机)和**服务端**(大 QMT 客户端内置 Python)。
|
||||
|
||||
### 第 1 步:客户端安装包
|
||||
|
||||
```bash
|
||||
pip install "xtquant-big-convert[redis]" # redis 传输(默认,推荐)
|
||||
# 或 zmq 同机低延迟:pip install xtquant-big-convert(基础版已含 pyzmq)
|
||||
```
|
||||
|
||||
> 没发布到 PyPI 的私有 fork 用源码安装:`git clone <repo> && cd xtquant_big_convert && pip install -e .[redis]`
|
||||
|
||||
### 第 2 步:把服务端文件同步到 QMT 的 python 目录
|
||||
|
||||
需要拷 4 项到大 QMT 的 `python` 目录(如 `D:\国金证券QMT交易端\python\`):
|
||||
|
||||
```
|
||||
bigqmt_signal_trader/ (整个包,pip 装的在 site-packages 里)
|
||||
bigqmt_signal_trader_strategy.py
|
||||
bigqmt_signal_trader_redis_rpc_runtime.py
|
||||
BIGQMT_REDIS_DRYRUN.py (★ QMT 编辑器入口,GBK 编码)
|
||||
```
|
||||
|
||||
pip 安装后的文件位置可以用这条命令定位(输出目录里就有全部 4 项):
|
||||
|
||||
```bash
|
||||
python -c "import bigqmt_signal_trader_strategy as m, os; print(os.path.dirname(m.__file__))"
|
||||
```
|
||||
|
||||
> QMT 沙箱若拒绝 `import redis`(部分券商白名单拦截),改用仓库里的 `bigqmt_no_redis/` 无 redis 版本(自包含 ZMQ 传输)。
|
||||
|
||||
### 第 3 步:创建 QMT 端私有配置
|
||||
|
||||
在 QMT 的 `python` 目录创建 `bigqmt_signal_trader_local_config.py`(含账号密码,**不要提交 git**):
|
||||
|
||||
```python
|
||||
# coding: utf-8
|
||||
BIGQMT_ACCOUNT_ID = "资金账号"
|
||||
BIGQMT_REDIS_CONFIG = {
|
||||
"host": "Redis地址", "port": 6379, "db": 5, "password": "Redis密码",
|
||||
"rpc_allow_order_methods": False, # 下单开关,默认关闭;确认风控后改 True
|
||||
"rpc_process_in_listener": True,
|
||||
"rpc_listener_methods": ("*",),
|
||||
"rpc_background_threads": False, # 若切 zmq/mysql 传输必须改 True
|
||||
"schedule_adjust": True,
|
||||
"schedule_adjust_interval": "500nMilliSecond",
|
||||
}
|
||||
```
|
||||
|
||||
> 切 zmq:配置里加 `"transport": "zmq"` 并把 `rpc_background_threads` 改 `True`(QMT 端需装 pyzmq 19.0.2,Python 3.6 最后支持的版本)。
|
||||
|
||||
### 第 4 步:在 QMT 策略编辑器运行入口
|
||||
|
||||
QMT 策略编辑器里**只加载运行 `BIGQMT_REDIS_DRYRUN.py` 一个文件**(它自动 import 其余模块)。
|
||||
若 QMT 装在非默认路径且用 exec 方式加载,需改文件里 `_known_qmt_python_dir()` 的 fallback 路径。
|
||||
|
||||
启动成功标志(QMT 输出面板):
|
||||
|
||||
```
|
||||
[bigqmt_shell] local redis config loaded keys=[...]
|
||||
[bigqmt_shell] local account config loaded=True
|
||||
[bigqmt_rpc] started channel=bigqmt:rpc:req:你的账号
|
||||
[bigqmt_signal_trader] init ok
|
||||
```
|
||||
|
||||
### 第 5 步:客户端配置 + 验证
|
||||
|
||||
客户端用环境变量(或 `bigqmt_signal_trader_client_config.py`)指向同一套 Redis/账号:
|
||||
|
||||
```powershell
|
||||
$env:BIGQMT_ACCOUNT_ID="资金账号"
|
||||
$env:BIGQMT_REDIS_HOST="Redis地址"; $env:BIGQMT_REDIS_PORT="6379"
|
||||
$env:BIGQMT_REDIS_DB="5"; $env:BIGQMT_REDIS_PASSWORD="Redis密码"
|
||||
```
|
||||
|
||||
然后验证(redis ~13ms / zmq ~0.7ms 为正常):
|
||||
|
||||
```bash
|
||||
python scripts/qmt.py ping
|
||||
```
|
||||
|
||||
### 部署排错速查
|
||||
|
||||
| 现象 | 排查 |
|
||||
|------|------|
|
||||
| `ping` 超时 | 客户端/服务端 transport 不一致(一边 redis 一边 zmq);QMT 端服务没启动;Redis 地址/密码/db 不一致 |
|
||||
| QMT 面板报 `import redis` 被拒 | 换 `bigqmt_no_redis/` 无 redis 版本 |
|
||||
| 启动了但查询全空 | 账号没对上:服务端 `BIGQMT_ACCOUNT_ID` vs 客户端 `BIGQMT_ACCOUNT_ID`;QMT 需在实盘模式 |
|
||||
| 下单报 `ORDER_DISABLED` | 正常保护,服务端配置 `rpc_allow_order_methods` 改 `True` 才放行 |
|
||||
| 详细错误日志 | QMT python 目录下 `logs/bigqmt_*.log`(保留 7 天),排错首选 |
|
||||
|
||||
## 快速开始
|
||||
|
||||
### 第 0 步:确认连通性
|
||||
|
||||
```bash
|
||||
python scripts/qmt.py ping
|
||||
```
|
||||
|
||||
返回 `ok: true` 且 `latency_ms` 合理(redis ~13ms / zmq ~0.7ms)即表示服务端就绪。
|
||||
|
||||
### 第 1 步:一键快照(资产+持仓+委托+成交)
|
||||
|
||||
```bash
|
||||
python scripts/qmt.py snapshot
|
||||
```
|
||||
|
||||
一次 RPC 往返返回账户全景,适合快速了解当前状态。
|
||||
|
||||
## 命令速查
|
||||
|
||||
### 行情分析
|
||||
|
||||
| 命令 | 用途 | 示例 |
|
||||
|------|------|------|
|
||||
| `tick <codes...>` | 实时五档盘口 | `tick 600000.SH 000001.SZ` |
|
||||
| `kline <code>` | K线/历史行情 | `kline 600000.SH --period 1d --count 60 --dividend front` |
|
||||
| `instrument <code>` | 合约详情 | `instrument 600000.SH` |
|
||||
| `sector [name]` | 板块成分股/板块列表 | `sector "沪深A股"` |
|
||||
| `trading-dates` | 交易日历 | `trading-dates --count 10` |
|
||||
| `north` | 北向资金 | `north --period 1d` |
|
||||
| `longhubang <code>` | 龙虎榜 | `longhubang 600000.SH --count 5` |
|
||||
| `financial <codes...>` | 财务数据 | `financial 000001.SZ --tables Capital.CAPITAL` |
|
||||
| `download <codes...>` | 下载历史数据 | `download 600654.SH --period 1d --dividend front` |
|
||||
| `quote-subscribe <codes...>` | 实时全推订阅 | `quote-subscribe SH SZ --max 10` |
|
||||
|
||||
### 账户/持仓/委托
|
||||
|
||||
| 命令 | 用途 | 示例 |
|
||||
|------|------|------|
|
||||
| `account` | 账户资产 | `account` |
|
||||
| `positions [code]` | 持仓列表 | `positions` / `positions 600000.SH` |
|
||||
| `orders` | 今日委托 | `orders --cancelable` |
|
||||
| `trades` | 今日成交 | `trades` |
|
||||
| `snapshot` | 一键全景 | `snapshot` |
|
||||
|
||||
### 下单/撤单
|
||||
|
||||
| 命令 | 用途 | 示例 |
|
||||
|------|------|------|
|
||||
| `buy <code> <volume>` | 买入 | `buy 600000.SH 100 --price 7.50` |
|
||||
| `sell <code> <volume>` | 卖出 | `sell 600000.SH 100 --price 7.50` |
|
||||
| `cancel <order_id>` | 撤单 | `cancel 12345 --market SH` |
|
||||
|
||||
> 下单命令支持 `--dry-run`(只打印不下单)、`--latest`(最新价)、`--strategy`、`--remark`。
|
||||
|
||||
### 扩展查询(高频)
|
||||
|
||||
| 命令 | 用途 | 示例 |
|
||||
|------|------|------|
|
||||
| `holiday` | 节假日列表 | `holiday` |
|
||||
| `stock-name <code>` | 股票名称 | `stock-name 600000.SH` |
|
||||
| `instrument-type <code>` | 品种类型 | `instrument-type 600000.SH` |
|
||||
| `divid-factors <code>` | 除权除息因子 | `divid-factors 600000.SH` |
|
||||
| `market-times [market]` | 日内交易时段 | `market-times SH` |
|
||||
| `trading-calendar [market]` | 交易日历(含时段) | `trading-calendar SH` |
|
||||
| `option-list <code>` | 期权列表 | `option-list 510050.SH` |
|
||||
| `bsm-price ...` | BSM 期权定价 | `bsm-price C 3.0 2.8 0.03 0.3 30` |
|
||||
| `bsm-iv ...` | BSM 隐含波动率 | `bsm-iv C 3.0 2.8 0.25 0.03 30` |
|
||||
| `hkt-stats <code>` | 港股通统计 | `hkt-stats 600000.SH` |
|
||||
| `hkt-details <code>` | 港股通明细 | `hkt-details 600000.SH` |
|
||||
| `hkt-rate` | 港股通汇率 | `hkt-rate` |
|
||||
| `top10-holder <code>` | 十大股东 | `top10-holder 600000.SH` |
|
||||
| `holder-num <code>` | 股东户数 | `holder-num 600000.SH` |
|
||||
| `ipo` / `ipo-limit` | 新股数据/申购额度 | `ipo` |
|
||||
| `credit-assure` | 融资担保品合约 | `credit-assure` |
|
||||
| `credit-short` | 融券标的合约 | `credit-short` |
|
||||
| `credit-debt` | 负债合约 | `credit-debt` |
|
||||
| `his-st <code>` | 历史 ST 数据 | `his-st 600000.SH` |
|
||||
| `index-weight <index>` | 指数权重 | `index-weight 000300.SH` |
|
||||
| `industry <name>` | 行业成分 | `industry 银行` |
|
||||
| `sector-info [name]` | 板块详情 | `sector-info 沪深A股` |
|
||||
| `local-data <code>` | 本地缓存数据 | `local-data 600000.SH` |
|
||||
| `timetag2dt <ms>` | 毫秒时间戳转日期 | `timetag2dt 1751353200000` |
|
||||
| `dt2timetag <dt>` | 日期转毫秒时间戳 | `dt2timetag 20250701150000` |
|
||||
|
||||
### 通用 RPC(兜底所有方法)
|
||||
|
||||
`rpc <method> [json_params]` 可调用**任意白名单方法**(含未列出的,如 `get_l2_quote` / `call_formula` / `get_raw_financial_data` 等):
|
||||
|
||||
```bash
|
||||
python scripts/qmt.py rpc get_holidays
|
||||
python scripts/qmt.py rpc get_stock_name '{"stock":"600000.SH"}'
|
||||
python scripts/qmt.py rpc get_l2_quote '{"stock_code":"600000.SH","count":5}'
|
||||
python scripts/qmt.py rpc call_formula '{"formula_name":"MA","stock_code":"600000.SH","period":"1d"}'
|
||||
```
|
||||
|
||||
## 典型工作流
|
||||
|
||||
### 场景一:行情分析
|
||||
|
||||
分析某只股票的技术面:
|
||||
|
||||
```bash
|
||||
# 1. 看实时盘口
|
||||
python scripts/qmt.py tick 600000.SH
|
||||
|
||||
# 2. 拉最近 60 根日 K(前复权),输出含 MA5/MA20/MA60 统计
|
||||
python scripts/qmt.py kline 600000.SH --period 1d --count 60 --dividend front
|
||||
|
||||
# 3. 看合约详情(名称、上市日、最小变动价位等)
|
||||
python scripts/qmt.py instrument 600000.SH
|
||||
|
||||
# 4. 看近期龙虎榜
|
||||
python scripts/qmt.py longhubang 600000.SH --count 5
|
||||
```
|
||||
|
||||
### 场景二:持仓监控
|
||||
|
||||
```bash
|
||||
# 一键看全景
|
||||
python scripts/qmt.py snapshot
|
||||
|
||||
# 只看持仓(含浮动盈亏)
|
||||
python scripts/qmt.py positions
|
||||
|
||||
# 看可撤委托
|
||||
python scripts/qmt.py orders --cancelable
|
||||
```
|
||||
|
||||
### 场景三:下单交易
|
||||
|
||||
```bash
|
||||
# 0. 先看当前价
|
||||
python scripts/qmt.py tick 600000.SH
|
||||
|
||||
# 1. 干跑确认参数
|
||||
python scripts/qmt.py buy 600000.SH 100 --price 7.50 --dry-run
|
||||
|
||||
# 2. 真实下单(限价 7.50 买 100 股)
|
||||
python scripts/qmt.py buy 600000.SH 100 --price 7.50 --strategy my_strat
|
||||
|
||||
# 3. 确认委托进了系统
|
||||
python scripts/qmt.py orders
|
||||
|
||||
# 4. 需要时撤单
|
||||
python scripts/qmt.py cancel <order_sysid> --market SH
|
||||
```
|
||||
|
||||
### 场景四:批量行情分析
|
||||
|
||||
```bash
|
||||
# 同时看多只股票的盘口
|
||||
python scripts/qmt.py tick 600000.SH 000001.SZ 600519.SH
|
||||
|
||||
# 看板块成分股
|
||||
python scripts/qmt.py sector "沪深A股"
|
||||
|
||||
# 看北向资金流向
|
||||
python scripts/qmt.py north
|
||||
```
|
||||
|
||||
## 安全须知
|
||||
|
||||
1. **下单默认关闭**:服务端 `rpc_allow_order_methods` 默认 `False`。必须由人工在服务端配置中
|
||||
显式开启后才能下单,否则 `buy`/`sell`/`cancel` 会报 `ORDER_DISABLED` 错误。
|
||||
|
||||
2. **下单前先看价**:始终先用 `tick` 确认当前价格,避免下出明显不合理的委托。
|
||||
|
||||
3. **超时防重复**:如果 `buy`/`sell` 报 `ORDER_TIMEOUT`,委托可能已提交。**先用 `orders` 查询确认**,
|
||||
不要直接重试,避免重复下单。
|
||||
|
||||
4. **strategy_name 一致性**:下单时的 `--strategy` 和查询时的 `--strategy` 必须一致。
|
||||
查全部委托用 `orders --strategy ""`(空字符串=不过滤)。
|
||||
|
||||
5. **实盘模式**:QMT 必须运行在实盘模式(非模拟/模型交易)才能收到完整回报。
|
||||
|
||||
## 脚本说明
|
||||
|
||||
### scripts/qmt.py
|
||||
|
||||
统一 CLI 入口,包含以下子命令:
|
||||
|
||||
**基础查询**:
|
||||
- `ping` — 连通性检测(含延迟测量)
|
||||
- `account` — 查询账户资产(现金/冻结/总资产/市值)
|
||||
- `positions [code]` — 查询持仓(含浮动盈亏计算)
|
||||
- `orders [--cancelable] [--strategy ""]` — 查询今日委托(含语义化状态名)
|
||||
- `trades [--strategy ""]` — 查询今日成交
|
||||
- `snapshot` — 一键全景(资产+持仓+委托+成交)
|
||||
|
||||
**行情**:
|
||||
- `tick <codes...>` — 实时五档盘口(含涨跌幅计算)
|
||||
- `kline <code> [--period 1d] [--count N] [--dividend front]` — K线(含 MA5/20/60 统计)
|
||||
- `instrument <code>` — 合约详情
|
||||
- `sector [name]` — 板块成分股/板块列表
|
||||
- `trading-dates [--count N]` — 交易日历
|
||||
- `north [--period 1d]` — 北向资金
|
||||
- `longhubang <code> [--count N]` — 龙虎榜
|
||||
- `financial <codes...> [--tables T1,T2]` — 财务数据
|
||||
- `download <codes...>` — 下载历史数据到服务端
|
||||
- `quote-subscribe <codes...> [--max N] [--timeout S]` — 实时全推行情订阅
|
||||
|
||||
**扩展查询**:
|
||||
- `holiday` — 节假日列表
|
||||
- `stock-name <code>` — 股票名称
|
||||
- `instrument-type <code>` — 品种类型
|
||||
- `divid-factors <code>` — 除权除息因子
|
||||
- `market-times [market]` — 日内交易时段
|
||||
- `trading-calendar [market]` — 交易日历(含时段)
|
||||
- `option-list <code>` — 期权列表
|
||||
- `bsm-price` / `bsm-iv` — BSM 期权定价/隐含波动率
|
||||
- `hkt-stats` / `hkt-details` / `hkt-rate` — 港股通统计/明细/汇率
|
||||
- `top10-holder <code>` / `holder-num <code>` — 十大股东/股东户数
|
||||
- `ipo` / `ipo-limit` — 新股数据/申购额度
|
||||
- `credit-assure` / `credit-short` / `credit-debt` — 融资融券查询
|
||||
- `his-st <code>` — 历史 ST 数据
|
||||
- `index-weight <index>` — 指数权重
|
||||
- `industry <name>` — 行业成分
|
||||
- `sector-info [name]` — 板块详情
|
||||
- `local-data <code>` — 本地缓存数据
|
||||
- `timetag2dt` / `dt2timetag` — 时间戳转换
|
||||
|
||||
**交易**:
|
||||
- `buy <code> <volume> [--price P] [--latest]` — 买入下单
|
||||
- `sell <code> <volume> [--price P] [--latest]` — 卖出下单
|
||||
- `cancel <order_id> [--market SH]` — 撤单
|
||||
|
||||
**通用兜底**:
|
||||
- `rpc <method> [json_params]` — 调用任意白名单方法(未列出的方法都能这样调)
|
||||
|
||||
**配置自动发现**:脚本会自动把仓库 `src/` 加入 `sys.path`(开发模式直接运行,无需 pip install),并自动发现 QMT 的 python 目录(读 `local_config.py` 里的 transport 配置)。配置从环境变量(`BIGQMT_ACCOUNT_ID`/`BIGQMT_REDIS_HOST` 等)或配置文件读取。
|
||||
|
||||
**输出格式**:默认 JSON(`ok`/`data`/`ts` 三字段),加 `--table` 切换表格输出。错误返回 `ok: false` + `error`/`detail`/`code`,退出码 1。
|
||||
|
||||
## 参考
|
||||
|
||||
详细的 API 参数、返回值结构、常量定义和已知陷阱见 `references/api_reference.md`。
|
||||
当命令速查不够用时(如需要直接 RPC 调用、查看信用交易类型、了解回调系统等),查阅该文件。
|
||||
@@ -1,523 +0,0 @@
|
||||
# QMT API 参考手册
|
||||
|
||||
本文档是 `qmt-trader` skill 的完整 API 参考。当 SKILL.md 的速查不够用时,查阅本文件获取
|
||||
参数细节、返回值结构和已知陷阱。
|
||||
|
||||
---
|
||||
|
||||
## 1. 初始化与配置
|
||||
|
||||
### 配置来源(优先级从高到低)
|
||||
|
||||
1. **环境变量**
|
||||
| 变量 | 默认 | 说明 |
|
||||
|------|------|------|
|
||||
| `BIGQMT_ACCOUNT_ID` | — | 资金账号 |
|
||||
| `BIGQMT_REDIS_HOST` | `127.0.0.1` | Redis 地址 |
|
||||
| `BIGQMT_REDIS_PORT` | `6379` | Redis 端口 |
|
||||
| `BIGQMT_REDIS_DB` | `5` | Redis DB |
|
||||
| `BIGQMT_REDIS_PASSWORD` | — | Redis 密码 |
|
||||
| `BIGQMT_RPC_TRANSPORT` | `redis` | 传输方式 redis/zmq |
|
||||
| `BIGQMT_RPC_TIMEOUT_SECONDS` | `6.0` | RPC 超时 |
|
||||
|
||||
2. **配置文件** `bigqmt_signal_trader_client_config.py`(在 PYTHONPATH 中,gitignored)
|
||||
|
||||
3. **备选配置文件** `bigqmt_signal_trader_local_config.py`
|
||||
|
||||
### Python 初始化
|
||||
|
||||
```python
|
||||
from bigqmt_signal_trader.xtquant_compat import StockAccount, configure, xt_trader, xtdata
|
||||
|
||||
configure() # 从配置/环境变量初始化
|
||||
acc = StockAccount(xt_trader.client.account_id, "STOCK")
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 2. 行情数据 API
|
||||
|
||||
### 2.1 get_full_tick — 实时五档盘口
|
||||
|
||||
```python
|
||||
xtdata.get_full_tick(code_list)
|
||||
```
|
||||
|
||||
- **参数**: `code_list: list[str]`,如 `["000001.SZ", "600000.SH"]`;也支持整市场 `["SH"]`, `["SZ"]`
|
||||
- **返回**: `dict[code -> dict]`,每只含 `lastPrice`/`open`/`high`/`low`/`lastClose`/`volume`/`amount`/
|
||||
`bidPrice`(10档)/`askPrice`(10档)/`bidVol`/`askVol`/`time`/`stime`
|
||||
- **CLI**: `python qmt.py tick 600000.SH 000001.SZ`
|
||||
- **注意**: 整市场快照数据量大(5000+ 股),超时自动设 30 秒
|
||||
|
||||
### 2.2 get_market_data_ex — K线/历史行情
|
||||
|
||||
```python
|
||||
xtdata.get_market_data_ex(
|
||||
field_list=None, # ["close","open","high","low","volume","amount"] 或 None=全部
|
||||
stock_list=None, # ["000001.SZ"]
|
||||
period="1d", # "1d"/"1m"/"5m"/"15m"/"30m"/"60m"/"tick"
|
||||
start_time="", # "YYYYMMDD" 或 "YYYYMMDDHHMMSS"
|
||||
end_time="",
|
||||
count=-1, # -1=不限
|
||||
dividend_type="none", # "none"/"front"(前复权)/"back"(后复权)
|
||||
fill_data=True, # 是否填充缺失
|
||||
)
|
||||
```
|
||||
|
||||
- **返回**: `dict[code -> pandas.DataFrame]`,index 是时间戳字符串,列含 `time`(epoch ms)/`open`/`high`/`low`/`close`/`volume`/`amount`
|
||||
- **CLI**: `python qmt.py kline 600000.SH --period 1d --count 60 --dividend front`
|
||||
- **自愈**: 请求复权但服务端缺原始数据时(返回全 0),自动触发下载+重试
|
||||
- **陷阱**: 前/后复权必须先在服务端下载原始数据,否则返回全 0(已自愈但仍可能首次慢)
|
||||
|
||||
### 2.3 get_instrument_detail — 合约详情
|
||||
|
||||
```python
|
||||
xtdata.get_instrument_detail(stock_code) # 别名 get_instrumentdetail
|
||||
```
|
||||
|
||||
- **返回**: `dict`,含名称/上市日/合约乘数/最小变动价位等约 30 字段
|
||||
- **CLI**: `python qmt.py instrument 600000.SH`
|
||||
|
||||
### 2.4 get_stock_list_in_sector — 板块成分股
|
||||
|
||||
```python
|
||||
xtdata.get_stock_list_in_sector(sector_name) # 如 "沪深A股", "科创板", "创业板"
|
||||
```
|
||||
|
||||
- **返回**: `list[str]` 代码列表
|
||||
- **CLI**: `python qmt.py sector "沪深A股"`
|
||||
|
||||
### 2.5 get_sector_list — 板块列表
|
||||
|
||||
```python
|
||||
xtdata.get_sector_list()
|
||||
```
|
||||
|
||||
- **返回**: `list[str]`
|
||||
- **CLI**: `python qmt.py sector`
|
||||
- **注意**: 大 QMT 环境 fallback 返回 13 个常用板块名(非完整列表)
|
||||
|
||||
### 2.6 get_trading_dates — 交易日历
|
||||
|
||||
```python
|
||||
xtdata.get_trading_dates(market="SH", start_time="", end_time="", count=-1)
|
||||
```
|
||||
|
||||
- **CLI**: `python qmt.py trading-dates --count 10`
|
||||
|
||||
### 2.7 get_north_finance_change — 北向资金
|
||||
|
||||
```python
|
||||
xtdata.get_north_finance_change(period="1d")
|
||||
```
|
||||
|
||||
- **CLI**: `python qmt.py north`
|
||||
|
||||
### 2.8 get_longhubang — 龙虎榜
|
||||
|
||||
```python
|
||||
xtdata.get_longhubang(stock_list=["600000.SH"], start_time="", end_time="", count=5)
|
||||
```
|
||||
|
||||
- **返回**: `pandas.DataFrame`
|
||||
- **CLI**: `python qmt.py longhubang 600000.SH --count 5`
|
||||
|
||||
### 2.9 get_financial_data — 财务数据
|
||||
|
||||
```python
|
||||
xtdata.get_financial_data(
|
||||
stock_list=["000001.SZ"],
|
||||
table_list=["Capital.CAPITAL"], # 表名
|
||||
start_time="", end_time="",
|
||||
)
|
||||
```
|
||||
|
||||
- **CLI**: `python qmt.py financial 000001.SZ --tables Capital.CAPITAL`
|
||||
|
||||
### 2.10 download_history_data2 — 下载历史数据
|
||||
|
||||
```python
|
||||
xtdata.download_history_data2(
|
||||
stock_list=["600654.SH"], period="1d",
|
||||
start_time="20240101", dividend_type="front",
|
||||
)
|
||||
```
|
||||
|
||||
- **返回**: `{"finished": N, "total": M}`
|
||||
- **CLI**: `python qmt.py download 600654.SH --period 1d --start 20240101 --dividend front`
|
||||
|
||||
### 2.11 subscribe_whole_quote — 全推行情订阅
|
||||
|
||||
```python
|
||||
sub_id = xtdata.subscribe_whole_quote(["SH","SZ"], callback=on_quote)
|
||||
# ... 运行策略 ...
|
||||
xtdata.unsubscribe_quote(sub_id)
|
||||
```
|
||||
|
||||
- **机制**: 服务端真推送(非轮询),增量推送有变化的品种
|
||||
- **CLI**: `python qmt.py quote-subscribe SH SZ --max 10 --timeout 30`
|
||||
- **心跳**: 客户端 3 秒一次 keepalive,服务端重启后自动恢复
|
||||
|
||||
---
|
||||
|
||||
## 3. 账户/持仓/委托查询 API
|
||||
|
||||
### 3.1 query_stock_asset — 查询资产
|
||||
|
||||
```python
|
||||
asset = xt_trader.query_stock_asset(acc)
|
||||
```
|
||||
|
||||
- **返回属性**: `account_id` / `cash`(可用现金) / `frozen_cash` / `total_asset` / `market_value`
|
||||
- **CLI**: `python qmt.py account`
|
||||
- **容错**: RPC 失败时从 Redis 缓存 `bigqmt:positions:{account_id}` 读取
|
||||
|
||||
### 3.2 query_stock_positions — 查询全部持仓
|
||||
|
||||
```python
|
||||
positions = xt_trader.query_stock_positions(acc)
|
||||
```
|
||||
|
||||
- **返回属性**: `stock_code` / `stock_name` / `volume`(总持仓) / `can_use_volume`(可用) /
|
||||
`avg_price`(成本) / `price`(最新价) / `market_value` / `frozen_volume` / `yesterday_volume`
|
||||
- **CLI**: `python qmt.py positions [code]`
|
||||
|
||||
### 3.3 query_stock_position — 查询单只持仓
|
||||
|
||||
```python
|
||||
pos = xt_trader.query_stock_position(acc, "600000.SH")
|
||||
```
|
||||
|
||||
- **返回**: 单个对象或 `None`
|
||||
|
||||
### 3.4 query_stock_orders — 查询委托
|
||||
|
||||
```python
|
||||
orders = xt_trader.query_stock_orders(acc, cancelable_only=False, strategy_name="")
|
||||
```
|
||||
|
||||
- **返回属性**: `stock_code` / `order_type`(23=BUY,24=SELL) / `order_status` /
|
||||
`order_volume` / `traded_volume` / `price` / `order_sysid` / `order_remark`
|
||||
- **CLI**: `python qmt.py orders [--cancelable] [--strategy ""]`
|
||||
- **⚠️ strategy_name 陷阱**: 下单时的 strategy_name 必须和查询时一致。服务端默认 `""` 返回全部;
|
||||
客户端 `BigQmtXtTrader` 默认 `"bigqmt_signal_trader"`。用 `""` 查全部最安全。
|
||||
|
||||
### 3.5 query_stock_trades — 查询成交
|
||||
|
||||
```python
|
||||
trades = xt_trader.query_stock_trades(acc, strategy_name="")
|
||||
```
|
||||
|
||||
- **返回属性**: `stock_code` / `order_type` / `traded_volume` / `traded_price` /
|
||||
`traded_at` / `order_sysid` / `trade_id`
|
||||
- **CLI**: `python qmt.py trades`
|
||||
|
||||
### 3.6 委托状态码
|
||||
|
||||
| 值 | 常量 | 含义 |
|
||||
|----|------|------|
|
||||
| 48 | ORDER_UNREPORTED | 未申报 |
|
||||
| 49 | ORDER_WAIT_REPORTING | 等待申报 |
|
||||
| 50 | ORDER_REPORTED | 已申报 |
|
||||
| 51 | ORDER_REPORTED_CANCEL | 已申报撤单 |
|
||||
| 52 | ORDER_PARTSUCC_CANCEL | 部成撤单 |
|
||||
| 53 | ORDER_PART_CANCEL | 部撤 |
|
||||
| 54 | ORDER_CANCELED | 已撤 |
|
||||
| 55 | ORDER_PART_SUCC | 部分成交 |
|
||||
| 56 | ORDER_SUCCEEDED | 全部成交 |
|
||||
| 57 | ORDER_JUNK | 废单 |
|
||||
| 255 | ORDER_UNKNOWN | 未知 |
|
||||
|
||||
可撤状态: 49, 50, 55
|
||||
|
||||
---
|
||||
|
||||
## 4. 下单 API
|
||||
|
||||
### 4.1 order_stock — 同步下单
|
||||
|
||||
```python
|
||||
from bigqmt_signal_trader.xtquant_compat import STOCK_BUY, STOCK_SELL, FIX_PRICE, LATEST_PRICE
|
||||
|
||||
order_id = xt_trader.order_stock(
|
||||
acc, # StockAccount
|
||||
stock_code, # "600000.SH"
|
||||
order_type, # STOCK_BUY(23) / STOCK_SELL(24)
|
||||
order_volume, # int,委托数量
|
||||
price_type, # FIX_PRICE(11) / LATEST_PRICE(5)
|
||||
price, # float,限价单价格(最新价时传 0)
|
||||
strategy_name, # str
|
||||
order_remark, # str,user_order_id
|
||||
)
|
||||
```
|
||||
|
||||
- **返回**: `order_sys_id`(字符串) 或 `-1`(失败)
|
||||
- **CLI**: `python qmt.py buy 600000.SH 100 --price 7.50 [--strategy s] [--remark r]`
|
||||
- **CLI**: `python qmt.py sell 600000.SH 100 --price 7.50`
|
||||
- **⚠️ 权限**: 服务端默认 `rpc_allow_order_methods=False`,必须显式开启才能下单
|
||||
- **⚠️ 超时**: 超时后委托可能已提交,先查 `query_orders` 确认,避免重复下单
|
||||
|
||||
### 4.2 order_stock_async — 异步下单
|
||||
|
||||
```python
|
||||
seq = xt_trader.order_stock_async(acc, code, order_type, vol, price_type, price, strategy, remark)
|
||||
```
|
||||
|
||||
- **返回**: seq(结果通过 callback 回调)
|
||||
|
||||
### 4.3 order_stock_batch — 批量下单
|
||||
|
||||
```python
|
||||
results = xt_trader.order_stock_batch(acc, orders, batch_id="")
|
||||
# orders: list[dict],每项含 stock_code/action/volume/price/price_type/strategy_name
|
||||
```
|
||||
|
||||
- **上限**: 500 条/批
|
||||
|
||||
### 4.4 信用交易委托类型
|
||||
|
||||
| 常量 | 值 | 用途 |
|
||||
|------|-----|------|
|
||||
| CREDIT_BUY | 23 | 担保品买入 |
|
||||
| CREDIT_SELL | 24 | 担保品卖出 |
|
||||
| CREDIT_FIN_BUY | 27 | 融资买入 |
|
||||
| CREDIT_SLO_SELL | 28 | 融券卖出 |
|
||||
| CREDIT_BUY_SECU_REPAY | 29 | 买券还券 |
|
||||
| CREDIT_DIRECT_SECU_REPAY | 30 | 直接还券 |
|
||||
| CREDIT_SELL_SECU_REPAY | 31 | 卖券还款 |
|
||||
| CREDIT_DIRECT_CASH_REPAY | 32 | 直接还款 |
|
||||
|
||||
---
|
||||
|
||||
## 5. 撤单 API
|
||||
|
||||
### 5.1 cancel_order_stock_sysid
|
||||
|
||||
```python
|
||||
success = xt_trader.cancel_order_stock_sysid(acc, market, order_sysid)
|
||||
# market: "SH" / "SZ" / ""
|
||||
```
|
||||
|
||||
- **CLI**: `python qmt.py cancel <order_sysid> --market SH`
|
||||
|
||||
### 5.2 cancel_order_stock
|
||||
|
||||
```python
|
||||
success = xt_trader.cancel_order_stock(acc, order_id)
|
||||
# 等价于 cancel_order_stock_sysid(acc, "", order_id)
|
||||
```
|
||||
|
||||
---
|
||||
|
||||
## 6. 回调系统
|
||||
|
||||
```python
|
||||
from bigqmt_signal_trader.xtquant_compat import XtQuantTraderCallback
|
||||
|
||||
class MyCallback(XtQuantTraderCallback):
|
||||
def on_stock_order(self, order): ... # 委托变更
|
||||
def on_stock_trade(self, trade): ... # 成交推送
|
||||
def on_order_error(self, error): ... # 委托错误
|
||||
def on_cancel_error(self, error): ... # 撤单错误
|
||||
def on_order_stock_async_response(self, resp): ...
|
||||
def on_account_status(self, status): ...
|
||||
|
||||
xt_trader.register_callback(MyCallback())
|
||||
xt_trader.start()
|
||||
xt_trader.connect()
|
||||
xt_trader.subscribe(acc)
|
||||
```
|
||||
|
||||
事件推送通过 Redis pubsub 频道:
|
||||
- `bigqmt:exec:order:{account_id}`
|
||||
- `bigqmt:exec:trade:{account_id}`
|
||||
- `bigqmt:exec:order_error:{account_id}`
|
||||
- `bigqmt:exec:cancel_error:{account_id}`
|
||||
|
||||
---
|
||||
|
||||
## 7. 关键陷阱速查
|
||||
|
||||
### 7.1 strategy_name 不匹配
|
||||
- 下单用 `strategy_name="rpc_test"` → 查询用 `strategy_name="bigqmt_signal_trader"` → 返回空
|
||||
- **解决**: 查询时传 `strategy_name=""` 返回全部,或保持一致
|
||||
|
||||
### 7.2 下单静默失败
|
||||
- `passorder` 调用成功但委托没进系统(QMT 风控拒绝但没报错)
|
||||
- **解决**: 服务端下单后等 0.5 秒查 `query_orders` 确认;检查返回的 `server_error` 字段
|
||||
|
||||
### 7.3 复权 K 线返回全 0
|
||||
- 服务端缺原始数据时,前/后复权返回的 close 全是 0.0
|
||||
- **解决**: 先 `download_history_data2` 下载原始数据(客户端有自愈机制)
|
||||
|
||||
### 7.4 Transport 不匹配
|
||||
- 客户端 redis / 服务端 zmq → ping 超时
|
||||
- **解决**: 两端 `transport` 字段保持一致
|
||||
|
||||
### 7.5 QMT 必须运行在实盘模式
|
||||
- 模拟模式下委托进 QMT 界面但不在真实委托队列,`query_orders` 查不到
|
||||
- `order_stock` 返回 -1,触发 `on_order_error`
|
||||
|
||||
### 7.6 整市场快照数据量大
|
||||
- `get_full_tick(["SH"])` 返回 5000+ 股完整盘口
|
||||
- **解决**: 启用 `full_tick_cache` 或增大超时(已自动设 30 秒)
|
||||
|
||||
### 7.7 全推行情是增量的
|
||||
- `subscribe_whole_quote` 的大 QMT 回调只推有变化的品种
|
||||
- **解决**: 订阅成功后客户端自动调一次 `get_full_tick` 打底
|
||||
|
||||
### 7.8 下单超时与重复下单
|
||||
- `order_stock` 超时 → 委托可能已提交但没收到响应
|
||||
- **解决**: 超时后先查 `query_orders`/`query_trades` 确认状态,再决定是否重试
|
||||
|
||||
---
|
||||
|
||||
## 8. 常量速查
|
||||
|
||||
### 交易常量
|
||||
|
||||
| 常量 | 值 | 用途 |
|
||||
|------|-----|------|
|
||||
| STOCK_BUY | 23 | 股票买入 |
|
||||
| STOCK_SELL | 24 | 股票卖出 |
|
||||
| FIX_PRICE | 11 | 限价/指定价 |
|
||||
| LATEST_PRICE | 5 | 最新价 |
|
||||
| MARKET_PEER_PRICE_FIRST | 44 | 对手方最优价 |
|
||||
|
||||
### 账号类型
|
||||
|
||||
| 常量 | 值 |
|
||||
|------|-----|
|
||||
| FUTURE_ACCOUNT | 1 |
|
||||
| SECURITY_ACCOUNT | 2 |
|
||||
| CREDIT_ACCOUNT | 3 |
|
||||
| FUTURE_OPTION_ACCOUNT | 5 |
|
||||
| STOCK_OPTION_ACCOUNT | 6 |
|
||||
|
||||
### 期货委托类型(部分)
|
||||
|
||||
| 常量 | 值 | 用途 |
|
||||
|------|-----|------|
|
||||
| FUTURE_OPEN_LONG | 0 | 开多 |
|
||||
| FUTURE_CLOSE_LONG_TODAY | 2 | 平今多 |
|
||||
| FUTURE_OPEN_SHORT | 3 | 开空 |
|
||||
| FUTURE_CLOSE_SHORT_TODAY | 4 | 平今空 |
|
||||
| FUTURE_CLOSE_LONG_HISTORY | 6 | 平昨多 |
|
||||
| FUTURE_CLOSE_SHORT_HISTORY | 7 | 平昨空 |
|
||||
|
||||
---
|
||||
|
||||
## 9. 直接 RPC 调用(绕过兼容层)
|
||||
|
||||
当兼容层方法不够用时,可直接调 RPC:
|
||||
|
||||
```python
|
||||
from bigqmt_signal_trader.redis_rpc import call_redis_rpc
|
||||
import redis
|
||||
|
||||
r = redis.Redis(host="...", port=6379, db=5, password="...")
|
||||
resp = call_redis_rpc(r, "ACCOUNT_ID", "get_full_tick", {"codes": ["000001.SZ"]})
|
||||
print(resp["data"]["000001.SZ"]["lastPrice"])
|
||||
```
|
||||
|
||||
- **万能入口**: `xtdata.call_method("get_float_caps", stockcode="000001.SZ")`
|
||||
- **方法别名映射**:
|
||||
- `get_full_tick` → `get_ticks`
|
||||
- `get_instrument_detail` → `get_instrument`
|
||||
- `query_stock_asset` → `get_asset`
|
||||
- `query_stock_positions` → `get_positions`
|
||||
- `query_stock_orders` → `query_orders`
|
||||
- `query_stock_trades` → `query_trades`
|
||||
- `order_stock` → `submit_order`
|
||||
- `cancel_order_stock` → `cancel_order`
|
||||
|
||||
### RPC 响应结构
|
||||
|
||||
```json
|
||||
{
|
||||
"ok": true,
|
||||
"data": {...},
|
||||
"error": "",
|
||||
"server_error": "",
|
||||
"handled_at": "2024-07-01 15:00:00"
|
||||
}
|
||||
```
|
||||
|
||||
- `ok=true`: `data` 为方法返回值(DataFrame 已序列化,客户端自动还原 pandas 对象)
|
||||
- `ok=false`: `error` 为错误信息
|
||||
- `server_error`: 额外诊断(如 passorder 提交但委托未进系统)
|
||||
|
||||
---
|
||||
|
||||
## 10. 可用 RPC 方法白名单(117 个只读 + 3 个下单/撤单)
|
||||
|
||||
### 行情快照
|
||||
`get_ticks`/`get_full_tick`, `get_instrument`/`get_instrument_detail`, `get_instrument_type`,
|
||||
`get_stock_name`, `get_stock_type`, `get_last_close`, `get_last_volume`, `get_float_caps`,
|
||||
`get_total_share`, `get_turn_over_rate`, `get_weight_in_index`, `get_contract_multiplier`,
|
||||
`get_contract_expire_date`, `get_open_date`, `get_svol`, `get_bvol`, `get_risk_free_rate`,
|
||||
`is_stock_type`, `get_cb_info`
|
||||
|
||||
### K线/历史
|
||||
`get_market_data`, `get_market_data_ex`, `get_local_data`, `get_close_price`, `get_index_weight`
|
||||
|
||||
### L2 行情(需 L2 权限)
|
||||
`get_l2_quote`, `get_l2_order`, `get_l2_transaction`, `subscribe_l2thousand`
|
||||
|
||||
### 板块
|
||||
`get_stock_list_in_sector`, `get_sector_list`, `get_sector_info`, `create_sector`, `add_sector`, `remove_sector`
|
||||
|
||||
### 交易日历/时段
|
||||
`get_trading_dates`, `get_holidays`, `get_markets`, `get_market_last_trade_date`,
|
||||
`get_date_location`, `get_trading_calendar`, `get_trade_times`
|
||||
|
||||
### 数据下载
|
||||
`download_history_data`, `download_history_data2`, `download_holiday_data`,
|
||||
`download_etf_info`, `download_cb_data`, `download_history_contracts`,
|
||||
`download_index_weight`, `download_sector_data`
|
||||
|
||||
### 财务/因子
|
||||
`get_financial_data`, `download_financial_data`, `download_financial_data2`,
|
||||
`get_raw_financial_data`, `get_factor_data`
|
||||
|
||||
### ETF/期权/期货
|
||||
`get_etf_info`, `get_ipo_info`, `get_option_list`, `get_his_option_list`,
|
||||
`get_his_option_list_batch`, `get_option_detail_data`, `get_option_undl_data`,
|
||||
`get_option_undl`, `get_ETF_list`, `get_main_contract`, `get_his_contract_list`
|
||||
|
||||
### 期权定价
|
||||
`bsm_price`, `bsm_iv`, `get_option_iv`
|
||||
|
||||
### 龙虎榜/股东
|
||||
`get_longhubang`, `get_top10_share_holder`, `get_holder_num`, `get_turnover_rate`,
|
||||
`get_industry`, `get_his_st_data`, `get_his_index_data`
|
||||
|
||||
### 资金流
|
||||
`get_north_finance_change`, `get_hkt_statistics`, `get_hkt_details`, `get_hkt_exchange_rate`
|
||||
|
||||
### 因子/模型
|
||||
`call_formula`, `subscribe_formula`, `unsubscribe_formula`, `get_formula_result`, `gen_factor_index`
|
||||
|
||||
### 时间转换(纯本地)
|
||||
`datetime_to_timetag`, `timetag_to_datetime`
|
||||
|
||||
### 账户查询
|
||||
`get_asset`, `get_positions`, `query_stock_position`, `query_orders`, `query_trades`,
|
||||
`get_history_trade_detail_data`, `get_value_by_order_id`, `get_last_order_id`
|
||||
|
||||
### 融资融券(需两融权限)
|
||||
`get_assure_contract`, `get_enable_short_contract`, `get_unclosed_compacts`,
|
||||
`get_closed_compacts`, `get_debt_contract`
|
||||
|
||||
### 期权持仓
|
||||
`get_option_subject_position`, `get_comb_option`
|
||||
|
||||
### 持仓同步
|
||||
`sync_positions`
|
||||
|
||||
### 下单/撤单(需开启 rpc_allow_order_methods)
|
||||
`submit_order`/`order_stock`, `submit_orders_batch`/`order_stock_batch`,
|
||||
`cancel_order`/`cancel_order_stock`/`cancel_order_stock_sysid`
|
||||
|
||||
### 全推行情
|
||||
`subscribe_whole_quote`, `unsubscribe_whole_quote`, `quote_keepalive`
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,108 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Run the full test suite from a single entry point.
|
||||
|
||||
Groups tests by area and prints a clear per-group + total report. Optional
|
||||
live/API tests (need a running QMT + redis) are skipped by default.
|
||||
|
||||
Usage:
|
||||
python run_all_tests.py # all offline tests (default)
|
||||
python run_all_tests.py -v # verbose
|
||||
python run_all_tests.py --live # also run live RPC tests (needs QMT running)
|
||||
python run_all_tests.py --group signal_trader # only one group
|
||||
python run_all_tests.py --group backtest
|
||||
"""
|
||||
import argparse
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
ROOT = os.path.dirname(os.path.abspath(__file__))
|
||||
SRC = os.path.join(ROOT, "src")
|
||||
|
||||
# Test groups: (name, paths, requires_live)
|
||||
GROUPS = [
|
||||
("signal_trader", [os.path.join("tests", "bigqmt_signal_trader")], False),
|
||||
("backtest", [os.path.join("tests", "bigqmt_backtest")], False),
|
||||
]
|
||||
LIVE_GROUP = ("live_api", ["test_all_apis.py"], True)
|
||||
|
||||
|
||||
def run_group(name, paths, verbose):
|
||||
"""Run a pytest group, return (passed, failed, skipped, seconds)."""
|
||||
cmd = [sys.executable, "-m", "pytest"] + paths + ["-q" if not verbose else "-v"]
|
||||
t0 = time.time()
|
||||
proc = subprocess.run(cmd, cwd=ROOT, capture_output=True, text=True)
|
||||
elapsed = time.time() - t0
|
||||
out = (proc.stdout or "") + (proc.stderr or "")
|
||||
# Parse pytest summary like "290 passed in 8.5s" / "1 failed, 289 passed, 3 skipped"
|
||||
passed = failed = skipped = 0
|
||||
for line in out.splitlines():
|
||||
line = line.strip()
|
||||
if not any(tok in line for tok in ("passed", "failed", "skipped", "error")):
|
||||
continue
|
||||
for part in line.split(","):
|
||||
words = part.strip().split()
|
||||
if len(words) >= 2 and words[0].isdigit():
|
||||
num = int(words[0])
|
||||
if words[1].startswith("passed"):
|
||||
passed += num
|
||||
elif words[1].startswith("failed") or words[1].startswith("error"):
|
||||
failed += num
|
||||
elif words[1].startswith("skipped"):
|
||||
skipped += num
|
||||
return passed, failed, skipped, elapsed, out, proc.returncode
|
||||
|
||||
|
||||
def main():
|
||||
parser = argparse.ArgumentParser(description="Run all bigqmt tests")
|
||||
parser.add_argument("-v", "--verbose", action="store_true")
|
||||
parser.add_argument("--live", action="store_true", help="also run live RPC tests (needs QMT)")
|
||||
parser.add_argument("--group", help="run only this group (signal_trader/backtest/live_api)")
|
||||
args = parser.parse_args()
|
||||
|
||||
groups = list(GROUPS)
|
||||
if args.live:
|
||||
groups.append(LIVE_GROUP)
|
||||
if args.group:
|
||||
groups = [g for g in groups if g[0] == args.group]
|
||||
if not groups:
|
||||
print("Unknown group: %s (available: %s)" % (args.group, ", ".join(g[0] for g in GROUPS + [LIVE_GROUP])))
|
||||
return 1
|
||||
|
||||
print("=" * 70)
|
||||
print("Big QMT Bridge - 全量测试")
|
||||
print("=" * 70)
|
||||
|
||||
total_passed = total_failed = total_skipped = 0
|
||||
total_time = 0.0
|
||||
failed_groups = []
|
||||
for name, paths, needs_live in groups:
|
||||
print("\n--- %s ---" % name)
|
||||
passed, failed, skipped, elapsed, out, rc = run_group(name, paths, args.verbose)
|
||||
total_passed += passed
|
||||
total_failed += failed
|
||||
total_skipped += skipped
|
||||
total_time += elapsed
|
||||
status = "PASS" if failed == 0 and rc == 0 else "FAIL"
|
||||
print(" %s: %d passed, %d failed, %d skipped (%.1fs)" % (status, passed, failed, skipped, elapsed))
|
||||
if failed or rc != 0:
|
||||
failed_groups.append(name)
|
||||
if not args.verbose:
|
||||
# print the failing part of the output for visibility
|
||||
tail = "\n".join(out.splitlines()[-20:])
|
||||
print(tail)
|
||||
|
||||
print("\n" + "=" * 70)
|
||||
print("=== 汇总 ===")
|
||||
print("通过 %d / 失败 %d / 跳过 %d / 总计 %d" % (total_passed, total_failed, total_skipped, total_passed + total_failed + total_skipped))
|
||||
print("总耗时 %.1fs" % total_time)
|
||||
if failed_groups:
|
||||
print("失败分组: %s" % ", ".join(failed_groups))
|
||||
return 1
|
||||
print("全部通过 ✅")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -1,254 +0,0 @@
|
||||
#coding:gbk
|
||||
"""QMT bridge entry using the same file-loader pattern as qmt_realtime strategies.
|
||||
|
||||
Broker QMT strategy sandboxes may reject local package names through their
|
||||
normal ``import`` allowlist. The realtime QMT strategies in gupiao_ztfx load
|
||||
their colocated helpers through ``importlib.util.spec_from_file_location``.
|
||||
This entry applies path-based loading to the bridge package, including its
|
||||
internal relative imports, while leaving all standard-library and QMT imports
|
||||
untouched. This terminal's spec loader ignores custom builtins for nested
|
||||
package imports, so local bridge files are compiled explicitly after resolving
|
||||
their path.
|
||||
"""
|
||||
import builtins as _builtins
|
||||
import importlib as _importlib
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
|
||||
|
||||
_LOCAL_ROOTS = (
|
||||
"bigqmt_signal_trader",
|
||||
"bigqmt_signal_trader_strategy",
|
||||
"bigqmt_signal_trader_redis_rpc_runtime",
|
||||
"bigqmt_signal_trader_local_config",
|
||||
)
|
||||
_ORIGINAL_IMPORT = _builtins.__import__
|
||||
_ORIGINAL_IMPORT_MODULE = _importlib.import_module
|
||||
_ORIGINAL_RELOAD = _importlib.reload
|
||||
|
||||
|
||||
def _known_qmt_python_dir():
|
||||
# Find the QMT python dir from sys.path instead of a hardcoded path, so
|
||||
# the bridge loads regardless of broker install location or launch mode
|
||||
# (editor / paste-run / exec). Falls back to empty when not found.
|
||||
for p in sys.path:
|
||||
if p and r"\python" in p and os.path.isdir(p):
|
||||
return p
|
||||
return ""
|
||||
|
||||
|
||||
try:
|
||||
_SOURCE_ROOT = os.path.dirname(os.path.abspath(__file__))
|
||||
except Exception:
|
||||
_SOURCE_ROOT = _known_qmt_python_dir()
|
||||
if not _SOURCE_ROOT:
|
||||
_SOURCE_ROOT = _known_qmt_python_dir()
|
||||
|
||||
|
||||
def _is_local_module(name):
|
||||
return any(name == root or name.startswith(root + ".") for root in _LOCAL_ROOTS)
|
||||
|
||||
|
||||
def _resolve_name(name, module_globals, level):
|
||||
if not level:
|
||||
return name
|
||||
package = (module_globals or {}).get("__package__") or (module_globals or {}).get("__name__", "")
|
||||
if not package:
|
||||
raise ImportError("relative import without package")
|
||||
for unused in range(level - 1):
|
||||
if "." not in package:
|
||||
raise ImportError("relative import beyond top-level package")
|
||||
package = package.rsplit(".", 1)[0]
|
||||
return package + ("." + name if name else "")
|
||||
|
||||
|
||||
def _find_local_source(name):
|
||||
relative = name.replace(".", os.sep)
|
||||
dirs = []
|
||||
if _SOURCE_ROOT:
|
||||
dirs.append(_SOURCE_ROOT)
|
||||
for p in sys.path:
|
||||
if p and os.path.isdir(p) and p not in dirs:
|
||||
dirs.append(p)
|
||||
for d in dirs:
|
||||
package_init = os.path.join(d, relative, "__init__.py")
|
||||
if os.path.isfile(package_init):
|
||||
return package_init, True
|
||||
module_file = os.path.join(d, relative + ".py")
|
||||
if os.path.isfile(module_file):
|
||||
return module_file, False
|
||||
raise ModuleNotFoundError("local source not found: %s" % name, name=name)
|
||||
|
||||
|
||||
def _set_parent_attribute(name, module):
|
||||
if "." not in name:
|
||||
return
|
||||
parent_name, child_name = name.rsplit(".", 1)
|
||||
parent = _load_local_module(parent_name)
|
||||
setattr(parent, child_name, module)
|
||||
|
||||
|
||||
def _load_local_module(name):
|
||||
existing = sys.modules.get(name)
|
||||
if existing is not None:
|
||||
return existing
|
||||
source_path, is_package = _find_local_source(name)
|
||||
if "." in name:
|
||||
_load_local_module(name.rsplit(".", 1)[0])
|
||||
module = types.ModuleType(name)
|
||||
module.__file__ = source_path
|
||||
module.__package__ = name if is_package else name.rpartition(".")[0]
|
||||
if is_package:
|
||||
module.__path__ = [os.path.dirname(source_path)]
|
||||
module_builtins = dict(_builtins.__dict__)
|
||||
module_builtins["__import__"] = _local_import
|
||||
module.__dict__["__builtins__"] = module_builtins
|
||||
module.__dict__["__bigqmt_load_local_module"] = _load_local_module
|
||||
sys.modules[name] = module
|
||||
# QMT native allowlist rejects the root package eager exports.
|
||||
if name == "bigqmt_signal_trader":
|
||||
return module
|
||||
try:
|
||||
with open(source_path, "rb") as source_file:
|
||||
source = source_file.read()
|
||||
exec(compile(source, source_path, "exec"), module.__dict__)
|
||||
except Exception:
|
||||
sys.modules.pop(name, None)
|
||||
raise
|
||||
_set_parent_attribute(name, module)
|
||||
return module
|
||||
|
||||
|
||||
def _local_import(name, module_globals=None, module_locals=None, fromlist=(), level=0):
|
||||
absolute_name = _resolve_name(name, module_globals, level)
|
||||
if not _is_local_module(absolute_name):
|
||||
return _ORIGINAL_IMPORT(name, module_globals, module_locals, fromlist, level)
|
||||
module = _load_local_module(absolute_name)
|
||||
for child in fromlist or ():
|
||||
if child != "*":
|
||||
try:
|
||||
_load_local_module(absolute_name + "." + child)
|
||||
except ModuleNotFoundError:
|
||||
pass
|
||||
if fromlist:
|
||||
return module
|
||||
return _load_local_module(absolute_name.split(".", 1)[0])
|
||||
|
||||
|
||||
def _local_import_module(name, package=None):
|
||||
if _is_local_module(name):
|
||||
return _load_local_module(name)
|
||||
return _ORIGINAL_IMPORT_MODULE(name, package)
|
||||
|
||||
|
||||
def _local_reload(module):
|
||||
if _is_local_module(getattr(module, "__name__", "")):
|
||||
return _load_local_module(module.__name__)
|
||||
return _ORIGINAL_RELOAD(module)
|
||||
|
||||
|
||||
def _clear_local_modules():
|
||||
for name in list(sys.modules):
|
||||
if _is_local_module(name):
|
||||
sys.modules.pop(name, None)
|
||||
|
||||
|
||||
def _stop_previous_rpc_service():
|
||||
"""Release the previous QMT strategy's socket before clearing its module.
|
||||
|
||||
QMT can re-execute this entry in the same Python process. The old strategy
|
||||
module owns the RPC service and its ZMQ ROUTER socket, so dropping that
|
||||
module from ``sys.modules`` first would make the service unreachable and
|
||||
leave its port bound for the next strategy start.
|
||||
"""
|
||||
previous = sys.modules.get("bigqmt_signal_trader_strategy")
|
||||
reset = getattr(previous, "reset_app", None)
|
||||
if not callable(reset):
|
||||
return
|
||||
try:
|
||||
reset()
|
||||
print("[bigqmt_shell] previous rpc service stopped")
|
||||
except Exception as exc:
|
||||
# Continue the reload so a broken old instance does not prevent QMT
|
||||
# from reporting its normal startup error.
|
||||
print("[bigqmt_shell] previous rpc service stop failed: %s" % exc)
|
||||
|
||||
|
||||
_stop_previous_rpc_service()
|
||||
_clear_local_modules()
|
||||
_importlib.import_module = _local_import_module
|
||||
_importlib.reload = _local_reload
|
||||
print("[bigqmt_shell] importlib entry source_root=%s" % _SOURCE_ROOT)
|
||||
|
||||
|
||||
def _fallback_account_id():
|
||||
for name in ("BIGQMT_ACCOUNT_ID", "account", "account_id", "accountID"):
|
||||
value = globals().get(name)
|
||||
if value:
|
||||
return str(value)
|
||||
return ""
|
||||
|
||||
|
||||
try:
|
||||
_local_import("bigqmt_signal_trader.adapters.redis_common", globals(), fromlist=("*",))
|
||||
_local_import("bigqmt_signal_trader.redis_rpc", globals(), fromlist=("*",))
|
||||
_strategy = _local_import("bigqmt_signal_trader_strategy", globals(), fromlist=("*",))
|
||||
_strategy.reset_app()
|
||||
except Exception as bridge_preload_error:
|
||||
print("[bigqmt_shell] bridge preload failed: %s" % bridge_preload_error)
|
||||
|
||||
_runtime = _local_import("bigqmt_signal_trader_redis_rpc_runtime", globals(), fromlist=("*",))
|
||||
|
||||
|
||||
def _load_local_config():
|
||||
return _local_import("bigqmt_signal_trader_local_config", globals(), fromlist=("*",))
|
||||
|
||||
|
||||
try:
|
||||
_config = _load_local_config()
|
||||
BIGQMT_REDIS_CONFIG = getattr(_config, "BIGQMT_REDIS_CONFIG", {})
|
||||
print("[bigqmt_shell] local redis config loaded keys=%s" % sorted((BIGQMT_REDIS_CONFIG or {}).keys()))
|
||||
_runtime.configure_runtime_redis(BIGQMT_REDIS_CONFIG)
|
||||
except Exception as redis_config_error:
|
||||
print("[bigqmt_shell] local redis config load failed: %s" % redis_config_error)
|
||||
|
||||
try:
|
||||
_config = _load_local_config()
|
||||
BIGQMT_ACCOUNT_ID = getattr(_config, "BIGQMT_ACCOUNT_ID", "")
|
||||
print("[bigqmt_shell] local account config loaded=%s" % bool(BIGQMT_ACCOUNT_ID))
|
||||
_runtime.configure_runtime_account(BIGQMT_ACCOUNT_ID)
|
||||
except Exception as account_config_error:
|
||||
print("[bigqmt_shell] local account config load failed: %s" % account_config_error)
|
||||
account_id = _fallback_account_id()
|
||||
if account_id:
|
||||
_runtime.configure_runtime_account(account_id)
|
||||
|
||||
try:
|
||||
qmt_extra = {}
|
||||
for function_name in (
|
||||
"get_history_trade_detail_data", "get_value_by_order_id", "get_last_order_id",
|
||||
"get_ipo_data", "get_new_purchase_limit", "get_assure_contract",
|
||||
"get_enable_short_contract", "get_unclosed_compacts", "get_closed_compacts",
|
||||
"get_debt_contract", "get_option_subject_position", "get_comb_option",
|
||||
"get_hkt_exchange_rate",
|
||||
"download_history_data", "download_history_data2",
|
||||
):
|
||||
if function_name in globals():
|
||||
qmt_extra[function_name] = globals()[function_name]
|
||||
print("[bigqmt_shell] down_history_data bound=%s" % ("down_history_data" in qmt_extra))
|
||||
_runtime.bind_runtime_api(
|
||||
passorder_func=globals().get("passorder"),
|
||||
cancel_func=globals().get("cancel"),
|
||||
get_trade_detail_data_func=globals().get("get_trade_detail_data"),
|
||||
extra_funcs=qmt_extra or None,
|
||||
)
|
||||
except NameError:
|
||||
pass
|
||||
|
||||
|
||||
init = _runtime.init
|
||||
handlebar = _runtime.handlebar
|
||||
adjust = _runtime.adjust
|
||||
order_callback = _runtime.order_callback
|
||||
deal_callback = _runtime.deal_callback
|
||||
@@ -1,152 +0,0 @@
|
||||
#coding:gbk
|
||||
"""Isolated QMT backtest entry for external ZMQ strategies.
|
||||
|
||||
This file is ASCII-only. It loads only the bigqmt_backtest package and never
|
||||
loads or mutates the live bridge package.
|
||||
"""
|
||||
|
||||
import builtins as _builtins
|
||||
import os
|
||||
import sys
|
||||
import types
|
||||
|
||||
|
||||
BACKTEST_ZMQ_CONFIG = {
|
||||
"bind_endpoint": "tcp://127.0.0.1:16662",
|
||||
"run_id": "",
|
||||
"account_id": "",
|
||||
"account_type": "STOCK",
|
||||
"strategy_name": "ZMQ_BACKTEST",
|
||||
"combo_type": 1101,
|
||||
"quick_trade": 2,
|
||||
"market_price_type": 5,
|
||||
"limit_price_type": 11,
|
||||
"bar_wait_timeout_seconds": 60,
|
||||
"require_qmt_backtest": True,
|
||||
}
|
||||
|
||||
|
||||
_LOCAL_ROOT = "bigqmt_backtest"
|
||||
_ORIGINAL_IMPORT = _builtins.__import__
|
||||
|
||||
|
||||
def _known_qmt_python_dir():
|
||||
for p in sys.path:
|
||||
if p and r"\python" in p and os.path.isdir(p):
|
||||
return p
|
||||
return ""
|
||||
|
||||
|
||||
try:
|
||||
_SOURCE_ROOT = os.path.dirname(os.path.abspath(__file__))
|
||||
except Exception:
|
||||
_SOURCE_ROOT = _known_qmt_python_dir()
|
||||
if not _SOURCE_ROOT:
|
||||
_SOURCE_ROOT = _known_qmt_python_dir()
|
||||
|
||||
|
||||
def _is_local(name):
|
||||
return name == _LOCAL_ROOT or name.startswith(_LOCAL_ROOT + ".")
|
||||
|
||||
|
||||
def _resolve_name(name, module_globals, level):
|
||||
if not level:
|
||||
return name
|
||||
package = (module_globals or {}).get("__package__") or ""
|
||||
if not package:
|
||||
raise ImportError("relative import without package")
|
||||
for unused in range(level - 1):
|
||||
package = package.rsplit(".", 1)[0]
|
||||
return package + (("." + name) if name else "")
|
||||
|
||||
|
||||
def _find_source(name):
|
||||
relative = name.replace(".", os.sep)
|
||||
dirs = []
|
||||
if _SOURCE_ROOT:
|
||||
dirs.append(_SOURCE_ROOT)
|
||||
for p in sys.path:
|
||||
if p and os.path.isdir(p) and p not in dirs:
|
||||
dirs.append(p)
|
||||
for d in dirs:
|
||||
package_init = os.path.join(d, relative, "__init__.py")
|
||||
if os.path.isfile(package_init):
|
||||
return package_init, True
|
||||
module_file = os.path.join(d, relative + ".py")
|
||||
if os.path.isfile(module_file):
|
||||
return module_file, False
|
||||
raise ModuleNotFoundError("local source not found: %s" % name, name=name)
|
||||
|
||||
|
||||
def _load_local_module(name):
|
||||
existing = sys.modules.get(name)
|
||||
if existing is not None:
|
||||
return existing
|
||||
source_path, is_package = _find_source(name)
|
||||
if "." in name:
|
||||
_load_local_module(name.rsplit(".", 1)[0])
|
||||
module = types.ModuleType(name)
|
||||
module.__file__ = source_path
|
||||
module.__package__ = name if is_package else name.rpartition(".")[0]
|
||||
if is_package:
|
||||
module.__path__ = [os.path.dirname(source_path)]
|
||||
module_builtins = dict(_builtins.__dict__)
|
||||
module_builtins["__import__"] = _local_import
|
||||
module.__dict__["__builtins__"] = module_builtins
|
||||
sys.modules[name] = module
|
||||
if name == _LOCAL_ROOT:
|
||||
return module
|
||||
try:
|
||||
with open(source_path, "rb") as source_file:
|
||||
source = source_file.read()
|
||||
exec(compile(source, source_path, "exec"), module.__dict__)
|
||||
except Exception:
|
||||
sys.modules.pop(name, None)
|
||||
raise
|
||||
if "." in name:
|
||||
parent_name, child_name = name.rsplit(".", 1)
|
||||
setattr(_load_local_module(parent_name), child_name, module)
|
||||
return module
|
||||
|
||||
|
||||
def _local_import(name, module_globals=None, module_locals=None, fromlist=(), level=0):
|
||||
absolute_name = _resolve_name(name, module_globals, level)
|
||||
if not _is_local(absolute_name):
|
||||
return _ORIGINAL_IMPORT(name, module_globals, module_locals, fromlist, level)
|
||||
module = _load_local_module(absolute_name)
|
||||
for child in fromlist or ():
|
||||
if child != "*":
|
||||
try:
|
||||
_load_local_module(absolute_name + "." + child)
|
||||
except ModuleNotFoundError:
|
||||
pass
|
||||
if fromlist:
|
||||
return module
|
||||
return _load_local_module(absolute_name.split(".", 1)[0])
|
||||
|
||||
|
||||
for _name in sorted(
|
||||
[name for name in list(sys.modules) if _is_local(name)],
|
||||
key=lambda item: item.count("."),
|
||||
reverse=True,
|
||||
):
|
||||
sys.modules.pop(_name, None)
|
||||
|
||||
|
||||
_runtime = _load_local_module("bigqmt_backtest.qmt_runtime")
|
||||
_runtime.configure(**BACKTEST_ZMQ_CONFIG)
|
||||
_runtime.bind_qmt_api(
|
||||
passorder_func=globals().get("passorder") or getattr(_builtins, "passorder", None),
|
||||
cancel_func=globals().get("cancel") or getattr(_builtins, "cancel", None),
|
||||
get_trade_detail_data_func=(
|
||||
globals().get("get_trade_detail_data")
|
||||
or getattr(_builtins, "get_trade_detail_data", None)
|
||||
),
|
||||
)
|
||||
|
||||
init = _runtime.init
|
||||
handlebar = _runtime.handlebar
|
||||
order_callback = _runtime.order_callback
|
||||
deal_callback = _runtime.deal_callback
|
||||
stop = _runtime.stop
|
||||
after_backtest = _runtime.after_backtest
|
||||
@@ -1,23 +0,0 @@
|
||||
"""Isolated ZMQ bridge for QMT-native and standalone backtests.
|
||||
|
||||
This package deliberately does not import ``bigqmt_signal_trader``. The live
|
||||
bridge and both backtest backends therefore have separate module state,
|
||||
identities, and order gateways. QMT-native mode never uses the local broker.
|
||||
"""
|
||||
|
||||
from .client import BacktestZmqClient
|
||||
from .data_feed import CsvBarFeed, InMemoryBarFeed
|
||||
from .engine import BacktestConfig, BacktestEngine
|
||||
from .protocol import BacktestBridgeProtocol
|
||||
|
||||
|
||||
__all__ = [
|
||||
"BacktestBridgeProtocol",
|
||||
"BacktestConfig",
|
||||
"BacktestEngine",
|
||||
"BacktestZmqClient",
|
||||
"CsvBarFeed",
|
||||
"InMemoryBarFeed",
|
||||
]
|
||||
|
||||
__version__ = "1.0.0"
|
||||
@@ -1,5 +0,0 @@
|
||||
from .server import main
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,358 +0,0 @@
|
||||
"""A-share simulated broker used only by the standalone backtest runtime."""
|
||||
|
||||
import json
|
||||
from decimal import Decimal
|
||||
|
||||
from .models import (
|
||||
BacktestFill,
|
||||
BacktestOrder,
|
||||
Position,
|
||||
ZERO,
|
||||
decimal_value,
|
||||
json_number,
|
||||
money,
|
||||
normalize_symbol,
|
||||
round_price,
|
||||
)
|
||||
|
||||
|
||||
ACTIVE_ORDER_STATUSES = ("PENDING", "PARTIALLY_FILLED")
|
||||
|
||||
|
||||
class SimulatedBroker(object):
|
||||
def __init__(self, config):
|
||||
self.config = config
|
||||
self.cash = money(config.initial_cash)
|
||||
self.positions = {}
|
||||
for symbol, payload in dict(config.initial_positions or {}).items():
|
||||
data = dict(payload or {})
|
||||
position = Position(
|
||||
symbol,
|
||||
quantity=data.get("quantity", data.get("volume", 0)),
|
||||
available=data.get("available"),
|
||||
today_buy=data.get("today_buy", 0),
|
||||
avg_cost=data.get("avg_cost", data.get("cost", 0)),
|
||||
)
|
||||
if position.quantity > 0:
|
||||
self.positions[position.symbol] = position
|
||||
self.orders_list = []
|
||||
self.fills_list = []
|
||||
self._client_order_ids = {}
|
||||
self._client_order_fingerprints = {}
|
||||
self._order_sequence = 0
|
||||
self._fill_sequence = 0
|
||||
self._trading_date = None
|
||||
self.total_fees = ZERO
|
||||
self.turnover = ZERO
|
||||
|
||||
def _new_order(self, payload, frame_index, submitted_at):
|
||||
self._order_sequence += 1
|
||||
return BacktestOrder(
|
||||
order_id="bt-order-%06d" % self._order_sequence,
|
||||
client_order_id=payload.get("client_order_id"),
|
||||
symbol=payload.get("symbol"),
|
||||
side=payload.get("side"),
|
||||
quantity=payload.get("quantity"),
|
||||
order_type=payload.get("order_type", "MARKET"),
|
||||
limit_price=payload.get("limit_price", payload.get("price")),
|
||||
submitted_index=frame_index,
|
||||
submitted_at=submitted_at,
|
||||
time_in_force=payload.get("time_in_force", self.config.time_in_force),
|
||||
)
|
||||
|
||||
def _reject(self, order, reason):
|
||||
order.status = "REJECTED"
|
||||
order.reject_reason = str(reason)
|
||||
return order
|
||||
|
||||
def _reserved_sell(self, symbol):
|
||||
return sum(
|
||||
order.remaining
|
||||
for order in self.orders_list
|
||||
if order.symbol == symbol and order.side == "SELL" and order.status in ACTIVE_ORDER_STATUSES
|
||||
)
|
||||
|
||||
def submit(self, payload, frame_index, submitted_at):
|
||||
payload = dict(payload or {})
|
||||
client_order_id = str(payload.get("client_order_id") or "")
|
||||
if client_order_id and client_order_id in self._client_order_ids:
|
||||
limit_value = payload.get("limit_price", payload.get("price"))
|
||||
fingerprint_payload = {
|
||||
"symbol": normalize_symbol(payload.get("symbol")),
|
||||
"side": str(payload.get("side") or "").upper(),
|
||||
"quantity": int(payload.get("quantity") or 0),
|
||||
"order_type": str(payload.get("order_type") or "MARKET").upper(),
|
||||
"limit_price": None if limit_value in (None, "") else float(decimal_value(limit_value)),
|
||||
"time_in_force": str(payload.get("time_in_force", self.config.time_in_force)).upper(),
|
||||
}
|
||||
fingerprint = json.dumps(fingerprint_payload, sort_keys=True, separators=(",", ":"))
|
||||
if self._client_order_fingerprints.get(client_order_id) != fingerprint:
|
||||
raise ValueError("client_order_id reused with different order payload")
|
||||
return self._client_order_ids[client_order_id]
|
||||
try:
|
||||
symbol = normalize_symbol(payload.get("symbol"))
|
||||
side = str(payload.get("side") or "").upper()
|
||||
quantity = int(payload.get("quantity") or 0)
|
||||
order_type = str(payload.get("order_type") or "MARKET").upper()
|
||||
if side not in ("BUY", "SELL"):
|
||||
raise ValueError("side must be BUY or SELL")
|
||||
if quantity <= 0:
|
||||
raise ValueError("quantity must be positive")
|
||||
if order_type not in ("MARKET", "LIMIT"):
|
||||
raise ValueError("order_type must be MARKET or LIMIT")
|
||||
if order_type == "LIMIT" and decimal_value(payload.get("limit_price", payload.get("price"))) <= 0:
|
||||
raise ValueError("positive limit_price is required for LIMIT order")
|
||||
payload.update({"symbol": symbol, "side": side, "quantity": quantity, "order_type": order_type})
|
||||
order = self._new_order(payload, frame_index, submitted_at)
|
||||
except Exception as exc:
|
||||
self._order_sequence += 1
|
||||
order = BacktestOrder(
|
||||
"bt-order-%06d" % self._order_sequence,
|
||||
client_order_id,
|
||||
payload.get("symbol") or "UNKNOWN",
|
||||
payload.get("side") or "UNKNOWN",
|
||||
int(payload.get("quantity") or 0),
|
||||
payload.get("order_type") or "MARKET",
|
||||
payload.get("limit_price", payload.get("price")),
|
||||
frame_index,
|
||||
submitted_at,
|
||||
payload.get("time_in_force", self.config.time_in_force),
|
||||
)
|
||||
self._reject(order, "invalid_order:%s" % exc)
|
||||
self.orders_list.append(order)
|
||||
return order
|
||||
|
||||
self.orders_list.append(order)
|
||||
if client_order_id:
|
||||
self._client_order_ids[client_order_id] = order
|
||||
fingerprint_payload = {
|
||||
"symbol": order.symbol,
|
||||
"side": order.side,
|
||||
"quantity": order.quantity,
|
||||
"order_type": order.order_type,
|
||||
"limit_price": None if order.limit_price is None else float(order.limit_price),
|
||||
"time_in_force": order.time_in_force,
|
||||
}
|
||||
self._client_order_fingerprints[client_order_id] = json.dumps(
|
||||
fingerprint_payload, sort_keys=True, separators=(",", ":")
|
||||
)
|
||||
lot_size = self.config.lot_size
|
||||
if order.side == "BUY" and order.quantity % lot_size:
|
||||
return self._reject(order, "buy_quantity_not_round_lot")
|
||||
if order.side == "SELL":
|
||||
position = self.positions.get(order.symbol)
|
||||
available = 0 if position is None else max(position.available - self._reserved_sell(order.symbol) + order.quantity, 0)
|
||||
if available <= 0:
|
||||
return self._reject(order, "t_plus_one_unavailable")
|
||||
if order.quantity > available:
|
||||
return self._reject(order, "insufficient_sellable")
|
||||
if order.quantity % lot_size and order.quantity != available:
|
||||
return self._reject(order, "sell_quantity_not_round_lot")
|
||||
return order
|
||||
|
||||
def cancel(self, order_id):
|
||||
for order in self.orders_list:
|
||||
if order.order_id == str(order_id) or order.client_order_id == str(order_id):
|
||||
if order.status not in ACTIVE_ORDER_STATUSES:
|
||||
return order
|
||||
order.status = "CANCELLED"
|
||||
order.reject_reason = "cancelled_by_strategy"
|
||||
return order
|
||||
raise KeyError("order not found: %s" % order_id)
|
||||
|
||||
def _settle_trading_day(self, trading_date):
|
||||
if self._trading_date == trading_date:
|
||||
return
|
||||
if self._trading_date is not None:
|
||||
for position in self.positions.values():
|
||||
position.available = position.quantity
|
||||
position.today_buy = 0
|
||||
self._trading_date = trading_date
|
||||
|
||||
def _limits(self, order, bar):
|
||||
prev_close = decimal_value(bar.get("prev_close") or bar.get("close"))
|
||||
if bar.get("price_limit_rate") not in (None, ""):
|
||||
rate = decimal_value(bar.get("price_limit_rate"))
|
||||
else:
|
||||
pure = order.symbol.split(".", 1)[0]
|
||||
if order.symbol.endswith(".BJ"):
|
||||
rate = Decimal("0.30")
|
||||
elif pure.startswith(("300", "301", "688", "689")):
|
||||
rate = Decimal("0.20")
|
||||
else:
|
||||
rate = decimal_value(self.config.price_limit_rate)
|
||||
up_limit = bar.get("up_limit")
|
||||
down_limit = bar.get("down_limit")
|
||||
up_limit = round_price(order.symbol, up_limit if up_limit not in (None, "") else prev_close * (Decimal("1") + rate))
|
||||
down_limit = round_price(order.symbol, down_limit if down_limit not in (None, "") else prev_close * (Decimal("1") - rate))
|
||||
return up_limit, down_limit
|
||||
|
||||
def _match_price(self, order, bar):
|
||||
if bool(bar.get("suspended")) or float(bar.get("volume") or 0) <= 0:
|
||||
return None, "suspended_or_no_volume"
|
||||
open_price = round_price(order.symbol, bar["open"])
|
||||
high = round_price(order.symbol, bar["high"])
|
||||
low = round_price(order.symbol, bar["low"])
|
||||
up_limit, down_limit = self._limits(order, bar)
|
||||
if order.side == "BUY" and open_price == high == low == up_limit:
|
||||
return None, "limit_up_locked"
|
||||
if order.side == "SELL" and open_price == high == low == down_limit:
|
||||
return None, "limit_down_locked"
|
||||
if order.order_type == "LIMIT" and not down_limit <= order.limit_price <= up_limit:
|
||||
return None, "limit_price_outside_daily_range"
|
||||
if order.order_type == "MARKET":
|
||||
price = open_price
|
||||
elif order.side == "BUY":
|
||||
if low > order.limit_price:
|
||||
return None, "limit_not_crossed"
|
||||
price = min(open_price, order.limit_price)
|
||||
else:
|
||||
if high < order.limit_price:
|
||||
return None, "limit_not_crossed"
|
||||
price = max(open_price, order.limit_price)
|
||||
slip = decimal_value(self.config.slippage_bps) / Decimal("10000")
|
||||
if order.side == "BUY":
|
||||
price = min(round_price(order.symbol, price * (Decimal("1") + slip)), up_limit)
|
||||
else:
|
||||
price = max(round_price(order.symbol, price * (Decimal("1") - slip)), down_limit)
|
||||
return price, ""
|
||||
|
||||
def _fees(self, side, amount):
|
||||
rate = self.config.buy_commission_rate if side == "BUY" else self.config.sell_commission_rate
|
||||
commission = max(amount * decimal_value(rate), decimal_value(self.config.min_commission)) if rate else ZERO
|
||||
stamp = amount * decimal_value(self.config.stamp_tax_rate) if side == "SELL" else ZERO
|
||||
transfer = amount * decimal_value(self.config.transfer_fee_rate)
|
||||
return money(commission), money(stamp), money(transfer)
|
||||
|
||||
def _volume_cap(self, order, bar, used_volume=0):
|
||||
raw = int(float(bar.get("volume") or 0) * float(self.config.max_volume_participation))
|
||||
cap = max((raw // self.config.lot_size) * self.config.lot_size - int(used_volume), 0)
|
||||
return min(order.remaining, cap)
|
||||
|
||||
def _affordable_buy_quantity(self, quantity, price):
|
||||
quantity = (int(quantity) // self.config.lot_size) * self.config.lot_size
|
||||
while quantity > 0:
|
||||
amount = money(price * quantity)
|
||||
fees = sum(self._fees("BUY", amount), ZERO)
|
||||
if self.cash >= amount + fees:
|
||||
return quantity
|
||||
quantity -= self.config.lot_size
|
||||
return 0
|
||||
|
||||
def _apply_fill(self, order, quantity, price, frame_index, filled_at):
|
||||
amount = money(price * quantity)
|
||||
commission, stamp, transfer = self._fees(order.side, amount)
|
||||
self._fill_sequence += 1
|
||||
fill = BacktestFill(
|
||||
"bt-fill-%06d" % self._fill_sequence,
|
||||
order,
|
||||
quantity,
|
||||
price,
|
||||
commission,
|
||||
stamp,
|
||||
transfer,
|
||||
frame_index,
|
||||
filled_at,
|
||||
)
|
||||
fees = fill.total_fee
|
||||
position = self.positions.get(order.symbol)
|
||||
if order.side == "BUY":
|
||||
if position is None:
|
||||
position = Position(order.symbol)
|
||||
self.positions[order.symbol] = position
|
||||
old_cost = position.avg_cost * position.quantity
|
||||
self.cash = money(self.cash - amount - fees)
|
||||
position.quantity += quantity
|
||||
position.today_buy += quantity
|
||||
position.avg_cost = (old_cost + amount + fees) / position.quantity
|
||||
else:
|
||||
if position is None or position.available < quantity:
|
||||
raise RuntimeError("sellable quantity changed before fill")
|
||||
self.cash = money(self.cash + amount - fees)
|
||||
position.quantity -= quantity
|
||||
position.available -= quantity
|
||||
position.realized_pnl += amount - fees - position.avg_cost * quantity
|
||||
if position.quantity <= 0:
|
||||
self.positions.pop(order.symbol, None)
|
||||
order.filled_quantity += quantity
|
||||
order.status = "FILLED" if order.remaining == 0 else "PARTIALLY_FILLED"
|
||||
self.total_fees += fees
|
||||
self.turnover += amount
|
||||
self.fills_list.append(fill)
|
||||
return fill
|
||||
|
||||
def advance(self, frame_index, frame):
|
||||
trading_date = str(frame["datetime"])[:10]
|
||||
self._settle_trading_day(trading_date)
|
||||
fills = []
|
||||
used_volume = {}
|
||||
for order in self.orders_list:
|
||||
if order.status not in ACTIVE_ORDER_STATUSES or frame_index <= order.submitted_index:
|
||||
continue
|
||||
bar = frame["bars"].get(order.symbol)
|
||||
if bar is None:
|
||||
continue
|
||||
if order.time_in_force == "DAY" and str(order.submitted_at)[:10] != trading_date:
|
||||
order.status = "EXPIRED"
|
||||
order.reject_reason = "day_order_expired"
|
||||
continue
|
||||
order.last_attempt_index = frame_index
|
||||
price, reason = self._match_price(order, bar)
|
||||
if price is None:
|
||||
if order.time_in_force == "NEXT_BAR":
|
||||
order.status = "EXPIRED"
|
||||
order.reject_reason = reason
|
||||
continue
|
||||
quantity = self._volume_cap(order, bar, used_volume.get(order.symbol, 0))
|
||||
if quantity <= 0:
|
||||
reason = "volume_participation_exhausted"
|
||||
elif order.side == "BUY":
|
||||
quantity = self._affordable_buy_quantity(quantity, price)
|
||||
if quantity <= 0:
|
||||
reason = "insufficient_cash"
|
||||
else:
|
||||
position = self.positions.get(order.symbol)
|
||||
quantity = min(quantity, 0 if position is None else position.available)
|
||||
if quantity <= 0:
|
||||
reason = "t_plus_one_unavailable"
|
||||
if quantity > 0:
|
||||
fills.append(self._apply_fill(order, quantity, price, frame_index, frame["datetime"]))
|
||||
used_volume[order.symbol] = used_volume.get(order.symbol, 0) + quantity
|
||||
if order.time_in_force == "NEXT_BAR" and order.remaining > 0:
|
||||
if order.filled_quantity == 0:
|
||||
order.status = "EXPIRED"
|
||||
else:
|
||||
order.status = "PARTIALLY_FILLED_EXPIRED"
|
||||
order.reject_reason = reason or "next_bar_remaining_expired"
|
||||
return fills
|
||||
|
||||
def expire_open_orders(self, reason="backtest_finished"):
|
||||
for order in self.orders_list:
|
||||
if order.status in ACTIVE_ORDER_STATUSES:
|
||||
order.status = "EXPIRED"
|
||||
order.reject_reason = reason
|
||||
|
||||
def snapshot(self, bars):
|
||||
positions = {}
|
||||
market_value = ZERO
|
||||
for symbol in sorted(self.positions):
|
||||
position = self.positions[symbol]
|
||||
bar = bars.get(symbol) or {}
|
||||
mark = decimal_value(bar.get("close"), position.avg_cost)
|
||||
market_value += mark * position.quantity
|
||||
positions[symbol] = position.to_dict(mark)
|
||||
total_asset = money(self.cash + market_value)
|
||||
return {
|
||||
"cash": json_number(self.cash, 2),
|
||||
"market_value": json_number(money(market_value), 2),
|
||||
"total_asset": json_number(total_asset, 2),
|
||||
"positions": positions,
|
||||
"total_fees": json_number(money(self.total_fees), 2),
|
||||
"turnover": json_number(money(self.turnover), 2),
|
||||
}
|
||||
|
||||
def orders(self):
|
||||
return [order.to_dict() for order in self.orders_list]
|
||||
|
||||
def fills(self):
|
||||
return [fill.to_dict() for fill in self.fills_list]
|
||||
@@ -1,134 +0,0 @@
|
||||
"""External-strategy client SDK for the ZMQ backtest bridge."""
|
||||
|
||||
import json
|
||||
import uuid
|
||||
|
||||
|
||||
class BacktestRemoteError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class BacktestZmqClient(object):
|
||||
def __init__(
|
||||
self,
|
||||
endpoint,
|
||||
run_id,
|
||||
client_id="external-strategy",
|
||||
timeout_seconds=10.0,
|
||||
):
|
||||
self.endpoint = str(endpoint)
|
||||
self.run_id = str(run_id)
|
||||
self.client_id = str(client_id)
|
||||
self.timeout_seconds = float(timeout_seconds)
|
||||
self._context = None
|
||||
self._socket = None
|
||||
|
||||
def _connect(self):
|
||||
if self._socket is not None:
|
||||
return self._socket
|
||||
import zmq
|
||||
|
||||
self._context = zmq.Context.instance()
|
||||
self._socket = self._context.socket(zmq.REQ)
|
||||
self._socket.setsockopt(zmq.LINGER, 0)
|
||||
self._socket.connect(self.endpoint)
|
||||
return self._socket
|
||||
|
||||
def _reset_socket(self):
|
||||
if self._socket is not None:
|
||||
self._socket.close(linger=0)
|
||||
self._socket = None
|
||||
|
||||
def request(self, method, params=None, request_id=None):
|
||||
import zmq
|
||||
|
||||
request_id = str(request_id or uuid.uuid4().hex)
|
||||
envelope = {
|
||||
"schema_version": 1,
|
||||
"request_id": request_id,
|
||||
"run_id": self.run_id,
|
||||
"client_id": self.client_id,
|
||||
"method": str(method),
|
||||
"params": dict(params or {}),
|
||||
}
|
||||
socket = self._connect()
|
||||
socket.send(json.dumps(envelope, ensure_ascii=False, separators=(",", ":")).encode("utf-8"))
|
||||
poller = zmq.Poller()
|
||||
poller.register(socket, zmq.POLLIN)
|
||||
events = dict(poller.poll(int(self.timeout_seconds * 1000)))
|
||||
if socket not in events:
|
||||
self._reset_socket()
|
||||
raise TimeoutError("backtest ZMQ request timed out: %s" % method)
|
||||
response = json.loads(socket.recv().decode("utf-8"))
|
||||
if str(response.get("request_id") or "") != request_id:
|
||||
raise BacktestRemoteError("response request_id mismatch")
|
||||
if not response.get("ok"):
|
||||
raise BacktestRemoteError(str(response.get("error") or "remote request failed"))
|
||||
return response.get("data")
|
||||
|
||||
def ping(self):
|
||||
return self.request("ping")
|
||||
|
||||
def describe(self):
|
||||
data = self.request("describe")
|
||||
if not self.run_id and data.get("run_id"):
|
||||
self.run_id = str(data["run_id"])
|
||||
return data
|
||||
|
||||
def start(self):
|
||||
return self.request("start")
|
||||
|
||||
def next_bar(self):
|
||||
return self.request("next_bar")
|
||||
|
||||
def state(self):
|
||||
return self.request("state")
|
||||
|
||||
def submit_order(
|
||||
self,
|
||||
symbol,
|
||||
side,
|
||||
quantity,
|
||||
order_type="MARKET",
|
||||
limit_price=None,
|
||||
client_order_id="",
|
||||
time_in_force="NEXT_BAR",
|
||||
):
|
||||
params = {
|
||||
"symbol": symbol,
|
||||
"side": side,
|
||||
"quantity": int(quantity),
|
||||
"order_type": order_type,
|
||||
"client_order_id": client_order_id,
|
||||
"time_in_force": time_in_force,
|
||||
}
|
||||
if limit_price is not None:
|
||||
params["limit_price"] = limit_price
|
||||
return self.request("submit_order", params)
|
||||
|
||||
def cancel_order(self, order_id):
|
||||
return self.request("cancel_order", {"order_id": order_id})
|
||||
|
||||
def history(self, symbol, count=100, fields=None):
|
||||
params = {"symbol": symbol, "count": int(count)}
|
||||
if fields is not None:
|
||||
params["fields"] = list(fields)
|
||||
return self.request("history", params)
|
||||
|
||||
def orders(self):
|
||||
return self.request("orders")
|
||||
|
||||
def fills(self):
|
||||
return self.request("fills")
|
||||
|
||||
def finish(self):
|
||||
return self.request("finish")
|
||||
|
||||
def close(self):
|
||||
self._reset_socket()
|
||||
|
||||
def __enter__(self):
|
||||
return self
|
||||
|
||||
def __exit__(self, exc_type, exc_value, traceback):
|
||||
self.close()
|
||||
@@ -1,278 +0,0 @@
|
||||
"""Deterministic historical bar feeds for the backtest bridge."""
|
||||
|
||||
import csv
|
||||
import datetime as dt
|
||||
import hashlib
|
||||
import io
|
||||
import json
|
||||
import os
|
||||
import threading
|
||||
|
||||
from .models import normalize_symbol
|
||||
|
||||
|
||||
DATETIME_FIELDS = ("datetime", "timestamp", "time", "date", "stime")
|
||||
SYMBOL_FIELDS = ("symbol", "stock_code", "code", "stock")
|
||||
REQUIRED_PRICE_FIELDS = ("open", "high", "low", "close")
|
||||
|
||||
|
||||
def _first(row, names, default=None):
|
||||
for name in names:
|
||||
value = row.get(name)
|
||||
if value not in (None, ""):
|
||||
return value
|
||||
return default
|
||||
|
||||
|
||||
def parse_datetime(value):
|
||||
if isinstance(value, dt.datetime):
|
||||
return value.replace(tzinfo=None)
|
||||
if isinstance(value, dt.date):
|
||||
return dt.datetime.combine(value, dt.time())
|
||||
text = str(value or "").strip()
|
||||
if not text:
|
||||
raise ValueError("bar datetime is required")
|
||||
if text.isdigit():
|
||||
if len(text) == 8:
|
||||
return dt.datetime.strptime(text, "%Y%m%d")
|
||||
if len(text) == 14:
|
||||
return dt.datetime.strptime(text, "%Y%m%d%H%M%S")
|
||||
numeric = int(text)
|
||||
if numeric > 10 ** 12:
|
||||
numeric = numeric / 1000.0
|
||||
return dt.datetime.fromtimestamp(numeric)
|
||||
normalized = text.replace("T", " ").replace("Z", "").strip()
|
||||
from_isoformat = getattr(dt.datetime, "fromisoformat", None)
|
||||
if from_isoformat is not None:
|
||||
try:
|
||||
return from_isoformat(normalized).replace(tzinfo=None)
|
||||
except ValueError:
|
||||
pass
|
||||
for fmt in ("%Y-%m-%d %H:%M:%S", "%Y-%m-%d %H:%M", "%Y/%m/%d %H:%M:%S", "%Y-%m-%d"):
|
||||
try:
|
||||
return dt.datetime.strptime(normalized, fmt)
|
||||
except ValueError:
|
||||
continue
|
||||
raise ValueError("unsupported bar datetime: %s" % text)
|
||||
|
||||
|
||||
def _bool_value(value):
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
return str(value or "").strip().lower() in ("1", "true", "yes", "y", "on")
|
||||
|
||||
|
||||
def _optional_float(value):
|
||||
return None if value in (None, "") else float(value)
|
||||
|
||||
|
||||
def normalize_bar(row, default_symbol=""):
|
||||
timestamp = parse_datetime(_first(row, DATETIME_FIELDS))
|
||||
symbol = normalize_symbol(_first(row, SYMBOL_FIELDS, default_symbol))
|
||||
bar = {
|
||||
"datetime": timestamp.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"symbol": symbol,
|
||||
}
|
||||
for field in REQUIRED_PRICE_FIELDS:
|
||||
value = row.get(field)
|
||||
if value in (None, ""):
|
||||
raise ValueError("%s is required for %s at %s" % (field, symbol, bar["datetime"]))
|
||||
bar[field] = float(value)
|
||||
if bar[field] <= 0:
|
||||
raise ValueError("%s must be positive for %s at %s" % (field, symbol, bar["datetime"]))
|
||||
if bar["high"] < max(bar["open"], bar["close"], bar["low"]):
|
||||
raise ValueError("bar high is inconsistent for %s at %s" % (symbol, bar["datetime"]))
|
||||
if bar["low"] > min(bar["open"], bar["close"], bar["high"]):
|
||||
raise ValueError("bar low is inconsistent for %s at %s" % (symbol, bar["datetime"]))
|
||||
bar["volume"] = float(row.get("volume") or 0)
|
||||
bar["amount"] = float(row.get("amount") or 0)
|
||||
bar["prev_close"] = _optional_float(row.get("prev_close"))
|
||||
bar["up_limit"] = _optional_float(row.get("up_limit"))
|
||||
bar["down_limit"] = _optional_float(row.get("down_limit"))
|
||||
bar["suspended"] = _bool_value(row.get("suspended"))
|
||||
if row.get("price_limit_rate") not in (None, ""):
|
||||
bar["price_limit_rate"] = float(row["price_limit_rate"])
|
||||
return timestamp, bar
|
||||
|
||||
|
||||
class InMemoryBarFeed(object):
|
||||
def __init__(self, rows, source="memory", data_hash=None, default_symbol=""):
|
||||
normalized = []
|
||||
for row in rows:
|
||||
timestamp, bar = normalize_bar(dict(row), default_symbol=default_symbol)
|
||||
normalized.append((timestamp, bar))
|
||||
normalized.sort(key=lambda item: (item[0], item[1]["symbol"]))
|
||||
seen = set()
|
||||
frames = []
|
||||
current_timestamp = None
|
||||
current_bars = None
|
||||
previous_close = {}
|
||||
for timestamp, bar in normalized:
|
||||
identity = (timestamp, bar["symbol"])
|
||||
if identity in seen:
|
||||
raise ValueError("duplicate bar for %s at %s" % (bar["symbol"], bar["datetime"]))
|
||||
seen.add(identity)
|
||||
if bar["prev_close"] is None:
|
||||
bar["prev_close"] = previous_close.get(bar["symbol"])
|
||||
previous_close[bar["symbol"]] = bar["close"]
|
||||
if current_timestamp != timestamp:
|
||||
current_timestamp = timestamp
|
||||
current_bars = {}
|
||||
frames.append(
|
||||
{
|
||||
"datetime": timestamp.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"bars": current_bars,
|
||||
}
|
||||
)
|
||||
current_bars[bar["symbol"]] = bar
|
||||
if not frames:
|
||||
raise ValueError("historical data is empty")
|
||||
self._frames = frames
|
||||
self.source = str(source)
|
||||
if data_hash:
|
||||
self.data_hash = str(data_hash)
|
||||
else:
|
||||
payload = json.dumps(frames, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
||||
self.data_hash = hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||
|
||||
def __len__(self):
|
||||
return len(self._frames)
|
||||
|
||||
def frame(self, index):
|
||||
frame = self._frames[int(index)]
|
||||
return {"datetime": frame["datetime"], "bars": {key: dict(value) for key, value in frame["bars"].items()}}
|
||||
|
||||
def history(self, symbol, end_index, count=100, fields=None):
|
||||
symbol = normalize_symbol(symbol)
|
||||
end_index = min(int(end_index), len(self._frames) - 1)
|
||||
count = max(int(count or 0), 0)
|
||||
result = []
|
||||
for index in range(0, end_index + 1):
|
||||
bar = self._frames[index]["bars"].get(symbol)
|
||||
if bar is None:
|
||||
continue
|
||||
if fields:
|
||||
item = {"datetime": bar["datetime"], "symbol": symbol}
|
||||
for field in fields:
|
||||
if field in bar:
|
||||
item[str(field)] = bar[field]
|
||||
else:
|
||||
item = dict(bar)
|
||||
result.append(item)
|
||||
return result[-count:] if count else []
|
||||
|
||||
|
||||
class CsvBarFeed(InMemoryBarFeed):
|
||||
def __init__(self, path, default_symbol="", encoding="utf-8-sig"):
|
||||
absolute = os.path.abspath(path)
|
||||
with open(absolute, "rb") as handle:
|
||||
raw = handle.read()
|
||||
digest = hashlib.sha256(raw).hexdigest()
|
||||
text = raw.decode(encoding)
|
||||
reader = csv.DictReader(io.StringIO(text, newline=""))
|
||||
rows = list(reader)
|
||||
super(CsvBarFeed, self).__init__(
|
||||
rows,
|
||||
source=absolute,
|
||||
data_hash=digest,
|
||||
default_symbol=default_symbol,
|
||||
)
|
||||
|
||||
|
||||
class StreamingBarFeed(object):
|
||||
"""Thread-safe feed populated by QMT ``handlebar`` callbacks.
|
||||
|
||||
The external strategy can only read through ``frame``/``history`` with an
|
||||
engine-controlled end index, so bars already captured from QMT but not yet
|
||||
advanced to remain inaccessible.
|
||||
"""
|
||||
|
||||
def __init__(self, source="qmt_native_backtest"):
|
||||
self.source = str(source)
|
||||
self._frames = []
|
||||
self._seen = set()
|
||||
self._previous_close = {}
|
||||
self._condition = threading.Condition()
|
||||
self.closed = False
|
||||
|
||||
@property
|
||||
def data_hash(self):
|
||||
with self._condition:
|
||||
payload = json.dumps(
|
||||
self._frames,
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
)
|
||||
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||
|
||||
def __len__(self):
|
||||
with self._condition:
|
||||
return len(self._frames)
|
||||
|
||||
def append(self, row, default_symbol=""):
|
||||
timestamp, bar = normalize_bar(dict(row), default_symbol=default_symbol)
|
||||
identity = (timestamp, bar["symbol"])
|
||||
with self._condition:
|
||||
if identity in self._seen:
|
||||
return False
|
||||
if self.closed:
|
||||
raise RuntimeError("streaming feed is closed")
|
||||
self._seen.add(identity)
|
||||
if bar["prev_close"] is None:
|
||||
bar["prev_close"] = self._previous_close.get(bar["symbol"])
|
||||
self._previous_close[bar["symbol"]] = bar["close"]
|
||||
timestamp_text = timestamp.strftime("%Y-%m-%d %H:%M:%S")
|
||||
if self._frames and self._frames[-1]["datetime"] == timestamp_text:
|
||||
self._frames[-1]["bars"][bar["symbol"]] = bar
|
||||
elif self._frames and self._frames[-1]["datetime"] > timestamp_text:
|
||||
raise ValueError("streaming bars must be appended chronologically")
|
||||
else:
|
||||
self._frames.append({"datetime": timestamp_text, "bars": {bar["symbol"]: bar}})
|
||||
self._condition.notify_all()
|
||||
return True
|
||||
|
||||
def close(self):
|
||||
with self._condition:
|
||||
self.closed = True
|
||||
self._condition.notify_all()
|
||||
|
||||
def wait_for_index(self, index, timeout_seconds=None):
|
||||
index = int(index)
|
||||
with self._condition:
|
||||
if len(self._frames) > index:
|
||||
return True
|
||||
self._condition.wait_for(
|
||||
lambda: len(self._frames) > index or self.closed,
|
||||
timeout=timeout_seconds,
|
||||
)
|
||||
return len(self._frames) > index
|
||||
|
||||
def frame(self, index):
|
||||
with self._condition:
|
||||
frame = self._frames[int(index)]
|
||||
return {
|
||||
"datetime": frame["datetime"],
|
||||
"bars": {key: dict(value) for key, value in frame["bars"].items()},
|
||||
}
|
||||
|
||||
def history(self, symbol, end_index, count=100, fields=None):
|
||||
symbol = normalize_symbol(symbol)
|
||||
count = max(int(count or 0), 0)
|
||||
with self._condition:
|
||||
end_index = min(int(end_index), len(self._frames) - 1)
|
||||
frames = self._frames[: end_index + 1]
|
||||
result = []
|
||||
for frame in frames:
|
||||
bar = frame["bars"].get(symbol)
|
||||
if bar is None:
|
||||
continue
|
||||
if fields:
|
||||
item = {"datetime": bar["datetime"], "symbol": symbol}
|
||||
for field in fields:
|
||||
if field in bar:
|
||||
item[str(field)] = bar[field]
|
||||
else:
|
||||
item = dict(bar)
|
||||
result.append(item)
|
||||
return result[-count:] if count else []
|
||||
@@ -1,372 +0,0 @@
|
||||
"""Deterministic bar-by-bar backtest engine."""
|
||||
|
||||
import csv
|
||||
import datetime as dt
|
||||
import hashlib
|
||||
import json
|
||||
import math
|
||||
import os
|
||||
|
||||
from .broker import ACTIVE_ORDER_STATUSES, SimulatedBroker
|
||||
from .models import decimal_value, normalize_symbol
|
||||
|
||||
|
||||
ENGINE_VERSION = "1.0.0"
|
||||
PROTOCOL_VERSION = 1
|
||||
|
||||
|
||||
class BacktestConfig(object):
|
||||
def __init__(
|
||||
self,
|
||||
run_id,
|
||||
output_dir,
|
||||
initial_cash=1000000,
|
||||
initial_positions=None,
|
||||
buy_commission_rate=0.0003,
|
||||
sell_commission_rate=0.0003,
|
||||
min_commission=5,
|
||||
stamp_tax_rate=0.0005,
|
||||
transfer_fee_rate=0.00001,
|
||||
slippage_bps=0,
|
||||
max_volume_participation=0.1,
|
||||
price_limit_rate=0.10,
|
||||
lot_size=100,
|
||||
time_in_force="NEXT_BAR",
|
||||
seed=0,
|
||||
strategy_name="external_zmq_strategy",
|
||||
parameters=None,
|
||||
fee_schedule="a_share_2023_08_28",
|
||||
market_rules_version="a_share_v1",
|
||||
):
|
||||
self.run_id = str(run_id or "").strip()
|
||||
if not self.run_id:
|
||||
raise ValueError("run_id is required")
|
||||
self.output_dir = os.path.abspath(output_dir)
|
||||
self.initial_cash = decimal_value(initial_cash)
|
||||
if self.initial_cash < 0:
|
||||
raise ValueError("initial_cash cannot be negative")
|
||||
self.initial_positions = dict(initial_positions or {})
|
||||
self.buy_commission_rate = decimal_value(buy_commission_rate)
|
||||
self.sell_commission_rate = decimal_value(sell_commission_rate)
|
||||
self.min_commission = decimal_value(min_commission)
|
||||
self.stamp_tax_rate = decimal_value(stamp_tax_rate)
|
||||
self.transfer_fee_rate = decimal_value(transfer_fee_rate)
|
||||
self.slippage_bps = decimal_value(slippage_bps)
|
||||
self.max_volume_participation = float(max_volume_participation)
|
||||
if not 0 < self.max_volume_participation <= 1:
|
||||
raise ValueError("max_volume_participation must be in (0, 1]")
|
||||
self.price_limit_rate = decimal_value(price_limit_rate)
|
||||
self.lot_size = int(lot_size)
|
||||
if self.lot_size <= 0:
|
||||
raise ValueError("lot_size must be positive")
|
||||
self.time_in_force = str(time_in_force or "NEXT_BAR").upper()
|
||||
if self.time_in_force not in ("NEXT_BAR", "DAY"):
|
||||
raise ValueError("time_in_force must be NEXT_BAR or DAY")
|
||||
self.seed = int(seed)
|
||||
self.strategy_name = str(strategy_name or "external_zmq_strategy")
|
||||
self.parameters = dict(parameters or {})
|
||||
self.fee_schedule = str(fee_schedule or "custom")
|
||||
self.market_rules_version = str(market_rules_version or "custom")
|
||||
|
||||
def to_dict(self, include_paths=True, include_identity=True):
|
||||
payload = {
|
||||
"initial_cash": float(self.initial_cash),
|
||||
"initial_positions": self.initial_positions,
|
||||
"buy_commission_rate": float(self.buy_commission_rate),
|
||||
"sell_commission_rate": float(self.sell_commission_rate),
|
||||
"min_commission": float(self.min_commission),
|
||||
"stamp_tax_rate": float(self.stamp_tax_rate),
|
||||
"transfer_fee_rate": float(self.transfer_fee_rate),
|
||||
"slippage_bps": float(self.slippage_bps),
|
||||
"max_volume_participation": self.max_volume_participation,
|
||||
"price_limit_rate": float(self.price_limit_rate),
|
||||
"lot_size": self.lot_size,
|
||||
"time_in_force": self.time_in_force,
|
||||
"seed": self.seed,
|
||||
"strategy_name": self.strategy_name,
|
||||
"parameters": self.parameters,
|
||||
"fee_schedule": self.fee_schedule,
|
||||
"market_rules_version": self.market_rules_version,
|
||||
}
|
||||
if include_identity:
|
||||
payload["run_id"] = self.run_id
|
||||
if include_paths:
|
||||
payload["output_dir"] = self.output_dir
|
||||
return payload
|
||||
|
||||
|
||||
class BacktestEngine(object):
|
||||
def __init__(self, feed, config):
|
||||
self.feed = feed
|
||||
self.config = config
|
||||
self.created_at = dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
self.started = False
|
||||
self.finished = False
|
||||
self.current_index = -1
|
||||
self.current_frame = None
|
||||
self.last_fills = []
|
||||
self.equity_curve = []
|
||||
self.position_rows = []
|
||||
self._result = None
|
||||
self.broker = SimulatedBroker(config)
|
||||
|
||||
def _record_state(self):
|
||||
snapshot = self.broker.snapshot(self.current_frame["bars"])
|
||||
equity_row = {
|
||||
"frame_index": self.current_index,
|
||||
"datetime": self.current_frame["datetime"],
|
||||
"cash": snapshot["cash"],
|
||||
"market_value": snapshot["market_value"],
|
||||
"total_asset": snapshot["total_asset"],
|
||||
}
|
||||
self.equity_curve.append(equity_row)
|
||||
for symbol, position in snapshot["positions"].items():
|
||||
row = {"frame_index": self.current_index, "datetime": self.current_frame["datetime"]}
|
||||
row.update(position)
|
||||
self.position_rows.append(row)
|
||||
|
||||
def start(self):
|
||||
if self.started:
|
||||
return self.state()
|
||||
self.started = True
|
||||
self.current_index = 0
|
||||
self.current_frame = self.feed.frame(0)
|
||||
self.broker._settle_trading_day(self.current_frame["datetime"][:10])
|
||||
self._record_state()
|
||||
return self.state()
|
||||
|
||||
def _require_started(self):
|
||||
if not self.started:
|
||||
raise RuntimeError("backtest has not started")
|
||||
if self.finished:
|
||||
raise RuntimeError("backtest is already finished")
|
||||
|
||||
def submit_order(self, payload):
|
||||
self._require_started()
|
||||
return self.broker.submit(
|
||||
payload,
|
||||
frame_index=self.current_index,
|
||||
submitted_at=self.current_frame["datetime"],
|
||||
).to_dict()
|
||||
|
||||
def cancel_order(self, order_id):
|
||||
self._require_started()
|
||||
return self.broker.cancel(order_id).to_dict()
|
||||
|
||||
def next_bar(self):
|
||||
self._require_started()
|
||||
if self.current_index >= len(self.feed) - 1:
|
||||
return self.state()
|
||||
self.current_index += 1
|
||||
self.current_frame = self.feed.frame(self.current_index)
|
||||
self.last_fills = self.broker.advance(self.current_index, self.current_frame)
|
||||
self._record_state()
|
||||
return self.state()
|
||||
|
||||
def history(self, symbol, count=100, fields=None):
|
||||
self._require_started()
|
||||
return self.feed.history(normalize_symbol(symbol), self.current_index, count=count, fields=fields)
|
||||
|
||||
def orders(self):
|
||||
return self.broker.orders()
|
||||
|
||||
def fills(self):
|
||||
return self.broker.fills()
|
||||
|
||||
def state(self):
|
||||
if not self.started:
|
||||
return {
|
||||
"run_id": self.config.run_id,
|
||||
"started": False,
|
||||
"finished": self.finished,
|
||||
"done": False,
|
||||
"frame_index": -1,
|
||||
"frame_count": len(self.feed),
|
||||
}
|
||||
portfolio = self.broker.snapshot(self.current_frame["bars"])
|
||||
return {
|
||||
"run_id": self.config.run_id,
|
||||
"started": True,
|
||||
"finished": self.finished,
|
||||
"done": self.current_index >= len(self.feed) - 1,
|
||||
"frame_index": self.current_index,
|
||||
"frame_count": len(self.feed),
|
||||
"datetime": self.current_frame["datetime"],
|
||||
"bars": {key: dict(value) for key, value in self.current_frame["bars"].items()},
|
||||
"fills": [fill.to_dict() for fill in self.last_fills],
|
||||
"cash": portfolio["cash"],
|
||||
"market_value": portfolio["market_value"],
|
||||
"total_asset": portfolio["total_asset"],
|
||||
"positions": portfolio["positions"],
|
||||
"total_fees": portfolio["total_fees"],
|
||||
"turnover": portfolio["turnover"],
|
||||
}
|
||||
|
||||
def _metrics(self):
|
||||
assets = [float(row["total_asset"]) for row in self.equity_curve]
|
||||
initial = assets[0] if assets else float(self.config.initial_cash)
|
||||
final = assets[-1] if assets else initial
|
||||
peak = None
|
||||
max_drawdown = 0.0
|
||||
for value in assets:
|
||||
peak = value if peak is None else max(peak, value)
|
||||
if peak > 0:
|
||||
max_drawdown = min(max_drawdown, value / peak - 1.0)
|
||||
dates = sorted(set(row["datetime"][:10] for row in self.equity_curve))
|
||||
total_return = 0.0 if initial == 0 else final / initial - 1.0
|
||||
annualized = None
|
||||
if len(dates) > 1 and initial > 0 and final > 0:
|
||||
annualized = math.pow(final / initial, 252.0 / len(dates)) - 1.0
|
||||
filled_orders = len([order for order in self.broker.orders_list if order.filled_quantity > 0])
|
||||
rejected_orders = len([order for order in self.broker.orders_list if order.status == "REJECTED"])
|
||||
return {
|
||||
"initial_total_asset": round(initial, 2),
|
||||
"final_total_asset": round(final, 2),
|
||||
"total_return": round(total_return, 10),
|
||||
"annualized_return": None if annualized is None else round(annualized, 10),
|
||||
"max_drawdown": round(-max_drawdown, 10),
|
||||
"trading_days": len(dates),
|
||||
"bar_count": len(self.equity_curve),
|
||||
"order_count": len(self.broker.orders_list),
|
||||
"filled_order_count": filled_orders,
|
||||
"rejected_order_count": rejected_orders,
|
||||
"fill_count": len(self.broker.fills_list),
|
||||
"total_fees": round(float(self.broker.total_fees), 2),
|
||||
"turnover": round(float(self.broker.turnover), 2),
|
||||
}
|
||||
|
||||
def _signature_payload(self, metrics):
|
||||
orders = []
|
||||
for item in self.orders():
|
||||
clean = dict(item)
|
||||
clean.pop("client_order_id", None)
|
||||
orders.append(clean)
|
||||
return {
|
||||
"engine_version": ENGINE_VERSION,
|
||||
"data_hash": self.feed.data_hash,
|
||||
"config": self.config.to_dict(include_paths=False, include_identity=False),
|
||||
"orders": orders,
|
||||
"fills": self.fills(),
|
||||
"equity": self.equity_curve,
|
||||
"metrics": metrics,
|
||||
}
|
||||
|
||||
def finish(self):
|
||||
if self._result is not None:
|
||||
return dict(self._result)
|
||||
if not self.started:
|
||||
self.start()
|
||||
self.broker.expire_open_orders()
|
||||
metrics = self._metrics()
|
||||
signature_json = json.dumps(
|
||||
self._signature_payload(metrics), ensure_ascii=False, sort_keys=True, separators=(",", ":")
|
||||
)
|
||||
signature = hashlib.sha256(signature_json.encode("utf-8")).hexdigest()
|
||||
self.finished = True
|
||||
final_state = self.state()
|
||||
self._result = {
|
||||
"schema_version": 1,
|
||||
"engine_version": ENGINE_VERSION,
|
||||
"run_id": self.config.run_id,
|
||||
"strategy_name": self.config.strategy_name,
|
||||
"data_hash": self.feed.data_hash,
|
||||
"deterministic_signature": signature,
|
||||
"metrics": metrics,
|
||||
"final_state": final_state,
|
||||
}
|
||||
self._write_artifacts()
|
||||
return dict(self._result)
|
||||
|
||||
@staticmethod
|
||||
def _write_json(path, payload):
|
||||
with open(path, "w", encoding="utf-8", newline="\n") as handle:
|
||||
json.dump(payload, handle, ensure_ascii=False, sort_keys=True, indent=2)
|
||||
handle.write("\n")
|
||||
|
||||
@staticmethod
|
||||
def _write_csv(path, rows, fieldnames):
|
||||
with open(path, "w", encoding="utf-8-sig", newline="") as handle:
|
||||
writer = csv.DictWriter(handle, fieldnames=fieldnames, extrasaction="ignore")
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
def _write_artifacts(self):
|
||||
output_dir = self.config.output_dir
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
meta = {
|
||||
"schema_version": 1,
|
||||
"engine_version": ENGINE_VERSION,
|
||||
"protocol_version": PROTOCOL_VERSION,
|
||||
"run_id": self.config.run_id,
|
||||
"created_at": self.created_at,
|
||||
"data_source": self.feed.source,
|
||||
"data_hash": self.feed.data_hash,
|
||||
"frame_count": len(self.feed),
|
||||
"config": self.config.to_dict(),
|
||||
"live_ready": False,
|
||||
"execution_channel": "backtest_zmq_only",
|
||||
}
|
||||
self._write_json(os.path.join(output_dir, "meta.json"), meta)
|
||||
self._write_json(os.path.join(output_dir, "result.json"), self._result)
|
||||
self._write_csv(
|
||||
os.path.join(output_dir, "orders.csv"),
|
||||
self.orders(),
|
||||
(
|
||||
"order_id", "client_order_id", "symbol", "side", "quantity", "filled_quantity",
|
||||
"remaining_quantity", "order_type", "limit_price", "submitted_index", "submitted_at",
|
||||
"time_in_force", "status", "reject_reason",
|
||||
),
|
||||
)
|
||||
self._write_csv(
|
||||
os.path.join(output_dir, "fills.csv"),
|
||||
self.fills(),
|
||||
(
|
||||
"fill_id", "order_id", "client_order_id", "symbol", "side", "quantity", "price",
|
||||
"amount", "commission", "stamp_tax", "transfer_fee", "total_fee", "filled_index", "filled_at",
|
||||
),
|
||||
)
|
||||
self._write_csv(
|
||||
os.path.join(output_dir, "equity.csv"),
|
||||
self.equity_curve,
|
||||
("frame_index", "datetime", "cash", "market_value", "total_asset"),
|
||||
)
|
||||
self._write_csv(
|
||||
os.path.join(output_dir, "positions.csv"),
|
||||
self.position_rows,
|
||||
(
|
||||
"frame_index", "datetime", "symbol", "quantity", "available", "today_buy", "avg_cost",
|
||||
"realized_pnl", "mark_price", "market_value",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class StreamingBacktestEngine(BacktestEngine):
|
||||
"""Backtest engine whose bars arrive from a QMT backtest callback thread."""
|
||||
|
||||
def __init__(self, feed, config, bar_wait_timeout_seconds=60.0):
|
||||
super(StreamingBacktestEngine, self).__init__(feed, config)
|
||||
self.bar_wait_timeout_seconds = float(bar_wait_timeout_seconds)
|
||||
|
||||
def start(self):
|
||||
if not self.started and not self.feed.wait_for_index(0, self.bar_wait_timeout_seconds):
|
||||
raise TimeoutError("timed out waiting for the first QMT backtest bar")
|
||||
return super(StreamingBacktestEngine, self).start()
|
||||
|
||||
def next_bar(self):
|
||||
self._require_started()
|
||||
target = self.current_index + 1
|
||||
if not self.feed.wait_for_index(target, self.bar_wait_timeout_seconds):
|
||||
if self.feed.closed:
|
||||
return self.state()
|
||||
raise TimeoutError("timed out waiting for QMT backtest bar index %d" % target)
|
||||
self.current_index = target
|
||||
self.current_frame = self.feed.frame(target)
|
||||
self.last_fills = self.broker.advance(self.current_index, self.current_frame)
|
||||
self._record_state()
|
||||
return self.state()
|
||||
|
||||
def state(self):
|
||||
state = super(StreamingBacktestEngine, self).state()
|
||||
if state.get("started"):
|
||||
state["done"] = bool(self.feed.closed and self.current_index >= len(self.feed) - 1)
|
||||
return state
|
||||
@@ -1,179 +0,0 @@
|
||||
"""Backtest-only domain models with JSON-safe serialization."""
|
||||
|
||||
from decimal import Decimal, ROUND_HALF_UP
|
||||
|
||||
|
||||
ZERO = Decimal("0")
|
||||
MONEY_QUANT = Decimal("0.01")
|
||||
|
||||
|
||||
def decimal_value(value, default="0"):
|
||||
if value in (None, ""):
|
||||
value = default
|
||||
if isinstance(value, Decimal):
|
||||
return value
|
||||
return Decimal(str(value))
|
||||
|
||||
|
||||
def money(value):
|
||||
return decimal_value(value).quantize(MONEY_QUANT, rounding=ROUND_HALF_UP)
|
||||
|
||||
|
||||
def json_number(value, digits=None):
|
||||
if value is None:
|
||||
return None
|
||||
number = float(value)
|
||||
return round(number, digits) if digits is not None else number
|
||||
|
||||
|
||||
def normalize_symbol(value):
|
||||
text = str(value or "").strip().upper()
|
||||
if not text:
|
||||
raise ValueError("symbol is required")
|
||||
if "." in text:
|
||||
pure, market = text.rsplit(".", 1)
|
||||
if pure and market in ("SH", "SZ", "BJ"):
|
||||
return "%s.%s" % (pure, market)
|
||||
return text
|
||||
if text.isdigit() and len(text) == 6:
|
||||
if text.startswith(("4", "8")):
|
||||
return text + ".BJ"
|
||||
if text.startswith(("5", "6", "9")):
|
||||
return text + ".SH"
|
||||
return text + ".SZ"
|
||||
return text
|
||||
|
||||
|
||||
def price_precision(symbol):
|
||||
pure = normalize_symbol(symbol).split(".", 1)[0]
|
||||
return 3 if pure.startswith(("15", "16", "50", "51", "52", "56", "58")) else 2
|
||||
|
||||
|
||||
def price_quant(symbol):
|
||||
return Decimal("0.001") if price_precision(symbol) == 3 else Decimal("0.01")
|
||||
|
||||
|
||||
def round_price(symbol, value):
|
||||
return decimal_value(value).quantize(price_quant(symbol), rounding=ROUND_HALF_UP)
|
||||
|
||||
|
||||
class Position(object):
|
||||
def __init__(self, symbol, quantity=0, available=None, today_buy=0, avg_cost=0, realized_pnl=0):
|
||||
self.symbol = normalize_symbol(symbol)
|
||||
self.quantity = int(quantity or 0)
|
||||
self.available = self.quantity if available is None else int(available or 0)
|
||||
self.today_buy = int(today_buy or 0)
|
||||
self.avg_cost = decimal_value(avg_cost)
|
||||
self.realized_pnl = decimal_value(realized_pnl)
|
||||
|
||||
def to_dict(self, mark_price=None):
|
||||
market_value = None if mark_price is None else money(decimal_value(mark_price) * self.quantity)
|
||||
return {
|
||||
"symbol": self.symbol,
|
||||
"quantity": self.quantity,
|
||||
"available": self.available,
|
||||
"today_buy": self.today_buy,
|
||||
"avg_cost": json_number(self.avg_cost, 6),
|
||||
"realized_pnl": json_number(self.realized_pnl, 2),
|
||||
"mark_price": json_number(mark_price, 6),
|
||||
"market_value": json_number(market_value, 2),
|
||||
}
|
||||
|
||||
|
||||
class BacktestOrder(object):
|
||||
def __init__(
|
||||
self,
|
||||
order_id,
|
||||
client_order_id,
|
||||
symbol,
|
||||
side,
|
||||
quantity,
|
||||
order_type,
|
||||
limit_price,
|
||||
submitted_index,
|
||||
submitted_at,
|
||||
time_in_force="NEXT_BAR",
|
||||
):
|
||||
self.order_id = str(order_id)
|
||||
self.client_order_id = str(client_order_id or "")
|
||||
self.symbol = normalize_symbol(symbol)
|
||||
self.side = str(side).upper()
|
||||
self.quantity = int(quantity)
|
||||
self.filled_quantity = 0
|
||||
self.order_type = str(order_type).upper()
|
||||
self.limit_price = decimal_value(limit_price) if limit_price not in (None, "") else None
|
||||
self.submitted_index = int(submitted_index)
|
||||
self.submitted_at = str(submitted_at)
|
||||
self.time_in_force = str(time_in_force or "NEXT_BAR").upper()
|
||||
self.status = "PENDING"
|
||||
self.reject_reason = ""
|
||||
self.last_attempt_index = None
|
||||
|
||||
@property
|
||||
def remaining(self):
|
||||
return max(self.quantity - self.filled_quantity, 0)
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
"order_id": self.order_id,
|
||||
"client_order_id": self.client_order_id,
|
||||
"symbol": self.symbol,
|
||||
"side": self.side,
|
||||
"quantity": self.quantity,
|
||||
"filled_quantity": self.filled_quantity,
|
||||
"remaining_quantity": self.remaining,
|
||||
"order_type": self.order_type,
|
||||
"limit_price": json_number(self.limit_price, 6),
|
||||
"submitted_index": self.submitted_index,
|
||||
"submitted_at": self.submitted_at,
|
||||
"time_in_force": self.time_in_force,
|
||||
"status": self.status,
|
||||
"reject_reason": self.reject_reason,
|
||||
}
|
||||
|
||||
|
||||
class BacktestFill(object):
|
||||
def __init__(
|
||||
self,
|
||||
fill_id,
|
||||
order,
|
||||
quantity,
|
||||
price,
|
||||
commission,
|
||||
stamp_tax,
|
||||
transfer_fee,
|
||||
filled_index,
|
||||
filled_at,
|
||||
):
|
||||
self.fill_id = str(fill_id)
|
||||
self.order_id = order.order_id
|
||||
self.client_order_id = order.client_order_id
|
||||
self.symbol = order.symbol
|
||||
self.side = order.side
|
||||
self.quantity = int(quantity)
|
||||
self.price = decimal_value(price)
|
||||
self.amount = money(self.price * self.quantity)
|
||||
self.commission = money(commission)
|
||||
self.stamp_tax = money(stamp_tax)
|
||||
self.transfer_fee = money(transfer_fee)
|
||||
self.total_fee = money(self.commission + self.stamp_tax + self.transfer_fee)
|
||||
self.filled_index = int(filled_index)
|
||||
self.filled_at = str(filled_at)
|
||||
|
||||
def to_dict(self):
|
||||
return {
|
||||
"fill_id": self.fill_id,
|
||||
"order_id": self.order_id,
|
||||
"client_order_id": self.client_order_id,
|
||||
"symbol": self.symbol,
|
||||
"side": self.side,
|
||||
"quantity": self.quantity,
|
||||
"price": json_number(self.price, 6),
|
||||
"amount": json_number(self.amount, 2),
|
||||
"commission": json_number(self.commission, 2),
|
||||
"stamp_tax": json_number(self.stamp_tax, 2),
|
||||
"transfer_fee": json_number(self.transfer_fee, 2),
|
||||
"total_fee": json_number(self.total_fee, 2),
|
||||
"filled_index": self.filled_index,
|
||||
"filled_at": self.filled_at,
|
||||
}
|
||||
@@ -1,137 +0,0 @@
|
||||
"""Versioned request/response protocol for the backtest-only ZMQ bridge."""
|
||||
|
||||
import datetime as dt
|
||||
import json
|
||||
|
||||
|
||||
SCHEMA_VERSION = 1
|
||||
|
||||
|
||||
class BacktestBridgeProtocol(object):
|
||||
def __init__(self, engine, request_cache_size=10000):
|
||||
self.engine = engine
|
||||
self.request_cache_size = int(request_cache_size)
|
||||
self.client_id = None
|
||||
self._responses = {}
|
||||
self._request_fingerprints = {}
|
||||
self._response_order = []
|
||||
|
||||
def _response(self, request, ok, data=None, error=""):
|
||||
execution_backend = str(getattr(self.engine, "execution_backend", "LOCAL_SIM"))
|
||||
return {
|
||||
"schema_version": SCHEMA_VERSION,
|
||||
"request_id": str(request.get("request_id") or ""),
|
||||
"run_id": self.engine.config.run_id,
|
||||
"client_id": str(request.get("client_id") or ""),
|
||||
"method": str(request.get("method") or ""),
|
||||
"ok": bool(ok),
|
||||
"data": data,
|
||||
"error": str(error or ""),
|
||||
"execution_mode": "QMT_BACKTEST" if execution_backend == "QMT_NATIVE" else "BACKTEST",
|
||||
"execution_backend": execution_backend,
|
||||
"live_ready": False,
|
||||
"handled_at": dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _fingerprint(request):
|
||||
return json.dumps(request, ensure_ascii=False, sort_keys=True, separators=(",", ":"))
|
||||
|
||||
def _remember(self, request_id, request, response):
|
||||
self._responses[request_id] = response
|
||||
self._request_fingerprints[request_id] = self._fingerprint(request)
|
||||
self._response_order.append(request_id)
|
||||
while len(self._response_order) > self.request_cache_size:
|
||||
oldest = self._response_order.pop(0)
|
||||
self._responses.pop(oldest, None)
|
||||
self._request_fingerprints.pop(oldest, None)
|
||||
|
||||
def _validate(self, request):
|
||||
if not isinstance(request, dict):
|
||||
raise ValueError("request must be a JSON object")
|
||||
if int(request.get("schema_version") or 0) != SCHEMA_VERSION:
|
||||
raise ValueError("unsupported schema_version")
|
||||
if not str(request.get("request_id") or ""):
|
||||
raise ValueError("request_id is required")
|
||||
method = str(request.get("method") or "").lower()
|
||||
requested_run_id = str(request.get("run_id") or "")
|
||||
discovery = method in ("ping", "describe") and not requested_run_id
|
||||
if not discovery and requested_run_id != self.engine.config.run_id:
|
||||
raise ValueError("run_id mismatch")
|
||||
if not str(request.get("client_id") or ""):
|
||||
raise ValueError("client_id is required")
|
||||
if not str(request.get("method") or ""):
|
||||
raise ValueError("method is required")
|
||||
|
||||
def _claim_or_check_client(self, request):
|
||||
client_id = str(request["client_id"])
|
||||
method = str(request["method"]).lower()
|
||||
if self.client_id is None and method == "start":
|
||||
self.client_id = client_id
|
||||
if method not in ("ping", "describe") and self.client_id != client_id:
|
||||
raise PermissionError("run is owned by another client_id")
|
||||
|
||||
def _dispatch(self, request):
|
||||
method = str(request["method"]).lower()
|
||||
params = dict(request.get("params") or {})
|
||||
if method == "ping":
|
||||
return {
|
||||
"status": "ok",
|
||||
"started": self.engine.started,
|
||||
"finished": self.engine.finished,
|
||||
}
|
||||
if method == "describe":
|
||||
execution_backend = str(getattr(self.engine, "execution_backend", "LOCAL_SIM"))
|
||||
return {
|
||||
"schema_version": SCHEMA_VERSION,
|
||||
"run_id": self.engine.config.run_id,
|
||||
"engine_version": str(getattr(self.engine, "engine_version", "1.0.0")),
|
||||
"execution_backend": execution_backend,
|
||||
"engine_owner": "QMT" if execution_backend == "QMT_NATIVE" else "LOCAL",
|
||||
"matching_owner": "QMT" if execution_backend == "QMT_NATIVE" else "LOCAL",
|
||||
"methods": [
|
||||
"ping", "describe", "start", "next_bar", "submit_order", "cancel_order",
|
||||
"state", "history", "orders", "fills", "finish",
|
||||
],
|
||||
"fill_timing": str(getattr(self.engine, "fill_timing", "next_symbol_bar")),
|
||||
"live_ready": False,
|
||||
}
|
||||
if method == "start":
|
||||
return self.engine.start()
|
||||
if method == "next_bar":
|
||||
return self.engine.next_bar()
|
||||
if method == "submit_order":
|
||||
return self.engine.submit_order(params)
|
||||
if method == "cancel_order":
|
||||
return self.engine.cancel_order(params.get("order_id"))
|
||||
if method == "state":
|
||||
return self.engine.state()
|
||||
if method == "history":
|
||||
return self.engine.history(
|
||||
params.get("symbol"),
|
||||
count=params.get("count", 100),
|
||||
fields=params.get("fields"),
|
||||
)
|
||||
if method == "orders":
|
||||
return self.engine.orders()
|
||||
if method == "fills":
|
||||
return self.engine.fills()
|
||||
if method == "finish":
|
||||
return self.engine.finish()
|
||||
raise ValueError("unsupported method: %s" % method)
|
||||
|
||||
def handle(self, request):
|
||||
request_id = str((request or {}).get("request_id") or "")
|
||||
if request_id and request_id in self._responses:
|
||||
if self._request_fingerprints.get(request_id) != self._fingerprint(request):
|
||||
return self._response(request, False, None, "request_id reused with different payload")
|
||||
return self._responses[request_id]
|
||||
try:
|
||||
self._validate(request)
|
||||
self._claim_or_check_client(request)
|
||||
response = self._response(request, True, self._dispatch(request))
|
||||
except Exception as exc:
|
||||
response = self._response(request or {}, False, None, "%s: %s" % (exc.__class__.__name__, exc))
|
||||
if request_id:
|
||||
self._remember(request_id, request, response)
|
||||
return response
|
||||
@@ -1,718 +0,0 @@
|
||||
"""QMT-native backtest service exposed to external strategies over ZMQ.
|
||||
|
||||
QMT remains the only backtest engine and matching system. The ZMQ listener
|
||||
thread only queues commands; every QMT API call is executed by ``handlebar`` on
|
||||
QMT's callback thread.
|
||||
"""
|
||||
|
||||
import datetime as dt
|
||||
import threading
|
||||
import uuid
|
||||
|
||||
from .data_feed import StreamingBarFeed, parse_datetime
|
||||
from .models import normalize_symbol
|
||||
from .protocol import BacktestBridgeProtocol
|
||||
from .zmq_server import ZmqBacktestServer
|
||||
|
||||
|
||||
_CONFIG = {}
|
||||
_QMT_API = {}
|
||||
_RUNTIME = None
|
||||
|
||||
|
||||
def configure(**kwargs):
|
||||
_CONFIG.update(kwargs)
|
||||
|
||||
|
||||
def bind_qmt_api(passorder_func=None, cancel_func=None, get_trade_detail_data_func=None):
|
||||
if passorder_func is not None:
|
||||
_QMT_API["passorder"] = passorder_func
|
||||
if cancel_func is not None:
|
||||
_QMT_API["cancel"] = cancel_func
|
||||
if get_trade_detail_data_func is not None:
|
||||
_QMT_API["get_trade_detail_data"] = get_trade_detail_data_func
|
||||
|
||||
|
||||
def _sequence(value):
|
||||
if value is None:
|
||||
return []
|
||||
if isinstance(value, dict):
|
||||
for item in value.values():
|
||||
result = _sequence(item)
|
||||
if result:
|
||||
return result
|
||||
return []
|
||||
if isinstance(value, (list, tuple)):
|
||||
return list(value)
|
||||
if hasattr(value, "tolist"):
|
||||
try:
|
||||
result = value.tolist()
|
||||
return result if isinstance(result, list) else [result]
|
||||
except Exception:
|
||||
pass
|
||||
if hasattr(value, "values"):
|
||||
try:
|
||||
return list(value.values)
|
||||
except Exception:
|
||||
pass
|
||||
return [value]
|
||||
|
||||
|
||||
def _last_value(value):
|
||||
values = _sequence(value)
|
||||
return values[-1] if values else None
|
||||
|
||||
|
||||
def _attr(value, names, default=None):
|
||||
for name in names:
|
||||
if isinstance(value, dict) and name in value:
|
||||
result = value.get(name)
|
||||
else:
|
||||
result = getattr(value, name, None)
|
||||
if result is not None:
|
||||
return result
|
||||
return default
|
||||
|
||||
|
||||
def _json_number(value):
|
||||
if value in (None, ""):
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def _full_symbol(row):
|
||||
code = str(_attr(row, ("m_strInstrumentID", "instrument_id", "stock_code", "symbol"), "") or "")
|
||||
market = str(_attr(row, ("m_strExchangeID", "exchange_id", "market"), "") or "").upper()
|
||||
if "." not in code and market in ("SH", "SZ", "BJ"):
|
||||
code = code + "." + market
|
||||
return normalize_symbol(code) if code else ""
|
||||
|
||||
|
||||
def _side_from_offset(value):
|
||||
try:
|
||||
return "BUY" if int(value or 0) == 48 else "SELL"
|
||||
except (TypeError, ValueError):
|
||||
return str(value or "")
|
||||
|
||||
|
||||
def _is_qmt_backtest(context):
|
||||
value = getattr(context, "do_back_test", None)
|
||||
if callable(value):
|
||||
try:
|
||||
value = value()
|
||||
except Exception:
|
||||
value = None
|
||||
if bool(value):
|
||||
return True
|
||||
for name in ("is_backtest", "is_back_test", "backtest"):
|
||||
value = getattr(context, name, None)
|
||||
if callable(value):
|
||||
try:
|
||||
value = value()
|
||||
except Exception:
|
||||
value = None
|
||||
if bool(value):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
class QmtBarExtractor(object):
|
||||
def __init__(self):
|
||||
self.previous_close = {}
|
||||
|
||||
@staticmethod
|
||||
def _symbol(context):
|
||||
raw = ""
|
||||
for name in ("stock", "symbol", "stockcode"):
|
||||
value = getattr(context, name, None)
|
||||
if value:
|
||||
raw = str(value)
|
||||
break
|
||||
if not raw:
|
||||
raise ValueError("QMT ContextInfo has no stock symbol")
|
||||
if "." not in raw:
|
||||
market = str(getattr(context, "market", "") or "").upper()
|
||||
if market in ("SH", "SZ", "BJ"):
|
||||
raw = raw + "." + market
|
||||
return normalize_symbol(raw)
|
||||
|
||||
@staticmethod
|
||||
def _timestamp(context):
|
||||
barpos = getattr(context, "barpos", getattr(context, "bar_index", None))
|
||||
getter = getattr(context, "get_bar_timetag", None)
|
||||
if callable(getter) and barpos is not None:
|
||||
return parse_datetime(getter(barpos)).strftime("%Y-%m-%d %H:%M:%S")
|
||||
for name in ("bar_time", "datetime", "timestamp"):
|
||||
value = getattr(context, name, None)
|
||||
if value not in (None, ""):
|
||||
return parse_datetime(value).strftime("%Y-%m-%d %H:%M:%S")
|
||||
raise ValueError("QMT ContextInfo has no deterministic bar timestamp")
|
||||
|
||||
@staticmethod
|
||||
def _periods(context):
|
||||
result = []
|
||||
for value in (getattr(context, "period", None), "1m", "1d"):
|
||||
text = str(value or "").strip()
|
||||
if text and text not in result:
|
||||
result.append(text)
|
||||
return result
|
||||
|
||||
def _history_value(self, context, field):
|
||||
getter = getattr(context, "get_history_data", None)
|
||||
if not callable(getter):
|
||||
return None
|
||||
for period in self._periods(context):
|
||||
for call in (
|
||||
lambda period=period: getter(1, period, field),
|
||||
lambda period=period: getter(field, 1, period),
|
||||
lambda: getter(field, 1),
|
||||
):
|
||||
try:
|
||||
value = _last_value(call())
|
||||
if value not in (None, ""):
|
||||
return value
|
||||
except Exception:
|
||||
continue
|
||||
return None
|
||||
|
||||
def _field(self, context, field, aliases=()):
|
||||
for name in (field,) + tuple(aliases):
|
||||
value = _last_value(getattr(context, name, None))
|
||||
if value not in (None, ""):
|
||||
return value
|
||||
for name in (field,) + tuple(aliases):
|
||||
value = self._history_value(context, name)
|
||||
if value not in (None, ""):
|
||||
return value
|
||||
return None
|
||||
|
||||
def extract(self, context):
|
||||
symbol = self._symbol(context)
|
||||
close = self._field(context, "close")
|
||||
row = {
|
||||
"datetime": self._timestamp(context),
|
||||
"symbol": symbol,
|
||||
"open": self._field(context, "open"),
|
||||
"high": self._field(context, "high"),
|
||||
"low": self._field(context, "low"),
|
||||
"close": close,
|
||||
"volume": self._field(context, "volume", ("vol",)) or 0,
|
||||
"amount": self._field(context, "amount") or 0,
|
||||
"prev_close": self._field(context, "prev_close", ("preClose", "lastClose")),
|
||||
}
|
||||
if row["prev_close"] in (None, ""):
|
||||
row["prev_close"] = self.previous_close.get(symbol)
|
||||
self.previous_close[symbol] = close
|
||||
return row
|
||||
|
||||
|
||||
class NativeSessionConfig(object):
|
||||
def __init__(self, run_id, strategy_name, account_id):
|
||||
self.run_id = str(run_id)
|
||||
self.strategy_name = str(strategy_name)
|
||||
self.account_id = str(account_id)
|
||||
|
||||
|
||||
class QmtNativeBacktestSession(object):
|
||||
"""Engine-shaped adapter whose actual engine and broker are both QMT."""
|
||||
|
||||
engine_version = "qmt-native-1.0.0"
|
||||
execution_backend = "QMT_NATIVE"
|
||||
fill_timing = "qmt_native_matching"
|
||||
|
||||
def __init__(self, config=None, qmt_api=None):
|
||||
options = dict(config or {})
|
||||
run_id = str(options.get("run_id") or ("qmt-native-" + dt.datetime.now().strftime("%Y%m%d-%H%M%S")))
|
||||
self.config = NativeSessionConfig(
|
||||
run_id=run_id,
|
||||
strategy_name=options.get("strategy_name") or "ZMQ_BACKTEST",
|
||||
account_id=options.get("account_id") or "",
|
||||
)
|
||||
self.account_type = str(options.get("account_type") or "STOCK")
|
||||
self.combo_type = int(options.get("combo_type") or 1101)
|
||||
self.quick_trade = int(options.get("quick_trade") if options.get("quick_trade") is not None else 2)
|
||||
self.market_price_type = int(options.get("market_price_type") or 5)
|
||||
self.limit_price_type = int(options.get("limit_price_type") or 11)
|
||||
self.bar_wait_timeout = float(options.get("bar_wait_timeout_seconds") or 60.0)
|
||||
self.require_backtest = bool(options.get("require_qmt_backtest", True))
|
||||
self.qmt_api = dict(qmt_api or {})
|
||||
self.feed = StreamingBarFeed(source="qmt_native_backtest")
|
||||
self.extractor = QmtBarExtractor()
|
||||
self.created_at = dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
self.started = False
|
||||
self.finished = False
|
||||
self.qmt_completed = False
|
||||
self.current_index = -1
|
||||
self.current_frame = None
|
||||
self._released_index = -1
|
||||
self._condition = threading.Condition()
|
||||
self._pending_commands = []
|
||||
self._orders = {}
|
||||
self._fills = {}
|
||||
self._current_frame_fills = []
|
||||
self._published_fill_keys = set()
|
||||
self._positions = {}
|
||||
self._asset = {"cash": None, "total_asset": None}
|
||||
self._context = None
|
||||
self._failure = ""
|
||||
|
||||
def bind_context(self, context):
|
||||
if self.require_backtest and not _is_qmt_backtest(context):
|
||||
raise RuntimeError("QMT native bridge refused to run outside QMT backtest mode")
|
||||
if not self.config.account_id:
|
||||
raise RuntimeError("account_id is required for the QMT native backtest bridge")
|
||||
self._context = context
|
||||
if self.config.account_id and hasattr(context, "set_account"):
|
||||
context.set_account(self.config.account_id)
|
||||
|
||||
def _require_api(self, name):
|
||||
func = self.qmt_api.get(name)
|
||||
if func is None:
|
||||
raise RuntimeError("QMT runtime API is unavailable: %s" % name)
|
||||
return func
|
||||
|
||||
def _require_started(self):
|
||||
if not self.started:
|
||||
raise RuntimeError("external strategy has not attached")
|
||||
if self.finished:
|
||||
raise RuntimeError("external strategy session is already finished")
|
||||
|
||||
def start(self):
|
||||
with self._condition:
|
||||
if not self.started:
|
||||
self.started = True
|
||||
self._condition.notify_all()
|
||||
if self.current_index < 0 and not self.qmt_completed:
|
||||
ready = self._condition.wait_for(
|
||||
lambda: self.current_index >= 0 or self.qmt_completed or bool(self._failure),
|
||||
timeout=self.bar_wait_timeout,
|
||||
)
|
||||
if not ready:
|
||||
raise TimeoutError("timed out waiting for QMT's first backtest bar")
|
||||
if self._failure:
|
||||
raise RuntimeError(self._failure)
|
||||
return self._state_unlocked()
|
||||
|
||||
def _queue_command(self, command):
|
||||
with self._condition:
|
||||
self._require_started()
|
||||
if self.qmt_completed or self.current_index < 0:
|
||||
raise RuntimeError("QMT backtest has no active bar")
|
||||
if self._released_index >= self.current_index:
|
||||
raise RuntimeError("current QMT bar has already been released")
|
||||
command = dict(command)
|
||||
command["frame_index"] = self.current_index
|
||||
self._pending_commands.append(command)
|
||||
return command
|
||||
|
||||
def submit_order(self, payload):
|
||||
payload = dict(payload or {})
|
||||
side = str(payload.get("side") or "").upper()
|
||||
if side not in ("BUY", "SELL"):
|
||||
raise ValueError("side must be BUY or SELL")
|
||||
quantity = int(payload.get("quantity") or 0)
|
||||
if quantity <= 0:
|
||||
raise ValueError("quantity must be positive")
|
||||
symbol = normalize_symbol(payload.get("symbol"))
|
||||
order_type = str(payload.get("order_type") or "MARKET").upper()
|
||||
if order_type not in ("MARKET", "LIMIT"):
|
||||
raise ValueError("order_type must be MARKET or LIMIT")
|
||||
limit_price = payload.get("limit_price")
|
||||
if order_type == "LIMIT" and limit_price in (None, ""):
|
||||
raise ValueError("limit_price is required for LIMIT order")
|
||||
client_order_id = str(payload.get("client_order_id") or ("zmq:" + uuid.uuid4().hex[:20]))
|
||||
record = {
|
||||
"order_id": client_order_id,
|
||||
"client_order_id": client_order_id,
|
||||
"symbol": symbol,
|
||||
"side": side,
|
||||
"quantity": quantity,
|
||||
"filled_quantity": 0,
|
||||
"order_type": order_type,
|
||||
"limit_price": None if limit_price in (None, "") else float(limit_price),
|
||||
"status": "QUEUED",
|
||||
"reject_reason": "",
|
||||
"submitted_index": self.current_index,
|
||||
"submitted_at": (self.current_frame or {}).get("datetime", ""),
|
||||
"execution_backend": self.execution_backend,
|
||||
}
|
||||
with self._condition:
|
||||
self._orders[client_order_id] = record
|
||||
self._queue_command({"kind": "submit", "client_order_id": client_order_id})
|
||||
return dict(record)
|
||||
|
||||
def cancel_order(self, order_id):
|
||||
order_id = str(order_id or "").strip()
|
||||
if not order_id:
|
||||
raise ValueError("order_id is required")
|
||||
command = self._queue_command({"kind": "cancel", "order_id": order_id})
|
||||
return {"order_id": order_id, "status": "CANCEL_QUEUED", "frame_index": command["frame_index"]}
|
||||
|
||||
def _execute_submit(self, command, context):
|
||||
client_order_id = command["client_order_id"]
|
||||
with self._condition:
|
||||
record = dict(self._orders[client_order_id])
|
||||
if not self.config.account_id:
|
||||
raise RuntimeError("account_id is required for QMT native passorder")
|
||||
passorder = self._require_api("passorder")
|
||||
side = record["side"]
|
||||
order_type = record["order_type"]
|
||||
price_type = self.limit_price_type if order_type == "LIMIT" else self.market_price_type
|
||||
price = float(record["limit_price"] or 0)
|
||||
result = passorder(
|
||||
23 if side == "BUY" else 24,
|
||||
self.combo_type,
|
||||
self.config.account_id,
|
||||
record["symbol"],
|
||||
price_type,
|
||||
price,
|
||||
int(record["quantity"]),
|
||||
self.config.strategy_name,
|
||||
self.quick_trade,
|
||||
client_order_id,
|
||||
context,
|
||||
)
|
||||
with self._condition:
|
||||
target = self._orders[client_order_id]
|
||||
target["status"] = "SUBMITTED"
|
||||
if result not in (None, ""):
|
||||
target["qmt_order_id"] = str(result)
|
||||
|
||||
def _execute_cancel(self, command, context):
|
||||
cancel = self._require_api("cancel")
|
||||
order_id = command["order_id"]
|
||||
with self._condition:
|
||||
record = self._orders.get(order_id)
|
||||
qmt_order_id = (
|
||||
(record or {}).get("qmt_order_id")
|
||||
or (record or {}).get("order_id")
|
||||
or order_id
|
||||
)
|
||||
result = cancel(qmt_order_id, self.config.account_id, self.account_type, context)
|
||||
with self._condition:
|
||||
record = self._orders.get(order_id)
|
||||
if record is not None:
|
||||
record["status"] = "CANCEL_SUBMITTED" if result is not False else "CANCEL_REJECTED"
|
||||
|
||||
def _execute_commands(self, commands, context):
|
||||
for command in commands:
|
||||
try:
|
||||
if command["kind"] == "submit":
|
||||
self._execute_submit(command, context)
|
||||
elif command["kind"] == "cancel":
|
||||
self._execute_cancel(command, context)
|
||||
except Exception as exc:
|
||||
key = command.get("client_order_id") or command.get("order_id")
|
||||
with self._condition:
|
||||
record = self._orders.get(key)
|
||||
if record is not None:
|
||||
record["status"] = "REJECTED"
|
||||
record["reject_reason"] = "%s: %s" % (exc.__class__.__name__, exc)
|
||||
print("[bigqmt_backtest] QMT command failed kind=%s error=%s" % (command.get("kind"), exc))
|
||||
|
||||
def on_bar(self, context):
|
||||
self.bind_context(context)
|
||||
self._refresh_qmt_state()
|
||||
row = self.extractor.extract(context)
|
||||
appended = self.feed.append(row)
|
||||
if not appended:
|
||||
return False
|
||||
with self._condition:
|
||||
self.current_index = len(self.feed) - 1
|
||||
self.current_frame = self.feed.frame(self.current_index)
|
||||
new_fill_keys = [key for key in self._fills if key not in self._published_fill_keys]
|
||||
self._current_frame_fills = [dict(self._fills[key]) for key in new_fill_keys]
|
||||
self._published_fill_keys.update(new_fill_keys)
|
||||
index = self.current_index
|
||||
self._condition.notify_all()
|
||||
released = self._condition.wait_for(
|
||||
lambda: self._released_index >= index or self.finished or bool(self._failure),
|
||||
timeout=self.bar_wait_timeout,
|
||||
)
|
||||
if not released:
|
||||
self._failure = "external strategy timed out on QMT bar index %d" % index
|
||||
self._condition.notify_all()
|
||||
raise TimeoutError(self._failure)
|
||||
commands = [item for item in self._pending_commands if item.get("frame_index") == index]
|
||||
self._pending_commands = [item for item in self._pending_commands if item.get("frame_index") != index]
|
||||
self._execute_commands(commands, context)
|
||||
self._refresh_qmt_state()
|
||||
print(
|
||||
"[bigqmt_backtest] QMT native bar released index=%d datetime=%s symbol=%s commands=%d"
|
||||
% (index, row["datetime"], row["symbol"], len(commands))
|
||||
)
|
||||
return True
|
||||
|
||||
def next_bar(self):
|
||||
with self._condition:
|
||||
self._require_started()
|
||||
previous = self.current_index
|
||||
if self.qmt_completed:
|
||||
return self._state_unlocked()
|
||||
self._released_index = max(self._released_index, previous)
|
||||
self._condition.notify_all()
|
||||
ready = self._condition.wait_for(
|
||||
lambda: self.current_index > previous or self.qmt_completed or bool(self._failure),
|
||||
timeout=self.bar_wait_timeout,
|
||||
)
|
||||
if not ready:
|
||||
raise TimeoutError("timed out waiting for QMT backtest bar after index %d" % previous)
|
||||
if self._failure:
|
||||
raise RuntimeError(self._failure)
|
||||
return self._state_unlocked()
|
||||
|
||||
def history(self, symbol, count=100, fields=None):
|
||||
with self._condition:
|
||||
self._require_started()
|
||||
end_index = self.current_index
|
||||
return self.feed.history(symbol, end_index, count=count, fields=fields)
|
||||
|
||||
def orders(self):
|
||||
with self._condition:
|
||||
return [dict(value) for value in self._orders.values()]
|
||||
|
||||
def fills(self):
|
||||
with self._condition:
|
||||
return [dict(value) for value in self._fills.values()]
|
||||
|
||||
def _state_unlocked(self):
|
||||
if self.current_frame is None:
|
||||
return {
|
||||
"run_id": self.config.run_id,
|
||||
"started": self.started,
|
||||
"finished": self.finished,
|
||||
"done": self.qmt_completed,
|
||||
"frame_index": -1,
|
||||
"frame_count": len(self.feed),
|
||||
"execution_backend": self.execution_backend,
|
||||
}
|
||||
return {
|
||||
"run_id": self.config.run_id,
|
||||
"started": self.started,
|
||||
"finished": self.finished,
|
||||
"done": self.qmt_completed,
|
||||
"frame_index": self.current_index,
|
||||
"frame_count": len(self.feed),
|
||||
"datetime": self.current_frame["datetime"],
|
||||
"bars": {key: dict(value) for key, value in self.current_frame["bars"].items()},
|
||||
"fills": [dict(value) for value in self._current_frame_fills],
|
||||
"cash": self._asset.get("cash"),
|
||||
"total_asset": self._asset.get("total_asset"),
|
||||
"positions": {key: dict(value) for key, value in self._positions.items()},
|
||||
"execution_backend": self.execution_backend,
|
||||
"qmt_completed": self.qmt_completed,
|
||||
"failure": self._failure,
|
||||
}
|
||||
|
||||
def state(self):
|
||||
with self._condition:
|
||||
return self._state_unlocked()
|
||||
|
||||
def finish(self):
|
||||
with self._condition:
|
||||
if self.finished:
|
||||
return self._result_unlocked()
|
||||
self._released_index = max(self._released_index, self.current_index)
|
||||
self.finished = True
|
||||
self._condition.notify_all()
|
||||
return self._result_unlocked()
|
||||
|
||||
def _result_unlocked(self):
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"engine_version": self.engine_version,
|
||||
"run_id": self.config.run_id,
|
||||
"strategy_name": self.config.strategy_name,
|
||||
"execution_backend": self.execution_backend,
|
||||
"qmt_completed": self.qmt_completed,
|
||||
"order_count": len(self._orders),
|
||||
"fill_count": len(self._fills),
|
||||
"final_state": self._state_unlocked(),
|
||||
"result_owner": "QMT",
|
||||
}
|
||||
|
||||
def on_qmt_stop(self):
|
||||
self.feed.close()
|
||||
with self._condition:
|
||||
new_fill_keys = [key for key in self._fills if key not in self._published_fill_keys]
|
||||
self._current_frame_fills = [dict(self._fills[key]) for key in new_fill_keys]
|
||||
self._published_fill_keys.update(new_fill_keys)
|
||||
self.qmt_completed = True
|
||||
self._condition.notify_all()
|
||||
print("[bigqmt_backtest] QMT native backtest completed bars=%d" % len(self.feed))
|
||||
|
||||
def on_order(self, order):
|
||||
item = {
|
||||
"order_id": str(_attr(order, ("m_strOrderSysID", "order_sys_id", "order_id"), "") or ""),
|
||||
"client_order_id": str(_attr(order, ("m_strRemark", "remark", "user_order_id"), "") or ""),
|
||||
"symbol": _full_symbol(order),
|
||||
"side": _side_from_offset(_attr(order, ("m_nOffsetFlag", "offset_flag"), 0)),
|
||||
"quantity": int(_attr(order, ("m_nVolumeTotalOriginal", "volume", "quantity"), 0) or 0),
|
||||
"filled_quantity": int(_attr(order, ("m_nVolumeTraded", "traded_volume", "filled_quantity"), 0) or 0),
|
||||
"price": _json_number(_attr(order, ("m_dLimitPrice", "m_dPrice", "price"))),
|
||||
"status": str(_attr(order, ("m_nOrderStatus", "status"), "") or ""),
|
||||
}
|
||||
key = item["client_order_id"] or item["order_id"] or ("order:" + uuid.uuid4().hex)
|
||||
with self._condition:
|
||||
existing = self._orders.get(key, {})
|
||||
existing.update(item)
|
||||
self._orders[key] = existing
|
||||
self._condition.notify_all()
|
||||
return dict(existing)
|
||||
|
||||
def on_trade(self, trade):
|
||||
item = {
|
||||
"fill_id": str(_attr(trade, ("m_strTradeID", "trade_id", "fill_id"), "") or ""),
|
||||
"order_id": str(_attr(trade, ("m_strOrderSysID", "order_sys_id", "order_id"), "") or ""),
|
||||
"client_order_id": str(_attr(trade, ("m_strRemark", "remark", "user_order_id"), "") or ""),
|
||||
"symbol": _full_symbol(trade),
|
||||
"side": _side_from_offset(_attr(trade, ("m_nOffsetFlag", "offset_flag"), 0)),
|
||||
"quantity": int(_attr(trade, ("m_nVolume", "volume", "quantity"), 0) or 0),
|
||||
"price": _json_number(_attr(trade, ("m_dPrice", "m_dTradePrice", "price"))),
|
||||
"filled_at": str(_attr(trade, ("m_strTradeTime", "trade_time", "filled_at"), "") or ""),
|
||||
}
|
||||
key = item["fill_id"] or "%s:%s:%s" % (item["order_id"], item["quantity"], item["price"])
|
||||
with self._condition:
|
||||
self._fills[key] = item
|
||||
self._condition.notify_all()
|
||||
return dict(item)
|
||||
|
||||
def _query(self, detail_type):
|
||||
query = self.qmt_api.get("get_trade_detail_data")
|
||||
if query is None or not self.config.account_id:
|
||||
return []
|
||||
calls = []
|
||||
if detail_type in ("ORDER", "DEAL", "TRADE"):
|
||||
calls.append(lambda: query(
|
||||
self.config.account_id, self.account_type, detail_type, self.config.strategy_name
|
||||
))
|
||||
calls.append(lambda: query(self.config.account_id, self.account_type, detail_type))
|
||||
last_error = None
|
||||
for call in calls:
|
||||
try:
|
||||
return list(call() or [])
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
print("[bigqmt_backtest] QMT query failed type=%s error=%s" % (detail_type, last_error))
|
||||
return []
|
||||
|
||||
def _refresh_qmt_state(self):
|
||||
positions = {}
|
||||
for row in self._query("POSITION"):
|
||||
symbol = _full_symbol(row)
|
||||
if not symbol:
|
||||
continue
|
||||
positions[symbol] = {
|
||||
"symbol": symbol,
|
||||
"quantity": int(_attr(row, ("m_nVolume", "volume", "quantity"), 0) or 0),
|
||||
"available": int(_attr(row, ("m_nCanUseVolume", "available", "can_use_volume"), 0) or 0),
|
||||
"avg_cost": _json_number(_attr(row, ("m_dOpenPrice", "m_dCostPrice", "cost", "avg_cost"))),
|
||||
}
|
||||
asset_rows = self._query("ACCOUNT") or self._query("ASSET")
|
||||
asset = {"cash": None, "total_asset": None}
|
||||
if asset_rows:
|
||||
row = asset_rows[0]
|
||||
asset = {
|
||||
"cash": _json_number(_attr(row, ("m_dAvailable", "m_dAvailableCash", "available_cash", "cash"))),
|
||||
"total_asset": _json_number(_attr(row, ("m_dBalance", "m_dAsset", "total_asset", "asset"))),
|
||||
}
|
||||
order_rows = self._query("ORDER")
|
||||
trade_rows = self._query("DEAL") or self._query("TRADE")
|
||||
with self._condition:
|
||||
self._positions = positions
|
||||
self._asset = asset
|
||||
for row in order_rows:
|
||||
self.on_order(row)
|
||||
for row in trade_rows:
|
||||
self.on_trade(row)
|
||||
|
||||
|
||||
class QmtBacktestBridgeRuntime(object):
|
||||
def __init__(self, config=None, qmt_api=None):
|
||||
config = dict(config or {})
|
||||
bind_endpoint = str(config.pop("bind_endpoint", "tcp://127.0.0.1:16662"))
|
||||
self.engine = QmtNativeBacktestSession(config=config, qmt_api=qmt_api)
|
||||
self.protocol = BacktestBridgeProtocol(self.engine)
|
||||
self.server = ZmqBacktestServer(self.protocol, endpoint=bind_endpoint, exit_on_finish=True)
|
||||
self.server_thread = None
|
||||
|
||||
def start(self, context):
|
||||
if self.server_thread is not None:
|
||||
return
|
||||
self.engine.bind_context(context)
|
||||
self.server_thread = threading.Thread(
|
||||
target=self.server.serve_forever,
|
||||
name="bigqmt-native-backtest-zmq",
|
||||
daemon=True,
|
||||
)
|
||||
self.server_thread.start()
|
||||
if not self.server.wait_until_ready(5.0) or not self.server.actual_endpoint:
|
||||
raise RuntimeError("QMT native backtest ZMQ service failed to bind")
|
||||
print(
|
||||
"[bigqmt_backtest] QMT native service started run_id=%s endpoint=%s account=%s live_ready=False"
|
||||
% (self.engine.config.run_id, self.server.actual_endpoint, self.engine.config.account_id)
|
||||
)
|
||||
|
||||
def on_bar(self, context):
|
||||
return self.engine.on_bar(context)
|
||||
|
||||
def on_order(self, order):
|
||||
return self.engine.on_order(order)
|
||||
|
||||
def on_trade(self, trade):
|
||||
return self.engine.on_trade(trade)
|
||||
|
||||
def on_qmt_stop(self):
|
||||
self.engine.on_qmt_stop()
|
||||
|
||||
def stop_server(self):
|
||||
self.server.stop()
|
||||
|
||||
|
||||
def reset_runtime():
|
||||
global _RUNTIME
|
||||
if _RUNTIME is not None:
|
||||
_RUNTIME.stop_server()
|
||||
_RUNTIME = None
|
||||
|
||||
|
||||
def get_runtime():
|
||||
return _RUNTIME
|
||||
|
||||
|
||||
def init(ContextInfo):
|
||||
global _RUNTIME
|
||||
reset_runtime()
|
||||
_RUNTIME = QmtBacktestBridgeRuntime(_CONFIG, _QMT_API)
|
||||
_RUNTIME.start(ContextInfo)
|
||||
return _RUNTIME
|
||||
|
||||
|
||||
def handlebar(ContextInfo):
|
||||
if _RUNTIME is None:
|
||||
init(ContextInfo)
|
||||
return _RUNTIME.on_bar(ContextInfo)
|
||||
|
||||
|
||||
def order_callback(ContextInfo, orderInfo):
|
||||
if _RUNTIME is not None:
|
||||
return _RUNTIME.on_order(orderInfo)
|
||||
return None
|
||||
|
||||
|
||||
def deal_callback(ContextInfo, dealInfo):
|
||||
if _RUNTIME is not None:
|
||||
return _RUNTIME.on_trade(dealInfo)
|
||||
return None
|
||||
|
||||
|
||||
def stop(ContextInfo=None):
|
||||
if _RUNTIME is not None:
|
||||
_RUNTIME.on_qmt_stop()
|
||||
|
||||
|
||||
def after_backtest(ContextInfo=None):
|
||||
return stop(ContextInfo)
|
||||
@@ -1,80 +0,0 @@
|
||||
"""Command-line entry for the standalone ZMQ backtest bridge."""
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import uuid
|
||||
|
||||
from .data_feed import CsvBarFeed
|
||||
from .engine import BacktestConfig, BacktestEngine
|
||||
from .protocol import BacktestBridgeProtocol
|
||||
from .zmq_server import ZmqBacktestServer
|
||||
|
||||
|
||||
def _load_config(path):
|
||||
if not path:
|
||||
return {}
|
||||
with open(os.path.abspath(path), encoding="utf-8") as handle:
|
||||
payload = json.load(handle)
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("config JSON must be an object")
|
||||
return payload
|
||||
|
||||
|
||||
def _parser():
|
||||
parser = argparse.ArgumentParser(description="Standalone ZMQ backtest bridge")
|
||||
parser.add_argument("--data", required=True, help="UTF-8 CSV historical bar file")
|
||||
parser.add_argument("--config", default="", help="Optional UTF-8 JSON backtest config")
|
||||
parser.add_argument("--run-id", default="", help="Unique run identity")
|
||||
parser.add_argument("--output-dir", default="", help="Artifact directory")
|
||||
parser.add_argument("--bind", default="tcp://127.0.0.1:16661", help="ZMQ REP bind endpoint")
|
||||
parser.add_argument("--default-symbol", default="", help="Used when CSV has no symbol column")
|
||||
parser.add_argument("--initial-cash", type=float, default=None)
|
||||
parser.add_argument("--slippage-bps", type=float, default=None)
|
||||
parser.add_argument("--max-volume-participation", type=float, default=None)
|
||||
parser.add_argument("--keep-running", action="store_true", help="Do not stop server after finish")
|
||||
return parser
|
||||
|
||||
|
||||
def build_engine(args):
|
||||
payload = _load_config(args.config)
|
||||
run_id = str(args.run_id or payload.pop("run_id", "") or ("bt-" + uuid.uuid4().hex[:12]))
|
||||
output_dir = args.output_dir or payload.pop("output_dir", "") or os.path.join("backtest_runs", run_id)
|
||||
if args.initial_cash is not None:
|
||||
payload["initial_cash"] = args.initial_cash
|
||||
if args.slippage_bps is not None:
|
||||
payload["slippage_bps"] = args.slippage_bps
|
||||
if args.max_volume_participation is not None:
|
||||
payload["max_volume_participation"] = args.max_volume_participation
|
||||
config = BacktestConfig(run_id=run_id, output_dir=output_dir, **payload)
|
||||
feed = CsvBarFeed(args.data, default_symbol=args.default_symbol)
|
||||
return BacktestEngine(feed, config)
|
||||
|
||||
|
||||
def main(argv=None):
|
||||
args = _parser().parse_args(argv)
|
||||
engine = build_engine(args)
|
||||
server = ZmqBacktestServer(
|
||||
BacktestBridgeProtocol(engine),
|
||||
endpoint=args.bind,
|
||||
exit_on_finish=not args.keep_running,
|
||||
)
|
||||
startup = {
|
||||
"event": "backtest_bridge_starting",
|
||||
"run_id": engine.config.run_id,
|
||||
"bind": args.bind,
|
||||
"data_hash": engine.feed.data_hash,
|
||||
"output_dir": engine.config.output_dir,
|
||||
"live_ready": False,
|
||||
}
|
||||
print(json.dumps(startup, ensure_ascii=False, sort_keys=True), flush=True)
|
||||
try:
|
||||
server.serve_forever()
|
||||
except KeyboardInterrupt:
|
||||
server.stop()
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -1,75 +0,0 @@
|
||||
"""Small external-strategy contract and a synchronous ZMQ runner."""
|
||||
|
||||
|
||||
class StrategyContext(object):
|
||||
def __init__(self, client):
|
||||
self.client = client
|
||||
self.state = None
|
||||
|
||||
@property
|
||||
def now(self):
|
||||
return None if self.state is None else self.state.get("datetime")
|
||||
|
||||
@property
|
||||
def cash(self):
|
||||
return 0 if self.state is None else self.state.get("cash", 0)
|
||||
|
||||
@property
|
||||
def positions(self):
|
||||
return {} if self.state is None else self.state.get("positions", {})
|
||||
|
||||
def history(self, symbol, count=100, fields=None):
|
||||
return self.client.history(symbol, count=count, fields=fields)
|
||||
|
||||
|
||||
class ExternalStrategyRunner(object):
|
||||
"""Drive a user strategy without exposing future bars.
|
||||
|
||||
Strategy methods are optional:
|
||||
|
||||
* ``on_start(context)``
|
||||
* ``on_bar(context, bars) -> iterable[order dict]``
|
||||
* ``on_fill(context, fill)``
|
||||
* ``on_finish(context, result)``
|
||||
"""
|
||||
|
||||
def __init__(self, client, strategy):
|
||||
self.client = client
|
||||
self.strategy = strategy
|
||||
self.context = StrategyContext(client)
|
||||
|
||||
def _call(self, name, *args):
|
||||
callback = getattr(self.strategy, name, None)
|
||||
return callback(*args) if callback is not None else None
|
||||
|
||||
def _apply_orders(self, orders):
|
||||
for order in list(orders or []):
|
||||
payload = dict(order)
|
||||
self.client.submit_order(
|
||||
symbol=payload["symbol"],
|
||||
side=payload["side"],
|
||||
quantity=payload["quantity"],
|
||||
order_type=payload.get("order_type", "MARKET"),
|
||||
limit_price=payload.get("limit_price"),
|
||||
client_order_id=payload.get("client_order_id", ""),
|
||||
time_in_force=payload.get("time_in_force", "NEXT_BAR"),
|
||||
)
|
||||
|
||||
def run(self):
|
||||
if not self.client.run_id:
|
||||
self.client.describe()
|
||||
state = self.client.start()
|
||||
self.context.state = state
|
||||
self._call("on_start", self.context)
|
||||
while True:
|
||||
for fill in state.get("fills", []):
|
||||
self._call("on_fill", self.context, fill)
|
||||
orders = self._call("on_bar", self.context, state.get("bars", {}))
|
||||
self._apply_orders(orders)
|
||||
if state.get("done"):
|
||||
break
|
||||
state = self.client.next_bar()
|
||||
self.context.state = state
|
||||
result = self.client.finish()
|
||||
self._call("on_finish", self.context, result)
|
||||
return result
|
||||
@@ -1,65 +0,0 @@
|
||||
"""Minimal REQ/REP ZMQ server for one isolated backtest run."""
|
||||
|
||||
import json
|
||||
import threading
|
||||
|
||||
|
||||
class ZmqBacktestServer(object):
|
||||
def __init__(self, protocol, endpoint="tcp://127.0.0.1:16661", exit_on_finish=False, poll_ms=100):
|
||||
self.protocol = protocol
|
||||
self.endpoint = str(endpoint)
|
||||
self.exit_on_finish = bool(exit_on_finish)
|
||||
self.poll_ms = int(poll_ms)
|
||||
self._stop_event = threading.Event()
|
||||
self._ready_event = threading.Event()
|
||||
self.actual_endpoint = None
|
||||
|
||||
def wait_until_ready(self, timeout_seconds=None):
|
||||
return self._ready_event.wait(timeout_seconds)
|
||||
|
||||
def stop(self):
|
||||
self._stop_event.set()
|
||||
|
||||
def serve_forever(self):
|
||||
import zmq
|
||||
|
||||
context = zmq.Context.instance()
|
||||
socket = context.socket(zmq.REP)
|
||||
socket.setsockopt(zmq.LINGER, 0)
|
||||
socket.setsockopt(zmq.RCVHWM, 1000)
|
||||
socket.setsockopt(zmq.SNDHWM, 1000)
|
||||
try:
|
||||
if self.endpoint.endswith(":0"):
|
||||
base = self.endpoint.rsplit(":", 1)[0]
|
||||
port = socket.bind_to_random_port(base)
|
||||
self.actual_endpoint = "%s:%d" % (base, port)
|
||||
else:
|
||||
socket.bind(self.endpoint)
|
||||
self.actual_endpoint = self.endpoint
|
||||
self._ready_event.set()
|
||||
poller = zmq.Poller()
|
||||
poller.register(socket, zmq.POLLIN)
|
||||
while not self._stop_event.is_set():
|
||||
events = dict(poller.poll(self.poll_ms))
|
||||
if socket not in events:
|
||||
continue
|
||||
try:
|
||||
request = json.loads(socket.recv().decode("utf-8"))
|
||||
response = self.protocol.handle(request)
|
||||
except Exception as exc:
|
||||
response = {
|
||||
"schema_version": 1,
|
||||
"request_id": "",
|
||||
"run_id": self.protocol.engine.config.run_id,
|
||||
"client_id": "",
|
||||
"method": "",
|
||||
"ok": False,
|
||||
"data": None,
|
||||
"error": "%s: %s" % (exc.__class__.__name__, exc),
|
||||
}
|
||||
socket.send(json.dumps(response, ensure_ascii=False, separators=(",", ":")).encode("utf-8"))
|
||||
if self.exit_on_finish and self.protocol.engine.finished:
|
||||
break
|
||||
finally:
|
||||
self._ready_event.set()
|
||||
socket.close(linger=0)
|
||||
@@ -1,50 +0,0 @@
|
||||
# bigqmt_signal_trader
|
||||
|
||||
大 QMT 信号交易包的核心骨架。当前版本只完成可替换包边界和 dry-run 运行入口,不会发送真实委托。
|
||||
|
||||
## 已完成
|
||||
|
||||
- `TradeSignal`、`OrderRequest`、`PositionSnapshot`、`AccountSnapshot` 等核心数据模型。
|
||||
- `SignalSource`、`MarketDataProvider`、`PositionProvider`、`OrderGateway`、`PositionSyncSink`、`StateStore` 等替换接口。
|
||||
- `SignalTradingApp.tick()` 编排流程:
|
||||
1. 读取信号。
|
||||
2. 原子 claim。
|
||||
3. 读取持仓。
|
||||
4. 计算买卖数量。
|
||||
5. 生成价格。
|
||||
6. 调用可替换 `OrderGateway`。
|
||||
7. 写回状态。
|
||||
8. 同步持仓快照。
|
||||
- `DryRunOrderGateway`:记录委托请求,不调用真实 `passorder`。
|
||||
- `bigqmt_signal_trader_strategy.py`:大 QMT 运行文件骨架,响应 `init`、`adjust`、`order_callback`、`deal_callback`、`sync_positions`。
|
||||
|
||||
## 当前安全状态
|
||||
|
||||
默认 `adapter_factory.build_app()` 使用:
|
||||
|
||||
- 空信号源。
|
||||
- 空行情源。
|
||||
- 空持仓源。
|
||||
- dry-run 下单 gateway。
|
||||
- no-op 状态存储。
|
||||
- 内存持仓同步 sink。
|
||||
|
||||
因此即使大 QMT 加载该运行文件,也不会真实下单。
|
||||
|
||||
## 后续接入顺序
|
||||
|
||||
1. 实现 `BigQmtMarketDataProvider` 和 `BigQmtPositionProvider`。
|
||||
2. 实现 `BigQmtOrderGateway(passorder/cancel/get_trade_detail_data)`。
|
||||
3. 实现 Redis Stream / MySQL outbox 信号源。
|
||||
4. 实现 Redis / MySQL 状态写回。
|
||||
5. 实现 Redis / MySQL 持仓同步 sink。
|
||||
6. dry-run 跑通后,再按账户灰度切换真实下单。
|
||||
|
||||
## 测试
|
||||
|
||||
```powershell
|
||||
cd <REPO_ROOT>
|
||||
python -m unittest discover -s tests\bigqmt_signal_trader
|
||||
```
|
||||
|
||||
|
||||
@@ -1,32 +0,0 @@
|
||||
"""可替换的大 QMT 信号下单包核心模块。"""
|
||||
|
||||
__version__ = "0.2.0"
|
||||
|
||||
from .app import SignalTradingApp
|
||||
from .models import (
|
||||
AccountSnapshot,
|
||||
AssetSnapshot,
|
||||
OrderRequest,
|
||||
OrderSubmitResult,
|
||||
PositionSnapshot,
|
||||
SignalAction,
|
||||
SignalStatus,
|
||||
TradeSignal,
|
||||
)
|
||||
from .xtquant_compat import BigQmtRpcClient, BigQmtXtData, BigQmtXtTrader
|
||||
|
||||
__all__ = [
|
||||
"AccountSnapshot",
|
||||
"AssetSnapshot",
|
||||
"BigQmtRpcClient",
|
||||
"BigQmtXtData",
|
||||
"BigQmtXtTrader",
|
||||
"OrderRequest",
|
||||
"OrderSubmitResult",
|
||||
"PositionSnapshot",
|
||||
"SignalAction",
|
||||
"SignalStatus",
|
||||
"SignalTradingApp",
|
||||
"TradeSignal",
|
||||
"__version__",
|
||||
]
|
||||
@@ -1,159 +0,0 @@
|
||||
"""根据配置装配 SignalTradingApp。
|
||||
|
||||
当前第一版只提供安全的空信号 + dry-run 默认实现,后续再接 Redis/MySQL/大 QMT
|
||||
真实 adapter。这样大 QMT 运行文件可以先加载和响应调度,不会误发真实委托。
|
||||
"""
|
||||
|
||||
from .adapters.order_dryrun import DryRunOrderGateway
|
||||
from .app import SignalTradingApp
|
||||
from .models import AssetSnapshot
|
||||
|
||||
|
||||
class EmptySignalSource:
|
||||
def fetch(self, account_id, limit):
|
||||
return []
|
||||
|
||||
def ack(self, signal):
|
||||
return None
|
||||
|
||||
|
||||
class EmptyMarketDataProvider:
|
||||
def get_ticks(self, codes):
|
||||
return {}
|
||||
|
||||
def get_instrument(self, code):
|
||||
return {}
|
||||
|
||||
|
||||
class EmptyPositionProvider:
|
||||
def get_positions(self, account_id):
|
||||
return {}
|
||||
|
||||
def get_asset(self, account_id):
|
||||
return AssetSnapshot(account_id=account_id, cash=None, total_asset=None)
|
||||
|
||||
|
||||
class NoopPositionSyncSink:
|
||||
def __init__(self):
|
||||
self.snapshots = []
|
||||
|
||||
def publish(self, snapshot):
|
||||
self.snapshots.append(snapshot)
|
||||
|
||||
|
||||
class NoopStateStore:
|
||||
def claim(self, signal, consumer_id):
|
||||
return False
|
||||
|
||||
def mark_submitted(self, signal_id, result):
|
||||
return None
|
||||
|
||||
def mark_finished(self, signal_id, status, message=""):
|
||||
return None
|
||||
|
||||
|
||||
def _config_bool(value, default=False):
|
||||
if value is None:
|
||||
return default
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
return str(value).strip().lower() in ("1", "true", "yes", "y", "on")
|
||||
|
||||
|
||||
def build_app(context_info=None, config=None):
|
||||
config = config or {}
|
||||
mode = str(config.get("mode") or "dryrun").lower()
|
||||
account_id = config.get("account_id", "default")
|
||||
source_type = str(config.get("signal_source_type") or config.get("source_type") or "").lower()
|
||||
state_type = str(config.get("state_store_type") or "").lower()
|
||||
position_sync_type = str(config.get("position_sync_type") or "").lower()
|
||||
|
||||
signal_source = config.get("signal_source")
|
||||
market_data = config.get("market_data")
|
||||
position_provider = config.get("position_provider")
|
||||
order_gateway = config.get("order_gateway")
|
||||
position_sync_sink = config.get("position_sync_sink")
|
||||
state_store = config.get("state_store")
|
||||
redis_client = config.get("redis_client")
|
||||
|
||||
if source_type == "redis" or state_type == "redis" or position_sync_type == "redis":
|
||||
from .adapters.redis_common import build_redis_client
|
||||
|
||||
redis_client = redis_client or build_redis_client(config.get("redis") or {})
|
||||
|
||||
if source_type == "redis":
|
||||
from .adapters.signal_redis import RedisStreamSignalSource
|
||||
|
||||
redis_cfg = config.get("redis") or {}
|
||||
signal_source = signal_source or RedisStreamSignalSource(
|
||||
redis_client=redis_client,
|
||||
stream_key_template=redis_cfg.get("stream_key_template", "bigqmt:signals:{account_id}"),
|
||||
group_name=redis_cfg.get("group_name", "bigqmt-signal-trader"),
|
||||
consumer_name=redis_cfg.get("consumer_name", "bigqmt-consumer"),
|
||||
block_ms=int(redis_cfg.get("block_ms", 0)),
|
||||
)
|
||||
|
||||
if state_type == "redis" or (source_type == "redis" and state_store is None):
|
||||
from .adapters.state_redis import RedisStateStore
|
||||
|
||||
redis_cfg = config.get("redis") or {}
|
||||
state_store = state_store or RedisStateStore(
|
||||
redis_client=redis_client,
|
||||
account_id=account_id,
|
||||
claim_key_template=redis_cfg.get("claim_key_template", "bigqmt:signal_claim:{account_id}:{signal_id}"),
|
||||
status_key_template=redis_cfg.get("status_key_template", "bigqmt:signal_status:{account_id}:{signal_id}"),
|
||||
claim_ttl_seconds=int(redis_cfg.get("claim_ttl_seconds", 3600)),
|
||||
status_ttl_seconds=int(redis_cfg.get("status_ttl_seconds", 86400)),
|
||||
)
|
||||
|
||||
if position_sync_type == "redis":
|
||||
from .adapters.position_sync_redis import RedisPositionSyncSink
|
||||
|
||||
redis_cfg = config.get("redis") or {}
|
||||
position_sync_sink = position_sync_sink or RedisPositionSyncSink(
|
||||
redis_client=redis_client,
|
||||
key_template=redis_cfg.get("position_key_template", "bigqmt:positions:{account_id}"),
|
||||
event_stream_template=redis_cfg.get("position_event_stream_template", "bigqmt:position_events:{account_id}"),
|
||||
ttl_seconds=int(redis_cfg.get("position_ttl_seconds", 120)),
|
||||
publish_events=_config_bool(redis_cfg.get("position_publish_events"), True),
|
||||
)
|
||||
|
||||
if mode == "bigqmt":
|
||||
from .adapters.market_bigqmt import BigQmtMarketDataProvider
|
||||
from .adapters.order_bigqmt import BigQmtOrderGateway
|
||||
from .adapters.position_bigqmt import BigQmtPositionProvider
|
||||
|
||||
qmt_api = config.get("qmt_api") or {}
|
||||
get_trade_detail_data_func = qmt_api.get("get_trade_detail_data")
|
||||
market_data = market_data or BigQmtMarketDataProvider(context_info, qmt_api=qmt_api)
|
||||
position_provider = position_provider or BigQmtPositionProvider(
|
||||
get_trade_detail_data_func=get_trade_detail_data_func,
|
||||
account_type=config.get("account_type", "STOCK"),
|
||||
)
|
||||
# passorder / cancel need the RAW QMT ContextInfo as their last arg -- QMT's
|
||||
# injected passorder reads internals off it (e.g. .request_id). Our runtime
|
||||
# wrapper (BigQmtRuntimeAdapter) doesn't have those, so unwrap it here.
|
||||
raw_context_info = getattr(context_info, "context_info", context_info)
|
||||
order_gateway = order_gateway or BigQmtOrderGateway(
|
||||
context_info=raw_context_info,
|
||||
account_id=account_id,
|
||||
passorder_func=qmt_api.get("passorder"),
|
||||
cancel_func=qmt_api.get("cancel"),
|
||||
get_trade_detail_data_func=get_trade_detail_data_func,
|
||||
account_type=config.get("account_type", "STOCK"),
|
||||
combo_type=int(config.get("combo_type", 1101)),
|
||||
price_type=int(config.get("order_price_type", 11)),
|
||||
quick_trade=int(config.get("quick_trade", 2)),
|
||||
)
|
||||
|
||||
return SignalTradingApp(
|
||||
account_id=account_id,
|
||||
signal_source=signal_source or EmptySignalSource(),
|
||||
market_data=market_data or EmptyMarketDataProvider(),
|
||||
position_provider=position_provider or EmptyPositionProvider(),
|
||||
order_gateway=order_gateway or DryRunOrderGateway(),
|
||||
position_sync_sink=position_sync_sink or NoopPositionSyncSink(),
|
||||
state_store=state_store or NoopStateStore(),
|
||||
consumer_id=config.get("consumer_id", "bigqmt-signal-trader"),
|
||||
fetch_limit=config.get("fetch_limit", 20),
|
||||
)
|
||||
@@ -1 +0,0 @@
|
||||
"""具体外部系统 adapter。"""
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,254 +0,0 @@
|
||||
"""Big QMT order gateway.
|
||||
|
||||
The passorder signature follows src/api/qmt_jq_trade.
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
|
||||
from ..code_utils import normalize_stock_code
|
||||
from ..models import CancelResult, OrderSnapshot, OrderSubmitResult, SignalAction, TradeSnapshot
|
||||
from .position_bigqmt import _attr, _full_code
|
||||
|
||||
|
||||
PRICE_TYPE_ALIASES = {
|
||||
"LIMIT": 11,
|
||||
"FIX_PRICE": 11,
|
||||
"LATEST_PRICE": 5,
|
||||
"MARKET_PEER_PRICE_FIRST": 44,
|
||||
"MARKET_SH_CONVERT_5_LIMIT": 43,
|
||||
"MARKET_SZ_CONVERT_5_CANCEL": 47,
|
||||
}
|
||||
|
||||
|
||||
def _action_from_offset_flag(offset_flag):
|
||||
return SignalAction.BUY.value if int(offset_flag or 0) == 48 else SignalAction.SELL.value
|
||||
|
||||
|
||||
# 报单时间。大 QMT 的 ORDER 行把日期和时间分成两个字段, MiniQMT 的
|
||||
# XtOrder.order_time 是 Unix 秒, 所以要拼接后转换。成交那条路径早就读了
|
||||
# m_strTradeTime, 委托这边一直漏掉 (issue #48)。
|
||||
_ORDER_DATE_FIELDS = ("m_strInsertDate", "m_strOrderDate", "insert_date", "order_date")
|
||||
_ORDER_TIME_FIELDS = ("m_strInsertTime", "m_strOrderTime", "insert_time", "order_time")
|
||||
|
||||
# 取不到时打印该行实际有哪些 m_*, 每进程一次。字段名无法离线核实,
|
||||
# 猜一个然后静默返回 0 正是订单方向那个 bug 的成因。
|
||||
_missing_order_time_reported = []
|
||||
|
||||
|
||||
def _report_missing_order_time(row):
|
||||
if _missing_order_time_reported:
|
||||
return
|
||||
_missing_order_time_reported.append(True)
|
||||
try:
|
||||
available = sorted(n for n in dir(row) if n.startswith("m_"))
|
||||
except Exception:
|
||||
available = []
|
||||
print(
|
||||
"[bigqmt_order] order_time not found (tried %s / %s); ORDER row exposes: %s"
|
||||
% (", ".join(_ORDER_DATE_FIELDS), ", ".join(_ORDER_TIME_FIELDS),
|
||||
", ".join(available) or "<none>")
|
||||
)
|
||||
|
||||
|
||||
def _order_time_seconds(row):
|
||||
"""把 ORDER 行的报单日期+时间转成 Unix 秒, 拿不到返回 0。
|
||||
|
||||
容忍几种实际会遇到的写法: 日期 '20260819' 或 '2026-08-19',
|
||||
时间 '093015'、'09:30:15' 或 '09:30:15.123'。已经是数字时间戳的直接用
|
||||
(毫秒会被归一到秒)。
|
||||
"""
|
||||
raw_time = _attr(row, _ORDER_TIME_FIELDS)
|
||||
raw_date = _attr(row, _ORDER_DATE_FIELDS)
|
||||
if raw_time is None and raw_date is None:
|
||||
_report_missing_order_time(row)
|
||||
return 0
|
||||
|
||||
# 已是数字: 当成时间戳 (>1e11 视为毫秒)。
|
||||
if isinstance(raw_time, (int, float)) and not isinstance(raw_time, bool):
|
||||
value = float(raw_time)
|
||||
if value > 1e11:
|
||||
value /= 1000.0
|
||||
if value > 1e8: # 像时间戳而不是 093015 这种时分秒
|
||||
return int(value)
|
||||
|
||||
date_text = "".join(ch for ch in str(raw_date or "") if ch.isdigit())
|
||||
time_text = "".join(ch for ch in str(raw_time or "") if ch.isdigit())
|
||||
if not date_text or len(date_text) < 8:
|
||||
return 0
|
||||
time_text = (time_text + "000000")[:6] # 补齐到 HHMMSS, 丢掉毫秒
|
||||
try:
|
||||
import time as _time
|
||||
|
||||
parsed = _time.strptime(date_text[:8] + time_text, "%Y%m%d%H%M%S")
|
||||
return int(_time.mktime(parsed))
|
||||
except Exception:
|
||||
return 0
|
||||
|
||||
|
||||
def _price_type_value(value, default):
|
||||
if value is None or value == "":
|
||||
return int(default)
|
||||
try:
|
||||
return int(value)
|
||||
except (TypeError, ValueError):
|
||||
text = str(value).strip().upper()
|
||||
return int(PRICE_TYPE_ALIASES.get(text, default))
|
||||
|
||||
|
||||
class BigQmtOrderGateway:
|
||||
def __init__(
|
||||
self,
|
||||
context_info,
|
||||
account_id="",
|
||||
passorder_func=None,
|
||||
cancel_func=None,
|
||||
get_trade_detail_data_func=None,
|
||||
account_type="STOCK",
|
||||
combo_type=1101,
|
||||
price_type=11,
|
||||
quick_trade=2,
|
||||
):
|
||||
self.context_info = context_info
|
||||
self.account_id = account_id
|
||||
self.passorder = passorder_func
|
||||
self.cancel_func = cancel_func
|
||||
self.get_trade_detail_data = get_trade_detail_data_func
|
||||
self.account_type = account_type
|
||||
self.combo_type = combo_type
|
||||
self.price_type = price_type
|
||||
self.quick_trade = quick_trade
|
||||
|
||||
def _require_passorder(self):
|
||||
if self.passorder is None:
|
||||
raise RuntimeError("passorder is not available in Big QMT runtime")
|
||||
return self.passorder
|
||||
|
||||
def _require_cancel(self):
|
||||
if self.cancel_func is None:
|
||||
raise RuntimeError("cancel is not available in Big QMT runtime")
|
||||
return self.cancel_func
|
||||
|
||||
def _require_query_func(self):
|
||||
if self.get_trade_detail_data is None:
|
||||
raise RuntimeError("get_trade_detail_data is not available in Big QMT runtime")
|
||||
return self.get_trade_detail_data
|
||||
|
||||
@staticmethod
|
||||
def build_user_order_id(signal_id):
|
||||
text = str(signal_id or "")
|
||||
digest = hashlib.sha1(text.encode("utf-8")).hexdigest()[:10]
|
||||
return "bq:%s:%s" % (digest, text[:30])
|
||||
|
||||
def submit(self, request):
|
||||
passorder = self._require_passorder()
|
||||
action = str(request.action).upper()
|
||||
if action == SignalAction.BUY.value:
|
||||
op_type = 23
|
||||
elif action == SignalAction.SELL.value:
|
||||
op_type = 24
|
||||
else:
|
||||
raise ValueError("unsupported order action: %s" % request.action)
|
||||
|
||||
user_order_id = str(request.remark or "").strip() or self.build_user_order_id(request.signal_id)
|
||||
account_id = request.account_id or self.account_id
|
||||
passorder(
|
||||
op_type,
|
||||
self.combo_type,
|
||||
account_id,
|
||||
normalize_stock_code(request.stock_code),
|
||||
_price_type_value(request.price_type, self.price_type),
|
||||
float(request.price),
|
||||
int(request.volume),
|
||||
request.strategy_name,
|
||||
self.quick_trade,
|
||||
user_order_id,
|
||||
self.context_info,
|
||||
)
|
||||
return OrderSubmitResult(
|
||||
status="SUBMITTED",
|
||||
user_order_id=user_order_id,
|
||||
order_sys_id=None,
|
||||
message="passorder submitted",
|
||||
)
|
||||
|
||||
def cancel(self, order_ref):
|
||||
cancel_func = self._require_cancel()
|
||||
ok = cancel_func(order_ref.order_sys_id, self.account_id, self.account_type, self.context_info)
|
||||
return CancelResult(success=bool(ok), message="" if ok else "cancel returned false")
|
||||
|
||||
def query_orders(self, account_id, strategy_name):
|
||||
try:
|
||||
return self.query_orders_strict(account_id, strategy_name)
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def query_orders_strict(self, account_id, strategy_name):
|
||||
query = self._require_query_func()
|
||||
rows = query(account_id, self.account_type, "ORDER", strategy_name) or []
|
||||
result = []
|
||||
for row in rows:
|
||||
result.append(
|
||||
OrderSnapshot(
|
||||
order_sys_id=str(_attr(row, ("m_strOrderSysID", "order_sys_id"), "") or ""),
|
||||
user_order_id=str(_attr(row, ("m_strRemark", "user_order_id", "remark"), "") or ""),
|
||||
stock_code=_full_code(
|
||||
_attr(row, ("m_strInstrumentID", "instrument_id", "stock_code")),
|
||||
_attr(row, ("m_strExchangeID", "exchange_id", "market")),
|
||||
),
|
||||
action=_action_from_offset_flag(_attr(row, ("m_nOffsetFlag", "offset_flag"), 0)),
|
||||
volume=int(_attr(row, ("m_nVolumeTotalOriginal", "volume"), 0) or 0),
|
||||
traded_volume=int(_attr(row, ("m_nVolumeTraded", "traded_volume"), 0) or 0),
|
||||
status=str(_attr(row, ("m_nOrderStatus", "status"), "") or ""),
|
||||
price=float(_attr(row, ("m_dLimitPrice", "m_dPrice", "price"), 0.0) or 0.0),
|
||||
strategy_name=str(_attr(row, ("m_strStrategyName", "strategy_name"), "") or ""),
|
||||
remark=str(_attr(row, ("m_strRemark", "remark"), "") or ""),
|
||||
order_time=_order_time_seconds(row),
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
def query_trades(self, account_id, strategy_name):
|
||||
try:
|
||||
return self.query_trades_strict(account_id, strategy_name)
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def query_trades_strict(self, account_id, strategy_name):
|
||||
query = self._require_query_func()
|
||||
rows = []
|
||||
last_error = None
|
||||
for detail_type in ("DEAL", "TRADE"):
|
||||
try:
|
||||
if str(strategy_name or "").strip():
|
||||
rows = query(account_id, self.account_type, detail_type, strategy_name) or []
|
||||
else:
|
||||
rows = query(account_id, self.account_type, detail_type) or []
|
||||
if rows:
|
||||
break
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
if not rows and last_error is not None:
|
||||
raise last_error
|
||||
result = []
|
||||
for row in rows:
|
||||
result.append(
|
||||
TradeSnapshot(
|
||||
trade_id=str(_attr(row, ("m_strTradeID", "trade_id"), "") or ""),
|
||||
order_sys_id=str(_attr(row, ("m_strOrderSysID", "order_sys_id"), "") or ""),
|
||||
stock_code=_full_code(
|
||||
_attr(row, ("m_strInstrumentID", "instrument_id", "stock_code")),
|
||||
_attr(row, ("m_strExchangeID", "exchange_id", "market")),
|
||||
),
|
||||
action=_action_from_offset_flag(_attr(row, ("m_nOffsetFlag", "offset_flag"), 0)),
|
||||
volume=int(_attr(row, ("m_nVolume", "volume"), 0) or 0),
|
||||
price=float(_attr(row, ("m_dPrice", "m_dTradePrice", "price"), 0.0) or 0.0),
|
||||
traded_at=str(_attr(row, ("m_strTradeTime", "trade_time", "traded_at"), "") or ""),
|
||||
user_order_id=str(_attr(row, ("m_strRemark", "user_order_id", "remark"), "") or ""),
|
||||
)
|
||||
)
|
||||
return result
|
||||
|
||||
def query_submission_identities_strict(self, account_id, strategy_name):
|
||||
orders = self.query_orders_strict(account_id, strategy_name)
|
||||
trades = self.query_trades_strict(account_id, strategy_name)
|
||||
return orders, trades
|
||||
@@ -1,30 +0,0 @@
|
||||
"""不发真实委托的下单 gateway,用于联调和回放。"""
|
||||
|
||||
import hashlib
|
||||
|
||||
from ..models import OrderSubmitResult
|
||||
|
||||
|
||||
class DryRunOrderGateway:
|
||||
def __init__(self):
|
||||
self.submitted = []
|
||||
self.cancelled = []
|
||||
|
||||
def submit(self, request):
|
||||
self.submitted.append(request)
|
||||
digest = hashlib.sha1(request.signal_id.encode("utf-8")).hexdigest()[:10]
|
||||
return OrderSubmitResult(
|
||||
status="DRY_RUN",
|
||||
user_order_id=f"dryrun:bq:{digest}:{request.signal_id}",
|
||||
order_sys_id=None,
|
||||
)
|
||||
|
||||
def cancel(self, order_ref):
|
||||
self.cancelled.append(order_ref)
|
||||
return None
|
||||
|
||||
def query_orders(self, account_id, strategy_name):
|
||||
return []
|
||||
|
||||
def query_trades(self, account_id, strategy_name):
|
||||
return []
|
||||
@@ -1,146 +0,0 @@
|
||||
"""Big QMT position and asset adapters."""
|
||||
|
||||
from ..code_utils import normalize_stock_code
|
||||
from ..models import AssetSnapshot, PositionSnapshot
|
||||
|
||||
|
||||
def _attr(obj, names, default=None):
|
||||
for name in names:
|
||||
if hasattr(obj, name):
|
||||
value = getattr(obj, name)
|
||||
if value is not None:
|
||||
return value
|
||||
return default
|
||||
|
||||
|
||||
def _float_or_none(value):
|
||||
if value is None:
|
||||
return None
|
||||
try:
|
||||
return float(value)
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
# Candidate ThinkTrader field names on the ACCOUNT row of get_trade_detail_data.
|
||||
# The MiniQMT SDK only documents the normalized name (XtAsset.frozen_cash); the
|
||||
# big QMT ACCOUNT struct is a different surface and brokers vary, so probe the
|
||||
# plausible spellings the way cash/total_asset already do.
|
||||
_FROZEN_CASH_FIELDS = (
|
||||
"m_dFrozenCash",
|
||||
"m_dFrozen",
|
||||
"m_dFrozenBalance",
|
||||
"m_dFrozenMargin",
|
||||
"frozen_cash",
|
||||
"frozen",
|
||||
)
|
||||
_MARKET_VALUE_FIELDS = (
|
||||
"m_dInstrumentValue",
|
||||
"m_dStockValue",
|
||||
"m_dMarketValue",
|
||||
"market_value",
|
||||
)
|
||||
|
||||
# Printed once per process when the frozen field is not found, listing what the
|
||||
# row actually carries. Guessing a field name and shipping it unverified is how
|
||||
# the order-direction bug happened; this makes the real name self-reporting.
|
||||
_missing_field_reported = set()
|
||||
|
||||
|
||||
def _report_missing_field(label, row, candidates):
|
||||
if label in _missing_field_reported:
|
||||
return
|
||||
_missing_field_reported.add(label)
|
||||
try:
|
||||
available = sorted(name for name in dir(row) if name.startswith("m_"))
|
||||
except Exception:
|
||||
available = []
|
||||
print(
|
||||
"[bigqmt_asset] %s not found (tried %s); ACCOUNT row exposes: %s"
|
||||
% (label, ", ".join(candidates), ", ".join(available) or "<none>")
|
||||
)
|
||||
|
||||
|
||||
def _full_code(instrument_id, exchange_id):
|
||||
code = str(instrument_id or "").strip().upper()
|
||||
market = str(exchange_id or "").strip().upper()
|
||||
if "." in code:
|
||||
return normalize_stock_code(code)
|
||||
if market in ("SH", "SZ"):
|
||||
return normalize_stock_code("%s.%s" % (code, market))
|
||||
return normalize_stock_code(code)
|
||||
|
||||
|
||||
class BigQmtPositionProvider:
|
||||
def __init__(self, get_trade_detail_data_func, account_type="STOCK"):
|
||||
self.get_trade_detail_data = get_trade_detail_data_func
|
||||
self.account_type = account_type
|
||||
|
||||
def _require_query_func(self):
|
||||
if self.get_trade_detail_data is None:
|
||||
raise RuntimeError("get_trade_detail_data is not available in Big QMT runtime")
|
||||
return self.get_trade_detail_data
|
||||
|
||||
def get_positions(self, account_id):
|
||||
query = self._require_query_func()
|
||||
# QMT's get_trade_detail_data can raise on POSITION queries in some
|
||||
# states (e.g. context not bound). Degrade to empty like get_asset does.
|
||||
try:
|
||||
rows = query(account_id, self.account_type, "POSITION") or []
|
||||
except Exception:
|
||||
return {}
|
||||
positions = {}
|
||||
for row in rows:
|
||||
code = _full_code(
|
||||
_attr(row, ("m_strInstrumentID", "instrument_id", "stock_code")),
|
||||
_attr(row, ("m_strExchangeID", "exchange_id", "market")),
|
||||
)
|
||||
positions[code] = PositionSnapshot(
|
||||
stock_code=code,
|
||||
volume=int(_attr(row, ("m_nVolume", "volume"), 0) or 0),
|
||||
available=int(_attr(row, ("m_nCanUseVolume", "available", "can_use_volume"), 0) or 0),
|
||||
cost=float(_attr(row, ("m_dOpenPrice", "m_dCostPrice", "cost"), 0.0) or 0.0),
|
||||
stock_name=str(_attr(row, ("m_strInstrumentName", "stock_name"), "") or ""),
|
||||
market_value=_float_or_none(_attr(row, ("m_dMarketValue", "m_dInstrumentValue", "market_value"))),
|
||||
price=_float_or_none(_attr(row, ("m_dLastPrice", "m_dSettlementPrice", "price", "last_price"))),
|
||||
open_price=_float_or_none(_attr(row, ("m_dOpenPrice", "m_dCostPrice", "open_price", "cost"))),
|
||||
frozen_volume=int(_attr(row, ("m_nFrozenVolume", "frozen_volume"), 0) or 0),
|
||||
on_road_volume=int(_attr(row, ("m_nOnRoadVolume", "on_road_volume"), 0) or 0),
|
||||
yesterday_volume=int(_attr(row, ("m_nYesterdayVolume", "yesterday_volume"), 0) or 0),
|
||||
direction=int(_attr(row, ("m_nDirection", "direction"), 48) or 48),
|
||||
)
|
||||
return positions
|
||||
|
||||
def get_asset(self, account_id):
|
||||
query = self._require_query_func()
|
||||
rows = []
|
||||
for detail_type in ("ACCOUNT", "ASSET"):
|
||||
try:
|
||||
rows = query(account_id, self.account_type, detail_type) or []
|
||||
if rows:
|
||||
break
|
||||
except Exception:
|
||||
rows = []
|
||||
if not rows:
|
||||
return AssetSnapshot(account_id=account_id, cash=None, total_asset=None)
|
||||
|
||||
row = rows[0]
|
||||
cash = _attr(row, ("m_dAvailable", "m_dAvailableCash", "available_cash", "cash"))
|
||||
total_asset = _attr(row, ("m_dBalance", "m_dAsset", "total_asset", "asset"))
|
||||
frozen_cash = _attr(row, _FROZEN_CASH_FIELDS)
|
||||
market_value = _attr(row, _MARKET_VALUE_FIELDS)
|
||||
if frozen_cash is None:
|
||||
_report_missing_field("frozen_cash", row, _FROZEN_CASH_FIELDS)
|
||||
if market_value is None and cash is not None and total_asset is not None:
|
||||
# Derive only as a last resort. Without frozen_cash this overstates
|
||||
# market value by the frozen amount, so subtract it when known.
|
||||
market_value = float(total_asset) - float(cash)
|
||||
if frozen_cash is not None:
|
||||
market_value -= float(frozen_cash)
|
||||
return AssetSnapshot(
|
||||
account_id=account_id,
|
||||
cash=float(cash) if cash is not None else None,
|
||||
total_asset=float(total_asset) if total_asset is not None else None,
|
||||
frozen_cash=float(frozen_cash) if frozen_cash is not None else None,
|
||||
market_value=float(market_value) if market_value is not None else None,
|
||||
)
|
||||
-65
@@ -1,65 +0,0 @@
|
||||
"""Publish Big QMT position snapshots to Redis."""
|
||||
|
||||
import datetime as _dt
|
||||
import json
|
||||
|
||||
|
||||
class RedisPositionSyncSink:
|
||||
def __init__(
|
||||
self,
|
||||
redis_client,
|
||||
key_template="bigqmt:positions:{account_id}",
|
||||
event_stream_template="bigqmt:position_events:{account_id}",
|
||||
ttl_seconds=120,
|
||||
publish_events=True,
|
||||
):
|
||||
self.redis = redis_client
|
||||
self.key_template = key_template
|
||||
self.event_stream_template = event_stream_template
|
||||
self.ttl_seconds = int(ttl_seconds)
|
||||
self.publish_events = bool(publish_events)
|
||||
|
||||
@staticmethod
|
||||
def _time_text(value):
|
||||
if isinstance(value, _dt.datetime):
|
||||
return value.strftime("%Y-%m-%d %H:%M:%S")
|
||||
return str(value)
|
||||
|
||||
def _snapshot_to_dict(self, snapshot):
|
||||
return {
|
||||
"account_id": snapshot.account_id,
|
||||
"reason": snapshot.reason,
|
||||
"updated_at": self._time_text(snapshot.updated_at),
|
||||
"asset": {
|
||||
"cash": snapshot.asset.cash,
|
||||
"total_asset": snapshot.asset.total_asset,
|
||||
# Carried so the client's cached-asset fallback exposes the same
|
||||
# fields as a live query_stock_asset.
|
||||
"frozen_cash": getattr(snapshot.asset, "frozen_cash", None),
|
||||
"market_value": getattr(snapshot.asset, "market_value", None),
|
||||
},
|
||||
"positions": {
|
||||
code: {
|
||||
"stock_code": position.stock_code,
|
||||
"volume": position.volume,
|
||||
"available": position.available,
|
||||
"cost": position.cost,
|
||||
"stock_name": position.stock_name,
|
||||
}
|
||||
for code, position in snapshot.positions.items()
|
||||
},
|
||||
}
|
||||
|
||||
def publish(self, snapshot):
|
||||
payload = json.dumps(self._snapshot_to_dict(snapshot), ensure_ascii=False)
|
||||
key = self.key_template.format(account_id=snapshot.account_id)
|
||||
if self.ttl_seconds > 0:
|
||||
self.redis.setex(key, self.ttl_seconds, payload)
|
||||
else:
|
||||
self.redis.set(key, payload)
|
||||
if self.publish_events:
|
||||
stream_key = self.event_stream_template.format(account_id=snapshot.account_id)
|
||||
# Cap the stream to prevent unbounded memory growth. Order/trade
|
||||
# events already use maxlen=2000; position events were missing it,
|
||||
# causing 4.2GB+ streams in production (issue #21).
|
||||
self.redis.xadd(stream_key, {"payload": payload}, maxlen=2000, approximate=True)
|
||||
@@ -1,56 +0,0 @@
|
||||
"""Redis client helpers for Big QMT signal trader."""
|
||||
|
||||
import os
|
||||
|
||||
|
||||
def _float_or_none(value, default=None):
|
||||
if value is None:
|
||||
return default
|
||||
if value == "":
|
||||
return default
|
||||
text = str(value).strip()
|
||||
if text.lower() in ("none", "null"):
|
||||
return None
|
||||
return float(value)
|
||||
|
||||
|
||||
def build_redis_client(config=None):
|
||||
config = config or {}
|
||||
try:
|
||||
import redis
|
||||
except ImportError as exc: # pragma: no cover
|
||||
raise RuntimeError("redis package is required when Redis adapters are enabled") from exc
|
||||
|
||||
url = config.get("url") or os.environ.get("BIGQMT_REDIS_URL")
|
||||
if url:
|
||||
return redis.Redis.from_url(
|
||||
url,
|
||||
socket_connect_timeout=_float_or_none(config.get("socket_connect_timeout", 1.5), 1.5),
|
||||
socket_timeout=_float_or_none(config.get("socket_timeout", 1.5), 1.5),
|
||||
)
|
||||
|
||||
host = config.get("host") or os.environ.get("BIGQMT_REDIS_HOST") or "127.0.0.1"
|
||||
port = int(config.get("port") or os.environ.get("BIGQMT_REDIS_PORT") or 6379)
|
||||
db = int(config.get("db") or os.environ.get("BIGQMT_REDIS_DB") or 5)
|
||||
username = config.get("username") or os.environ.get("BIGQMT_REDIS_USERNAME") or None
|
||||
password = config.get("password") or os.environ.get("BIGQMT_REDIS_PASSWORD") or None
|
||||
return redis.Redis(
|
||||
host=host,
|
||||
port=port,
|
||||
db=db,
|
||||
username=username,
|
||||
password=password,
|
||||
socket_connect_timeout=_float_or_none(config.get("socket_connect_timeout", 1.5), 1.5),
|
||||
socket_timeout=_float_or_none(config.get("socket_timeout", 1.5), 1.5),
|
||||
health_check_interval=int(config.get("health_check_interval", 30)),
|
||||
)
|
||||
|
||||
|
||||
def decode_text(value):
|
||||
if isinstance(value, bytes):
|
||||
return value.decode("utf-8")
|
||||
return str(value)
|
||||
|
||||
|
||||
def redis_mapping_to_text(mapping):
|
||||
return {decode_text(key): decode_text(value) for key, value in (mapping or {}).items()}
|
||||
@@ -1,115 +0,0 @@
|
||||
"""Redis Stream signal source.
|
||||
|
||||
Redis is only a transport here. It must never call passorder or inspect QMT.
|
||||
"""
|
||||
|
||||
import datetime as _dt
|
||||
import json
|
||||
|
||||
from ..models import TradeSignal
|
||||
from .redis_common import decode_text, redis_mapping_to_text
|
||||
|
||||
|
||||
DEFAULT_STREAM_KEY_TEMPLATE = "bigqmt:signals:{account_id}"
|
||||
DEFAULT_GROUP = "bigqmt-signal-trader"
|
||||
|
||||
|
||||
def _json_default(value):
|
||||
if isinstance(value, (_dt.datetime, _dt.date)):
|
||||
return value.strftime("%Y-%m-%d %H:%M:%S")
|
||||
return str(value)
|
||||
|
||||
|
||||
def _coerce_scalar(value):
|
||||
text = decode_text(value).strip()
|
||||
lowered = text.lower()
|
||||
if lowered == "true":
|
||||
return True
|
||||
if lowered == "false":
|
||||
return False
|
||||
if lowered in ("none", "null"):
|
||||
return None
|
||||
return text
|
||||
|
||||
|
||||
def parse_stream_payload(fields):
|
||||
text_fields = redis_mapping_to_text(fields)
|
||||
payload_text = text_fields.get("payload") or text_fields.get("data")
|
||||
if payload_text:
|
||||
payload = json.loads(payload_text)
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("Redis stream payload must be a JSON object")
|
||||
return payload
|
||||
return {decode_text(key): _coerce_scalar(value) for key, value in fields.items()}
|
||||
|
||||
|
||||
class RedisStreamSignalSource:
|
||||
def __init__(
|
||||
self,
|
||||
redis_client,
|
||||
stream_key_template=DEFAULT_STREAM_KEY_TEMPLATE,
|
||||
group_name=DEFAULT_GROUP,
|
||||
consumer_name="bigqmt-consumer",
|
||||
block_ms=0,
|
||||
):
|
||||
self.redis = redis_client
|
||||
self.stream_key_template = stream_key_template
|
||||
self.group_name = group_name
|
||||
self.consumer_name = consumer_name
|
||||
self.block_ms = int(block_ms or 0)
|
||||
self._stream_ids_by_signal_id = {}
|
||||
self._created_groups = set()
|
||||
|
||||
def _stream_key(self, account_id):
|
||||
return self.stream_key_template.format(account_id=account_id)
|
||||
|
||||
def _ensure_group(self, stream_key):
|
||||
if stream_key in self._created_groups:
|
||||
return
|
||||
try:
|
||||
self.redis.xgroup_create(stream_key, self.group_name, id="0-0", mkstream=True)
|
||||
except Exception as exc:
|
||||
if "BUSYGROUP" not in str(exc):
|
||||
raise
|
||||
self._created_groups.add(stream_key)
|
||||
|
||||
def fetch(self, account_id, limit):
|
||||
stream_key = self._stream_key(account_id)
|
||||
self._ensure_group(stream_key)
|
||||
kwargs = {
|
||||
"groupname": self.group_name,
|
||||
"consumername": self.consumer_name,
|
||||
"streams": {stream_key: ">"},
|
||||
"count": int(limit),
|
||||
}
|
||||
if self.block_ms > 0:
|
||||
kwargs["block"] = self.block_ms
|
||||
rows = self.redis.xreadgroup(**kwargs) or []
|
||||
signals = []
|
||||
for _, entries in rows:
|
||||
for stream_id, fields in entries:
|
||||
payload = parse_stream_payload(fields)
|
||||
signal = TradeSignal.from_dict(payload)
|
||||
self._stream_ids_by_signal_id[signal.signal_id] = (stream_key, stream_id)
|
||||
signals.append(signal)
|
||||
return signals
|
||||
|
||||
def ack(self, signal):
|
||||
ref = self._stream_ids_by_signal_id.pop(signal.signal_id, None)
|
||||
if not ref:
|
||||
return None
|
||||
stream_key, stream_id = ref
|
||||
return self.redis.xack(stream_key, self.group_name, stream_id)
|
||||
|
||||
|
||||
def push_trade_signal(redis_client, payload, account_id=None, stream_key_template=DEFAULT_STREAM_KEY_TEMPLATE):
|
||||
if isinstance(payload, TradeSignal):
|
||||
account_id = account_id or payload.account_id
|
||||
raw_payload = dict(payload.raw_payload)
|
||||
else:
|
||||
raw_payload = dict(payload)
|
||||
account_id = account_id or raw_payload.get("account_id")
|
||||
if not account_id:
|
||||
raise ValueError("account_id is required")
|
||||
stream_key = stream_key_template.format(account_id=account_id)
|
||||
return redis_client.xadd(stream_key, {"payload": json.dumps(raw_payload, ensure_ascii=False, default=_json_default)})
|
||||
@@ -1,90 +0,0 @@
|
||||
"""Redis signal state store."""
|
||||
|
||||
import datetime as _dt
|
||||
|
||||
|
||||
class RedisStateStore:
|
||||
def __init__(
|
||||
self,
|
||||
redis_client,
|
||||
account_id="default",
|
||||
claim_key_template="bigqmt:signal_claim:{account_id}:{signal_id}",
|
||||
status_key_template="bigqmt:signal_status:{account_id}:{signal_id}",
|
||||
claim_ttl_seconds=3600,
|
||||
status_ttl_seconds=86400,
|
||||
):
|
||||
self.redis = redis_client
|
||||
self.account_id = account_id
|
||||
self.claim_key_template = claim_key_template
|
||||
self.status_key_template = status_key_template
|
||||
self.claim_ttl_seconds = int(claim_ttl_seconds)
|
||||
self.status_ttl_seconds = int(status_ttl_seconds)
|
||||
self._accounts_by_signal_id = {}
|
||||
|
||||
def _account_for(self, signal_id):
|
||||
return self._accounts_by_signal_id.get(signal_id) or self.account_id
|
||||
|
||||
def _claim_key(self, account_id, signal_id):
|
||||
return self.claim_key_template.format(account_id=account_id, signal_id=signal_id)
|
||||
|
||||
def _status_key(self, account_id, signal_id):
|
||||
return self.status_key_template.format(account_id=account_id, signal_id=signal_id)
|
||||
|
||||
@staticmethod
|
||||
def _now_text():
|
||||
return _dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
def _write_status(self, account_id, signal_id, mapping):
|
||||
key = self._status_key(account_id, signal_id)
|
||||
fields = {
|
||||
"signal_id": signal_id,
|
||||
"account_id": account_id,
|
||||
"updated_at": self._now_text(),
|
||||
}
|
||||
fields.update({k: "" if v is None else str(v) for k, v in mapping.items()})
|
||||
self.redis.hset(key, mapping=fields)
|
||||
if self.status_ttl_seconds > 0:
|
||||
self.redis.expire(key, self.status_ttl_seconds)
|
||||
|
||||
def claim(self, signal, consumer_id):
|
||||
account_id = signal.account_id or self.account_id
|
||||
self._accounts_by_signal_id[signal.signal_id] = account_id
|
||||
key = self._claim_key(account_id, signal.signal_id)
|
||||
ok = self.redis.set(key, consumer_id, nx=True, ex=self.claim_ttl_seconds)
|
||||
if ok:
|
||||
self._write_status(
|
||||
account_id,
|
||||
signal.signal_id,
|
||||
{
|
||||
"status": "CLAIMED",
|
||||
"consumer_id": consumer_id,
|
||||
"stock_code": signal.stock_code,
|
||||
"action": signal.action.value,
|
||||
"message": "",
|
||||
},
|
||||
)
|
||||
return bool(ok)
|
||||
|
||||
def mark_submitted(self, signal_id, result):
|
||||
account_id = self._account_for(signal_id)
|
||||
self._write_status(
|
||||
account_id,
|
||||
signal_id,
|
||||
{
|
||||
"status": result.status,
|
||||
"user_order_id": result.user_order_id,
|
||||
"order_sys_id": result.order_sys_id,
|
||||
"message": result.message,
|
||||
},
|
||||
)
|
||||
|
||||
def mark_finished(self, signal_id, status, message=""):
|
||||
account_id = self._account_for(signal_id)
|
||||
self._write_status(
|
||||
account_id,
|
||||
signal_id,
|
||||
{
|
||||
"status": status,
|
||||
"message": message,
|
||||
},
|
||||
)
|
||||
@@ -1,98 +0,0 @@
|
||||
"""信号交易应用编排层。"""
|
||||
|
||||
import datetime as _dt
|
||||
|
||||
from .models import AccountSnapshot, OrderRequest, SignalAction
|
||||
from .price_engine import build_order_price
|
||||
from .risk_guard import validate_signal
|
||||
|
||||
|
||||
class SignalTradingApp:
|
||||
def __init__(
|
||||
self,
|
||||
account_id,
|
||||
signal_source,
|
||||
market_data,
|
||||
position_provider,
|
||||
order_gateway,
|
||||
position_sync_sink,
|
||||
state_store,
|
||||
consumer_id="bigqmt-signal-trader",
|
||||
fetch_limit=20,
|
||||
):
|
||||
self.account_id = account_id
|
||||
self.signal_source = signal_source
|
||||
self.market_data = market_data
|
||||
self.position_provider = position_provider
|
||||
self.order_gateway = order_gateway
|
||||
self.position_sync_sink = position_sync_sink
|
||||
self.state_store = state_store
|
||||
self.consumer_id = consumer_id
|
||||
self.fetch_limit = int(fetch_limit)
|
||||
|
||||
def tick(self, now=None):
|
||||
now = now or _dt.datetime.now()
|
||||
signals = self.signal_source.fetch(self.account_id, self.fetch_limit)
|
||||
positions = self.position_provider.get_positions(self.account_id)
|
||||
|
||||
for signal in signals:
|
||||
if not self.state_store.claim(signal, self.consumer_id):
|
||||
continue
|
||||
try:
|
||||
self._handle_signal(signal, now, positions)
|
||||
except Exception as exc:
|
||||
self.state_store.mark_finished(signal.signal_id, "FAILED", str(exc))
|
||||
self.signal_source.ack(signal)
|
||||
|
||||
self.sync_positions("tick", now=now)
|
||||
|
||||
def _handle_signal(self, signal, now, positions):
|
||||
decision = validate_signal(signal, now, positions)
|
||||
if not decision.allowed:
|
||||
self.state_store.mark_finished(signal.signal_id, "SKIPPED", decision.reason)
|
||||
self.signal_source.ack(signal)
|
||||
return
|
||||
|
||||
price = build_order_price(
|
||||
self.market_data,
|
||||
decision.stock_code,
|
||||
signal.action.value,
|
||||
price_type=signal.price_type,
|
||||
fixed_price=signal.price,
|
||||
)
|
||||
request = OrderRequest(
|
||||
signal_id=signal.signal_id,
|
||||
account_id=signal.account_id,
|
||||
action=signal.action.value,
|
||||
stock_code=decision.stock_code,
|
||||
volume=decision.volume,
|
||||
price=price,
|
||||
price_type="LIMIT",
|
||||
strategy_name=signal.strategy_name,
|
||||
remark=signal.remark,
|
||||
)
|
||||
result = self.order_gateway.submit(request)
|
||||
self.state_store.mark_submitted(signal.signal_id, result)
|
||||
self.signal_source.ack(signal)
|
||||
|
||||
def on_init(self, runtime):
|
||||
return None
|
||||
|
||||
def on_order_event(self, event):
|
||||
return None
|
||||
|
||||
def on_trade_event(self, event):
|
||||
self.sync_positions("trade_event")
|
||||
|
||||
def sync_positions(self, reason, now=None):
|
||||
now = now or _dt.datetime.now()
|
||||
asset = self.position_provider.get_asset(self.account_id)
|
||||
positions = self.position_provider.get_positions(self.account_id)
|
||||
snapshot = AccountSnapshot(
|
||||
account_id=self.account_id,
|
||||
asset=asset,
|
||||
positions=positions,
|
||||
reason=reason,
|
||||
updated_at=now,
|
||||
)
|
||||
self.position_sync_sink.publish(snapshot)
|
||||
@@ -1,48 +0,0 @@
|
||||
"""证券代码标准化和委托数量处理。"""
|
||||
|
||||
import re
|
||||
|
||||
|
||||
_DIGIT_CODE_RE = re.compile(r"^\d{6}$")
|
||||
|
||||
|
||||
def normalize_stock_code(code):
|
||||
text = str(code or "").strip().upper()
|
||||
if not text:
|
||||
return ""
|
||||
if text.startswith("SH") and _DIGIT_CODE_RE.match(text[2:]):
|
||||
return f"{text[2:]}.SH"
|
||||
if text.startswith("SZ") and _DIGIT_CODE_RE.match(text[2:]):
|
||||
return f"{text[2:]}.SZ"
|
||||
if text.endswith(".SH") or text.endswith(".SZ"):
|
||||
prefix = text[:6]
|
||||
if _DIGIT_CODE_RE.match(prefix):
|
||||
return text
|
||||
if _DIGIT_CODE_RE.match(text):
|
||||
market = "SH" if text.startswith(("5", "6")) else "SZ"
|
||||
return f"{text}.{market}"
|
||||
raise ValueError(f"invalid stock code: {code}")
|
||||
|
||||
|
||||
def min_lot(stock_code):
|
||||
normalized = normalize_stock_code(stock_code)
|
||||
pure = normalized.split(".")[0]
|
||||
return 200 if pure.startswith("688") else 100
|
||||
|
||||
|
||||
def round_buy_volume(stock_code, amount):
|
||||
lot = min_lot(stock_code)
|
||||
value = int(amount or 0)
|
||||
if value <= 0:
|
||||
return 0
|
||||
return (value // lot) * lot
|
||||
|
||||
|
||||
def round_sell_volume(stock_code, amount, sell_all=False):
|
||||
value = int(amount or 0)
|
||||
if value <= 0:
|
||||
return 0
|
||||
if sell_all:
|
||||
return value
|
||||
lot = min_lot(stock_code)
|
||||
return (value // lot) * lot
|
||||
@@ -1,81 +0,0 @@
|
||||
"""可替换 adapter 的接口定义。"""
|
||||
|
||||
import datetime as _dt
|
||||
from typing import Dict, List
|
||||
|
||||
try:
|
||||
from typing import Protocol
|
||||
except ImportError: # pragma: no cover
|
||||
from typing_extensions import Protocol
|
||||
|
||||
from .models import (
|
||||
AccountSnapshot,
|
||||
AssetSnapshot,
|
||||
CancelResult,
|
||||
OrderRef,
|
||||
OrderRequest,
|
||||
OrderSnapshot,
|
||||
OrderSubmitResult,
|
||||
PositionSnapshot,
|
||||
TradeSignal,
|
||||
TradeSnapshot,
|
||||
)
|
||||
|
||||
|
||||
class SignalSource(Protocol):
|
||||
def fetch(self, account_id: str, limit: int) -> List[TradeSignal]:
|
||||
...
|
||||
|
||||
def ack(self, signal: TradeSignal) -> None:
|
||||
...
|
||||
|
||||
|
||||
class MarketDataProvider(Protocol):
|
||||
def get_ticks(self, codes: List[str]) -> Dict[str, dict]:
|
||||
...
|
||||
|
||||
def get_instrument(self, code: str) -> dict:
|
||||
...
|
||||
|
||||
|
||||
class PositionProvider(Protocol):
|
||||
def get_positions(self, account_id: str) -> Dict[str, PositionSnapshot]:
|
||||
...
|
||||
|
||||
def get_asset(self, account_id: str) -> AssetSnapshot:
|
||||
...
|
||||
|
||||
|
||||
class OrderGateway(Protocol):
|
||||
def submit(self, request: OrderRequest) -> OrderSubmitResult:
|
||||
...
|
||||
|
||||
def cancel(self, order_ref: OrderRef) -> CancelResult:
|
||||
...
|
||||
|
||||
def query_orders(self, account_id: str, strategy_name: str) -> List[OrderSnapshot]:
|
||||
...
|
||||
|
||||
def query_trades(self, account_id: str, strategy_name: str) -> List[TradeSnapshot]:
|
||||
...
|
||||
|
||||
|
||||
class PositionSyncSink(Protocol):
|
||||
def publish(self, snapshot: AccountSnapshot) -> None:
|
||||
...
|
||||
|
||||
|
||||
class StateStore(Protocol):
|
||||
def claim(self, signal: TradeSignal, consumer_id: str) -> bool:
|
||||
...
|
||||
|
||||
def mark_submitted(self, signal_id: str, result: OrderSubmitResult) -> None:
|
||||
...
|
||||
|
||||
def mark_finished(self, signal_id: str, status: str, message: str = "") -> None:
|
||||
...
|
||||
|
||||
|
||||
class RuntimeAdapter(Protocol):
|
||||
def now(self) -> _dt.datetime:
|
||||
...
|
||||
@@ -1,252 +0,0 @@
|
||||
"""Async, chunked download jobs for Big QMT.
|
||||
|
||||
A client submits a download job (fire-and-forget) into a Redis queue and polls
|
||||
its status. The Big QMT strategy thread drains one job at a time and downloads a
|
||||
bounded slice of symbols per tick (``chunk_size`` symbols, capped by a wall-clock
|
||||
budget), so a long ``download_history_data2`` never blocks the strategy thread /
|
||||
RPC pump. Historical bars land in the Big QMT machine's local store; clients then
|
||||
read them back with fast ``get_local_data`` / ``get_market_data`` calls.
|
||||
|
||||
Redis layout (per account). All stored VALUES are digit-free encoded (see _enc)
|
||||
so the QMT terminal's redis compliance filter never trips on stock codes in the
|
||||
job data the pump reads back:
|
||||
- ``bigqmt:dljob:pending:{account_id}`` list of pending job ids (RPUSH/LPOP)
|
||||
- ``bigqmt:dljob:item:{account_id}:{job_id}`` encoded job blob incl. progress
|
||||
- ``bigqmt:dljob:active:{account_id}`` id of the job being processed now
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
import uuid
|
||||
|
||||
|
||||
# Dedicated "bigqmt:dljob:*" namespace for the client<->pump protocol.
|
||||
QUEUE_KEY_TEMPLATE = "bigqmt:dljob:pending:{account_id}"
|
||||
JOB_KEY_TEMPLATE = "bigqmt:dljob:item:{account_id}:{job_id}"
|
||||
CURRENT_KEY_TEMPLATE = "bigqmt:dljob:active:{account_id}"
|
||||
|
||||
DEFAULT_JOB_TTL_SECONDS = 3600
|
||||
DEFAULT_CHUNK_SIZE = 10
|
||||
DEFAULT_MAX_WALL_SECONDS = 0.5
|
||||
|
||||
# Terminal + in-flight states.
|
||||
PENDING = "pending"
|
||||
RUNNING = "running"
|
||||
DONE = "done"
|
||||
FAILED = "failed"
|
||||
_ACTIVE_STATES = (PENDING, RUNNING)
|
||||
|
||||
|
||||
def queue_key(account_id):
|
||||
return QUEUE_KEY_TEMPLATE.format(account_id=str(account_id or ""))
|
||||
|
||||
|
||||
def job_key(account_id, job_id):
|
||||
return JOB_KEY_TEMPLATE.format(account_id=str(account_id or ""), job_id=str(job_id or ""))
|
||||
|
||||
|
||||
def current_key(account_id):
|
||||
return CURRENT_KEY_TEMPLATE.format(account_id=str(account_id or ""))
|
||||
|
||||
|
||||
def _text(value):
|
||||
if value is None:
|
||||
return ""
|
||||
if isinstance(value, bytes):
|
||||
return value.decode("utf-8")
|
||||
return str(value)
|
||||
|
||||
|
||||
# The 国金证券 QMT terminal ships a redis client whose check_response() raises
|
||||
# "Sensitive Data Detected, Forbidden!" whenever a Redis *response* contains a
|
||||
# stock-code + operation-code DIGIT pattern (a brokerage control against trading
|
||||
# signals flowing through Redis). The pump runs inside that terminal and must read
|
||||
# job data (which contains stock codes) back from Redis. So every value the pump
|
||||
# reads is stored as a DIGIT-FREE token: hex-encode, then shift digits 0-9 -> the
|
||||
# letters g-p, making the stored value all letters (a-p). The stock-code regex
|
||||
# requires digits, so it can never match. Reversible; writes are never filtered
|
||||
# (only responses are), so only read-back values need this.
|
||||
_DIGIT_TO_ALPHA = str.maketrans("0123456789", "ghijklmnop")
|
||||
_ALPHA_TO_DIGIT = str.maketrans("ghijklmnop", "0123456789")
|
||||
|
||||
|
||||
def _enc(text_value):
|
||||
return _text(text_value).encode("utf-8").hex().translate(_DIGIT_TO_ALPHA)
|
||||
|
||||
|
||||
def _dec(token):
|
||||
text = _text(token)
|
||||
if not text:
|
||||
return None
|
||||
try:
|
||||
return bytes.fromhex(text.translate(_ALPHA_TO_DIGIT)).decode("utf-8")
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def submit_download_job(
|
||||
redis_client,
|
||||
account_id,
|
||||
stock_list,
|
||||
period,
|
||||
method="download_history_data2",
|
||||
start_time="",
|
||||
end_time="",
|
||||
incrementally=None,
|
||||
chunk_size=DEFAULT_CHUNK_SIZE,
|
||||
job_ttl_seconds=DEFAULT_JOB_TTL_SECONDS,
|
||||
):
|
||||
"""Queue a download job and return its initial status dict (non-blocking)."""
|
||||
codes = [str(code) for code in (stock_list or []) if str(code or "").strip()]
|
||||
if not codes:
|
||||
raise ValueError("stock_list is required for a download job")
|
||||
job_id = uuid.uuid4().hex[:16]
|
||||
now = time.time()
|
||||
job = {
|
||||
"job_id": job_id,
|
||||
"method": str(method or "download_history_data2"),
|
||||
"stock_list": codes,
|
||||
"period": period,
|
||||
"start_time": start_time or "",
|
||||
"end_time": end_time or "",
|
||||
"incrementally": incrementally,
|
||||
"chunk_size": int(chunk_size or DEFAULT_CHUNK_SIZE),
|
||||
"total": len(codes),
|
||||
"done": 0,
|
||||
"state": PENDING,
|
||||
"error": "",
|
||||
"created_at_ts": now,
|
||||
"updated_at_ts": now,
|
||||
}
|
||||
ttl = int(max(1, job_ttl_seconds))
|
||||
redis_client.setex(job_key(account_id, job_id), ttl, _enc(json.dumps(job, ensure_ascii=False)))
|
||||
redis_client.rpush(queue_key(account_id), _enc(job_id))
|
||||
try:
|
||||
redis_client.expire(queue_key(account_id), ttl)
|
||||
except Exception:
|
||||
pass
|
||||
return job
|
||||
|
||||
|
||||
def read_download_status(redis_client, account_id, job_id):
|
||||
"""Return the current job status dict, or None if unknown/expired."""
|
||||
decoded = _dec(redis_client.get(job_key(account_id, job_id)))
|
||||
if not decoded:
|
||||
return None
|
||||
try:
|
||||
job = json.loads(decoded)
|
||||
except Exception:
|
||||
return None
|
||||
return job if isinstance(job, dict) else None
|
||||
|
||||
|
||||
def wait_download_job(
|
||||
redis_client,
|
||||
account_id,
|
||||
job_id,
|
||||
wait_seconds=600.0,
|
||||
poll_interval_seconds=0.5,
|
||||
):
|
||||
"""Block (client-side only) until the job reaches a terminal state or timeout."""
|
||||
deadline = time.time() + max(0.0, float(wait_seconds))
|
||||
while True:
|
||||
status = read_download_status(redis_client, account_id, job_id)
|
||||
if status and status.get("state") in (DONE, FAILED):
|
||||
return status
|
||||
if time.time() >= deadline:
|
||||
return status
|
||||
time.sleep(max(0.05, float(poll_interval_seconds)))
|
||||
|
||||
|
||||
def _write_job(redis_client, account_id, job, job_ttl_seconds):
|
||||
job["updated_at_ts"] = time.time()
|
||||
ttl = int(max(1, job_ttl_seconds))
|
||||
redis_client.setex(job_key(account_id, job["job_id"]), ttl, _enc(json.dumps(job, ensure_ascii=False)))
|
||||
|
||||
|
||||
def _acquire_current_job(redis_client, account_id):
|
||||
ckey = current_key(account_id)
|
||||
current_id = _dec(redis_client.get(ckey))
|
||||
if current_id:
|
||||
job = read_download_status(redis_client, account_id, current_id)
|
||||
if job and job.get("state") in _ACTIVE_STATES:
|
||||
return job
|
||||
# Stale pointer (job done/failed/expired): drop it and pick the next one.
|
||||
redis_client.delete(ckey)
|
||||
while True:
|
||||
job_id = _dec(redis_client.lpop(queue_key(account_id)))
|
||||
if not job_id:
|
||||
return None
|
||||
job = read_download_status(redis_client, account_id, job_id)
|
||||
if job and job.get("state") in _ACTIVE_STATES:
|
||||
redis_client.set(ckey, _enc(job_id))
|
||||
return job
|
||||
# Skip unknown/expired/finished ids left in the queue.
|
||||
|
||||
|
||||
def _download_chunk(market_data, method, chunk, period, start_time, end_time, incrementally):
|
||||
if method == "download_history_data":
|
||||
for code in chunk:
|
||||
market_data.download_history_data(code, period, start_time, end_time, incrementally)
|
||||
else:
|
||||
market_data.download_history_data2(chunk, period, start_time, end_time, incrementally)
|
||||
|
||||
|
||||
def pump_download_jobs(
|
||||
redis_client,
|
||||
market_data,
|
||||
account_id,
|
||||
chunk_size=DEFAULT_CHUNK_SIZE,
|
||||
max_wall_seconds=DEFAULT_MAX_WALL_SECONDS,
|
||||
job_ttl_seconds=DEFAULT_JOB_TTL_SECONDS,
|
||||
):
|
||||
"""Advance the active download job by a bounded slice. Call once per tick.
|
||||
|
||||
Downloads at least one chunk (so progress is always made) and keeps going
|
||||
until the wall-clock budget is spent. Returns a small status summary, or None
|
||||
when there is no active job. Runs on the caller (strategy) thread.
|
||||
"""
|
||||
job = _acquire_current_job(redis_client, account_id)
|
||||
if job is None:
|
||||
return None
|
||||
stock_list = job.get("stock_list") or []
|
||||
total = int(job.get("total") or len(stock_list))
|
||||
done = int(job.get("done") or 0)
|
||||
step = int(job.get("chunk_size") or chunk_size or DEFAULT_CHUNK_SIZE)
|
||||
if step <= 0:
|
||||
step = DEFAULT_CHUNK_SIZE
|
||||
method = str(job.get("method") or "download_history_data2")
|
||||
period = job.get("period")
|
||||
start_time = job.get("start_time") or ""
|
||||
end_time = job.get("end_time") or ""
|
||||
incrementally = job.get("incrementally")
|
||||
|
||||
started_at = time.time()
|
||||
processed_this_tick = 0
|
||||
try:
|
||||
while done < total:
|
||||
# Always run one chunk; only the budget check (after the first) can
|
||||
# stop the tick, so a single heavy chunk is the smallest block unit.
|
||||
if max_wall_seconds and processed_this_tick and (time.time() - started_at) > float(max_wall_seconds):
|
||||
break
|
||||
chunk = stock_list[done:done + step]
|
||||
_download_chunk(market_data, method, chunk, period, start_time, end_time, incrementally)
|
||||
done += len(chunk)
|
||||
processed_this_tick += len(chunk)
|
||||
except Exception as exc:
|
||||
job["state"] = FAILED
|
||||
job["error"] = "%s: %s" % (exc.__class__.__name__, exc)
|
||||
job["done"] = done
|
||||
_write_job(redis_client, account_id, job, job_ttl_seconds)
|
||||
redis_client.delete(current_key(account_id))
|
||||
return {"job_id": job["job_id"], "state": FAILED, "done": done, "total": total, "error": job["error"]}
|
||||
|
||||
job["done"] = done
|
||||
if done >= total:
|
||||
job["state"] = DONE
|
||||
_write_job(redis_client, account_id, job, job_ttl_seconds)
|
||||
redis_client.delete(current_key(account_id))
|
||||
return {"job_id": job["job_id"], "state": DONE, "done": done, "total": total}
|
||||
job["state"] = RUNNING
|
||||
_write_job(redis_client, account_id, job, job_ttl_seconds)
|
||||
return {"job_id": job["job_id"], "state": RUNNING, "done": done, "total": total}
|
||||
@@ -1,449 +0,0 @@
|
||||
"""Real-time order/trade (execution) event push over Redis.
|
||||
|
||||
Big QMT fires ``order_callback(ContextInfo, orderInfo)`` and
|
||||
``deal_callback(ContextInfo, dealInfo)`` inside the strategy process. We normalize
|
||||
the QMT order/deal object (ThinkTrader ``m_*`` fields) into a plain dict and
|
||||
publish it to a Redis channel, so clients receive ``on_stock_order`` /
|
||||
``on_stock_trade`` callbacks in real time (MiniQMT style) instead of polling.
|
||||
|
||||
Channels (also used as capped streams for short replay, xadd + publish):
|
||||
- ``bigqmt:order_events:{account_id}``
|
||||
- ``bigqmt:trade_events:{account_id}``
|
||||
|
||||
The normalized field names match ``BigQmtXtTrader._order_from_dict`` /
|
||||
``_trade_from_dict`` so the client can shape them straight into MiniQMT objects.
|
||||
"""
|
||||
|
||||
import json
|
||||
import time
|
||||
|
||||
|
||||
ORDER_CHANNEL_TEMPLATE = "bigqmt:order_events:{account_id}"
|
||||
TRADE_CHANNEL_TEMPLATE = "bigqmt:trade_events:{account_id}"
|
||||
ORDER_ERROR_CHANNEL_TEMPLATE = "bigqmt:order_error_events:{account_id}"
|
||||
CANCEL_ERROR_CHANNEL_TEMPLATE = "bigqmt:cancel_error_events:{account_id}"
|
||||
ORDER_IDENTITY_KEY_TEMPLATE = "bigqmt:order_identity:{account_id}:{user_order_id}"
|
||||
|
||||
EVENT_ORDER = "order"
|
||||
EVENT_TRADE = "trade"
|
||||
EVENT_ORDER_ERROR = "order_error"
|
||||
EVENT_CANCEL_ERROR = "cancel_error"
|
||||
|
||||
# ThinkTrader enum_EEntrustBS (买卖方向, the m_nDirection field), universal across
|
||||
# 股票/期货/期权. Ref: https://dict.thinktrader.net/innerApi/enum_constants.html
|
||||
ENTRUST_BUY = 48 # 买入 / 多
|
||||
ENTRUST_SELL = 49 # 卖出 / 空
|
||||
ENTRUST_PLEDGE_IN = 81 # 质押入库
|
||||
ENTRUST_PLEDGE_OUT = 66 # 质押出库
|
||||
|
||||
# enum_EEntrustBS (买卖方向, the m_nDirection field), per QMT enum docs.
|
||||
# 48=买, 49=卖. Universal across 股票/期货/期权.
|
||||
#
|
||||
# Real-world findings from live COrderDetail/CDealDetail callbacks
|
||||
# (diagnosed via exec_events_debug_raw_fields=True, 2026-07-29):
|
||||
# QMT returns m_nDirection=48 **unconditionally** — even for sell orders.
|
||||
# m_nOffsetFlag correctly reflects direction (48=买, 49=卖 for stocks).
|
||||
# m_nOpType correctly reflects direction (23=买, 24=卖) on orders.
|
||||
# query_orders uses m_nOffsetFlag and works correctly in production.
|
||||
#
|
||||
# Therefore _extract_direction uses an arbitration chain:
|
||||
# Preferred: m_nOffsetFlag (most reliable in live callbacks, matches query_orders)
|
||||
# Fallback: m_nDirection (traditional EEntrustBS; can be stuck at 48 in calls)
|
||||
# Arbiter: when direction≠offset (futures: sell+open=49+48),
|
||||
# consult m_nOpType (23/24) to resolve the conflict; for trades
|
||||
# (no m_nOpType) trust m_nOffsetFlag (QMT docs confirm stock
|
||||
# direction=offset).
|
||||
# Last: order_type (MiniQMT STOCK_BUY=23 / STOCK_SELL=24) and plain text
|
||||
# Unknown -> "" (the raw value is always preserved so callers can refine).
|
||||
OFFSET_OPEN = 48
|
||||
OFFSET_CLOSE = 49
|
||||
OFFSET_CLOSE_TODAY = 51
|
||||
OFFSET_CLOSE_YESTERDAY = 52
|
||||
|
||||
_BUY_DIRECTIONS = {ENTRUST_BUY, str(ENTRUST_BUY), OFFSET_OPEN, str(OFFSET_OPEN), 23, "23", "BUY", "buy", "B"}
|
||||
_SELL_DIRECTIONS = {ENTRUST_SELL, str(ENTRUST_SELL), OFFSET_CLOSE, str(OFFSET_CLOSE), OFFSET_CLOSE_TODAY, str(OFFSET_CLOSE_TODAY), OFFSET_CLOSE_YESTERDAY, str(OFFSET_CLOSE_YESTERDAY), 24, "24", "SELL", "sell", "S"}
|
||||
|
||||
|
||||
def order_channel(account_id):
|
||||
return ORDER_CHANNEL_TEMPLATE.format(account_id=str(account_id or ""))
|
||||
|
||||
|
||||
def trade_channel(account_id):
|
||||
return TRADE_CHANNEL_TEMPLATE.format(account_id=str(account_id or ""))
|
||||
|
||||
|
||||
def order_error_channel(account_id):
|
||||
return ORDER_ERROR_CHANNEL_TEMPLATE.format(account_id=str(account_id or ""))
|
||||
|
||||
|
||||
def cancel_error_channel(account_id):
|
||||
return CANCEL_ERROR_CHANNEL_TEMPLATE.format(account_id=str(account_id or ""))
|
||||
|
||||
|
||||
def order_identity_key(account_id, user_order_id):
|
||||
return ORDER_IDENTITY_KEY_TEMPLATE.format(
|
||||
account_id=str(account_id or ""),
|
||||
user_order_id=str(user_order_id or ""),
|
||||
)
|
||||
|
||||
|
||||
def _attr(obj, names, default=None):
|
||||
for name in names:
|
||||
if isinstance(obj, dict):
|
||||
if name in obj and obj[name] is not None:
|
||||
return obj[name]
|
||||
else:
|
||||
value = getattr(obj, name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
return default
|
||||
|
||||
|
||||
def _action_from_direction(direction):
|
||||
if direction in _BUY_DIRECTIONS:
|
||||
return "BUY"
|
||||
if direction in _SELL_DIRECTIONS:
|
||||
return "SELL"
|
||||
return ""
|
||||
|
||||
|
||||
def _is_buy(val):
|
||||
v = int(val)
|
||||
return v in _BUY_DIRECTIONS
|
||||
|
||||
|
||||
def _is_sell(val):
|
||||
v = int(val)
|
||||
return v in _SELL_DIRECTIONS
|
||||
|
||||
|
||||
def _conflict_resolve(d_val, o_val, obj):
|
||||
"""When m_nDirection and m_nOffsetFlag disagree, arbitrate via m_nOpType.
|
||||
|
||||
Live diagnosis confirms:
|
||||
- Stock sell: direction=48(buy), offset=49(sell), op_type=24(sell) → sell
|
||||
- Futures sell+open: direction=49(sell), offset=48(open), op_type=24(sell) → sell
|
||||
- Futures buy+close: direction=48(buy), offset=49(close), op_type=23(buy) → buy
|
||||
|
||||
Returns a resolved value, or None if no arbiter can decide.
|
||||
"""
|
||||
op = _attr(obj, ["m_nOpType", "op_type", "order_type"])
|
||||
if op is not None:
|
||||
try:
|
||||
op_int = int(op)
|
||||
if op_int in _BUY_DIRECTIONS:
|
||||
return d_val if _is_buy(d_val) else o_val if _is_buy(o_val) else op
|
||||
if op_int in _SELL_DIRECTIONS:
|
||||
return d_val if _is_sell(d_val) else o_val if _is_sell(o_val) else op
|
||||
except (TypeError, ValueError):
|
||||
if op in _BUY_DIRECTIONS:
|
||||
return d_val if _is_buy(d_val) else o_val if _is_buy(o_val) else op
|
||||
if op in _SELL_DIRECTIONS:
|
||||
return d_val if _is_sell(d_val) else o_val if _is_sell(o_val) else op
|
||||
# no arbiter — trust offset (QMT docs confirm stock direction=offset)
|
||||
return o_val
|
||||
|
||||
|
||||
def _extract_direction(obj):
|
||||
"""Extract buy/sell direction, matching query_orders' reliable logic.
|
||||
|
||||
Priority chain (documented with live-diagnosis justification):
|
||||
1. m_nOffsetFlag — most reliable in live callbacks (matches query_orders)
|
||||
2. m_nDirection — traditional EEntrustBS (can be stuck at 48)
|
||||
3. Arbitration: when direction≠offset, consult m_nOpType (orders: 23/24)
|
||||
to resolve correctly for both stocks AND futures.
|
||||
4. m_nOpType / order_type — last resort fallback.
|
||||
|
||||
The raw value is always returned (even pledge=81) so callers can inspect it;
|
||||
_action_from_direction maps only known buy/sell values, leaving others "".
|
||||
|
||||
References
|
||||
----------
|
||||
- Live diagnosis 2026-07-29 (COrderDetail/CDealDetail):
|
||||
m_nDirection=48 unconditionally, m_nOffsetFlag=48(buy)/49(sell) correct,
|
||||
m_nOpType=23(buy)/24(sell) correct (orders only).
|
||||
- QMT enum docs: enum_EEntrustBS (48=买,49=卖), enum_EOffset_Flag_Type
|
||||
(48=开仓,49=平仓). For stocks direction=offset; for futures they differ.
|
||||
- query_orders uses m_nOffsetFlag and works correctly in production.
|
||||
"""
|
||||
offset = _attr(obj, ["m_nOffsetFlag", "offset_flag"])
|
||||
direction = _attr(obj, ["m_nDirection", "direction"])
|
||||
|
||||
# 1. offset alone — use it directly (matches query_orders)
|
||||
if offset is not None and direction is None:
|
||||
try:
|
||||
o = int(offset)
|
||||
if o in _BUY_DIRECTIONS or o in _SELL_DIRECTIONS:
|
||||
return offset
|
||||
except (TypeError, ValueError):
|
||||
if offset in _BUY_DIRECTIONS or offset in _SELL_DIRECTIONS:
|
||||
return offset
|
||||
|
||||
# 2. direction alone — use it
|
||||
if direction is not None and offset is None:
|
||||
try:
|
||||
d = int(direction)
|
||||
if d in _BUY_DIRECTIONS or d in _SELL_DIRECTIONS:
|
||||
return direction
|
||||
if d != 0:
|
||||
return direction
|
||||
except (TypeError, ValueError):
|
||||
if direction in _BUY_DIRECTIONS or direction in _SELL_DIRECTIONS:
|
||||
return direction
|
||||
return direction
|
||||
|
||||
# 3. both present
|
||||
if direction is not None and offset is not None:
|
||||
try:
|
||||
d = int(direction)
|
||||
o = int(offset)
|
||||
d_valid = (d in _BUY_DIRECTIONS or d in _SELL_DIRECTIONS)
|
||||
o_valid = (o in _BUY_DIRECTIONS or o in _SELL_DIRECTIONS)
|
||||
|
||||
if d_valid and o_valid:
|
||||
if d == o:
|
||||
return direction # agree → use either
|
||||
# disagree → arbitrate via m_nOpType
|
||||
return _conflict_resolve(d, o, obj)
|
||||
|
||||
if d_valid and not o_valid:
|
||||
return direction
|
||||
if o_valid and not d_valid:
|
||||
return offset
|
||||
# neither valid — fall through
|
||||
except (TypeError, ValueError):
|
||||
pass
|
||||
|
||||
# 4. last resort: m_nOpType / order_type
|
||||
return _attr(obj, ["m_nOpType", "op_type", "order_type"])
|
||||
|
||||
|
||||
# Fields we care about when diagnosing a direction misread. Anything starting
|
||||
# with "m_" is captured automatically; these are the MiniQMT-style names that
|
||||
# do not match that prefix.
|
||||
_RAW_SNAPSHOT_EXTRA_FIELDS = (
|
||||
"stock_code",
|
||||
"order_type",
|
||||
"op_type",
|
||||
"direction",
|
||||
"offset_flag",
|
||||
"order_status",
|
||||
"order_volume",
|
||||
"traded_volume",
|
||||
"price",
|
||||
"order_id",
|
||||
"order_sysid",
|
||||
"order_sys_id",
|
||||
"trade_id",
|
||||
"traded_id",
|
||||
"strategy_name",
|
||||
"strategyName",
|
||||
"user_order_id",
|
||||
"order_remark",
|
||||
"remark",
|
||||
)
|
||||
|
||||
|
||||
def raw_field_snapshot(obj, max_repr=120):
|
||||
"""Capture every readable field of a live QMT callback object.
|
||||
|
||||
Direction extraction relies on understanding what ``m_nDirection``,
|
||||
``m_nOffsetFlag`` and ``m_nOpType`` carry in live callbacks. This dumps
|
||||
every readable field so one live order settles the question.
|
||||
|
||||
Returns ``{name: "<type> <value>"}``. Never raises: a callback that dies
|
||||
while being diagnosed would be worse than no diagnosis.
|
||||
"""
|
||||
snapshot = {}
|
||||
try:
|
||||
if isinstance(obj, dict):
|
||||
names = list(obj.keys())
|
||||
else:
|
||||
names = [name for name in dir(obj) if name.startswith("m_")]
|
||||
names.extend(_RAW_SNAPSHOT_EXTRA_FIELDS)
|
||||
except Exception:
|
||||
return {"__error__": "dir() failed"}
|
||||
seen = set()
|
||||
for name in names:
|
||||
key = str(name)
|
||||
if key in seen or key.startswith("__"):
|
||||
continue
|
||||
seen.add(key)
|
||||
try:
|
||||
if isinstance(obj, dict):
|
||||
if key not in obj:
|
||||
continue
|
||||
value = obj[key]
|
||||
else:
|
||||
if not hasattr(obj, key):
|
||||
continue
|
||||
value = getattr(obj, key)
|
||||
if callable(value):
|
||||
continue
|
||||
text = repr(value)
|
||||
if len(text) > max_repr:
|
||||
text = text[:max_repr] + "..."
|
||||
snapshot[key] = "%s %s" % (type(value).__name__, text)
|
||||
except Exception as exc: # noqa: BLE001 - diagnostics must not break callbacks
|
||||
snapshot[key] = "<unreadable: %s>" % exc.__class__.__name__
|
||||
return snapshot
|
||||
|
||||
|
||||
def format_raw_snapshot(kind, obj):
|
||||
"""One-line, GBK-safe rendering of :func:`raw_field_snapshot` for the QMT panel."""
|
||||
snapshot = raw_field_snapshot(obj)
|
||||
parts = ["%s=%s" % (name, snapshot[name]) for name in sorted(snapshot)]
|
||||
return "[bigqmt_exec_raw] %s type=%s %s" % (
|
||||
kind,
|
||||
type(obj).__name__,
|
||||
" | ".join(parts) or "<no fields>",
|
||||
)
|
||||
|
||||
|
||||
def normalize_order_event(order, account_id=""):
|
||||
"""Build a JSON-able order event dict from a Big QMT orderInfo object."""
|
||||
direction = _extract_direction(order)
|
||||
return {
|
||||
"event_type": EVENT_ORDER,
|
||||
"account_id": str(_attr(order, ["m_strAccountID", "account_id"], account_id) or account_id or ""),
|
||||
"stock_code": str(_attr(order, ["m_strInstrumentID", "stock_code", "m_strInstrument"], "") or ""),
|
||||
"order_sys_id": str(_attr(order, ["m_strOrderSysID", "order_sys_id", "order_sysid", "order_id"], "") or ""),
|
||||
"order_volume": _attr(order, ["m_nVolumeTotal", "order_volume", "volume"]),
|
||||
"traded_volume": _attr(order, ["m_nVolumeTraded", "traded_volume"]),
|
||||
"price": _attr(order, ["m_dLimitPrice", "price", "limit_price"]),
|
||||
"status": _attr(order, ["m_nOrderStatus", "order_status", "status"]),
|
||||
"direction": direction,
|
||||
"action": _action_from_direction(direction),
|
||||
"offset_flag": _attr(order, ["m_nOffsetFlag", "offset_flag"]),
|
||||
"strategy_name": str(_attr(order, ["strategyName", "m_strStrategyName", "strategy_name"], "") or ""),
|
||||
"remark": str(_attr(order, ["m_strRemark", "order_remark", "remark", "user_order_id"], "") or ""),
|
||||
"user_order_id": str(_attr(order, ["m_strRemark", "user_order_id", "order_remark", "remark"], "") or ""),
|
||||
"opt_name": str(_attr(order, ["m_strOptName", "opt_name"], "") or ""),
|
||||
"created_at": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"created_at_ts": time.time(),
|
||||
}
|
||||
|
||||
|
||||
def remember_order_identity(redis_client, account_id, user_order_id, strategy_name="", stock_code="", ttl_seconds=86400):
|
||||
user_order_id = str(user_order_id or "").strip()
|
||||
if not user_order_id or redis_client is None:
|
||||
return None
|
||||
payload = {
|
||||
"account_id": str(account_id or ""),
|
||||
"user_order_id": user_order_id,
|
||||
"strategy_name": str(strategy_name or ""),
|
||||
"stock_code": str(stock_code or ""),
|
||||
"created_at_ts": time.time(),
|
||||
}
|
||||
try:
|
||||
redis_client.setex(
|
||||
order_identity_key(account_id, user_order_id),
|
||||
int(ttl_seconds or 86400),
|
||||
json.dumps(payload, ensure_ascii=False, default=str),
|
||||
)
|
||||
except Exception:
|
||||
pass
|
||||
return payload
|
||||
|
||||
|
||||
def enrich_order_identity(redis_client, account_id, event):
|
||||
if redis_client is None or not isinstance(event, dict):
|
||||
return event
|
||||
user_order_id = str(event.get("user_order_id") or event.get("remark") or "").strip()
|
||||
if not user_order_id:
|
||||
return event
|
||||
try:
|
||||
raw = redis_client.get(order_identity_key(account_id, user_order_id))
|
||||
except Exception:
|
||||
raw = None
|
||||
if not raw:
|
||||
return event
|
||||
try:
|
||||
identity = json.loads(raw.decode("utf-8") if isinstance(raw, (bytes, bytearray)) else str(raw))
|
||||
except Exception:
|
||||
return event
|
||||
if not event.get("strategy_name") and identity.get("strategy_name"):
|
||||
event["strategy_name"] = str(identity.get("strategy_name") or "")
|
||||
if not event.get("stock_code") and identity.get("stock_code"):
|
||||
event["stock_code"] = str(identity.get("stock_code") or "")
|
||||
return event
|
||||
|
||||
|
||||
def normalize_trade_event(trade, account_id=""):
|
||||
"""Build a JSON-able trade (成交) event dict from a Big QMT dealInfo object."""
|
||||
direction = _extract_direction(trade)
|
||||
return {
|
||||
"event_type": EVENT_TRADE,
|
||||
"account_id": str(_attr(trade, ["m_strAccountID", "account_id"], account_id) or account_id or ""),
|
||||
"stock_code": str(_attr(trade, ["m_strInstrumentID", "stock_code"], "") or ""),
|
||||
"order_sys_id": str(_attr(trade, ["m_strOrderSysID", "order_sys_id", "order_sysid", "order_id"], "") or ""),
|
||||
"trade_id": str(_attr(trade, ["m_strTradeID", "trade_id"], "") or ""),
|
||||
"volume": _attr(trade, ["m_nVolume", "volume", "traded_volume"]),
|
||||
"price": _attr(trade, ["m_dPrice", "price", "traded_price"]),
|
||||
"amount": _attr(trade, ["m_dTradeAmount", "amount"]),
|
||||
"commission": _attr(trade, ["m_dComssion", "m_dCommission", "commission"]),
|
||||
"direction": direction,
|
||||
"action": _action_from_direction(direction),
|
||||
"offset_flag": _attr(trade, ["m_nOffsetFlag", "offset_flag"]),
|
||||
"traded_at": str(_attr(trade, ["m_strTradeTime", "traded_at", "trade_time"], "") or ""),
|
||||
"created_at": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"created_at_ts": time.time(),
|
||||
}
|
||||
|
||||
|
||||
def _publish(redis_client, channel, event, maxlen=2000):
|
||||
raw = json.dumps(event, ensure_ascii=False, default=str)
|
||||
try:
|
||||
redis_client.xadd(channel, {"payload": raw}, maxlen=maxlen, approximate=True)
|
||||
except Exception:
|
||||
pass
|
||||
redis_client.publish(channel, raw)
|
||||
return event
|
||||
|
||||
|
||||
def publish_order_event(redis_client, account_id, event):
|
||||
return _publish(redis_client, order_channel(account_id), event)
|
||||
|
||||
|
||||
def publish_trade_event(redis_client, account_id, event):
|
||||
return _publish(redis_client, trade_channel(account_id), event)
|
||||
|
||||
|
||||
def publish_order_error_event(redis_client, account_id, event):
|
||||
return _publish(redis_client, order_error_channel(account_id), event)
|
||||
|
||||
|
||||
def publish_cancel_error_event(redis_client, account_id, event):
|
||||
return _publish(redis_client, cancel_error_channel(account_id), event)
|
||||
|
||||
|
||||
def normalize_order_error_event(order_error, account_id=""):
|
||||
"""Build a JSON-able order-error event dict (废单/拒单).
|
||||
|
||||
QMT order callbacks carry the failed order via m_strOrderSysID / error info.
|
||||
MiniQMT's on_order_error receives an XtOrderError with error_id/error_msg.
|
||||
"""
|
||||
return {
|
||||
"event_type": EVENT_ORDER_ERROR,
|
||||
"account_id": str(_attr(order_error, ["m_strAccountID", "account_id"], account_id) or account_id or ""),
|
||||
"stock_code": str(_attr(order_error, ["m_strInstrumentID", "stock_code"], "") or ""),
|
||||
"order_sys_id": str(_attr(order_error, ["m_strOrderSysID", "order_sys_id", "order_sysid", "order_id"], "") or ""),
|
||||
"error_id": _attr(order_error, ["m_nErrorID", "error_id", "m_nOrderStatus"]),
|
||||
"error_msg": str(_attr(order_error, ["m_strErrorMsg", "error_msg", "m_strMsg"], "") or ""),
|
||||
"created_at": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"created_at_ts": time.time(),
|
||||
}
|
||||
|
||||
|
||||
def normalize_cancel_error_event(cancel_error, account_id=""):
|
||||
"""Build a JSON-able cancel-error event dict (撤单失败)."""
|
||||
return {
|
||||
"event_type": EVENT_CANCEL_ERROR,
|
||||
"account_id": str(_attr(cancel_error, ["m_strAccountID", "account_id"], account_id) or account_id or ""),
|
||||
"stock_code": str(_attr(cancel_error, ["m_strInstrumentID", "stock_code"], "") or ""),
|
||||
"order_sys_id": str(_attr(cancel_error, ["m_strOrderSysID", "order_sys_id", "order_sysid", "order_id"], "") or ""),
|
||||
"error_id": _attr(cancel_error, ["m_nErrorID", "error_id"]),
|
||||
"error_msg": str(_attr(cancel_error, ["m_strErrorMsg", "error_msg", "m_strMsg"], "") or ""),
|
||||
"created_at": time.strftime("%Y-%m-%d %H:%M:%S"),
|
||||
"created_at_ts": time.time(),
|
||||
}
|
||||
@@ -1,718 +0,0 @@
|
||||
"""Direct client for the Big QMT FormulaServer RPC (default port 58600).
|
||||
|
||||
Why this exists
|
||||
---------------
|
||||
The RPC bridge in :mod:`redis_rpc` routes every read through the QMT *strategy*
|
||||
process: client -> redis/zmq -> QMT python thread -> ContextInfo -> back. That
|
||||
costs ~13ms (redis) or ~0.7ms-with-500ms-GIL-spikes (zmq), and every read
|
||||
competes for the QMT main-thread GIL against the strategy itself.
|
||||
|
||||
FormulaServer is the C++ quote/reference-data service inside the same QMT
|
||||
terminal, listening on the port named in ``config/formulaserver/formulaserver.ini``
|
||||
(``[server_formula] address``, default 58600). QMT ships its own client for it at
|
||||
``bin.x64/Lib/site-packages/qmt_api``. Talking to it directly bypasses the
|
||||
strategy process entirely: measured p50 **0.07ms**, and zero GIL contention.
|
||||
|
||||
What it can and cannot do
|
||||
-------------------------
|
||||
FormulaServer serves market/reference data ONLY. Every account, position, order
|
||||
and trade method answers ``ErrorID 200005 未找到该服务``, as do ``getFullTick``
|
||||
and ``getQuote``. So this is a read fast-path, never a replacement for the RPC
|
||||
bridge — trading, account queries and 五档 snapshots stay on it.
|
||||
|
||||
Deliberately NOT routed here, despite FormulaServer exposing something similar:
|
||||
|
||||
* ``get_trading_dates`` — FormulaServer wants a *stock code* (``000001.SZ``);
|
||||
passing a market (``SH``) silently returns ``[]``. Our callers pass markets.
|
||||
* ``get_divid_factors`` / ``get_risk_free_rate`` — parameter semantics differ
|
||||
(range vs single date, index vs timetag). A wrong calendar or dividend factor
|
||||
is worse than a slow one.
|
||||
* Adjusted bars — see :func:`_market_data_params`; ``dividendType`` appears to
|
||||
be ignored by the server, so only unadjusted requests are routed.
|
||||
|
||||
Every failure here is non-fatal: :class:`FormulaServerRouter` reports the method
|
||||
as unroutable and the caller falls back to the normal RPC path.
|
||||
"""
|
||||
|
||||
import os
|
||||
import socket
|
||||
import struct
|
||||
import threading
|
||||
import time
|
||||
import zlib
|
||||
|
||||
|
||||
DEFAULT_HOST = "127.0.0.1"
|
||||
DEFAULT_PORT = 58600
|
||||
DEFAULT_TIMEOUT_SECONDS = 3.0
|
||||
# After a transport failure, stop trying for this long so a dead/absent
|
||||
# FormulaServer costs one timeout rather than one per call.
|
||||
DEFAULT_FAILURE_COOLDOWN_SECONDS = 30.0
|
||||
|
||||
NET_CMD_RPC = 3
|
||||
COMPRESS_ZLIB = 1
|
||||
COMPRESS_DOUBLE_ZLIB = 2
|
||||
|
||||
# FormulaServer's "method not found" code. Distinct from a transport failure:
|
||||
# it means the server is healthy and simply does not implement the call.
|
||||
ERROR_METHOD_NOT_FOUND = 200005
|
||||
|
||||
|
||||
class FormulaServerError(RuntimeError):
|
||||
"""FormulaServer answered with a non-zero status (bad params, no such method)."""
|
||||
|
||||
def __init__(self, message, error_id=None):
|
||||
RuntimeError.__init__(self, message)
|
||||
self.error_id = error_id
|
||||
|
||||
|
||||
class FormulaServerUnavailable(RuntimeError):
|
||||
"""The FormulaServer could not be reached (connect/IO/protocol failure)."""
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# BSON codec
|
||||
# ---------------------------------------------------------------------------
|
||||
# FormulaServer frames BSON documents. pymongo's ``bson`` is used when present
|
||||
# (faster, battle-tested); otherwise the minimal codec below covers the types
|
||||
# this wire actually carries. Keeping a built-in path means an external client
|
||||
# needs no pymongo just to read market data.
|
||||
|
||||
def _load_bson():
|
||||
for module_name in ("bson", "xtquant.xtbson.bson36"):
|
||||
try:
|
||||
module = __import__(module_name, fromlist=["BSON"])
|
||||
except Exception:
|
||||
continue
|
||||
if hasattr(module, "BSON"):
|
||||
return module
|
||||
return None
|
||||
|
||||
|
||||
_BSON = _load_bson()
|
||||
|
||||
|
||||
def _encode_document(pairs):
|
||||
body = b"".join(_encode_element(str(key), value) for key, value in pairs)
|
||||
return struct.pack("<i", len(body) + 5) + body + b"\x00"
|
||||
|
||||
|
||||
def _encode_element(name, value):
|
||||
key = name.encode("utf-8") + b"\x00"
|
||||
if value is None:
|
||||
return b"\x0a" + key
|
||||
# bool before int: bool is an int subclass.
|
||||
if isinstance(value, bool):
|
||||
return b"\x08" + key + (b"\x01" if value else b"\x00")
|
||||
if isinstance(value, int):
|
||||
if -2147483648 <= value <= 2147483647:
|
||||
return b"\x10" + key + struct.pack("<i", value)
|
||||
return b"\x12" + key + struct.pack("<q", value)
|
||||
if isinstance(value, float):
|
||||
return b"\x01" + key + struct.pack("<d", value)
|
||||
if isinstance(value, bytes):
|
||||
return b"\x05" + key + struct.pack("<i", len(value)) + b"\x00" + value
|
||||
if isinstance(value, str):
|
||||
raw = value.encode("utf-8") + b"\x00"
|
||||
return b"\x02" + key + struct.pack("<i", len(raw)) + raw
|
||||
if isinstance(value, (list, tuple)):
|
||||
return b"\x04" + key + _encode_document(
|
||||
(str(index), item) for index, item in enumerate(value)
|
||||
)
|
||||
if isinstance(value, dict):
|
||||
return b"\x03" + key + _encode_document(value.items())
|
||||
raise TypeError("cannot BSON-encode %s" % type(value).__name__)
|
||||
|
||||
|
||||
def _decode_document(data, pos):
|
||||
size = struct.unpack_from("<i", data, pos)[0]
|
||||
end = pos + size
|
||||
pos += 4
|
||||
out = {}
|
||||
while pos < end - 1:
|
||||
type_byte = data[pos] if isinstance(data[pos], int) else ord(data[pos])
|
||||
pos += 1
|
||||
terminator = data.index(b"\x00", pos)
|
||||
name = data[pos:terminator].decode("utf-8", "replace")
|
||||
pos = terminator + 1
|
||||
out[name], pos = _decode_element(type_byte, data, pos)
|
||||
return out, end
|
||||
|
||||
|
||||
def _decode_element(type_byte, data, pos):
|
||||
if type_byte == 0x01:
|
||||
return struct.unpack_from("<d", data, pos)[0], pos + 8
|
||||
if type_byte == 0x02:
|
||||
length = struct.unpack_from("<i", data, pos)[0]
|
||||
pos += 4
|
||||
return data[pos:pos + length - 1].decode("utf-8", "replace"), pos + length
|
||||
if type_byte == 0x03:
|
||||
return _decode_document(data, pos)
|
||||
if type_byte == 0x04:
|
||||
doc, end = _decode_document(data, pos)
|
||||
try:
|
||||
ordered = sorted(doc, key=lambda key: int(key))
|
||||
except (TypeError, ValueError):
|
||||
ordered = sorted(doc)
|
||||
return [doc[key] for key in ordered], end
|
||||
if type_byte == 0x05:
|
||||
length = struct.unpack_from("<i", data, pos)[0]
|
||||
pos += 5 # int32 length + 1 subtype byte
|
||||
return data[pos:pos + length], pos + length
|
||||
if type_byte == 0x08:
|
||||
flag = data[pos] if isinstance(data[pos], int) else ord(data[pos])
|
||||
return bool(flag), pos + 1
|
||||
if type_byte in (0x09, 0x12):
|
||||
return struct.unpack_from("<q", data, pos)[0], pos + 8
|
||||
if type_byte == 0x0A:
|
||||
return None, pos
|
||||
if type_byte == 0x10:
|
||||
return struct.unpack_from("<i", data, pos)[0], pos + 4
|
||||
if type_byte == 0x11:
|
||||
return struct.unpack_from("<Q", data, pos)[0], pos + 8
|
||||
raise ValueError("unsupported BSON type byte 0x%02x" % type_byte)
|
||||
|
||||
|
||||
def bson_encode(document):
|
||||
if _BSON is not None:
|
||||
return _BSON.BSON.encode(document)
|
||||
return _encode_document(document.items())
|
||||
|
||||
|
||||
def bson_decode(payload):
|
||||
if _BSON is not None:
|
||||
return _BSON.BSON(payload).decode()
|
||||
return _decode_document(payload, 0)[0]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Address discovery
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
def read_formulaserver_port(qmt_root):
|
||||
"""Read ``[server_formula] address`` from a QMT install's formulaserver.ini.
|
||||
|
||||
``qmt_root`` is the terminal directory (the one holding ``bin.x64`` and
|
||||
``config``). Returns the port int, or None when the file is absent or
|
||||
unparsable — callers then fall back to :data:`DEFAULT_PORT`.
|
||||
"""
|
||||
if not qmt_root:
|
||||
return None
|
||||
path = os.path.join(str(qmt_root), "config", "formulaserver", "formulaserver.ini")
|
||||
try:
|
||||
try:
|
||||
import configparser
|
||||
except ImportError: # pragma: no cover - py2 safety net
|
||||
import ConfigParser as configparser
|
||||
parser = configparser.ConfigParser()
|
||||
if not parser.read(path):
|
||||
return None
|
||||
address = parser.get("server_formula", "address")
|
||||
except Exception:
|
||||
return None
|
||||
if ":" not in str(address):
|
||||
return None
|
||||
try:
|
||||
return int(str(address).rsplit(":", 1)[1])
|
||||
except (TypeError, ValueError):
|
||||
return None
|
||||
|
||||
|
||||
def resolve_address(config=None):
|
||||
"""Resolve (host, port) for the FormulaServer.
|
||||
|
||||
Priority: explicit ``host``/``port`` > ``formulaserver.ini`` under
|
||||
``qmt_root`` > ``BIGQMT_FORMULA_HOST``/``BIGQMT_FORMULA_PORT`` > defaults.
|
||||
The address binds ``0.0.0.0`` in QMT's shipped config, so a remote client
|
||||
can reach it too when the firewall allows.
|
||||
"""
|
||||
config = dict(config or {})
|
||||
host = str(config.get("host") or os.environ.get("BIGQMT_FORMULA_HOST") or DEFAULT_HOST)
|
||||
port = config.get("port")
|
||||
if not port:
|
||||
port = read_formulaserver_port(config.get("qmt_root"))
|
||||
if not port:
|
||||
port = os.environ.get("BIGQMT_FORMULA_PORT")
|
||||
try:
|
||||
port = int(port)
|
||||
except (TypeError, ValueError):
|
||||
port = DEFAULT_PORT
|
||||
return host, port
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Client
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
class FormulaServerClient(object):
|
||||
"""Thread-safe BSON-over-TCP client for FormulaServer.
|
||||
|
||||
One socket is shared under a lock. FormulaServer matches responses by
|
||||
sequence number, so concurrent use of a single socket would require
|
||||
demultiplexing; serializing is simpler and, at 0.07ms per call, ample.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
host=DEFAULT_HOST,
|
||||
port=DEFAULT_PORT,
|
||||
timeout_seconds=DEFAULT_TIMEOUT_SECONDS,
|
||||
print_prefix="[bigqmt_formula]",
|
||||
):
|
||||
self.host = str(host or DEFAULT_HOST)
|
||||
self.port = int(port or DEFAULT_PORT)
|
||||
self.timeout_seconds = float(timeout_seconds or DEFAULT_TIMEOUT_SECONDS)
|
||||
self.print_prefix = print_prefix
|
||||
self._lock = threading.Lock()
|
||||
self._socket = None
|
||||
self._seq = 0
|
||||
|
||||
# -- wire ------------------------------------------------------------
|
||||
def _connect_locked(self):
|
||||
if self._socket is not None:
|
||||
return self._socket
|
||||
try:
|
||||
sock = socket.create_connection((self.host, self.port), self.timeout_seconds)
|
||||
sock.settimeout(self.timeout_seconds)
|
||||
except Exception as exc:
|
||||
raise FormulaServerUnavailable(
|
||||
"connect %s:%s failed: %s" % (self.host, self.port, exc)
|
||||
)
|
||||
self._socket = sock
|
||||
return sock
|
||||
|
||||
def _close_locked(self):
|
||||
if self._socket is not None:
|
||||
try:
|
||||
self._socket.close()
|
||||
except Exception:
|
||||
pass
|
||||
self._socket = None
|
||||
|
||||
def close(self):
|
||||
with self._lock:
|
||||
self._close_locked()
|
||||
|
||||
def _recv_exactly(self, sock, length):
|
||||
chunks = []
|
||||
remaining = length
|
||||
while remaining > 0:
|
||||
more = sock.recv(remaining)
|
||||
if not more:
|
||||
raise FormulaServerUnavailable("socket closed mid-message")
|
||||
chunks.append(more)
|
||||
remaining -= len(more)
|
||||
return b"".join(chunks)
|
||||
|
||||
def _request_locked(self, func, params):
|
||||
sock = self._connect_locked()
|
||||
self._seq += 1
|
||||
seq = self._seq
|
||||
body = bson_encode({"func": str(func), "params": dict(params or {})})
|
||||
tag = ((seq >> 32) & 0x0F) << 8
|
||||
packet = struct.pack(
|
||||
"!IIHH%ds" % len(body),
|
||||
len(body) + 12,
|
||||
seq & 0xFFFFFFFF,
|
||||
NET_CMD_RPC,
|
||||
tag,
|
||||
body,
|
||||
)
|
||||
try:
|
||||
sock.sendall(packet)
|
||||
except Exception as exc:
|
||||
raise FormulaServerUnavailable("send failed: %s" % exc)
|
||||
# Subscription pushes share the socket; skip anything that is not our seq.
|
||||
while True:
|
||||
try:
|
||||
header = self._recv_exactly(sock, 4)
|
||||
pack_len = struct.unpack_from("!I", header, 0)[0]
|
||||
rest = self._recv_exactly(sock, pack_len - 4)
|
||||
except FormulaServerUnavailable:
|
||||
raise
|
||||
except Exception as exc:
|
||||
raise FormulaServerUnavailable("recv failed: %s" % exc)
|
||||
raw = header + rest
|
||||
try:
|
||||
got_seq, _cmd, got_tag, payload_bytes = struct.unpack_from(
|
||||
"!IHH%ds" % (pack_len - 12), raw, 4
|
||||
)
|
||||
if (got_tag & 7) in (COMPRESS_ZLIB, COMPRESS_DOUBLE_ZLIB):
|
||||
payload_bytes = zlib.decompress(payload_bytes)
|
||||
got_seq = ((got_tag >> 8) & 0x0F) << 32 | got_seq
|
||||
payload = bson_decode(payload_bytes)
|
||||
except Exception as exc:
|
||||
raise FormulaServerUnavailable("decode failed: %s" % exc)
|
||||
if got_seq != seq:
|
||||
continue
|
||||
if payload.get("status") == 0:
|
||||
return payload.get("params")
|
||||
detail = payload.get("params")
|
||||
error_id = None
|
||||
if isinstance(detail, dict):
|
||||
error_id = detail.get("ErrorID")
|
||||
raise FormulaServerError(
|
||||
"%s failed: %r" % (func, detail), error_id=error_id
|
||||
)
|
||||
|
||||
def request(self, func, params=None):
|
||||
"""Call ``func`` and return its ``params`` payload.
|
||||
|
||||
Retries once on a transport failure, since QMT restarts (or an idle
|
||||
socket reaped by the server) show up as a dead socket on first use.
|
||||
"""
|
||||
with self._lock:
|
||||
try:
|
||||
return self._request_locked(func, params)
|
||||
except FormulaServerError:
|
||||
raise
|
||||
except FormulaServerUnavailable:
|
||||
self._close_locked()
|
||||
return self._request_locked(func, params)
|
||||
|
||||
def ping(self):
|
||||
"""Cheap liveness probe. True when FormulaServer answers at all.
|
||||
|
||||
A ``FormulaServerError`` still counts as alive — the server replied.
|
||||
"""
|
||||
try:
|
||||
self.request("getLastVolume", {"stockCode": "000001.SZ"})
|
||||
return True
|
||||
except FormulaServerError:
|
||||
return True
|
||||
except FormulaServerUnavailable:
|
||||
return False
|
||||
|
||||
def __repr__(self):
|
||||
return "<FormulaServerClient %s:%s>" % (self.host, self.port)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Method mapping
|
||||
# ---------------------------------------------------------------------------
|
||||
# FormulaServer misspells two instrument fields relative to the xtdata SDK
|
||||
# (``FloatVolume``/``TotalVolume``). Downstream code reads the SDK spelling, so
|
||||
# alias them rather than let the lookup silently miss.
|
||||
_INSTRUMENT_ALIASES = (
|
||||
("FloatVolumn", "FloatVolume"),
|
||||
("TotalVolumn", "TotalVolume"),
|
||||
)
|
||||
|
||||
|
||||
def _first(params, names, default=None):
|
||||
for name in names:
|
||||
if name in params and params[name] is not None:
|
||||
return params[name]
|
||||
return default
|
||||
|
||||
|
||||
def _as_list(value):
|
||||
if value is None:
|
||||
return []
|
||||
if isinstance(value, str):
|
||||
return [value]
|
||||
return list(value)
|
||||
|
||||
|
||||
def _require_code(params, names):
|
||||
code = _first(params, names)
|
||||
text = str(code or "").strip()
|
||||
if not text:
|
||||
raise ValueError("a stock code is required (one of %s)" % ", ".join(names))
|
||||
return text
|
||||
|
||||
|
||||
def _instrument_params(params):
|
||||
return {"strOptionCode": _require_code(params, ("code", "stock_code", "stockcode"))}
|
||||
|
||||
|
||||
def _instrument_result(raw, params):
|
||||
detail = (raw or {}).get("result")
|
||||
if not isinstance(detail, dict):
|
||||
return detail or {}
|
||||
out = dict(detail)
|
||||
for wire_name, sdk_name in _INSTRUMENT_ALIASES:
|
||||
if wire_name in out and sdk_name not in out:
|
||||
out[sdk_name] = out[wire_name]
|
||||
return out
|
||||
|
||||
|
||||
def _scalar_result(raw, params):
|
||||
return (raw or {}).get("result")
|
||||
|
||||
|
||||
def _list_result(raw, params):
|
||||
return (raw or {}).get("result") or []
|
||||
|
||||
|
||||
def _last_volume_params(params):
|
||||
return {"stockCode": _require_code(params, ("stock", "code", "stock_code", "stockcode"))}
|
||||
|
||||
|
||||
def _total_share_params(params):
|
||||
return {"stockCode": _require_code(params, ("stockcode", "code", "stock_code", "stock"))}
|
||||
|
||||
|
||||
def _contract_multiplier_params(params):
|
||||
return {"contractCode": _require_code(params, ("stockcode", "code", "stock_code", "contract_code"))}
|
||||
|
||||
|
||||
def _main_contract_params(params):
|
||||
return {"codeMarket": _require_code(params, ("code_market", "codeMarket", "code"))}
|
||||
|
||||
|
||||
def _sector_params(params):
|
||||
name = str(_first(params, ("sector_name", "sectorName", "sector"), "") or "").strip()
|
||||
if not name:
|
||||
raise ValueError("sector_name is required")
|
||||
# ContextInfo's real_timetag defaults to -1; FormulaServer's realtime
|
||||
# defaults to 0. Both return identical constituents (verified), so normalize
|
||||
# the sentinel rather than forward a value the server never documents.
|
||||
realtime = _first(params, ("real_timetag", "realtime"), 0)
|
||||
try:
|
||||
realtime = int(realtime)
|
||||
except (TypeError, ValueError):
|
||||
realtime = 0
|
||||
if realtime < 0:
|
||||
realtime = 0
|
||||
return {"sectorName": name, "realtime": realtime}
|
||||
|
||||
|
||||
def _weight_in_index_params(params):
|
||||
index_code = _first(params, ("mtkindexcode", "index_code", "indexCode"))
|
||||
stock_code = _first(params, ("stockcode", "stock_code", "code"))
|
||||
if not index_code or not stock_code:
|
||||
raise ValueError("mtkindexcode and stockcode are required")
|
||||
return {"indexCode": str(index_code), "stockCode": str(stock_code)}
|
||||
|
||||
|
||||
def _market_data_params(params):
|
||||
fields = _as_list(_first(params, ("field_list", "fields"), None))
|
||||
codes = _as_list(_first(params, ("stock_list", "stock_code", "stockCodes"), None))
|
||||
if not fields or not codes:
|
||||
raise ValueError("field_list and stock_list are required")
|
||||
dividend_type = str(params.get("dividend_type") or "none").lower()
|
||||
# FormulaServer returns byte-identical bars for dividendType none/front, so
|
||||
# adjustment is not applied here. Serving an adjusted request from this path
|
||||
# would silently hand back unadjusted prices — refuse and let RPC answer.
|
||||
if dividend_type not in ("", "none"):
|
||||
raise ValueError("adjusted bars (dividend_type=%s) are not served here" % dividend_type)
|
||||
period = str(params.get("period") or "1d")
|
||||
count = params.get("count", -1)
|
||||
try:
|
||||
count = int(count)
|
||||
except (TypeError, ValueError):
|
||||
count = -1
|
||||
return {
|
||||
"fields": [str(field) for field in fields],
|
||||
"stockCodes": [str(code) for code in codes],
|
||||
"startTime": str(params.get("start_time") or ""),
|
||||
"endTime": str(params.get("end_time") or ""),
|
||||
"period": period,
|
||||
"dividendType": "none",
|
||||
"count": count,
|
||||
}
|
||||
|
||||
|
||||
def _market_data_result(raw, params):
|
||||
"""Translate FormulaServer's flat bar list into the RPC path's payload.
|
||||
|
||||
Wire shape is ``[code, [time, [field, value, ...], time, [...]], code, ...]``.
|
||||
We emit the same ``__bigqmt_type__: DataFrame`` envelope the QMT-side adapter
|
||||
builds, so the client's ``_restore_jsonable`` rebuilds identical DataFrames
|
||||
whichever path answered.
|
||||
"""
|
||||
flat = (raw or {}).get("result") or []
|
||||
fields = [str(field) for field in (_first(params, ("field_list", "fields"), None) or [])]
|
||||
columns = list(fields)
|
||||
if columns and "stime" not in columns:
|
||||
columns.insert(0, "stime")
|
||||
|
||||
parsed = {}
|
||||
for index in range(0, len(flat) - 1, 2):
|
||||
code = str(flat[index])
|
||||
timeline = flat[index + 1] or []
|
||||
records = []
|
||||
for offset in range(0, len(timeline) - 1, 2):
|
||||
stamp = timeline[offset]
|
||||
pairs = timeline[offset + 1] or []
|
||||
record = {"stime": stamp}
|
||||
for cursor in range(0, len(pairs) - 1, 2):
|
||||
record[str(pairs[cursor])] = pairs[cursor + 1]
|
||||
records.append(record)
|
||||
parsed[code] = records
|
||||
|
||||
requested = [str(code) for code in _as_list(_first(params, ("stock_list", "stock_code"), None))]
|
||||
for code in parsed:
|
||||
if code not in requested:
|
||||
requested.append(code)
|
||||
return {
|
||||
code: {
|
||||
"__bigqmt_type__": "DataFrame",
|
||||
"columns": columns,
|
||||
"records": parsed.get(code) or [],
|
||||
}
|
||||
for code in requested
|
||||
}
|
||||
|
||||
|
||||
# our RPC method -> (FormulaServer func, param builder, result adapter)
|
||||
METHOD_MAP = {
|
||||
"get_instrument": ("getInstrumentDetail", _instrument_params, _instrument_result),
|
||||
"get_instrumentdetail": ("getInstrumentDetail", _instrument_params, _instrument_result),
|
||||
"get_instrument_detail": ("getInstrumentDetail", _instrument_params, _instrument_result),
|
||||
"get_last_volume": ("getLastVolume", _last_volume_params, _scalar_result),
|
||||
"get_total_share": ("getTotalShare", _total_share_params, _scalar_result),
|
||||
"get_contract_multiplier": ("getContractMultiplier", _contract_multiplier_params, _scalar_result),
|
||||
"get_main_contract": ("getMainContract", _main_contract_params, _scalar_result),
|
||||
"get_weight_in_index": ("getWeightInIndex", _weight_in_index_params, _scalar_result),
|
||||
"get_stock_list_in_sector": ("getStockListInSector", _sector_params, _list_result),
|
||||
"get_market_data_ex": ("getMarketData", _market_data_params, _market_data_result),
|
||||
}
|
||||
|
||||
SUPPORTED_METHODS = tuple(sorted(METHOD_MAP))
|
||||
|
||||
|
||||
class Unroutable(Exception):
|
||||
"""This call cannot be served by FormulaServer — use the RPC bridge."""
|
||||
|
||||
|
||||
class FormulaServerRouter(object):
|
||||
"""Routes supported read methods to FormulaServer, or declines.
|
||||
|
||||
:meth:`call` raises :class:`Unroutable` for anything it cannot serve —
|
||||
method not mapped, params that do not translate, server down, feature
|
||||
disabled. Callers treat that as "fall back to RPC".
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
client=None,
|
||||
enabled=True,
|
||||
methods=None,
|
||||
failure_cooldown_seconds=DEFAULT_FAILURE_COOLDOWN_SECONDS,
|
||||
print_prefix="[bigqmt_formula]",
|
||||
config=None,
|
||||
):
|
||||
self.enabled = bool(enabled)
|
||||
self.print_prefix = print_prefix
|
||||
self.failure_cooldown_seconds = float(
|
||||
failure_cooldown_seconds or DEFAULT_FAILURE_COOLDOWN_SECONDS
|
||||
)
|
||||
if client is None and self.enabled:
|
||||
host, port = resolve_address(config)
|
||||
timeout = float((config or {}).get("timeout_seconds") or DEFAULT_TIMEOUT_SECONDS)
|
||||
client = FormulaServerClient(
|
||||
host=host, port=port, timeout_seconds=timeout, print_prefix=print_prefix
|
||||
)
|
||||
self.client = client
|
||||
if methods:
|
||||
self.methods = set(str(name) for name in methods) & set(METHOD_MAP)
|
||||
else:
|
||||
self.methods = set(METHOD_MAP)
|
||||
self._unavailable_until = 0.0
|
||||
self._announced = False
|
||||
# Methods the server itself rejected as unimplemented — never retried.
|
||||
self._unimplemented = set()
|
||||
self.hits = 0
|
||||
self.misses = 0
|
||||
|
||||
def _available(self):
|
||||
if not self.enabled or self.client is None:
|
||||
return False
|
||||
return time.time() >= self._unavailable_until
|
||||
|
||||
def _mark_unavailable(self, reason):
|
||||
self._unavailable_until = time.time() + self.failure_cooldown_seconds
|
||||
print(
|
||||
"%s unavailable, falling back to RPC for %.0fs: %s"
|
||||
% (self.print_prefix, self.failure_cooldown_seconds, reason)
|
||||
)
|
||||
|
||||
def supports(self, method):
|
||||
return (
|
||||
str(method) in self.methods
|
||||
and str(method) not in self._unimplemented
|
||||
and self._available()
|
||||
)
|
||||
|
||||
def call(self, method, params=None):
|
||||
"""Serve ``method`` from FormulaServer, or raise :class:`Unroutable`."""
|
||||
method = str(method)
|
||||
if not self.supports(method):
|
||||
raise Unroutable(method)
|
||||
func, build_params, adapt_result = METHOD_MAP[method]
|
||||
try:
|
||||
wire_params = build_params(dict(params or {}))
|
||||
except Exception as exc:
|
||||
# Params that do not translate are a per-call condition, not a
|
||||
# server fault — do not trip the breaker.
|
||||
self.misses += 1
|
||||
raise Unroutable("%s: %s" % (method, exc))
|
||||
try:
|
||||
raw = self.client.request(func, wire_params)
|
||||
except FormulaServerError as exc:
|
||||
self.misses += 1
|
||||
if exc.error_id == ERROR_METHOD_NOT_FOUND:
|
||||
self._unimplemented.add(method)
|
||||
print(
|
||||
"%s %s not implemented by this terminal, using RPC from now on"
|
||||
% (self.print_prefix, method)
|
||||
)
|
||||
raise Unroutable("%s: %s" % (method, exc))
|
||||
except FormulaServerUnavailable as exc:
|
||||
self.misses += 1
|
||||
self._mark_unavailable(str(exc))
|
||||
raise Unroutable("%s: %s" % (method, exc))
|
||||
except Exception as exc:
|
||||
self.misses += 1
|
||||
self._mark_unavailable("%s: %s" % (exc.__class__.__name__, exc))
|
||||
raise Unroutable("%s: %s" % (method, exc))
|
||||
try:
|
||||
result = adapt_result(raw, dict(params or {}))
|
||||
except Exception as exc:
|
||||
self.misses += 1
|
||||
raise Unroutable("%s: result adaptation failed: %s" % (method, exc))
|
||||
self.hits += 1
|
||||
if not self._announced:
|
||||
self._announced = True
|
||||
print(
|
||||
"%s active at %s:%s (%d methods routed direct)"
|
||||
% (self.print_prefix, self.client.host, self.client.port, len(self.methods))
|
||||
)
|
||||
return result
|
||||
|
||||
def stats(self):
|
||||
return {
|
||||
"enabled": self.enabled,
|
||||
"hits": self.hits,
|
||||
"misses": self.misses,
|
||||
"available": self._available(),
|
||||
"unimplemented": sorted(self._unimplemented),
|
||||
"methods": sorted(self.methods),
|
||||
}
|
||||
|
||||
def close(self):
|
||||
if self.client is not None:
|
||||
self.client.close()
|
||||
|
||||
|
||||
def build_router(config=None, print_prefix="[bigqmt_formula]"):
|
||||
"""Build a router from a ``formula_server`` config dict.
|
||||
|
||||
Recognised keys: ``enabled`` (default True), ``host``, ``port``,
|
||||
``qmt_root``, ``timeout_seconds``, ``methods``, ``failure_cooldown_seconds``.
|
||||
``enabled=False`` yields a router that declines everything, so callers need
|
||||
no None checks.
|
||||
"""
|
||||
config = dict(config or {})
|
||||
enabled = config.get("enabled", True)
|
||||
if isinstance(enabled, str):
|
||||
enabled = enabled.strip().lower() not in ("0", "false", "no", "off")
|
||||
return FormulaServerRouter(
|
||||
enabled=bool(enabled),
|
||||
methods=config.get("methods"),
|
||||
failure_cooldown_seconds=config.get("failure_cooldown_seconds"),
|
||||
print_prefix=print_prefix,
|
||||
config=config,
|
||||
)
|
||||
@@ -1,219 +0,0 @@
|
||||
"""Demand-driven Redis cache for Big QMT full tick snapshots."""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import pickle
|
||||
import time
|
||||
|
||||
from .code_utils import normalize_stock_code
|
||||
|
||||
|
||||
MARKET_CODES = {"SH", "SZ", "BJ", "HK"}
|
||||
DEMAND_KEY_TEMPLATE = "bigqmt:full_tick:demand:{account_id}"
|
||||
CACHE_KEY_TEMPLATE = "bigqmt:full_tick:cache:{account_id}:{request_id}"
|
||||
|
||||
|
||||
def _decode_text(value):
|
||||
if isinstance(value, bytes):
|
||||
return value.decode("utf-8")
|
||||
return str(value)
|
||||
|
||||
|
||||
def _loads_json(value):
|
||||
return json.loads(_decode_text(value))
|
||||
|
||||
|
||||
def normalize_full_tick_codes(codes):
|
||||
normalized = []
|
||||
seen = set()
|
||||
for code in codes or []:
|
||||
text = str(code or "").strip().upper()
|
||||
if not text:
|
||||
continue
|
||||
if text in MARKET_CODES:
|
||||
item = text
|
||||
else:
|
||||
item = normalize_stock_code(text)
|
||||
if item not in seen:
|
||||
seen.add(item)
|
||||
normalized.append(item)
|
||||
return sorted(normalized)
|
||||
|
||||
|
||||
def full_tick_request_id(codes):
|
||||
normalized = normalize_full_tick_codes(codes)
|
||||
digest = hashlib.sha1("|".join(normalized).encode("utf-8")).hexdigest()
|
||||
return digest[:20]
|
||||
|
||||
|
||||
def full_tick_demand_key(account_id):
|
||||
return DEMAND_KEY_TEMPLATE.format(account_id=str(account_id or ""))
|
||||
|
||||
|
||||
def full_tick_cache_key(account_id, codes=None, request_id=None):
|
||||
rid = str(request_id or full_tick_request_id(codes or []))
|
||||
return CACHE_KEY_TEMPLATE.format(account_id=str(account_id or ""), request_id=rid)
|
||||
|
||||
|
||||
def _dump_snapshot(payload):
|
||||
return pickle.dumps(payload, protocol=4)
|
||||
|
||||
|
||||
def _load_snapshot(raw):
|
||||
if not raw:
|
||||
return None
|
||||
try:
|
||||
return pickle.loads(raw)
|
||||
except Exception:
|
||||
try:
|
||||
return _loads_json(raw)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def request_full_tick_cache(redis_client, account_id, codes, demand_ttl_seconds=10, cache_ttl_seconds=10):
|
||||
normalized = normalize_full_tick_codes(codes)
|
||||
if not normalized:
|
||||
raise ValueError("full tick codes are required")
|
||||
now = time.time()
|
||||
request_id = full_tick_request_id(normalized)
|
||||
payload = {
|
||||
"request_id": request_id,
|
||||
"codes": normalized,
|
||||
"requested_at_ts": now,
|
||||
"expires_at_ts": now + float(demand_ttl_seconds),
|
||||
"cache_ttl_seconds": float(cache_ttl_seconds),
|
||||
}
|
||||
key = full_tick_demand_key(account_id)
|
||||
redis_client.hset(key, request_id, json.dumps(payload, ensure_ascii=False, sort_keys=True))
|
||||
try:
|
||||
redis_client.expire(key, max(30, int(float(demand_ttl_seconds) * 3)))
|
||||
except Exception:
|
||||
pass
|
||||
return payload
|
||||
|
||||
|
||||
def write_full_tick_cache(redis_client, account_id, codes, data, cache_ttl_seconds=10):
|
||||
normalized = normalize_full_tick_codes(codes)
|
||||
request_id = full_tick_request_id(normalized)
|
||||
now = time.time()
|
||||
payload = {
|
||||
"request_id": request_id,
|
||||
"codes": normalized,
|
||||
"updated_at_ts": now,
|
||||
"updated_at": time.strftime("%Y-%m-%d %H:%M:%S", time.localtime(now)),
|
||||
"data": data or {},
|
||||
}
|
||||
key = full_tick_cache_key(account_id, request_id=request_id)
|
||||
ttl = int(max(1, float(cache_ttl_seconds)))
|
||||
redis_client.setex(key, ttl, _dump_snapshot(payload))
|
||||
return payload
|
||||
|
||||
|
||||
def read_full_tick_cache(redis_client, account_id, codes, max_age_seconds=10):
|
||||
normalized = normalize_full_tick_codes(codes)
|
||||
key = full_tick_cache_key(account_id, codes=normalized)
|
||||
snapshot = _load_snapshot(redis_client.get(key))
|
||||
if not isinstance(snapshot, dict):
|
||||
return None
|
||||
if normalize_full_tick_codes(snapshot.get("codes") or []) != normalized:
|
||||
return None
|
||||
updated_at = float(snapshot.get("updated_at_ts") or 0)
|
||||
if updated_at <= 0:
|
||||
return None
|
||||
if time.time() - updated_at > float(max_age_seconds):
|
||||
return None
|
||||
data = snapshot.get("data")
|
||||
return data if isinstance(data, dict) else None
|
||||
|
||||
|
||||
def wait_full_tick_cache(redis_client, account_id, codes, max_age_seconds=10, wait_seconds=3.5, poll_interval_seconds=0.2):
|
||||
deadline = time.time() + max(0.0, float(wait_seconds))
|
||||
while True:
|
||||
data = read_full_tick_cache(redis_client, account_id, codes, max_age_seconds=max_age_seconds)
|
||||
if data is not None:
|
||||
return data
|
||||
if time.time() >= deadline:
|
||||
return None
|
||||
time.sleep(max(0.05, float(poll_interval_seconds)))
|
||||
|
||||
|
||||
def iter_active_full_tick_demands(redis_client, account_id, demand_ttl_seconds=10, max_requests=8):
|
||||
key = full_tick_demand_key(account_id)
|
||||
raw_mapping = redis_client.hgetall(key) or {}
|
||||
now = time.time()
|
||||
active = []
|
||||
for field, raw_payload in list(raw_mapping.items()):
|
||||
field_text = _decode_text(field)
|
||||
try:
|
||||
payload = _loads_json(raw_payload)
|
||||
except Exception:
|
||||
try:
|
||||
redis_client.hdel(key, field_text)
|
||||
except Exception:
|
||||
pass
|
||||
continue
|
||||
expires_at = float(payload.get("expires_at_ts") or 0)
|
||||
if expires_at <= now:
|
||||
try:
|
||||
redis_client.hdel(key, field_text)
|
||||
except Exception:
|
||||
pass
|
||||
continue
|
||||
codes = normalize_full_tick_codes(payload.get("codes") or [])
|
||||
if not codes:
|
||||
try:
|
||||
redis_client.hdel(key, field_text)
|
||||
except Exception:
|
||||
pass
|
||||
continue
|
||||
payload["codes"] = codes
|
||||
payload["cache_ttl_seconds"] = float(payload.get("cache_ttl_seconds") or demand_ttl_seconds)
|
||||
active.append(payload)
|
||||
active.sort(key=lambda item: float(item.get("requested_at_ts") or 0), reverse=True)
|
||||
return active[: int(max_requests)]
|
||||
|
||||
|
||||
def _demand_is_market(codes):
|
||||
return any(str(code).strip().upper() in MARKET_CODES for code in codes or [])
|
||||
|
||||
|
||||
def refresh_full_tick_cache(
|
||||
redis_client,
|
||||
context_info,
|
||||
account_id,
|
||||
demand_ttl_seconds=10,
|
||||
cache_ttl_seconds=10,
|
||||
max_requests=8,
|
||||
kind=None,
|
||||
max_wall_seconds=None,
|
||||
):
|
||||
"""Refresh cached snapshots for active demands.
|
||||
|
||||
``kind`` selects which demands to refresh: ``None`` (all), ``"symbol"``
|
||||
(only symbol-list demands), or ``"market"`` (only whole-market demands such
|
||||
as SH/SZ/BJ/HK). ``max_wall_seconds`` caps how long one refresh round may run
|
||||
on the caller (strategy) thread; the in-flight demand always completes and at
|
||||
least one demand is always refreshed before the budget can cut the round.
|
||||
"""
|
||||
started_at = time.time()
|
||||
refreshed = 0
|
||||
for demand in iter_active_full_tick_demands(
|
||||
redis_client,
|
||||
account_id,
|
||||
demand_ttl_seconds=demand_ttl_seconds,
|
||||
max_requests=max_requests,
|
||||
):
|
||||
codes = demand.get("codes") or []
|
||||
is_market = _demand_is_market(codes)
|
||||
if kind == "symbol" and is_market:
|
||||
continue
|
||||
if kind == "market" and not is_market:
|
||||
continue
|
||||
if max_wall_seconds and refreshed and (time.time() - started_at) > float(max_wall_seconds):
|
||||
break
|
||||
tick_data = context_info.get_full_tick(codes) or {}
|
||||
ttl = demand.get("cache_ttl_seconds") or cache_ttl_seconds
|
||||
write_full_tick_cache(redis_client, account_id, codes, tick_data, cache_ttl_seconds=ttl)
|
||||
refreshed += 1
|
||||
return refreshed
|
||||
@@ -1,202 +0,0 @@
|
||||
"""Client-side local cache for Big QMT market data.
|
||||
|
||||
Pull bars from Big QMT once over RPC, persist them on the client, then read them
|
||||
back with ``get_local_data`` without touching Big QMT again — for offline / local
|
||||
analysis. One file per (period, dividend_type, code); incremental merge + dedupe
|
||||
by time. Default storage is Parquet (columnar, compressed, cross-language); falls
|
||||
back to pickle when pyarrow is unavailable. A cache written in one format is read
|
||||
+ migrated transparently if the configured format changes.
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
|
||||
# Candidate time-column names produced by the RPC market-data path.
|
||||
_TIME_COLS = ("stime", "time", "index", "date", "datetime", "timetag")
|
||||
|
||||
|
||||
def _time_col(df):
|
||||
cols = list(getattr(df, "columns", []))
|
||||
for name in _TIME_COLS:
|
||||
if name in cols:
|
||||
return name
|
||||
return None
|
||||
|
||||
|
||||
def _pad_end(value):
|
||||
text = str(value)
|
||||
return text + "9" * (14 - len(text)) if 0 < len(text) < 14 else text
|
||||
|
||||
|
||||
def _drop_placeholder_rows(df):
|
||||
"""Big QMT fills dates it has no local data for with all-zero rows. A real bar
|
||||
never has close/open == 0, so drop those placeholders — the cache should hold
|
||||
only real bars, not 0-fill padding."""
|
||||
for col in ("close", "open", "price", "lastPrice"):
|
||||
if col in getattr(df, "columns", []):
|
||||
try:
|
||||
return df[df[col] != 0].reset_index(drop=True)
|
||||
except Exception:
|
||||
return df
|
||||
return df
|
||||
|
||||
|
||||
def _pyarrow_available():
|
||||
try:
|
||||
import pyarrow # noqa: F401
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def _resolve_format(fmt):
|
||||
fmt = str(fmt or "auto").lower()
|
||||
if fmt in ("parquet", "pq"):
|
||||
return "parquet"
|
||||
if fmt in ("pkl", "pickle"):
|
||||
return "pkl"
|
||||
# auto / unknown
|
||||
return "parquet" if _pyarrow_available() else "pkl"
|
||||
|
||||
|
||||
class LocalMarketCache:
|
||||
def __init__(self, cache_dir=None, fmt="auto"):
|
||||
self.cache_dir = str(cache_dir or os.path.join(os.path.expanduser("~"), ".bigqmt_cache"))
|
||||
self.fmt = _resolve_format(fmt)
|
||||
|
||||
def _ext(self):
|
||||
return ".parquet" if self.fmt == "parquet" else ".pkl"
|
||||
|
||||
def path(self, code, period, dividend_type="none"):
|
||||
safe_code = str(code or "").replace("/", "_").replace("\\", "_")
|
||||
div = str(dividend_type or "none")
|
||||
return os.path.join(self.cache_dir, str(period or "1d"), div, safe_code + self._ext())
|
||||
|
||||
def _existing_path(self, code, period, dividend_type):
|
||||
"""Return the on-disk file for this key in the configured format, else the
|
||||
other format (so switching format still finds + migrates the old cache)."""
|
||||
primary = self.path(code, period, dividend_type)
|
||||
if os.path.isfile(primary):
|
||||
return primary
|
||||
base = primary[: -len(self._ext())]
|
||||
for ext in (".parquet", ".pkl"):
|
||||
alt = base + ext
|
||||
if os.path.isfile(alt):
|
||||
return alt
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _read_file(path):
|
||||
import pandas as pd
|
||||
|
||||
# Read by actual file extension (an existing cache may be either format).
|
||||
if path.endswith(".pkl"):
|
||||
return pd.read_pickle(path)
|
||||
return pd.read_parquet(path)
|
||||
|
||||
def _write_file(self, df, path):
|
||||
# Write in the configured format regardless of the path (the temp file ends
|
||||
# with ".tmp", not the format extension).
|
||||
if self.fmt == "parquet":
|
||||
df.to_parquet(path, index=False)
|
||||
else:
|
||||
df.to_pickle(path)
|
||||
|
||||
def write(self, code, period, df, dividend_type="none"):
|
||||
"""Merge ``df`` into the cache for (code, period, dividend_type).
|
||||
|
||||
Dedupe is by time keeping the LAST write, so re-pulling a range overwrites
|
||||
stale values — which is exactly what front-adjusted (前复权) data needs after
|
||||
a new dividend re-scales history. Returns total rows stored.
|
||||
"""
|
||||
if df is None or not hasattr(df, "shape") or df.shape[0] == 0:
|
||||
return 0
|
||||
import pandas as pd
|
||||
|
||||
incoming = _drop_placeholder_rows(df.copy())
|
||||
primary = self.path(code, period, dividend_type)
|
||||
existing = self._existing_path(code, period, dividend_type)
|
||||
if incoming.shape[0] == 0:
|
||||
# Nothing real to add (all 0-fill placeholders); keep existing cache.
|
||||
if existing:
|
||||
try:
|
||||
return self._read_file(existing).shape[0]
|
||||
except Exception:
|
||||
return 0
|
||||
return 0
|
||||
directory = os.path.dirname(primary)
|
||||
if directory and not os.path.isdir(directory):
|
||||
os.makedirs(directory, exist_ok=True)
|
||||
merged = incoming
|
||||
tcol = _time_col(merged)
|
||||
if existing:
|
||||
try:
|
||||
old = self._read_file(existing)
|
||||
merged = pd.concat([old, merged], ignore_index=True)
|
||||
except Exception:
|
||||
pass
|
||||
if tcol and tcol in merged.columns:
|
||||
merged = merged.drop_duplicates(subset=[tcol], keep="last").sort_values(tcol).reset_index(drop=True)
|
||||
else:
|
||||
merged = merged.drop_duplicates().reset_index(drop=True)
|
||||
# Atomic-ish write (temp + replace) so a crash mid-write can't corrupt the file.
|
||||
tmp = primary + ".tmp"
|
||||
self._write_file(merged, tmp)
|
||||
os.replace(tmp, primary)
|
||||
# Migrated from the other format? drop the stale file.
|
||||
if existing and existing != primary:
|
||||
try:
|
||||
os.remove(existing)
|
||||
except Exception:
|
||||
pass
|
||||
return merged.shape[0]
|
||||
|
||||
def read(self, code, period, start_time="", end_time="", count=-1, dividend_type="none"):
|
||||
"""Return the cached DataFrame for (code, period, dividend_type), filtered."""
|
||||
existing = self._existing_path(code, period, dividend_type)
|
||||
if not existing:
|
||||
return None
|
||||
try:
|
||||
df = self._read_file(existing)
|
||||
except Exception:
|
||||
return None
|
||||
tcol = _time_col(df)
|
||||
if tcol and tcol in df.columns:
|
||||
series = df[tcol].astype(str)
|
||||
if start_time:
|
||||
df = df[series >= str(start_time)]
|
||||
if end_time:
|
||||
df = df[series <= _pad_end(end_time)]
|
||||
df = df.sort_values(tcol).reset_index(drop=True)
|
||||
try:
|
||||
n = int(count)
|
||||
except (TypeError, ValueError):
|
||||
n = -1
|
||||
if n > 0 and df.shape[0] > n:
|
||||
df = df.tail(n).reset_index(drop=True)
|
||||
return df
|
||||
|
||||
def covered(self, code, period, dividend_type="none"):
|
||||
"""Return (first_time, last_time, rows) for the cache, or None if empty."""
|
||||
df = self.read(code, period, dividend_type=dividend_type)
|
||||
if df is None or df.shape[0] == 0:
|
||||
return None
|
||||
tcol = _time_col(df)
|
||||
if not tcol:
|
||||
return (None, None, df.shape[0])
|
||||
series = df[tcol].astype(str)
|
||||
return (series.iloc[0], series.iloc[-1], df.shape[0])
|
||||
|
||||
def stats(self):
|
||||
"""Return (files, periods) currently cached across all dividend types."""
|
||||
files = 0
|
||||
periods = set()
|
||||
if os.path.isdir(self.cache_dir):
|
||||
for root, _dirs, fnames in os.walk(self.cache_dir):
|
||||
cached = [f for f in fnames if f.endswith(".parquet") or f.endswith(".pkl")]
|
||||
if cached:
|
||||
files += len(cached)
|
||||
rel = os.path.relpath(root, self.cache_dir)
|
||||
periods.add(rel.split(os.sep)[0] if rel != "." else rel)
|
||||
return files, sorted(periods)
|
||||
@@ -1,171 +0,0 @@
|
||||
"""File-based logging for the Big QMT bridge.
|
||||
|
||||
Diagnostics used to be print()-only, which is lost when the QMT output panel
|
||||
scrolls or the terminal restarts. This module wires Python's stdlib ``logging``
|
||||
to a rotating file so errors survive restarts and can be reviewed after a
|
||||
crash.
|
||||
|
||||
Usage (any module, both server and client):
|
||||
|
||||
from bigqmt_signal_trader.logging_setup import get_logger
|
||||
log = get_logger("rpc")
|
||||
log.info("started")
|
||||
log.error("download failed: %s", exc)
|
||||
|
||||
Behavior:
|
||||
- Log directory resolves to ``<qmt_python_dir>/logs`` when running inside QMT
|
||||
(found via a sys.path entry ending in ``\\python``), else ``~/.cache/bigqmt/logs``.
|
||||
- Rotates at midnight into ``bigqmt.log.YYYY-MM-DD`` backups, keeping the last
|
||||
7 days by default (override with env BIGQMT_LOG_RETENTION_DAYS).
|
||||
- Each record is also printed to stdout so the QMT output panel still shows it.
|
||||
- Thread-safe (logging is; the print side is best-effort wrapped).
|
||||
- Never raises: a logging failure must not bring down the strategy.
|
||||
- Opt out via env BIGQMT_LOG_ENABLED=0 / BIGQMT_LOG_TO_STDOUT=0.
|
||||
"""
|
||||
|
||||
import datetime as _dt
|
||||
import logging
|
||||
import logging.handlers
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
_LOGGER_NAME = "bigqmt"
|
||||
_initialized = False
|
||||
|
||||
|
||||
def _env_bool(name, default=True):
|
||||
value = os.environ.get(name)
|
||||
if value in (None, ""):
|
||||
return default
|
||||
return str(value).strip().lower() in ("1", "true", "yes", "y", "on")
|
||||
|
||||
|
||||
def _resolve_log_dir():
|
||||
"""Pick a writable log dir. Prefers the QMT python dir (the sys.path entry
|
||||
ending in ``\\python``) so logs sit beside the deployed strategy; falls back
|
||||
to a user cache dir otherwise. Deliberately does NOT use this package's own
|
||||
src/ directory (a repo checkout is not a writable runtime location)."""
|
||||
candidates = []
|
||||
for entry in sys.path:
|
||||
try:
|
||||
if entry and entry.endswith(r"\python") and os.path.isdir(entry):
|
||||
candidates.append(entry)
|
||||
except Exception:
|
||||
continue
|
||||
candidates.append(os.path.join(os.path.expanduser("~"), ".cache", "bigqmt"))
|
||||
for base in candidates:
|
||||
try:
|
||||
path = os.path.join(base, "logs")
|
||||
os.makedirs(path, exist_ok=True)
|
||||
# probe writability
|
||||
probe = os.path.join(path, ".write_test")
|
||||
with open(probe, "w"):
|
||||
pass
|
||||
try:
|
||||
os.remove(probe)
|
||||
except Exception:
|
||||
pass
|
||||
return path
|
||||
except Exception:
|
||||
continue
|
||||
return None
|
||||
|
||||
|
||||
class _SafeStreamHandler(logging.Handler):
|
||||
"""print() the record so the QMT output panel shows it; never raises."""
|
||||
|
||||
def emit(self, record):
|
||||
try:
|
||||
print(self.format(record))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _cleanup_old_logs(log_dir, retention_days):
|
||||
"""Delete rotated log files older than retention_days.
|
||||
|
||||
TimedRotatingFileHandler only prunes backups at rotation time; this sweeps
|
||||
stale files on startup too (e.g. after a weekend gap or a config change).
|
||||
"""
|
||||
try:
|
||||
cutoff = time.time() - retention_days * 86400
|
||||
for name in os.listdir(log_dir):
|
||||
if not (name.startswith("bigqmt") and name.endswith(".log") or ".log." in name):
|
||||
continue
|
||||
path = os.path.join(log_dir, name)
|
||||
try:
|
||||
if os.path.getmtime(path) < cutoff:
|
||||
os.remove(path)
|
||||
except Exception:
|
||||
continue
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _setup():
|
||||
global _initialized
|
||||
if _initialized:
|
||||
return
|
||||
_initialized = True
|
||||
logger = logging.getLogger(_LOGGER_NAME)
|
||||
logger.setLevel(logging.DEBUG)
|
||||
logger.propagate = False
|
||||
if not _env_bool("BIGQMT_LOG_ENABLED", True):
|
||||
logger.addHandler(logging.NullHandler())
|
||||
return
|
||||
|
||||
fmt = logging.Formatter(
|
||||
fmt="%(asctime)s [%(levelname)s] [%(name)s] %(message)s",
|
||||
datefmt="%Y-%m-%d %H:%M:%S",
|
||||
)
|
||||
|
||||
# File handler: rotate at midnight, keep the last 7 days only.
|
||||
log_dir = _resolve_log_dir()
|
||||
if log_dir is not None:
|
||||
try:
|
||||
fname = os.path.join(log_dir, "bigqmt.log")
|
||||
file_handler = logging.handlers.TimedRotatingFileHandler(
|
||||
fname,
|
||||
when="midnight",
|
||||
interval=1,
|
||||
backupCount=int(os.environ.get("BIGQMT_LOG_RETENTION_DAYS", 7)),
|
||||
encoding="utf-8",
|
||||
utc=False,
|
||||
)
|
||||
file_handler.setFormatter(fmt)
|
||||
logger.addHandler(file_handler)
|
||||
_cleanup_old_logs(log_dir, int(os.environ.get("BIGQMT_LOG_RETENTION_DAYS", 7)))
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Stdout handler so the QMT panel still shows logs.
|
||||
if _env_bool("BIGQMT_LOG_TO_STDOUT", True):
|
||||
stream = _SafeStreamHandler()
|
||||
stream.setLevel(logging.INFO)
|
||||
stream.setFormatter(fmt)
|
||||
logger.addHandler(stream)
|
||||
|
||||
if not logger.handlers:
|
||||
logger.addHandler(logging.NullHandler())
|
||||
|
||||
|
||||
def get_logger(name=""):
|
||||
"""Return a module logger under the shared ``bigqmt`` root.
|
||||
|
||||
``get_logger("rpc")`` -> logger named ``bigqmt.rpc``; the tag is shown in
|
||||
each log line so the old ``[bigqmt_rpc]`` prefixes remain visible.
|
||||
"""
|
||||
_setup()
|
||||
suffix = str(name or "").strip(".")
|
||||
full = _LOGGER_NAME if not suffix else "%s.%s" % (_LOGGER_NAME, suffix)
|
||||
return logging.getLogger(full)
|
||||
|
||||
|
||||
def log_file_path():
|
||||
"""Return the current log file path (or None if file logging is off)."""
|
||||
_setup()
|
||||
log_dir = _resolve_log_dir()
|
||||
if log_dir is None:
|
||||
return None
|
||||
return os.path.join(log_dir, "bigqmt.log")
|
||||
@@ -1,308 +0,0 @@
|
||||
"""交易信号、委托请求和账户快照的数据模型。"""
|
||||
|
||||
import datetime as _dt
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, List, Optional
|
||||
|
||||
|
||||
class SignalAction(str, Enum):
|
||||
BUY = "BUY"
|
||||
SELL = "SELL"
|
||||
CLEAR = "CLEAR"
|
||||
CANCEL = "CANCEL"
|
||||
|
||||
|
||||
class SignalStatus(str, Enum):
|
||||
PENDING = "PENDING"
|
||||
CLAIMED = "CLAIMED"
|
||||
SUBMITTED = "SUBMITTED"
|
||||
SKIPPED = "SKIPPED"
|
||||
FAILED = "FAILED"
|
||||
FILLED = "FILLED"
|
||||
|
||||
|
||||
def parse_datetime(value: Any, field_name: str) -> _dt.datetime:
|
||||
if isinstance(value, _dt.datetime):
|
||||
return value
|
||||
if isinstance(value, str) and value:
|
||||
try:
|
||||
return _dt.datetime.strptime(value, "%Y-%m-%d %H:%M:%S")
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"{field_name} must use format YYYY-MM-DD HH:MM:SS") from exc
|
||||
raise ValueError(f"{field_name} is required")
|
||||
|
||||
|
||||
def _optional_int(value: Any) -> Optional[int]:
|
||||
if value is None or value == "":
|
||||
return None
|
||||
return int(value)
|
||||
|
||||
|
||||
def _optional_float(value: Any) -> Optional[float]:
|
||||
if value is None or value == "":
|
||||
return None
|
||||
return float(value)
|
||||
|
||||
|
||||
def _bool_value(value: Any) -> bool:
|
||||
if isinstance(value, bool):
|
||||
return value
|
||||
if value is None or value == "":
|
||||
return False
|
||||
if isinstance(value, (int, float)):
|
||||
return bool(value)
|
||||
text = str(value).strip().lower()
|
||||
return text in ("1", "true", "yes", "y", "on")
|
||||
|
||||
|
||||
class TradeSignal:
|
||||
def __init__(
|
||||
self,
|
||||
signal_id,
|
||||
account_id,
|
||||
action,
|
||||
created_at,
|
||||
expire_at,
|
||||
schema_version,
|
||||
stock_code="",
|
||||
stock_name="",
|
||||
amount=None,
|
||||
percentage=None,
|
||||
price_type="AUTO_LIMIT",
|
||||
price=None,
|
||||
strategy_name="bigqmt_signal_trader",
|
||||
remark="",
|
||||
source="",
|
||||
source_type="auto",
|
||||
force=False,
|
||||
bypass_stop_buy=False,
|
||||
bypass_stop_sell=False,
|
||||
bypass_daily_limit=False,
|
||||
status=SignalStatus.PENDING,
|
||||
raw_payload=None,
|
||||
):
|
||||
self.signal_id = signal_id
|
||||
self.account_id = account_id
|
||||
self.action = action
|
||||
self.created_at = created_at
|
||||
self.expire_at = expire_at
|
||||
self.schema_version = schema_version
|
||||
self.stock_code = stock_code
|
||||
self.stock_name = stock_name
|
||||
self.amount = amount
|
||||
self.percentage = percentage
|
||||
self.price_type = price_type
|
||||
self.price = price
|
||||
self.strategy_name = strategy_name
|
||||
self.remark = remark
|
||||
self.source = source
|
||||
self.source_type = source_type
|
||||
self.force = force
|
||||
self.bypass_stop_buy = bypass_stop_buy
|
||||
self.bypass_stop_sell = bypass_stop_sell
|
||||
self.bypass_daily_limit = bypass_daily_limit
|
||||
self.status = status
|
||||
self.raw_payload = dict(raw_payload or {})
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, payload: Dict[str, Any]) -> "TradeSignal":
|
||||
required = ("signal_id", "account_id", "action", "created_at", "expire_at", "schema_version")
|
||||
for field_name in required:
|
||||
if payload.get(field_name) in (None, ""):
|
||||
raise ValueError(f"{field_name} is required")
|
||||
|
||||
try:
|
||||
action = SignalAction(str(payload["action"]).upper())
|
||||
except ValueError as exc:
|
||||
raise ValueError(f"unsupported action: {payload.get('action')}") from exc
|
||||
|
||||
amount = _optional_int(payload.get("amount"))
|
||||
percentage = _optional_float(payload.get("percentage"))
|
||||
stock_code = str(payload.get("stock_code") or "").strip().upper()
|
||||
|
||||
if action == SignalAction.BUY:
|
||||
if not stock_code:
|
||||
raise ValueError("stock_code is required for BUY")
|
||||
if amount is None or amount <= 0:
|
||||
raise ValueError("amount must be positive for BUY")
|
||||
elif action == SignalAction.SELL:
|
||||
if not stock_code:
|
||||
raise ValueError("stock_code is required for SELL")
|
||||
if amount is None and percentage is None:
|
||||
raise ValueError("amount or percentage is required for SELL")
|
||||
elif action == SignalAction.CLEAR and percentage is None:
|
||||
percentage = 100.0
|
||||
|
||||
return cls(
|
||||
signal_id=str(payload["signal_id"]),
|
||||
account_id=str(payload["account_id"]),
|
||||
action=action,
|
||||
stock_code=stock_code,
|
||||
stock_name=str(payload.get("stock_name") or ""),
|
||||
amount=amount,
|
||||
percentage=percentage,
|
||||
price_type=str(payload.get("price_type") or "AUTO_LIMIT").upper(),
|
||||
price=_optional_float(payload.get("price")),
|
||||
strategy_name=str(payload.get("strategy_name") or "bigqmt_signal_trader"),
|
||||
remark=str(payload.get("remark") or ""),
|
||||
source=str(payload.get("source") or ""),
|
||||
source_type=str(payload.get("source_type") or "auto"),
|
||||
force=_bool_value(payload.get("force", False)),
|
||||
bypass_stop_buy=_bool_value(payload.get("bypass_stop_buy", False)),
|
||||
bypass_stop_sell=_bool_value(payload.get("bypass_stop_sell", False)),
|
||||
bypass_daily_limit=_bool_value(payload.get("bypass_daily_limit", False)),
|
||||
created_at=parse_datetime(payload.get("created_at"), "created_at"),
|
||||
expire_at=parse_datetime(payload.get("expire_at"), "expire_at"),
|
||||
schema_version=int(payload["schema_version"]),
|
||||
raw_payload=dict(payload),
|
||||
)
|
||||
|
||||
def is_expired(self, now: _dt.datetime) -> bool:
|
||||
return now > self.expire_at
|
||||
|
||||
|
||||
class PositionSnapshot:
|
||||
def __init__(
|
||||
self,
|
||||
stock_code,
|
||||
volume,
|
||||
available,
|
||||
cost=0.0,
|
||||
stock_name="",
|
||||
market_value=None,
|
||||
price=None,
|
||||
open_price=None,
|
||||
frozen_volume=0,
|
||||
on_road_volume=0,
|
||||
yesterday_volume=None,
|
||||
direction=48,
|
||||
):
|
||||
self.stock_code = stock_code
|
||||
self.volume = volume
|
||||
self.available = available
|
||||
self.cost = cost
|
||||
self.stock_name = stock_name
|
||||
self.market_value = market_value
|
||||
self.price = price
|
||||
self.open_price = open_price
|
||||
self.frozen_volume = frozen_volume
|
||||
self.on_road_volume = on_road_volume
|
||||
self.yesterday_volume = yesterday_volume
|
||||
self.direction = direction
|
||||
|
||||
|
||||
class AssetSnapshot:
|
||||
"""Account funds, mirroring MiniQMT's ``XtAsset``.
|
||||
|
||||
Field names follow ``xtquant.xttype.XtAsset(account_id, cash, frozen_cash,
|
||||
market_value, total_asset)`` so ``query_stock_asset`` can hand callers the
|
||||
same attributes they get from MiniQMT.
|
||||
|
||||
``cash`` is 可用 (available), NOT the full 资金余额:
|
||||
``total_asset == cash + frozen_cash + market_value``. New fields are
|
||||
appended with None defaults so existing positional callers keep working,
|
||||
and None means "the terminal did not report it" — distinct from 0.0.
|
||||
"""
|
||||
|
||||
def __init__(self, account_id, cash=None, total_asset=None, frozen_cash=None, market_value=None):
|
||||
self.account_id = account_id
|
||||
self.cash = cash
|
||||
self.total_asset = total_asset
|
||||
self.frozen_cash = frozen_cash
|
||||
self.market_value = market_value
|
||||
|
||||
|
||||
class AccountSnapshot:
|
||||
def __init__(self, account_id, asset, positions, reason, updated_at):
|
||||
self.account_id = account_id
|
||||
self.asset = asset
|
||||
self.positions = positions
|
||||
self.reason = reason
|
||||
self.updated_at = updated_at
|
||||
|
||||
|
||||
class OrderRequest:
|
||||
def __init__(
|
||||
self,
|
||||
signal_id,
|
||||
account_id,
|
||||
action,
|
||||
stock_code,
|
||||
volume,
|
||||
price,
|
||||
price_type,
|
||||
strategy_name,
|
||||
remark="",
|
||||
):
|
||||
self.signal_id = signal_id
|
||||
self.account_id = account_id
|
||||
self.action = action
|
||||
self.stock_code = stock_code
|
||||
self.volume = volume
|
||||
self.price = price
|
||||
self.price_type = price_type
|
||||
self.strategy_name = strategy_name
|
||||
self.remark = remark
|
||||
|
||||
|
||||
class OrderSubmitResult:
|
||||
def __init__(self, status, user_order_id, order_sys_id=None, message=""):
|
||||
self.status = status
|
||||
self.user_order_id = user_order_id
|
||||
self.order_sys_id = order_sys_id
|
||||
self.message = message
|
||||
|
||||
|
||||
class OrderSnapshot:
|
||||
def __init__(
|
||||
self,
|
||||
order_sys_id,
|
||||
user_order_id,
|
||||
stock_code,
|
||||
action,
|
||||
volume,
|
||||
traded_volume,
|
||||
status,
|
||||
price=0.0,
|
||||
strategy_name="",
|
||||
remark="",
|
||||
order_time=0,
|
||||
):
|
||||
self.order_sys_id = order_sys_id
|
||||
self.user_order_id = user_order_id
|
||||
self.stock_code = stock_code
|
||||
self.action = action
|
||||
self.volume = volume
|
||||
self.traded_volume = traded_volume
|
||||
self.status = status
|
||||
self.price = price
|
||||
self.strategy_name = strategy_name
|
||||
self.remark = remark
|
||||
# 报单时间, Unix 秒 -- MiniQMT XtOrder.order_time 的语义。0 = 未上报。
|
||||
# 追加在末尾并给默认值, 保持既有位置参数调用不受影响。
|
||||
self.order_time = order_time
|
||||
|
||||
|
||||
class TradeSnapshot:
|
||||
def __init__(self, trade_id, order_sys_id, stock_code, action, volume, price,
|
||||
traded_at="", user_order_id=""):
|
||||
self.trade_id = trade_id
|
||||
self.order_sys_id = order_sys_id
|
||||
self.stock_code = stock_code
|
||||
self.action = action
|
||||
self.volume = volume
|
||||
self.price = price
|
||||
self.traded_at = traded_at
|
||||
self.user_order_id = user_order_id
|
||||
|
||||
|
||||
class OrderRef:
|
||||
def __init__(self, order_sys_id, user_order_id=""):
|
||||
self.order_sys_id = order_sys_id
|
||||
self.user_order_id = user_order_id
|
||||
|
||||
|
||||
class CancelResult:
|
||||
def __init__(self, success, message=""):
|
||||
self.success = success
|
||||
self.message = message
|
||||
@@ -1,53 +0,0 @@
|
||||
"""订单价格生成逻辑。"""
|
||||
|
||||
from .code_utils import normalize_stock_code
|
||||
|
||||
|
||||
def _price_precision(stock_code):
|
||||
pure = normalize_stock_code(stock_code).split(".")[0]
|
||||
return 3 if pure.startswith(("15", "16", "51", "52")) else 2
|
||||
|
||||
|
||||
def _second_level(values):
|
||||
if isinstance(values, (list, tuple)) and len(values) > 1:
|
||||
try:
|
||||
value = float(values[1])
|
||||
return value if value > 0 else 0.0
|
||||
except (TypeError, ValueError):
|
||||
return 0.0
|
||||
return 0.0
|
||||
|
||||
|
||||
def build_order_price(market_data, stock_code, action, price_type="AUTO_LIMIT", fixed_price=None):
|
||||
if str(price_type or "AUTO_LIMIT").upper() == "FIX_PRICE":
|
||||
if fixed_price is None:
|
||||
raise ValueError("fixed_price is required when price_type is FIX_PRICE")
|
||||
return float(fixed_price)
|
||||
|
||||
code = normalize_stock_code(stock_code)
|
||||
ticks = market_data.get_ticks([code])
|
||||
tick = ticks.get(code)
|
||||
if not tick:
|
||||
raise ValueError(f"missing tick data for {code}")
|
||||
|
||||
last_price = float(tick.get("lastPrice") or 0)
|
||||
if last_price <= 0:
|
||||
raise ValueError(f"invalid lastPrice for {code}")
|
||||
|
||||
instrument = market_data.get_instrument(code)
|
||||
if int(instrument.get("InstrumentStatus") or 0) > 0:
|
||||
raise ValueError(f"{code} is suspended")
|
||||
|
||||
precision = _price_precision(code)
|
||||
action_text = str(action).upper()
|
||||
if action_text == "BUY":
|
||||
up_stop = float(instrument.get("UpStopPrice") or last_price * 1.1)
|
||||
calculated = min(round(last_price * 1.002, precision), up_stop)
|
||||
ask2 = _second_level(tick.get("askPrice"))
|
||||
return round(ask2, precision) if ask2 and ask2 < calculated else calculated
|
||||
if action_text == "SELL":
|
||||
down_stop = float(instrument.get("DownStopPrice") or last_price * 0.9)
|
||||
calculated = max(round(last_price * 0.998, precision), down_stop)
|
||||
bid2 = _second_level(tick.get("bidPrice"))
|
||||
return round(bid2, precision) if bid2 and bid2 > calculated else calculated
|
||||
raise ValueError(f"unsupported action for price: {action}")
|
||||
@@ -1,484 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Start and stop a Big QMT terminal (issue #45).
|
||||
|
||||
Big QMT has to be restarted most mornings, and the login dialog is the reason
|
||||
it cannot simply be dropped into a scheduler. Two ways past it:
|
||||
|
||||
* **Passwordless (preferred)** -- ``XtMiniQmt.exe linkMini`` starts MiniQMT
|
||||
against an existing session with no dialog at all. This is what
|
||||
``免密登录qmt.bat`` does. No UI automation, so nothing here depends on a
|
||||
desktop being visible.
|
||||
* **Credential entry** -- for the full terminal (``XtItClient.exe``) the dialog
|
||||
is unavoidable. We drive it with ``win32api.SendMessage`` posted straight to
|
||||
the window handle, NOT with pyautogui/pywinauto. That distinction is the
|
||||
answer to the question in issue #45: pyautogui replays physical input at
|
||||
screen coordinates, so it needs the window focused and the desktop unlocked;
|
||||
SendMessage delivers to a handle and works on a background -- or locked --
|
||||
session, as long as the session still exists (an RDP disconnect is fine, a
|
||||
full logout is not).
|
||||
|
||||
Everything is scoped to one install directory. A machine here runs several QMT
|
||||
copies side by side, so an unscoped ``taskkill /im XtItClient.exe`` would take
|
||||
down someone else's trading session.
|
||||
|
||||
Waits are on observed readiness, never a fixed sleep: startup is complete when
|
||||
the FormulaServer port accepts a connection, which is also exactly what the
|
||||
rest of this package needs before it can do anything.
|
||||
|
||||
Windows only. ``psutil`` is used when importable, otherwise we shell out to
|
||||
``wmic``/``taskkill``.
|
||||
"""
|
||||
|
||||
import os
|
||||
import socket
|
||||
import subprocess
|
||||
import sys
|
||||
import time
|
||||
|
||||
from .logging_setup import get_logger
|
||||
|
||||
|
||||
log = get_logger("launcher")
|
||||
|
||||
# Processes a QMT install owns. miniquote/BrokerProxy/minibroker are children
|
||||
# that survive the main window and hold the ports we need to rebind.
|
||||
QMT_PROCESS_NAMES = (
|
||||
"XtItClient.exe",
|
||||
"XtMiniQmt.exe",
|
||||
"miniquote.exe",
|
||||
"BrokerProxy.exe",
|
||||
"minibroker.exe",
|
||||
)
|
||||
|
||||
# FormulaServer. Listening means the terminal is far enough up to answer.
|
||||
DEFAULT_READY_PORT = 58600
|
||||
|
||||
__all__ = [
|
||||
"QmtLauncherError",
|
||||
"close_qmt",
|
||||
"find_qmt_processes",
|
||||
"is_qmt_running",
|
||||
"open_qmt",
|
||||
"restart_qmt",
|
||||
"wait_until_ready",
|
||||
]
|
||||
|
||||
|
||||
class QmtLauncherError(RuntimeError):
|
||||
"""Launching or stopping a QMT terminal failed."""
|
||||
|
||||
|
||||
def _normalize_dir(path):
|
||||
if not path:
|
||||
return ""
|
||||
return os.path.normcase(os.path.normpath(os.path.abspath(str(path))))
|
||||
|
||||
|
||||
def resolve_install_dir(install_dir):
|
||||
"""Accept an install root, its bin.x64, or a path to an exe inside it.
|
||||
|
||||
Returns the normalized ``bin.x64`` directory, which is what process paths
|
||||
are compared against.
|
||||
"""
|
||||
path = str(install_dir or "").strip().strip('"').strip("'")
|
||||
if not path:
|
||||
raise QmtLauncherError("install_dir is required (QMT root, bin.x64, or an exe path)")
|
||||
if os.path.isfile(path) or path.lower().endswith(".exe"):
|
||||
path = os.path.dirname(path)
|
||||
normalized = os.path.normpath(os.path.abspath(path))
|
||||
if os.path.basename(normalized).lower() != "bin.x64":
|
||||
candidate = os.path.join(normalized, "bin.x64")
|
||||
if os.path.isdir(candidate):
|
||||
normalized = candidate
|
||||
return normalized
|
||||
|
||||
|
||||
# ---------------------------------------------------------------- discovery
|
||||
def _iter_processes_psutil():
|
||||
import psutil
|
||||
|
||||
for proc in psutil.process_iter(["pid", "name", "exe"]):
|
||||
try:
|
||||
info = proc.info
|
||||
yield int(info["pid"]), str(info.get("name") or ""), str(info.get("exe") or "")
|
||||
except (psutil.NoSuchProcess, psutil.AccessDenied, psutil.ZombieProcess):
|
||||
continue
|
||||
|
||||
|
||||
def _iter_processes_wmic():
|
||||
"""psutil-free fallback. wmic still ships on the Windows builds QMT runs on."""
|
||||
try:
|
||||
raw = subprocess.check_output(
|
||||
["wmic", "process", "get", "ProcessId,Name,ExecutablePath", "/format:csv"],
|
||||
stderr=subprocess.STDOUT,
|
||||
)
|
||||
except Exception as exc:
|
||||
raise QmtLauncherError(
|
||||
"cannot enumerate processes: psutil is not installed and wmic failed (%s)" % exc
|
||||
)
|
||||
text = raw.decode("utf-8", "replace") if isinstance(raw, bytes) else str(raw)
|
||||
for row in text.splitlines():
|
||||
parts = [p.strip() for p in row.split(",")]
|
||||
# CSV columns: Node,ExecutablePath,Name,ProcessId
|
||||
if len(parts) < 4 or parts[3].lower() in ("processid", ""):
|
||||
continue
|
||||
try:
|
||||
pid = int(parts[3])
|
||||
except ValueError:
|
||||
continue
|
||||
yield pid, parts[2], parts[1]
|
||||
|
||||
|
||||
def _iter_processes():
|
||||
try:
|
||||
import psutil # noqa: F401
|
||||
except ImportError:
|
||||
return _iter_processes_wmic()
|
||||
return _iter_processes_psutil()
|
||||
|
||||
|
||||
def find_qmt_processes(install_dir, names=QMT_PROCESS_NAMES):
|
||||
"""Return ``[(pid, name, exe), ...]`` for QMT processes under ``install_dir``.
|
||||
|
||||
A process with no readable exe path is skipped rather than guessed at: on a
|
||||
machine running several QMT copies, killing by name alone is how you take
|
||||
down the wrong account.
|
||||
"""
|
||||
target = _normalize_dir(resolve_install_dir(install_dir))
|
||||
wanted = set(str(n).lower() for n in names)
|
||||
found = []
|
||||
for pid, name, exe in _iter_processes():
|
||||
if name.lower() not in wanted or not exe:
|
||||
continue
|
||||
if _normalize_dir(os.path.dirname(exe)) == target:
|
||||
found.append((pid, name, exe))
|
||||
return found
|
||||
|
||||
|
||||
def is_qmt_running(install_dir):
|
||||
return bool(find_qmt_processes(install_dir))
|
||||
|
||||
|
||||
# ------------------------------------------------------------------ readiness
|
||||
def port_is_listening(port=DEFAULT_READY_PORT, host="127.0.0.1", timeout=1.0):
|
||||
sock = socket.socket()
|
||||
sock.settimeout(timeout)
|
||||
try:
|
||||
sock.connect((host, port))
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
finally:
|
||||
try:
|
||||
sock.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def wait_until_ready(port=DEFAULT_READY_PORT, host="127.0.0.1", timeout_seconds=180.0,
|
||||
poll_interval=2.0):
|
||||
"""Block until ``port`` accepts a connection. Returns seconds waited.
|
||||
|
||||
Raises :class:`QmtLauncherError` on timeout rather than returning False, so
|
||||
a scheduled restart fails loudly instead of letting the next step run
|
||||
against a terminal that never came up.
|
||||
"""
|
||||
deadline = time.time() + float(timeout_seconds)
|
||||
started = time.time()
|
||||
while time.time() < deadline:
|
||||
if port_is_listening(port, host):
|
||||
waited = time.time() - started
|
||||
log.info("qmt ready after %.1fs (%s:%d listening)", waited, host, port)
|
||||
return waited
|
||||
time.sleep(poll_interval)
|
||||
raise QmtLauncherError(
|
||||
"QMT did not become ready within %.0fs (%s:%d never listened)"
|
||||
% (timeout_seconds, host, port)
|
||||
)
|
||||
|
||||
|
||||
def wait_until_stopped(install_dir, timeout_seconds=60.0, poll_interval=1.0):
|
||||
deadline = time.time() + float(timeout_seconds)
|
||||
while time.time() < deadline:
|
||||
if not find_qmt_processes(install_dir):
|
||||
return True
|
||||
time.sleep(poll_interval)
|
||||
return False
|
||||
|
||||
|
||||
# --------------------------------------------------------------------- close
|
||||
def _terminate(pid, force=False):
|
||||
try:
|
||||
import psutil
|
||||
|
||||
proc = psutil.Process(pid)
|
||||
if force:
|
||||
proc.kill()
|
||||
else:
|
||||
proc.terminate()
|
||||
return True
|
||||
except ImportError:
|
||||
pass
|
||||
except Exception:
|
||||
return False
|
||||
cmd = ["taskkill", "/pid", str(pid)]
|
||||
if force:
|
||||
cmd.append("/f")
|
||||
try:
|
||||
subprocess.check_output(cmd, stderr=subprocess.STDOUT)
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
def close_qmt(install_dir, timeout_seconds=60.0, force_after_seconds=20.0):
|
||||
"""Stop every QMT process under ``install_dir``. Returns how many were stopped.
|
||||
|
||||
Asks politely first: the terminal flushes local data on a clean exit, and
|
||||
killing it outright is how the K-line store ends up truncated. Escalates to
|
||||
a hard kill only after ``force_after_seconds``.
|
||||
"""
|
||||
targets = find_qmt_processes(install_dir)
|
||||
if not targets:
|
||||
log.info("no QMT process under %s; nothing to close", install_dir)
|
||||
return 0
|
||||
|
||||
for pid, name, _exe in targets:
|
||||
log.info("closing %s (pid=%s)", name, pid)
|
||||
_terminate(pid, force=False)
|
||||
|
||||
if wait_until_stopped(install_dir, timeout_seconds=force_after_seconds):
|
||||
log.info("closed %d process(es) cleanly", len(targets))
|
||||
return len(targets)
|
||||
|
||||
remaining = find_qmt_processes(install_dir)
|
||||
log.warning("%d process(es) still alive after %.0fs; forcing",
|
||||
len(remaining), force_after_seconds)
|
||||
for pid, name, _exe in remaining:
|
||||
_terminate(pid, force=True)
|
||||
|
||||
grace = max(timeout_seconds - force_after_seconds, 5.0)
|
||||
if not wait_until_stopped(install_dir, timeout_seconds=grace):
|
||||
still = find_qmt_processes(install_dir)
|
||||
raise QmtLauncherError(
|
||||
"could not stop: %s" % ", ".join("%s(pid=%s)" % (n, p) for p, n, _ in still)
|
||||
)
|
||||
return len(targets)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------- open
|
||||
def _spawn(command, cwd=None, shell=False):
|
||||
log.info("launching: %s", command if isinstance(command, str) else " ".join(command))
|
||||
kwargs = {"cwd": cwd, "shell": shell,
|
||||
"stdout": subprocess.DEVNULL, "stderr": subprocess.DEVNULL}
|
||||
if os.name == "nt":
|
||||
# Detach so the terminal outlives this process -- otherwise a scheduled
|
||||
# task exiting takes QMT with it.
|
||||
detached = getattr(subprocess, "DETACHED_PROCESS", 0x00000008)
|
||||
new_group = getattr(subprocess, "CREATE_NEW_PROCESS_GROUP", 0x00000200)
|
||||
kwargs["creationflags"] = detached | new_group
|
||||
return subprocess.Popen(command, **kwargs)
|
||||
|
||||
|
||||
def open_qmt(install_dir, mode="auto", bat_path=None, exe_name=None,
|
||||
ready_port=DEFAULT_READY_PORT, ready_timeout_seconds=180.0,
|
||||
wait_ready=True, credentials=None, window_title_prefix=None):
|
||||
"""Start a QMT terminal under ``install_dir`` and wait until it answers.
|
||||
|
||||
``mode``:
|
||||
``"linkmini"`` -- ``XtMiniQmt.exe linkMini``, no login dialog.
|
||||
``"bat"`` -- run ``bat_path`` (e.g. 免密登录qmt.bat).
|
||||
``"exe"`` -- start ``exe_name`` (default XtItClient.exe) as-is; use
|
||||
when the terminal restores its own session.
|
||||
``"login"`` -- start the exe, then type credentials into the dialog.
|
||||
``"auto"`` -- bat if given, else linkmini if XtMiniQmt.exe exists,
|
||||
else exe.
|
||||
|
||||
``credentials`` (mode="login") is ``{"user": ..., "password": ...}``. Pass
|
||||
it from your local config or environment; never hardcode it, and note the
|
||||
values are typed into a window, so anything that can read that window can
|
||||
read them.
|
||||
"""
|
||||
bin_dir = resolve_install_dir(install_dir)
|
||||
if not os.path.isdir(bin_dir):
|
||||
raise QmtLauncherError("no such directory: %s" % bin_dir)
|
||||
|
||||
mode = str(mode or "auto").lower()
|
||||
if mode == "auto":
|
||||
if bat_path:
|
||||
mode = "bat"
|
||||
elif os.path.isfile(os.path.join(bin_dir, "XtMiniQmt.exe")):
|
||||
mode = "linkmini"
|
||||
else:
|
||||
mode = "exe"
|
||||
|
||||
if mode == "bat":
|
||||
if not bat_path or not os.path.isfile(bat_path):
|
||||
raise QmtLauncherError("bat_path is required for mode='bat': %r" % bat_path)
|
||||
_spawn([bat_path], cwd=os.path.dirname(bat_path), shell=True)
|
||||
elif mode == "linkmini":
|
||||
exe = os.path.join(bin_dir, "XtMiniQmt.exe")
|
||||
if not os.path.isfile(exe):
|
||||
raise QmtLauncherError("XtMiniQmt.exe not found in %s" % bin_dir)
|
||||
_spawn([exe, "linkMini"], cwd=bin_dir)
|
||||
elif mode in ("exe", "login"):
|
||||
exe = os.path.join(bin_dir, str(exe_name or "XtItClient.exe"))
|
||||
if not os.path.isfile(exe):
|
||||
raise QmtLauncherError("%s not found in %s" % (os.path.basename(exe), bin_dir))
|
||||
_spawn([exe], cwd=bin_dir)
|
||||
if mode == "login":
|
||||
_login_via_window(credentials or {}, window_title_prefix)
|
||||
else:
|
||||
raise QmtLauncherError("unknown mode %r (bat/linkmini/exe/login/auto)" % mode)
|
||||
|
||||
if not wait_ready:
|
||||
return 0.0
|
||||
return wait_until_ready(ready_port, timeout_seconds=ready_timeout_seconds)
|
||||
|
||||
|
||||
def _login_via_window(credentials, window_title_prefix=None, appear_timeout_seconds=90.0):
|
||||
"""Type credentials into the QMT login dialog via SendMessage.
|
||||
|
||||
Matches the window by title PREFIX. The reference implementation pinned the
|
||||
full title including a build number ("国金证券QMT交易端 1.0.0.29456"), which
|
||||
stops finding the window on the next terminal update.
|
||||
"""
|
||||
user = str(credentials.get("user") or credentials.get("account") or "")
|
||||
password = str(credentials.get("password") or "")
|
||||
if not user or not password:
|
||||
raise QmtLauncherError(
|
||||
"mode='login' needs credentials={'user':..., 'password':...}"
|
||||
)
|
||||
try:
|
||||
import win32api
|
||||
import win32con
|
||||
import win32gui
|
||||
except ImportError:
|
||||
raise QmtLauncherError(
|
||||
"mode='login' needs pywin32 (pip install pywin32); "
|
||||
"prefer mode='linkmini' or mode='bat', which need no UI automation"
|
||||
)
|
||||
|
||||
prefix = str(window_title_prefix or "QMT")
|
||||
|
||||
def _collect(hwnd, acc):
|
||||
if not win32gui.IsWindowVisible(hwnd):
|
||||
return
|
||||
title = win32gui.GetWindowText(hwnd) or ""
|
||||
if title.strip().startswith(prefix):
|
||||
acc.append(hwnd)
|
||||
|
||||
def _find():
|
||||
matches = []
|
||||
win32gui.EnumWindows(_collect, matches)
|
||||
return matches[0] if matches else None
|
||||
|
||||
deadline = time.time() + appear_timeout_seconds
|
||||
handle = None
|
||||
while time.time() < deadline:
|
||||
handle = _find()
|
||||
if handle:
|
||||
break
|
||||
time.sleep(2.0)
|
||||
if not handle:
|
||||
raise QmtLauncherError(
|
||||
"login window starting with %r did not appear within %.0fs"
|
||||
% (prefix, appear_timeout_seconds)
|
||||
)
|
||||
|
||||
def _send_text(text):
|
||||
for ch in str(text):
|
||||
win32api.SendMessage(handle, win32con.WM_KEYDOWN, ord(ch), 0)
|
||||
win32api.SendMessage(handle, win32con.WM_KEYUP, ord(ch), 0)
|
||||
time.sleep(0.2)
|
||||
|
||||
def _send_enter():
|
||||
win32api.SendMessage(handle, win32con.WM_KEYDOWN, win32con.VK_RETURN, 0)
|
||||
win32api.SendMessage(handle, win32con.WM_KEYUP, win32con.VK_RETURN, 0)
|
||||
time.sleep(1.0)
|
||||
|
||||
# Never log the values themselves.
|
||||
log.info("entering credentials into window %r", prefix)
|
||||
_send_text(user)
|
||||
_send_enter()
|
||||
_send_text(password)
|
||||
_send_enter()
|
||||
_send_enter()
|
||||
|
||||
|
||||
def restart_qmt(install_dir, settle_seconds=5.0, **open_kwargs):
|
||||
"""Close, wait for the ports to be released, then start again.
|
||||
|
||||
``settle_seconds`` matters: the FormulaServer and RPC sockets linger briefly
|
||||
after the process dies, and the ZMQ transport binds its configured port
|
||||
exactly (no scanning), so restarting too eagerly fails the rebind.
|
||||
"""
|
||||
closed = close_qmt(install_dir)
|
||||
if closed:
|
||||
time.sleep(settle_seconds)
|
||||
waited = open_qmt(install_dir, **open_kwargs)
|
||||
log.info("restart complete (closed=%d, ready in %.1fs)", closed, waited)
|
||||
return waited
|
||||
|
||||
|
||||
# ----------------------------------------------------------------------- CLI
|
||||
def main(argv=None):
|
||||
import argparse
|
||||
|
||||
parser = argparse.ArgumentParser(
|
||||
prog="python -m bigqmt_signal_trader.qmt_launcher",
|
||||
description="Start/stop a Big QMT terminal, scoped to one install directory.",
|
||||
)
|
||||
parser.add_argument("action", choices=("open", "close", "restart", "status"))
|
||||
parser.add_argument("--dir", required=True,
|
||||
help="QMT root, its bin.x64, or a path to an exe inside it")
|
||||
parser.add_argument("--mode", default="auto",
|
||||
choices=("auto", "bat", "linkmini", "exe", "login"))
|
||||
parser.add_argument("--bat", default=None, help="batch file for --mode bat")
|
||||
parser.add_argument("--exe", default=None, help="exe name for --mode exe/login")
|
||||
parser.add_argument("--port", type=int, default=DEFAULT_READY_PORT)
|
||||
parser.add_argument("--timeout", type=float, default=180.0)
|
||||
parser.add_argument("--no-wait", action="store_true")
|
||||
parser.add_argument("--title-prefix", default=None,
|
||||
help="login window title prefix (--mode login)")
|
||||
args = parser.parse_args(argv)
|
||||
|
||||
if args.action == "status":
|
||||
procs = find_qmt_processes(args.dir)
|
||||
if not procs:
|
||||
print("not running (%s)" % resolve_install_dir(args.dir))
|
||||
return 1
|
||||
for pid, name, exe in procs:
|
||||
print("%-16s pid=%-8s %s" % (name, pid, exe))
|
||||
print("ready port %d: %s" % (
|
||||
args.port, "listening" if port_is_listening(args.port) else "not listening"))
|
||||
return 0
|
||||
|
||||
credentials = None
|
||||
if args.mode == "login":
|
||||
# Read from the environment so a password never reaches argv, where it
|
||||
# would be visible to any process listing.
|
||||
credentials = {"user": os.environ.get("BIGQMT_LOGIN_USER", ""),
|
||||
"password": os.environ.get("BIGQMT_LOGIN_PASSWORD", "")}
|
||||
|
||||
try:
|
||||
if args.action == "close":
|
||||
print("closed %d process(es)" % close_qmt(args.dir))
|
||||
else:
|
||||
kwargs = dict(mode=args.mode, bat_path=args.bat, exe_name=args.exe,
|
||||
ready_port=args.port, ready_timeout_seconds=args.timeout,
|
||||
wait_ready=not args.no_wait, credentials=credentials,
|
||||
window_title_prefix=args.title_prefix)
|
||||
if args.action == "restart":
|
||||
restart_qmt(args.dir, **kwargs)
|
||||
else:
|
||||
open_qmt(args.dir, **kwargs)
|
||||
print("ok")
|
||||
except QmtLauncherError as exc:
|
||||
print("error: %s" % exc, file=sys.stderr)
|
||||
return 2
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
@@ -1,270 +0,0 @@
|
||||
"""Server→client whole-quote push channel.
|
||||
|
||||
The RPC transport is request/response only; whole-quote data needs the opposite
|
||||
direction — the server pushes each incremental tick batch to every client
|
||||
subscribed to that combination. This module provides one abstract channel with
|
||||
two interchangeable implementations:
|
||||
|
||||
* :class:`ZmqQuotePushChannel` — a ``PUB`` socket on the server, a ``SUB`` socket
|
||||
per client. Native to no-redis deployments. Fire-and-forget: a client that is
|
||||
down simply misses frames (acceptable for incremental quote pushes).
|
||||
* :class:`RedisQuotePushChannel` — redis ``publish``/``subscribe`` on a
|
||||
per-account, per-combination channel, for redis deployments.
|
||||
|
||||
Wire encoding is msgpack when available (smaller + faster for the
|
||||
``{code: {field: number}}`` payload shape), falling back to stdlib json so the
|
||||
channel stays usable without the optional dependency.
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
|
||||
try:
|
||||
import msgpack
|
||||
|
||||
_HAS_MSGPACK = True
|
||||
except Exception: # pragma: no cover - depends on optional dependency
|
||||
msgpack = None
|
||||
_HAS_MSGPACK = False
|
||||
|
||||
|
||||
def encode_push_payload(payload):
|
||||
"""Encode a push payload dict to bytes (msgpack preferred, json fallback)."""
|
||||
if _HAS_MSGPACK:
|
||||
return msgpack.packb(payload, use_bin_type=True)
|
||||
return json.dumps(payload, ensure_ascii=False, default=str).encode("utf-8")
|
||||
|
||||
|
||||
def decode_push_payload(blob):
|
||||
"""Inverse of :func:`encode_push_payload`. Accepts bytes or str.
|
||||
|
||||
Encoding is not symmetric across deployments: a server without msgpack
|
||||
falls back to json while a client with msgpack installed decodes with
|
||||
msgpack — ``msgpack.unpackb`` then raises ``ExtraData`` on the json text
|
||||
(its first byte ``{`` parses as an int, leaving trailing bytes). So try
|
||||
msgpack first, and fall back to json when the bytes are not a single
|
||||
valid msgpack object.
|
||||
"""
|
||||
if blob is None:
|
||||
return None
|
||||
if isinstance(blob, str):
|
||||
blob = blob.encode("utf-8")
|
||||
if _HAS_MSGPACK:
|
||||
try:
|
||||
return msgpack.unpackb(blob, raw=False)
|
||||
except Exception:
|
||||
pass
|
||||
return json.loads(blob.decode("utf-8"))
|
||||
|
||||
|
||||
class QuotePushChannel(object):
|
||||
"""Abstract push channel. Server side: ``start_publisher`` + ``publish``.
|
||||
Client side: ``start_subscriber(topics, on_msg)``. A single instance may act
|
||||
as publisher or subscriber depending on which start method is called."""
|
||||
|
||||
def start_publisher(self):
|
||||
raise NotImplementedError
|
||||
|
||||
def start_subscriber(self, topics, on_msg):
|
||||
raise NotImplementedError
|
||||
|
||||
def publish(self, topic, data):
|
||||
raise NotImplementedError
|
||||
|
||||
def stop(self):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class ZmqQuotePushChannel(QuotePushChannel):
|
||||
def __init__(self, bind_address=None, connect_address=None, context=None, print_prefix="[bigqmt_quote_push]"):
|
||||
self.bind_address = bind_address
|
||||
self.connect_address = connect_address
|
||||
self.print_prefix = print_prefix
|
||||
self._zmq = None
|
||||
self._context = context
|
||||
self._pub = None
|
||||
self._pub_lock = threading.Lock()
|
||||
self._sub = None
|
||||
self._sub_thread = None
|
||||
self._running = False
|
||||
|
||||
def _ensure_context(self):
|
||||
if self._zmq is None:
|
||||
import zmq
|
||||
|
||||
self._zmq = zmq
|
||||
if self._context is None:
|
||||
self._context = zmq.Context.instance()
|
||||
return self._zmq, self._context
|
||||
|
||||
# -- server side ---------------------------------------------------------
|
||||
def start_publisher(self):
|
||||
zmq, ctx = self._ensure_context()
|
||||
if not self.bind_address:
|
||||
raise ValueError("bind_address is required to start a publisher")
|
||||
self._pub = ctx.socket(zmq.PUB)
|
||||
self._pub.bind(self.bind_address)
|
||||
self._running = True
|
||||
|
||||
def publish(self, topic, data):
|
||||
payload = encode_push_payload({"combo_key": topic, "data": data})
|
||||
frame = [str(topic).encode("utf-8"), payload]
|
||||
# PUB socket is not thread-safe; serialize under the lock and read the
|
||||
# socket inside it so a concurrent stop() (which nulls _pub) can't hand
|
||||
# us a closed socket.
|
||||
with self._pub_lock:
|
||||
pub = self._pub
|
||||
if pub is None:
|
||||
return
|
||||
try:
|
||||
pub.send_multipart(frame)
|
||||
except Exception as exc:
|
||||
print("%s zmq publish failed: %s" % (self.print_prefix, exc))
|
||||
|
||||
# -- client side ---------------------------------------------------------
|
||||
def start_subscriber(self, topics, on_msg):
|
||||
zmq, ctx = self._ensure_context()
|
||||
if not self.connect_address:
|
||||
raise ValueError("connect_address is required to start a subscriber")
|
||||
sub = ctx.socket(zmq.SUB)
|
||||
sub.connect(self.connect_address)
|
||||
for topic in topics or []:
|
||||
sub.setsockopt(zmq.SUBSCRIBE, str(topic).encode("utf-8"))
|
||||
self._sub = sub
|
||||
self._running = True
|
||||
self._sub_thread = threading.Thread(
|
||||
target=self._sub_loop, args=(sub, on_msg), name="bigqmt-quote-push-sub", daemon=True
|
||||
)
|
||||
self._sub_thread.start()
|
||||
|
||||
def _sub_loop(self, sub, on_msg):
|
||||
# The SUB socket is owned by THIS thread; it must be closed HERE (in a
|
||||
# finally) and never from another thread. Closing a ZMQ socket cross-
|
||||
# thread trips a Windows signaler assertion and aborts the whole QMT
|
||||
# process (the "auto-exit" users hit).
|
||||
poller = self._zmq.Poller()
|
||||
poller.register(sub, self._zmq.POLLIN)
|
||||
try:
|
||||
while self._running:
|
||||
try:
|
||||
events = dict(poller.poll(200))
|
||||
except Exception:
|
||||
break
|
||||
if sub not in events:
|
||||
continue
|
||||
try:
|
||||
frames = sub.recv_multipart(self._zmq.NOBLOCK)
|
||||
except Exception:
|
||||
continue
|
||||
if len(frames) < 2:
|
||||
continue
|
||||
topic = frames[0].decode("utf-8", errors="ignore")
|
||||
data = decode_push_payload(frames[-1])
|
||||
payload_data = data.get("data") if isinstance(data, dict) else data
|
||||
try:
|
||||
on_msg(topic, payload_data)
|
||||
except Exception as exc:
|
||||
print("%s subscriber callback failed: %s" % (self.print_prefix, exc))
|
||||
finally:
|
||||
try:
|
||||
sub.close(linger=0)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def stop(self):
|
||||
# Signal the sub thread to exit and let IT close its own socket (see
|
||||
# _sub_loop). Closing the SUB socket from this (foreign) thread would
|
||||
# trip the Windows ZMQ signaler abort and crash QMT.
|
||||
self._running = False
|
||||
thread = self._sub_thread
|
||||
if thread is not None and thread.is_alive():
|
||||
thread.join(1.0)
|
||||
self._sub_thread = None
|
||||
self._sub = None
|
||||
# The PUB socket is only touched by publisher threads under _pub_lock;
|
||||
# null it first so a racing publish() sees None and bails, then close.
|
||||
with self._pub_lock:
|
||||
pub = self._pub
|
||||
self._pub = None
|
||||
if pub is not None:
|
||||
try:
|
||||
pub.close(linger=0)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
class RedisQuotePushChannel(QuotePushChannel):
|
||||
def __init__(self, redis_client, account_id="", channel_template="bigqmt:quote_push:{account_id}:{topic}", print_prefix="[bigqmt_quote_push]"):
|
||||
self.redis = redis_client
|
||||
self.account_id = str(account_id or "")
|
||||
self.channel_template = channel_template
|
||||
self.print_prefix = print_prefix
|
||||
self._running = False
|
||||
self._pubsub = None
|
||||
self._thread = None
|
||||
|
||||
def _channel(self, topic):
|
||||
return self.channel_template.format(account_id=self.account_id, topic=topic)
|
||||
|
||||
# -- server side ---------------------------------------------------------
|
||||
def start_publisher(self):
|
||||
# Redis publish needs no setup; present for interface symmetry.
|
||||
self._running = True
|
||||
|
||||
def publish(self, topic, data):
|
||||
payload = encode_push_payload({"combo_key": topic, "data": data})
|
||||
try:
|
||||
self.redis.publish(self._channel(topic), payload)
|
||||
except Exception as exc:
|
||||
print("%s redis publish failed: %s" % (self.print_prefix, exc))
|
||||
|
||||
# -- client side ---------------------------------------------------------
|
||||
def start_subscriber(self, topics, on_msg):
|
||||
self._running = True
|
||||
self._thread = threading.Thread(
|
||||
target=self._sub_loop, args=(list(topics or []), on_msg), name="bigqmt-quote-push-sub", daemon=True
|
||||
)
|
||||
self._thread.start()
|
||||
|
||||
def _sub_loop(self, topics, on_msg):
|
||||
# The pubsub connection is owned by THIS thread and closed HERE so a
|
||||
# concurrent stop() can't close it out from under us.
|
||||
pubsub = self.redis.pubsub(ignore_subscribe_messages=True)
|
||||
self._pubsub = pubsub
|
||||
channels = [self._channel(topic) for topic in topics]
|
||||
try:
|
||||
pubsub.subscribe(*channels)
|
||||
except Exception as exc:
|
||||
print("%s redis subscribe failed: %s" % (self.print_prefix, exc))
|
||||
return
|
||||
try:
|
||||
while self._running:
|
||||
try:
|
||||
message = pubsub.get_message(timeout=0.2)
|
||||
except Exception:
|
||||
break
|
||||
if not message or message.get("type") != "message":
|
||||
continue
|
||||
channel = message.get("channel")
|
||||
if isinstance(channel, bytes):
|
||||
channel = channel.decode("utf-8", errors="ignore")
|
||||
topic = str(channel).rsplit(":", 1)[-1]
|
||||
data = decode_push_payload(message.get("data"))
|
||||
payload_data = data.get("data") if isinstance(data, dict) else data
|
||||
try:
|
||||
on_msg(topic, payload_data)
|
||||
except Exception as exc:
|
||||
print("%s subscriber callback failed: %s" % (self.print_prefix, exc))
|
||||
finally:
|
||||
try:
|
||||
pubsub.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def stop(self):
|
||||
self._running = False
|
||||
thread = self._thread
|
||||
if thread is not None and thread.is_alive():
|
||||
thread.join(1.0)
|
||||
self._thread = None
|
||||
self._pubsub = None
|
||||
@@ -1,258 +0,0 @@
|
||||
"""Reference-counted whole-quote subscription manager (server side).
|
||||
|
||||
One big-QMT ``ContextInfo.subscribe_whole_quote`` subscription is shared by every
|
||||
client that asked for the same (normalized) code combination. The big-QMT
|
||||
subscription is only created for the first client of a combination and only torn
|
||||
down after the last client either unsubscribes or goes silent (keepalive timeout).
|
||||
|
||||
The manager talks to big QMT exclusively through a :class:`QuoteSourceAdapter`;
|
||||
it never touches ``ContextInfo`` directly so the real-environment wiring (method
|
||||
names / handle shape) stays isolated to the adapter.
|
||||
|
||||
Threading: ``subscribe``/``unsubscribe``/``keepalive`` run on the RPC thread,
|
||||
``reap_expired`` on the scheduler thread and ``on_push`` on big QMT's quote
|
||||
thread. Shared state is guarded by one re-entrant lock; calls out to the quote
|
||||
source and to the push publisher happen OUTSIDE the lock so a slow/blocking
|
||||
publish never stalls quote-thread state, and no callback can deadlock.
|
||||
"""
|
||||
|
||||
import threading
|
||||
|
||||
|
||||
def combo_key(code_list):
|
||||
"""Normalize a code list into an order-independent combination key.
|
||||
|
||||
Uppercases, strips whitespace, drops empties and duplicates, sorts. So
|
||||
``["SH","SZ"]``, ``["sz","sh"]`` and ``["SH","SH","SZ"]`` all map to
|
||||
``"SH,SZ"`` and share one big-QMT subscription.
|
||||
"""
|
||||
normalized = {str(code).strip().upper() for code in (code_list or []) if str(code or "").strip()}
|
||||
return ",".join(sorted(normalized))
|
||||
|
||||
|
||||
class QuoteSourceAdapter(object):
|
||||
"""Big-QMT whole-quote source. ContextInfo-backed implementation lives in the
|
||||
server runtime; tests substitute a fake. ``subscribe`` must return a handle
|
||||
usable by ``unsubscribe``."""
|
||||
|
||||
def subscribe(self, codes, on_push):
|
||||
raise NotImplementedError
|
||||
|
||||
def unsubscribe(self, handle):
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class ContextInfoQuoteSource(QuoteSourceAdapter):
|
||||
"""Real big-QMT source backed by the strategy's ``ContextInfo``.
|
||||
|
||||
Verified against the real environment: ``ContextInfo.subscribe_whole_quote(
|
||||
code_list, callback)`` returns an int subscription id (``< 0`` on failure) and
|
||||
pushes INCREMENTAL ``{code: tick}`` batches on a dedicated quote thread;
|
||||
``ContextInfo.unsubscribe_quote(sub_id)`` cancels it.
|
||||
"""
|
||||
|
||||
def __init__(self, context_info):
|
||||
self._context = context_info
|
||||
|
||||
def subscribe(self, codes, on_push):
|
||||
sub_id = self._context.subscribe_whole_quote(list(codes), callback=on_push)
|
||||
if sub_id is None or int(sub_id) < 0:
|
||||
raise RuntimeError("ContextInfo.subscribe_whole_quote failed for codes=%s" % (list(codes),))
|
||||
return int(sub_id)
|
||||
|
||||
def unsubscribe(self, handle):
|
||||
try:
|
||||
self._context.unsubscribe_quote(handle)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
class _Combo(object):
|
||||
__slots__ = ("key", "codes", "handle", "topic", "clients")
|
||||
|
||||
def __init__(self, key, codes, handle, topic):
|
||||
self.key = key
|
||||
self.codes = codes
|
||||
self.handle = handle
|
||||
self.topic = topic
|
||||
# (client_id, sub_id) -> last_seen. Sub-id granularity: one client may
|
||||
# hold several subscriptions to the same combination, and each one keeps
|
||||
# the shared big-QMT subscription alive independently.
|
||||
self.clients = {} # (client_id, sub_id) -> last_seen timestamp
|
||||
|
||||
|
||||
class QuoteSubscriptionManager(object):
|
||||
def __init__(self, source, heartbeat_timeout_seconds=30.0, time_func=None, on_push_publisher=None, push_endpoint=""):
|
||||
self._source = source
|
||||
self._heartbeat_timeout = float(heartbeat_timeout_seconds)
|
||||
self._now = time_func or _monotonic
|
||||
# Optional callable(topic, data) invoked when big QMT pushes a tick batch.
|
||||
# Wired to the QuotePushChannel in a later stage; None keeps dispatch inert.
|
||||
self._on_push_publisher = on_push_publisher
|
||||
# Advertised to clients in subscribe responses so a zmq subscriber knows
|
||||
# where to connect (redis subscribers derive the channel locally instead).
|
||||
self._push_endpoint = str(push_endpoint or "")
|
||||
self._lock = threading.RLock()
|
||||
self._combos = {} # combo_key -> _Combo
|
||||
self._sub_index = {} # (client_id, sub_id) -> combo_key
|
||||
|
||||
# -- subscription lifecycle ---------------------------------------------
|
||||
def subscribe(self, client_id, sub_id, code_list):
|
||||
"""Register (client_id, sub_id) against its combination; create the shared
|
||||
big-QMT subscription on first use. Idempotent for replayed subscribes."""
|
||||
client_id = str(client_id or "")
|
||||
sub_id = str(sub_id or "")
|
||||
key = combo_key(code_list)
|
||||
now = self._now()
|
||||
|
||||
with self._lock:
|
||||
combo = self._combos.get(key)
|
||||
if combo is None:
|
||||
codes = sorted({str(c).strip().upper() for c in (code_list or []) if str(c or "").strip()})
|
||||
# source.subscribe registers the on_push callback with big QMT; it
|
||||
# does not call back into the manager, so it is safe under the lock.
|
||||
handle = self._source.subscribe(codes, self._make_on_push(key))
|
||||
combo = _Combo(key, codes, handle, key)
|
||||
self._combos[key] = combo
|
||||
|
||||
combo.clients[(client_id, sub_id)] = now
|
||||
self._sub_index[(client_id, sub_id)] = key
|
||||
return {"combo_key": key, "topic": combo.topic, "push_endpoint": self._push_endpoint}
|
||||
|
||||
def unsubscribe(self, client_id, sub_id):
|
||||
"""Drop (client_id, sub_id); tear the big-QMT subscription down when the
|
||||
last subscription of the combination leaves. Unknown sub_ids are a no-op."""
|
||||
client_id = str(client_id or "")
|
||||
sub_id = str(sub_id or "")
|
||||
with self._lock:
|
||||
key = self._sub_index.pop((client_id, sub_id), None)
|
||||
if key is None:
|
||||
return
|
||||
handle_to_close = self._remove_subscription_locked(key, client_id, sub_id)
|
||||
self._close_source(handle_to_close)
|
||||
|
||||
def keepalive(self, client_id, sub_id):
|
||||
"""Refresh last_seen for (client_id, sub_id). Unknown sub_ids are a no-op."""
|
||||
client_id = str(client_id or "")
|
||||
key = self._sub_index.get((client_id, str(sub_id or "")))
|
||||
if key is None:
|
||||
return
|
||||
with self._lock:
|
||||
combo = self._combos.get(key)
|
||||
if combo is None:
|
||||
return
|
||||
combo.clients[(client_id, str(sub_id or ""))] = self._now()
|
||||
|
||||
# -- reaper ---------------------------------------------------------------
|
||||
def reap_expired(self, now=None):
|
||||
"""Remove subscriptions silent for longer than the keepalive timeout;
|
||||
tear down combos that end up empty. Returns the number reaped."""
|
||||
now = self._now() if now is None else now
|
||||
reaped = 0
|
||||
handles_to_close = []
|
||||
with self._lock:
|
||||
for key in list(self._combos.keys()):
|
||||
combo = self._combos.get(key)
|
||||
if combo is None:
|
||||
continue
|
||||
for (client_id, sub_id), last_seen in list(combo.clients.items()):
|
||||
if now - last_seen > self._heartbeat_timeout:
|
||||
self._sub_index.pop((client_id, sub_id), None)
|
||||
handle = self._remove_subscription_locked(key, client_id, sub_id)
|
||||
if handle is not None:
|
||||
handles_to_close.append(handle)
|
||||
reaped += 1
|
||||
for handle in handles_to_close:
|
||||
self._close_source(handle)
|
||||
return reaped
|
||||
|
||||
# -- internals -------------------------------------------------------------
|
||||
def _make_on_push(self, key):
|
||||
def on_push(data):
|
||||
publisher = self._on_push_publisher
|
||||
if publisher is None:
|
||||
return
|
||||
with self._lock:
|
||||
combo = self._combos.get(key)
|
||||
topic = combo.topic if combo is not None else None
|
||||
if topic is None:
|
||||
return
|
||||
# Publish outside the lock: it is network IO and must not stall the
|
||||
# quote thread or block reaper/RPC threads waiting on the lock.
|
||||
publisher(topic, data)
|
||||
|
||||
return on_push
|
||||
|
||||
def _remove_subscription_locked(self, key, client_id, sub_id):
|
||||
"""Remove one (client_id, sub_id) from a combo. If the combo has no
|
||||
subscriptions left, detach it and return its source handle for the
|
||||
caller to close OUTSIDE the lock; else return None. Caller must hold
|
||||
the lock."""
|
||||
combo = self._combos.get(key)
|
||||
if combo is None:
|
||||
return None
|
||||
combo.clients.pop((client_id, sub_id), None)
|
||||
if combo.clients:
|
||||
return None
|
||||
self._combos.pop(key, None)
|
||||
return combo.handle
|
||||
|
||||
def _close_source(self, handle):
|
||||
if handle is None:
|
||||
return
|
||||
try:
|
||||
self._source.unsubscribe(handle)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _monotonic():
|
||||
import time
|
||||
|
||||
return time.monotonic()
|
||||
|
||||
|
||||
def build_quote_subscription_service(
|
||||
context_info,
|
||||
transport_name="redis",
|
||||
account_id="",
|
||||
redis_client=None,
|
||||
zmq_bind_address=None,
|
||||
enabled=True,
|
||||
heartbeat_timeout_seconds=30.0,
|
||||
time_func=None,
|
||||
):
|
||||
"""Assemble the server-side whole-quote service: a ContextInfo-backed source,
|
||||
a push channel matching the RPC transport, and a QuoteSubscriptionManager
|
||||
wired so big-QMT pushes publish to the channel. Returns ``(manager, channel)``
|
||||
or ``None`` when disabled. The caller starts the channel publisher and feeds
|
||||
``manager.reap_expired`` from the scheduler loop."""
|
||||
if not enabled:
|
||||
return None
|
||||
from .quote_push_channel import RedisQuotePushChannel, ZmqQuotePushChannel
|
||||
|
||||
source = ContextInfoQuoteSource(context_info)
|
||||
transport_name = str(transport_name or "redis").lower()
|
||||
if transport_name == "zmq":
|
||||
bind_address = zmq_bind_address or _default_quote_push_zmq_bind(account_id)
|
||||
channel = ZmqQuotePushChannel(bind_address=bind_address)
|
||||
push_endpoint = bind_address
|
||||
else:
|
||||
channel = RedisQuotePushChannel(redis_client, account_id=account_id)
|
||||
push_endpoint = ""
|
||||
manager = QuoteSubscriptionManager(
|
||||
source,
|
||||
heartbeat_timeout_seconds=heartbeat_timeout_seconds,
|
||||
time_func=time_func,
|
||||
on_push_publisher=channel.publish,
|
||||
push_endpoint=push_endpoint,
|
||||
)
|
||||
return manager, channel
|
||||
|
||||
|
||||
def _default_quote_push_zmq_bind(account_id):
|
||||
"""Default server PUB bind address: loopback, RPC zmq port + 1 (client side
|
||||
derives the same host/port + 1 to connect)."""
|
||||
from .transports.zmq_transport import _default_zmq_port
|
||||
|
||||
return "tcp://0.0.0.0:%d" % (_default_zmq_port(account_id) + 1)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,45 +0,0 @@
|
||||
"""信号执行前的轻量风控和数量计算。"""
|
||||
|
||||
from .code_utils import normalize_stock_code, round_buy_volume, round_sell_volume
|
||||
from .models import SignalAction, TradeSignal
|
||||
|
||||
|
||||
class RiskDecision:
|
||||
def __init__(self, allowed, reason="", volume=0, stock_code=""):
|
||||
self.allowed = allowed
|
||||
self.reason = reason
|
||||
self.volume = volume
|
||||
self.stock_code = stock_code
|
||||
|
||||
|
||||
def build_trade_volume(signal: TradeSignal, positions):
|
||||
code = normalize_stock_code(signal.stock_code)
|
||||
if signal.action == SignalAction.BUY:
|
||||
return RiskDecision(True, volume=round_buy_volume(code, signal.amount), stock_code=code)
|
||||
|
||||
if signal.action == SignalAction.SELL:
|
||||
position = positions.get(code)
|
||||
if not position or position.available <= 0:
|
||||
return RiskDecision(False, "no_available_position", stock_code=code)
|
||||
if signal.amount is not None:
|
||||
raw_volume = min(int(signal.amount), int(position.available))
|
||||
sell_all = raw_volume == int(position.available)
|
||||
else:
|
||||
pct = float(signal.percentage or 100)
|
||||
raw_volume = int(int(position.available) * pct / 100.0)
|
||||
sell_all = pct >= 100
|
||||
volume = round_sell_volume(code, raw_volume, sell_all=sell_all)
|
||||
if volume <= 0:
|
||||
return RiskDecision(False, "volume_below_min_lot", stock_code=code)
|
||||
return RiskDecision(True, volume=volume, stock_code=code)
|
||||
|
||||
return RiskDecision(False, f"unsupported_action:{signal.action}", stock_code=code)
|
||||
|
||||
|
||||
def validate_signal(signal, now, positions):
|
||||
if signal.is_expired(now):
|
||||
return RiskDecision(False, "expired", stock_code=signal.stock_code)
|
||||
decision = build_trade_volume(signal, positions)
|
||||
if decision.allowed and decision.volume <= 0:
|
||||
return RiskDecision(False, "invalid_volume", stock_code=decision.stock_code)
|
||||
return decision
|
||||
@@ -1,71 +0,0 @@
|
||||
"""大 QMT 运行文件可复用的转发入口。"""
|
||||
|
||||
import datetime as _dt
|
||||
import traceback
|
||||
|
||||
from .logging_setup import get_logger
|
||||
|
||||
_log = get_logger("runner")
|
||||
|
||||
_APP = None
|
||||
|
||||
|
||||
def reset_app():
|
||||
global _APP
|
||||
_APP = None
|
||||
|
||||
|
||||
def get_app():
|
||||
return _APP
|
||||
|
||||
|
||||
def init_app(context_info, app_factory):
|
||||
global _APP
|
||||
_APP = app_factory(context_info)
|
||||
if hasattr(_APP, "on_init"):
|
||||
_APP.on_init(context_info)
|
||||
return _APP
|
||||
|
||||
|
||||
def tick_app(context_info, now=None):
|
||||
if _APP is None:
|
||||
return None
|
||||
now = now or _dt.datetime.now()
|
||||
try:
|
||||
return _APP.tick(now)
|
||||
except Exception:
|
||||
_log.error("tick_app failed:\n%s", traceback.format_exc())
|
||||
return None
|
||||
|
||||
|
||||
def forward_order_event(event):
|
||||
# Unguarded events reach QMT's order_callback, which stops the strategy on
|
||||
# raise. Guard like tick_app so a bad event (e.g. redis outage during the
|
||||
# position-sync publish) never stops the strategy.
|
||||
if _APP is None:
|
||||
return None
|
||||
try:
|
||||
return _APP.on_order_event(event)
|
||||
except Exception:
|
||||
_log.error("forward_order_event failed:\n%s", traceback.format_exc())
|
||||
return None
|
||||
|
||||
|
||||
def forward_trade_event(event):
|
||||
if _APP is None:
|
||||
return None
|
||||
try:
|
||||
return _APP.on_trade_event(event)
|
||||
except Exception:
|
||||
_log.error("forward_trade_event failed:\n%s", traceback.format_exc())
|
||||
return None
|
||||
|
||||
|
||||
def sync_positions_app(reason="manual"):
|
||||
if _APP is None:
|
||||
return None
|
||||
try:
|
||||
return _APP.sync_positions(reason)
|
||||
except Exception:
|
||||
_log.error("sync_positions_app failed:\n%s", traceback.format_exc())
|
||||
return None
|
||||
@@ -1,19 +0,0 @@
|
||||
"""大 QMT 运行环境适配器骨架。"""
|
||||
|
||||
import datetime as _dt
|
||||
|
||||
|
||||
class BigQmtRuntimeAdapter:
|
||||
def __init__(self, context_info):
|
||||
self.context_info = context_info
|
||||
|
||||
def now(self):
|
||||
return _dt.datetime.now()
|
||||
|
||||
@staticmethod
|
||||
def to_order_event(order):
|
||||
return order
|
||||
|
||||
@staticmethod
|
||||
def to_trade_event(trade):
|
||||
return trade
|
||||
@@ -1,16 +0,0 @@
|
||||
"""Pluggable transport layer for the BigQMT RPC bridge.
|
||||
|
||||
The :class:`~bigqmt_signal_trader.transports.base.RpcTransport` interface owns
|
||||
the wire: how a request dict travels from the client to the QMT server and how
|
||||
the response dict travels back. ``redis`` is the reference implementation; the
|
||||
same business layer (handlers / ``process_request`` / ``to_jsonable``) runs
|
||||
unchanged over any transport.
|
||||
|
||||
Select a transport with ``rpc.transport`` in the config (default ``"redis"``).
|
||||
See :mod:`~bigqmt_signal_trader.transports.factory`.
|
||||
"""
|
||||
|
||||
from .base import RpcTransport, TransportError, TransportTimeout
|
||||
from .factory import build_transport
|
||||
|
||||
__all__ = ["RpcTransport", "TransportError", "TransportTimeout", "build_transport"]
|
||||
@@ -1,121 +0,0 @@
|
||||
"""Abstract transport interface for the BigQMT RPC bridge.
|
||||
|
||||
A transport owns the request/response wire. The business layer (handlers,
|
||||
``process_request``, ``to_jsonable``, ``enqueue_payload``, ``drain_pending``)
|
||||
is transport-agnostic; it only deals with request/response dicts.
|
||||
|
||||
Two roles, one interface
|
||||
------------------------
|
||||
* **Client side** — :meth:`RpcTransport.send_request`: send a request dict and
|
||||
block for the matching response dict (matched by ``request_id``).
|
||||
* **Server side** — :meth:`RpcTransport.start_receiving` registers a callback
|
||||
``on_request(request_dict)`` invoked per inbound request; the callback returns
|
||||
the response dict. :meth:`RpcTransport.send_response` delivers a response
|
||||
back to the client that sent ``request_dict`` (reply routing info is read
|
||||
from the request).
|
||||
|
||||
The request dict always carries the existing envelope (``schema_version``,
|
||||
``request_id``, ``account_id``, ``method``, ``params``). It MAY carry reply
|
||||
routing hints (``reply_key``/``reply_channel``/``reply_list``/``ttl_seconds``);
|
||||
Redis uses them, other transports may ignore them and use native routing.
|
||||
"""
|
||||
|
||||
|
||||
class TransportError(RuntimeError):
|
||||
"""A transport failed (connection lost, encode error, etc.)."""
|
||||
|
||||
|
||||
class TransportTimeout(TimeoutError):
|
||||
"""A request did not complete within the timeout window."""
|
||||
|
||||
|
||||
class RpcTransport(object):
|
||||
"""Abstract request/response transport. Concrete implementations own the wire.
|
||||
|
||||
Subclasses MUST override :meth:`send_request`,
|
||||
:meth:`start_receiving`, :meth:`send_response`, and :meth:`stop`.
|
||||
"""
|
||||
|
||||
name = "abstract"
|
||||
|
||||
def __init__(self, account_id="", print_prefix="[bigqmt_rpc]"):
|
||||
self.account_id = str(account_id or "")
|
||||
self.print_prefix = print_prefix
|
||||
self._on_request = None
|
||||
self._running = False
|
||||
|
||||
# -- client side -------------------------------------------------------
|
||||
def send_request(self, request, timeout_seconds):
|
||||
"""Send a request dict and block for the response dict.
|
||||
|
||||
``request`` is the full request envelope. Returns the response dict
|
||||
(with ``request_id`` matching). Raises :class:`TransportTimeout` if no
|
||||
response arrives within ``timeout_seconds``.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
# -- server side -------------------------------------------------------
|
||||
def start_receiving(self, on_request):
|
||||
"""Begin accepting inbound requests on the server side.
|
||||
|
||||
``on_request(request_dict)`` is invoked per inbound request and MUST
|
||||
return the response dict. Implementations may spawn a background
|
||||
thread. Safe to call once per transport instance.
|
||||
"""
|
||||
self._on_request = on_request
|
||||
self._running = True
|
||||
|
||||
def send_response(self, request, response):
|
||||
"""Deliver ``response`` back to the client that sent ``request``.
|
||||
|
||||
Reply routing is read from ``request`` (e.g. ``reply_key`` /
|
||||
``reply_channel`` / ``reply_list`` for Redis, or a native peer handle
|
||||
for ZMQ). Must be safe to call from the request-handling callback.
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
def stop(self):
|
||||
"""Stop receiving and release any sockets/connections/threads."""
|
||||
self._running = False
|
||||
self._on_request = None
|
||||
|
||||
def deliver(self, request):
|
||||
"""Internal: invoke the registered ``on_request`` callback.
|
||||
|
||||
Concrete transports call this when an inbound request arrives. If the
|
||||
callback returns a non-None response dict, it is delivered back to the
|
||||
client via :meth:`send_response` automatically — so a callback only
|
||||
needs to ``return response``. Callbacks that send the response
|
||||
themselves (e.g. the Redis service path, which routes through
|
||||
``_publish_response``) should return ``None`` to suppress the auto-send.
|
||||
|
||||
Handler exceptions are turned into an ``ok=False`` response envelope so
|
||||
the receive loop keeps running.
|
||||
"""
|
||||
callback = self._on_request
|
||||
if callback is None:
|
||||
return None
|
||||
try:
|
||||
response = callback(request)
|
||||
except Exception as exc: # noqa: BLE001 - transport must survive
|
||||
import datetime as _dt
|
||||
|
||||
response = {
|
||||
"schema_version": 1,
|
||||
"request_id": str((request or {}).get("request_id") or ""),
|
||||
"account_id": str((request or {}).get("account_id") or self.account_id or ""),
|
||||
"method": str((request or {}).get("method") or ""),
|
||||
"ok": False,
|
||||
"data": None,
|
||||
"error": "%s: %s" % (exc.__class__.__name__, exc),
|
||||
"handled_at": _dt.datetime.now().strftime("%Y-%m-%d %H:%M:%S"),
|
||||
}
|
||||
if response is not None:
|
||||
try:
|
||||
self.send_response(request, response)
|
||||
except Exception:
|
||||
pass
|
||||
return response
|
||||
|
||||
def __repr__(self):
|
||||
return "<%s account_id=%r>" % (self.__class__.__name__, self.account_id)
|
||||
@@ -1,126 +0,0 @@
|
||||
"""Transport factory: pick a transport backend by name.
|
||||
|
||||
``build_transport(name, config, ...)`` returns a ready transport instance.
|
||||
``name`` is the ``rpc.transport`` config value (default ``"redis"``). Unknown
|
||||
names raise :class:`ValueError`. Optional dependencies (``zmq``, a mysql
|
||||
driver) are imported lazily; a missing dependency surfaces as a clear
|
||||
``ImportError`` only when that transport is actually selected.
|
||||
"""
|
||||
|
||||
from .base import RpcTransport
|
||||
from .redis_transport import RedisTransport
|
||||
|
||||
|
||||
KNOWN_TRANSPORTS = ("redis", "zmq", "mysql", "shm")
|
||||
|
||||
|
||||
def build_transport(
|
||||
name,
|
||||
config=None,
|
||||
account_id="",
|
||||
print_prefix="[bigqmt_rpc]",
|
||||
):
|
||||
"""Construct a transport by name.
|
||||
|
||||
``config`` is the ``rpc`` config dict. Each backend reads its own sub-keys
|
||||
(``config["zmq"]``, ``config["mysql"]``); Redis reads the legacy keys
|
||||
(``request_channel_template`` etc.) plus ``redis_client``/
|
||||
``response_redis_client`` that the caller may inject.
|
||||
"""
|
||||
config = dict(config or {})
|
||||
name = str(name or "redis").lower()
|
||||
|
||||
if name in ("redis", "", "default"):
|
||||
return _build_redis(config, account_id, print_prefix)
|
||||
if name == "zmq":
|
||||
return _build_zmq(config, account_id, print_prefix)
|
||||
if name == "mysql":
|
||||
return _build_mysql(config, account_id, print_prefix)
|
||||
if name == "shm":
|
||||
return _build_shm(config, account_id, print_prefix)
|
||||
raise ValueError(
|
||||
"unknown rpc transport %r (known: %s)" % (name, ", ".join(KNOWN_TRANSPORTS))
|
||||
)
|
||||
|
||||
|
||||
def _build_redis(config, account_id, print_prefix):
|
||||
redis_client = config.get("redis_client")
|
||||
if redis_client is None:
|
||||
from ..adapters.redis_common import build_redis_client
|
||||
|
||||
redis_config = dict(config.get("redis") or {})
|
||||
redis_client = build_redis_client(redis_config)
|
||||
response_redis_client = config.get("response_redis_client")
|
||||
if response_redis_client is None:
|
||||
response_redis_client = redis_client
|
||||
return RedisTransport(
|
||||
redis_client,
|
||||
account_id=account_id,
|
||||
response_redis_client=response_redis_client,
|
||||
request_channel_template=config.get(
|
||||
"request_channel_template", "bigqmt:rpc:req:{account_id}"
|
||||
),
|
||||
request_queue_template=config.get(
|
||||
"request_queue_template", "bigqmt:rpc:queue:{account_id}"
|
||||
),
|
||||
response_channel_template=config.get(
|
||||
"response_channel_template", "bigqmt:rpc:resp:{account_id}:{request_id}"
|
||||
),
|
||||
response_list_template=config.get(
|
||||
"response_list_template", "bigqmt:rpc:respq:{account_id}:{request_id}"
|
||||
),
|
||||
response_key_template=config.get(
|
||||
"response_key_template", "bigqmt:rpc:resp:{account_id}:{request_id}"
|
||||
),
|
||||
response_ttl_seconds=int(config.get("response_ttl_seconds", 60)),
|
||||
queue_poll_interval_seconds=float(config.get("queue_poll_interval_seconds", 0.02)),
|
||||
debug_log_limit=int(config.get("debug_log_limit", 0)),
|
||||
print_prefix=print_prefix,
|
||||
)
|
||||
|
||||
|
||||
def _build_zmq(config, account_id, print_prefix):
|
||||
from .zmq_transport import ZmqTransport
|
||||
|
||||
zmq_config = dict(config.get("zmq") or {})
|
||||
# Wire up service discovery: if the caller injected a redis_client (server
|
||||
# side) or provided redis connection settings, the ZMQ transport can
|
||||
# publish/look up the actual bound port when the default port is taken.
|
||||
discovery_client = zmq_config.get("discovery_redis_client")
|
||||
if discovery_client is None and config.get("redis_client") is not None:
|
||||
discovery_client = config.get("redis_client")
|
||||
zmq_config["discovery_redis_client"] = discovery_client
|
||||
if discovery_client is None and config.get("redis"):
|
||||
# Build a small client just for discovery from the redis config block.
|
||||
try:
|
||||
from ..adapters.redis_common import build_redis_client
|
||||
|
||||
discovery_client = build_redis_client(dict(config.get("redis") or {}))
|
||||
zmq_config["discovery_redis_client"] = discovery_client
|
||||
except Exception:
|
||||
pass
|
||||
return ZmqTransport.from_config(
|
||||
zmq_config,
|
||||
account_id=account_id,
|
||||
print_prefix=print_prefix,
|
||||
)
|
||||
|
||||
|
||||
def _build_mysql(config, account_id, print_prefix):
|
||||
from .mysql_transport import MysqlTransport
|
||||
|
||||
return MysqlTransport.from_config(
|
||||
config.get("mysql") or {},
|
||||
account_id=account_id,
|
||||
print_prefix=print_prefix,
|
||||
)
|
||||
|
||||
|
||||
def _build_shm(config, account_id, print_prefix):
|
||||
from .shm_transport import SharedMemoryTransport
|
||||
|
||||
return SharedMemoryTransport(
|
||||
account_id=account_id,
|
||||
print_prefix=print_prefix,
|
||||
**dict(config.get("shm") or {})
|
||||
)
|
||||
@@ -1,462 +0,0 @@
|
||||
"""MySQL transport for the BigQMT RPC bridge.
|
||||
|
||||
A compatibility-oriented backend for environments where Redis/ZMQ are
|
||||
unavailable but a relational DB is. Latency is dominated by polling cadence,
|
||||
so this is NOT a low-latency path (expect tens of ms); use it when the
|
||||
deployment constraints rule out the others.
|
||||
|
||||
Schema (auto-created on first connect)::
|
||||
|
||||
CREATE TABLE bigqmt_rpc_requests (
|
||||
request_id VARCHAR(64) PRIMARY KEY,
|
||||
account_id VARCHAR(64) NOT NULL,
|
||||
payload MEDIUMTEXT NOT NULL,
|
||||
created_at DOUBLE NOT NULL,
|
||||
claimed_at DOUBLE NULL,
|
||||
INDEX idx_account_created (account_id, created_at)
|
||||
);
|
||||
CREATE TABLE bigqmt_rpc_responses (
|
||||
request_id VARCHAR(64) PRIMARY KEY,
|
||||
payload MEDIUMTEXT NOT NULL,
|
||||
created_at DOUBLE NOT NULL
|
||||
);
|
||||
|
||||
The client ``INSERT``s a request row and polls ``bigqmt_rpc_responses`` by
|
||||
``request_id``; the server ``SELECT ... FOR UPDATE SKIP LOCKED`` (or a
|
||||
``claimed_at`` flag on older engines) claims a request, invokes the handler,
|
||||
then ``INSERT``s the response row. Rows are cleaned up lazily by TTL.
|
||||
|
||||
Any DB-API 2.0 driver works (pymysql, mysql-connector, sqlite3 for tests).
|
||||
Pass ``driver="pymysql"`` / ``"sqlite3"`` etc. via config.
|
||||
"""
|
||||
|
||||
import json
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from ..adapters.redis_common import decode_text
|
||||
from ..redis_rpc import encode_rpc_request_payload, decode_rpc_request_payload
|
||||
from .base import RpcTransport, TransportError, TransportTimeout
|
||||
|
||||
|
||||
REQUESTS_TABLE = "bigqmt_rpc_requests"
|
||||
RESPONSES_TABLE = "bigqmt_rpc_responses"
|
||||
|
||||
|
||||
_SCHEMA = [
|
||||
"""CREATE TABLE IF NOT EXISTS {requests} (
|
||||
request_id VARCHAR(64) PRIMARY KEY,
|
||||
account_id VARCHAR(64) NOT NULL,
|
||||
payload MEDIUMTEXT NOT NULL,
|
||||
created_at DOUBLE NOT NULL,
|
||||
claimed_at DOUBLE NULL
|
||||
)""",
|
||||
"""CREATE TABLE IF NOT EXISTS {responses} (
|
||||
request_id VARCHAR(64) PRIMARY KEY,
|
||||
payload MEDIUMTEXT NOT NULL,
|
||||
created_at DOUBLE NOT NULL
|
||||
)""",
|
||||
"CREATE INDEX IF NOT EXISTS idx_{requests}_account_created ON {requests} (account_id, created_at)",
|
||||
]
|
||||
|
||||
|
||||
def _loads(raw):
|
||||
if isinstance(raw, dict):
|
||||
return dict(raw)
|
||||
text = decode_text(raw)
|
||||
text = decode_rpc_request_payload(text)
|
||||
return json.loads(text)
|
||||
|
||||
|
||||
class MysqlTransport(RpcTransport):
|
||||
"""Polling-based transport over a relational DB.
|
||||
|
||||
Both client and server open a short-lived connection per operation to keep
|
||||
the implementation driver-agnostic and avoid cross-thread cursor state.
|
||||
For high throughput a connection pool would help; this backend targets
|
||||
compatibility, not throughput.
|
||||
"""
|
||||
|
||||
name = "mysql"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
driver="pymysql",
|
||||
connect_kwargs=None,
|
||||
requests_table=REQUESTS_TABLE,
|
||||
responses_table=RESPONSES_TABLE,
|
||||
account_id="",
|
||||
print_prefix="[bigqmt_rpc]",
|
||||
poll_interval_seconds=0.02,
|
||||
row_ttl_seconds=120,
|
||||
background_threads=True,
|
||||
pool_config=None,
|
||||
use_pool=True,
|
||||
):
|
||||
super(MysqlTransport, self).__init__(account_id=account_id, print_prefix=print_prefix)
|
||||
self.driver = driver
|
||||
self.connect_kwargs = dict(connect_kwargs or {})
|
||||
self.requests_table = requests_table
|
||||
self.responses_table = responses_table
|
||||
self.poll_interval_seconds = max(0.001, float(poll_interval_seconds))
|
||||
self.row_ttl_seconds = int(row_ttl_seconds)
|
||||
self.background_threads = bool(background_threads)
|
||||
self._thread = None
|
||||
self._schema_ready = False
|
||||
self.use_pool = bool(use_pool)
|
||||
self.pool_config = dict(pool_config or {})
|
||||
self._pool = None
|
||||
# paramstyle: mysql drivers use "format" (%s), sqlite3 uses "qmark" (?).
|
||||
# Resolved lazily on first connect.
|
||||
self._placeholder = None
|
||||
|
||||
def _resolve_placeholder(self, mod):
|
||||
style = getattr(mod, "paramstyle", "format")
|
||||
if style == "qmark":
|
||||
return "?"
|
||||
return "%s" # format / pyformat / default
|
||||
|
||||
def _ph(self):
|
||||
# Return the placeholder char (resolving lazily).
|
||||
if self._placeholder is None:
|
||||
try:
|
||||
mod = __import__(self.driver)
|
||||
self._placeholder = self._resolve_placeholder(mod)
|
||||
except ImportError:
|
||||
self._placeholder = "%s"
|
||||
return self._placeholder
|
||||
|
||||
def _sql(self, template):
|
||||
"""Render a SQL template: fill {t}/{requests}/{responses} table names
|
||||
and swap the standard ``%s`` placeholder for the driver's paramstyle."""
|
||||
return template.format(
|
||||
t=None, # not used; callers format table names themselves
|
||||
requests=self.requests_table,
|
||||
responses=self.responses_table,
|
||||
).replace("__PH__", self._ph())
|
||||
|
||||
@classmethod
|
||||
def from_config(cls, config, account_id="", print_prefix="[bigqmt_rpc]"):
|
||||
config = dict(config or {})
|
||||
driver = config.get("driver", "pymysql")
|
||||
connect_kwargs = dict(config.get("connect_kwargs") or {})
|
||||
# Allow flat keys (host/port/user/...) as a convenience.
|
||||
for key in ("host", "port", "user", "password", "database", "charset"):
|
||||
if key in config and key not in connect_kwargs:
|
||||
connect_kwargs[key] = config[key]
|
||||
return cls(
|
||||
driver=driver,
|
||||
connect_kwargs=connect_kwargs,
|
||||
requests_table=config.get("requests_table", REQUESTS_TABLE),
|
||||
responses_table=config.get("responses_table", RESPONSES_TABLE),
|
||||
account_id=config.get("account_id", account_id),
|
||||
print_prefix=print_prefix,
|
||||
poll_interval_seconds=float(config.get("poll_interval_seconds", 0.02)),
|
||||
row_ttl_seconds=int(config.get("row_ttl_seconds", 120)),
|
||||
background_threads=bool(config.get("background_threads", True)),
|
||||
pool_config=config.get("pool_config"),
|
||||
use_pool=bool(config.get("use_pool", True)),
|
||||
)
|
||||
|
||||
# -- driver access / connection pool ----------------------------------
|
||||
def _import_driver(self):
|
||||
try:
|
||||
return __import__(self.driver)
|
||||
except ImportError as exc: # pragma: no cover - depends on env
|
||||
raise TransportError(
|
||||
"db driver %r is required for the mysql transport: %s"
|
||||
% (self.driver, exc)
|
||||
)
|
||||
|
||||
def _build_pool(self):
|
||||
"""Create a DBUtils PooledDB backed by the configured driver.
|
||||
|
||||
Works with any DB-API 2.0 driver (pymysql, mysql.connector, sqlite3,
|
||||
...). Pool sizing comes from ``pool_config``; connection kwargs are
|
||||
forwarded to the driver's ``connect()``.
|
||||
"""
|
||||
try:
|
||||
from dbutils.pooled_db import PooledDB
|
||||
except ImportError as exc:
|
||||
raise TransportError(
|
||||
"DBUtils is required for the mysql transport connection pool: %s" % exc
|
||||
)
|
||||
driver = self._import_driver()
|
||||
cfg = dict(self.pool_config)
|
||||
# Sensible defaults for an RPC workload: small idle pool, modest cap,
|
||||
# reuse connections across threads. Callers override via pool_config.
|
||||
mincached = cfg.pop("mincached", 1)
|
||||
maxcached = cfg.pop("maxcached", 4)
|
||||
maxshared = cfg.pop("maxshared", 3)
|
||||
maxconnections = cfg.pop("maxconnections", 8)
|
||||
blocking = cfg.pop("blocking", True)
|
||||
maxusage = cfg.pop("maxusage", 0)
|
||||
reset = cfg.pop("reset", True)
|
||||
# Whatever remains in cfg is treated as extra creator kwargs (e.g.
|
||||
# ping, setsession) and merged under the connect kwargs.
|
||||
extra = cfg
|
||||
connect_kwargs = self._pooled_connect_args()
|
||||
connect_kwargs.update(extra)
|
||||
return PooledDB(
|
||||
creator=driver,
|
||||
mincached=mincached,
|
||||
maxcached=maxcached,
|
||||
maxshared=maxshared,
|
||||
maxconnections=maxconnections,
|
||||
blocking=blocking,
|
||||
maxusage=maxusage,
|
||||
reset=reset,
|
||||
**connect_kwargs
|
||||
)
|
||||
|
||||
def _pooled_connect_args(self):
|
||||
"""Return the kwargs to forward to the driver's connect().
|
||||
|
||||
Stripped of empty credential fields so drivers that reject empty
|
||||
username/password (e.g. pymysql with auth plugin) don't choke.
|
||||
"""
|
||||
cfg = dict(self.connect_kwargs)
|
||||
if not cfg.get("user") and "user" in cfg:
|
||||
cfg.pop("user")
|
||||
if not cfg.get("password") and "password" in cfg:
|
||||
cfg.pop("password")
|
||||
return cfg
|
||||
|
||||
def _connect(self):
|
||||
if not self.use_pool:
|
||||
return self._import_driver().connect(**self.connect_kwargs)
|
||||
if self._pool is None:
|
||||
self._pool = self._build_pool()
|
||||
# PooledDB.connection() hands out a pooled connection; calling .close()
|
||||
# on it returns it to the pool rather than closing the underlying socket.
|
||||
return self._pool.connection()
|
||||
|
||||
def _ensure_schema(self):
|
||||
if self._schema_ready:
|
||||
return
|
||||
ctx = {"requests": self.requests_table, "responses": self.responses_table}
|
||||
conn = self._connect()
|
||||
try:
|
||||
cur = conn.cursor()
|
||||
for stmt in _SCHEMA:
|
||||
try:
|
||||
cur.execute(stmt.format(**ctx))
|
||||
except Exception:
|
||||
# "CREATE INDEX IF NOT EXISTS" is not supported on some
|
||||
# MySQL versions; the index is an optimization, ignore failure.
|
||||
pass
|
||||
conn.commit()
|
||||
self._schema_ready = True
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _now(self):
|
||||
return time.time()
|
||||
|
||||
# -- client side ------------------------------------------------------
|
||||
def send_request(self, request, timeout_seconds):
|
||||
self._ensure_schema()
|
||||
request = dict(request)
|
||||
request.setdefault("request_id", uuid.uuid4().hex)
|
||||
request_id = str(request["request_id"])
|
||||
request.setdefault("account_id", self.account_id)
|
||||
payload = encode_rpc_request_payload(request)
|
||||
conn = self._connect()
|
||||
try:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
self._sql(
|
||||
"INSERT INTO {requests} (request_id, account_id, payload, created_at, claimed_at) "
|
||||
"VALUES (__PH__, __PH__, __PH__, __PH__, NULL)"
|
||||
),
|
||||
(request_id, str(request.get("account_id") or self.account_id), payload, self._now()),
|
||||
)
|
||||
conn.commit()
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
deadline = time.time() + float(timeout_seconds)
|
||||
while time.time() < deadline:
|
||||
conn = self._connect()
|
||||
try:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
self._sql("SELECT payload FROM {responses} WHERE request_id = __PH__"),
|
||||
(request_id,),
|
||||
)
|
||||
row = cur.fetchone()
|
||||
if row:
|
||||
payload = row[0]
|
||||
try:
|
||||
cur.execute(
|
||||
self._sql("DELETE FROM {responses} WHERE request_id = __PH__"),
|
||||
(request_id,),
|
||||
)
|
||||
conn.commit()
|
||||
except Exception:
|
||||
pass
|
||||
return _loads(payload)
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
time.sleep(self.poll_interval_seconds)
|
||||
raise TransportTimeout("mysql rpc timeout: %s" % request.get("method"))
|
||||
|
||||
# -- server side ------------------------------------------------------
|
||||
def start_receiving(self, on_request, background_threads=None):
|
||||
super(MysqlTransport, self).start_receiving(on_request)
|
||||
self._ensure_schema()
|
||||
if background_threads is None:
|
||||
background_threads = self.background_threads
|
||||
if not background_threads:
|
||||
print("%s mysql polling table=%s background_threads=False" % (
|
||||
self.print_prefix, self.requests_table))
|
||||
return
|
||||
self._thread = threading.Thread(
|
||||
target=self._poll_loop, name="bigqmt-mysql-rpc", daemon=True
|
||||
)
|
||||
self._thread.start()
|
||||
print("%s mysql started polling table=%s" % (
|
||||
self.print_prefix, self.requests_table))
|
||||
|
||||
def _poll_loop(self):
|
||||
while self._running:
|
||||
try:
|
||||
self._claim_and_handle_batch()
|
||||
except Exception as exc:
|
||||
if not self._running:
|
||||
break
|
||||
print("%s mysql poll failed: %s" % (self.print_prefix, exc))
|
||||
time.sleep(0.5)
|
||||
continue
|
||||
time.sleep(self.poll_interval_seconds)
|
||||
|
||||
def _claim_and_handle_batch(self, max_items=20):
|
||||
conn = self._connect()
|
||||
claimed = []
|
||||
try:
|
||||
cur = conn.cursor()
|
||||
# Claim rows: mark claimed_at so concurrent servers skip them.
|
||||
# Uses an atomic UPDATE ... WHERE claimed_at IS NULL with a cap.
|
||||
cur.execute(
|
||||
self._sql(
|
||||
"SELECT request_id, payload FROM {requests} WHERE account_id = __PH__ "
|
||||
"AND claimed_at IS NULL ORDER BY created_at ASC LIMIT __PH__"
|
||||
),
|
||||
(self.account_id, int(max_items)),
|
||||
)
|
||||
rows = cur.fetchall()
|
||||
now = self._now()
|
||||
for request_id, payload in rows:
|
||||
cur.execute(
|
||||
self._sql(
|
||||
"UPDATE {requests} SET claimed_at = __PH__ WHERE request_id = __PH__ "
|
||||
"AND claimed_at IS NULL"
|
||||
),
|
||||
(now, request_id),
|
||||
)
|
||||
if cur.rowcount > 0:
|
||||
claimed.append((request_id, payload))
|
||||
conn.commit()
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
for request_id, payload in claimed:
|
||||
try:
|
||||
request = _loads(payload)
|
||||
request["request_id"] = request_id
|
||||
except Exception as exc:
|
||||
print("%s mysql decode failed: %s" % (self.print_prefix, exc))
|
||||
self._delete_request(request_id)
|
||||
continue
|
||||
try:
|
||||
self.deliver(request)
|
||||
except Exception as exc:
|
||||
print("%s mysql deliver failed: %s" % (self.print_prefix, exc))
|
||||
self._delete_request(request_id)
|
||||
|
||||
def _delete_request(self, request_id):
|
||||
conn = self._connect()
|
||||
try:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
self._sql("DELETE FROM {requests} WHERE request_id = __PH__"),
|
||||
(request_id,),
|
||||
)
|
||||
conn.commit()
|
||||
except Exception:
|
||||
pass
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def send_response(self, request, response):
|
||||
request_id = str(
|
||||
response.get("request_id") or request.get("request_id") or ""
|
||||
)
|
||||
payload = encode_rpc_request_payload(response)
|
||||
# DELETE-then-INSERT is portable across MySQL and sqlite (avoids the
|
||||
# MySQL-only ON DUPLICATE KEY / REPLACE syntax). One connection, one txn.
|
||||
conn = self._connect()
|
||||
try:
|
||||
cur = conn.cursor()
|
||||
cur.execute(
|
||||
self._sql("DELETE FROM {responses} WHERE request_id = __PH__"),
|
||||
(request_id,),
|
||||
)
|
||||
cur.execute(
|
||||
self._sql(
|
||||
"INSERT INTO {responses} (request_id, payload, created_at) "
|
||||
"VALUES (__PH__, __PH__, __PH__)"
|
||||
),
|
||||
(request_id, payload, self._now()),
|
||||
)
|
||||
conn.commit()
|
||||
except Exception as exc:
|
||||
try:
|
||||
conn.rollback()
|
||||
except Exception:
|
||||
pass
|
||||
print("%s mysql response write failed: %s" % (self.print_prefix, exc))
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# -- non-background drain (strategy adjust thread) --------------------
|
||||
def drain_request_queue(self, max_items=20):
|
||||
if not self._running:
|
||||
return 0
|
||||
before = 0 # _claim_and_handle_batch handles its own count
|
||||
self._claim_and_handle_batch(max_items=max_items)
|
||||
return 0
|
||||
|
||||
def stop(self):
|
||||
super(MysqlTransport, self).stop()
|
||||
if self._thread is not None and self._thread.is_alive():
|
||||
self._thread.join(1.0)
|
||||
self._thread = None
|
||||
# Close the connection pool so background connections are released.
|
||||
if self._pool is not None:
|
||||
try:
|
||||
self._pool.close()
|
||||
except Exception:
|
||||
pass
|
||||
self._pool = None
|
||||
@@ -1,424 +0,0 @@
|
||||
"""Redis transport for the BigQMT RPC bridge.
|
||||
|
||||
This is the reference transport and the default. It preserves the exact wire
|
||||
behavior of the original ``RedisPubSubRpcService``:
|
||||
|
||||
* Client ``send_request``: ``RPUSH`` the (base64-obfuscated) request onto the
|
||||
per-account request queue, then ``BLPOP`` the per-request response list with
|
||||
a ``GET response_key`` fallback. A ``pubsub`` transport variant is kept for
|
||||
callers that pass ``transport="pubsub"`` to ``call_redis_rpc``.
|
||||
* Server receive: two background loops — a ``pubsub.subscribe`` loop and a
|
||||
``brpop`` queue loop. Either delivers inbound payloads to the registered
|
||||
``on_request`` callback.
|
||||
* Server ``send_response``: fan-out writes to ``reply_key`` (``SETEX``),
|
||||
``reply_list`` (``RPUSH`` + ``EXPIRE``) and ``reply_channel`` (``PUBLISH``).
|
||||
|
||||
The module-level :func:`call_redis_rpc` helper keeps its original signature and
|
||||
delegates here so existing callers and ``bench_latency.py`` are unchanged.
|
||||
"""
|
||||
|
||||
import threading
|
||||
import time
|
||||
import traceback
|
||||
import uuid
|
||||
|
||||
from ..adapters.redis_common import decode_text
|
||||
from ..redis_rpc import (
|
||||
decode_rpc_request_payload,
|
||||
encode_rpc_request_payload,
|
||||
)
|
||||
from .base import RpcTransport, TransportTimeout
|
||||
|
||||
import json # noqa: E402 (kept here so transport owns all wire encoding)
|
||||
|
||||
|
||||
REQUEST_CHANNEL_TEMPLATE = "bigqmt:rpc:req:{account_id}"
|
||||
REQUEST_QUEUE_TEMPLATE = "bigqmt:rpc:queue:{account_id}"
|
||||
RESPONSE_CHANNEL_TEMPLATE = "bigqmt:rpc:resp:{account_id}:{request_id}"
|
||||
RESPONSE_LIST_TEMPLATE = "bigqmt:rpc:respq:{account_id}:{request_id}"
|
||||
RESPONSE_KEY_TEMPLATE = "bigqmt:rpc:resp:{account_id}:{request_id}"
|
||||
|
||||
|
||||
def _format(template, account_id, request_id):
|
||||
if not template:
|
||||
return ""
|
||||
return template.format(account_id=account_id, request_id=request_id)
|
||||
|
||||
|
||||
def _loads(raw_payload):
|
||||
"""Decode a wire payload (bytes/str/dict) into a request dict."""
|
||||
if isinstance(raw_payload, dict):
|
||||
return dict(raw_payload)
|
||||
text = decode_text(raw_payload)
|
||||
text = decode_rpc_request_payload(text)
|
||||
payload = json.loads(text)
|
||||
if not isinstance(payload, dict):
|
||||
raise ValueError("rpc payload must be a json object")
|
||||
return payload
|
||||
|
||||
|
||||
def _is_redis_timeout(exc):
|
||||
name = exc.__class__.__name__.lower()
|
||||
module = getattr(exc.__class__, "__module__", "")
|
||||
text = str(exc).lower()
|
||||
return ("redis" in module and "timeout" in name) or "timeout reading from socket" in text
|
||||
|
||||
|
||||
class RedisTransport(RpcTransport):
|
||||
"""Redis-backed transport. Owns rpush/blpop/brpop/publish/setex."""
|
||||
|
||||
name = "redis"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
redis_client,
|
||||
account_id="",
|
||||
response_redis_client=None,
|
||||
request_channel_template=REQUEST_CHANNEL_TEMPLATE,
|
||||
request_queue_template=REQUEST_QUEUE_TEMPLATE,
|
||||
response_channel_template=RESPONSE_CHANNEL_TEMPLATE,
|
||||
response_list_template=RESPONSE_LIST_TEMPLATE,
|
||||
response_key_template=RESPONSE_KEY_TEMPLATE,
|
||||
response_ttl_seconds=60,
|
||||
queue_poll_interval_seconds=0.02,
|
||||
debug_log_limit=0,
|
||||
print_prefix="[bigqmt_rpc]",
|
||||
):
|
||||
super(RedisTransport, self).__init__(account_id=account_id, print_prefix=print_prefix)
|
||||
self.listen_redis = redis_client
|
||||
self.redis = response_redis_client or redis_client
|
||||
self.request_channel_template = request_channel_template
|
||||
self.request_queue_template = request_queue_template
|
||||
self.response_channel_template = response_channel_template
|
||||
self.response_list_template = response_list_template
|
||||
self.response_key_template = response_key_template
|
||||
self.response_ttl_seconds = int(response_ttl_seconds)
|
||||
self.queue_poll_interval_seconds = max(0.001, float(queue_poll_interval_seconds))
|
||||
self.debug_log_limit = int(debug_log_limit)
|
||||
self._received_count = 0
|
||||
self._published_count = 0
|
||||
self._pubsub = None
|
||||
self._thread = None
|
||||
self._queue_thread = None
|
||||
# Hooks so the service can observe/intercept received payloads (debug
|
||||
# logging, inline-vs-deferred dispatch). When None, the request is
|
||||
# delivered straight to the on_request callback.
|
||||
self.on_raw_payload = None
|
||||
|
||||
# -- properties mirroring the original service -------------------------
|
||||
@property
|
||||
def request_channel(self):
|
||||
return self.request_channel_template.format(account_id=self.account_id)
|
||||
|
||||
@property
|
||||
def request_queue(self):
|
||||
return self.request_queue_template.format(account_id=self.account_id)
|
||||
|
||||
def _response_clients(self):
|
||||
clients = [self.redis]
|
||||
if self.listen_redis is not self.redis:
|
||||
clients.append(self.listen_redis)
|
||||
return clients
|
||||
|
||||
# -- client side -------------------------------------------------------
|
||||
def send_request(self, request, timeout_seconds, transport="queue"):
|
||||
"""Send ``request`` and block for the response dict.
|
||||
|
||||
``transport`` selects the Redis sub-transport: ``"queue"`` (default,
|
||||
RPUSH+BLPOP) or ``"pubsub"`` (PUBLISH+subscribe). Kept for parity with
|
||||
the original ``call_redis_rpc`` signature.
|
||||
"""
|
||||
return _call_redis_rpc(
|
||||
self.listen_redis,
|
||||
self.account_id,
|
||||
request,
|
||||
timeout_seconds=float(timeout_seconds),
|
||||
transport=transport,
|
||||
request_channel_template=self.request_channel_template,
|
||||
request_queue_template=self.request_queue_template,
|
||||
response_channel_template=self.response_channel_template,
|
||||
response_list_template=self.response_list_template,
|
||||
response_key_template=self.response_key_template,
|
||||
)
|
||||
|
||||
# -- server side -------------------------------------------------------
|
||||
def start_receiving(self, on_request, background_threads=True):
|
||||
"""Spawn the pubsub + queue receive loops (unless ``background_threads``)."""
|
||||
super(RedisTransport, self).start_receiving(on_request)
|
||||
if not background_threads:
|
||||
print(
|
||||
"%s started queue=%s background_threads=False"
|
||||
% (self.print_prefix, self.request_queue)
|
||||
)
|
||||
return
|
||||
if (
|
||||
self._thread is not None
|
||||
and self._thread.is_alive()
|
||||
and self._queue_thread is not None
|
||||
and self._queue_thread.is_alive()
|
||||
):
|
||||
return
|
||||
self._thread = threading.Thread(
|
||||
target=self._listen_loop, name="bigqmt-redis-rpc", daemon=True
|
||||
)
|
||||
self._queue_thread = threading.Thread(
|
||||
target=self._queue_loop, name="bigqmt-redis-rpc-queue", daemon=True
|
||||
)
|
||||
self._thread.start()
|
||||
self._queue_thread.start()
|
||||
print(
|
||||
"%s started channel=%s queue=%s"
|
||||
% (self.print_prefix, self.request_channel, self.request_queue)
|
||||
)
|
||||
|
||||
def _listen_loop(self):
|
||||
while self._running:
|
||||
try:
|
||||
pubsub = self.listen_redis.pubsub(ignore_subscribe_messages=True)
|
||||
self._pubsub = pubsub
|
||||
pubsub.subscribe(self.request_channel)
|
||||
if self.debug_log_limit > 0:
|
||||
print(
|
||||
"%s subscribed channel=%s" % (self.print_prefix, self.request_channel)
|
||||
)
|
||||
while self._running:
|
||||
message = pubsub.get_message(timeout=1.0)
|
||||
if not self._running:
|
||||
break
|
||||
if not message or message.get("type") != "message":
|
||||
continue
|
||||
self._handle_received_payload(message.get("data"), "pubsub")
|
||||
except Exception:
|
||||
print(
|
||||
"%s listener failed:\n%s" % (self.print_prefix, traceback.format_exc())
|
||||
)
|
||||
time.sleep(1.0)
|
||||
finally:
|
||||
try:
|
||||
if self._pubsub is not None:
|
||||
self._pubsub.close()
|
||||
except Exception:
|
||||
pass
|
||||
self._pubsub = None
|
||||
|
||||
def _queue_loop(self):
|
||||
while self._running:
|
||||
try:
|
||||
if self.debug_log_limit > 0:
|
||||
print(
|
||||
"%s queue polling key=%s" % (self.print_prefix, self.request_queue)
|
||||
)
|
||||
while self._running:
|
||||
item = self.listen_redis.brpop(self.request_queue, timeout=1)
|
||||
if not self._running:
|
||||
break
|
||||
if not item:
|
||||
continue
|
||||
raw = (
|
||||
item[1]
|
||||
if isinstance(item, (list, tuple)) and len(item) >= 2
|
||||
else item
|
||||
)
|
||||
self._handle_received_payload(raw, "queue")
|
||||
except Exception:
|
||||
print(
|
||||
"%s queue listener failed:\n%s"
|
||||
% (self.print_prefix, traceback.format_exc())
|
||||
)
|
||||
time.sleep(1.0)
|
||||
|
||||
def _handle_received_payload(self, raw_payload, source):
|
||||
self._received_count += 1
|
||||
if self.on_raw_payload is not None:
|
||||
# Service wants to observe/intercept (e.g. debug log + dispatch fork).
|
||||
self.on_raw_payload(raw_payload, source)
|
||||
return
|
||||
# Default: decode and deliver straight to the registered callback.
|
||||
request = _loads(raw_payload)
|
||||
self.deliver(request)
|
||||
|
||||
def send_response(self, request, response):
|
||||
"""Fan out the response to reply_key/reply_list/reply_channel."""
|
||||
request_id = response.get("request_id") or request.get("request_id") or ""
|
||||
account_id = response.get("account_id") or request.get("account_id") or self.account_id
|
||||
payload = json.dumps(response, ensure_ascii=False)
|
||||
ttl_seconds = int(request.get("ttl_seconds") or self.response_ttl_seconds)
|
||||
response_key = request.get("reply_key") or _format(
|
||||
self.response_key_template, account_id, request_id
|
||||
)
|
||||
response_channel = request.get("reply_channel") or _format(
|
||||
self.response_channel_template, account_id, request_id
|
||||
)
|
||||
response_list = request.get("reply_list")
|
||||
if response_key:
|
||||
self._write_response_key(response_key, ttl_seconds, payload)
|
||||
if response_list:
|
||||
self._push_response_list(response_list, ttl_seconds, payload)
|
||||
if response_channel:
|
||||
self._publish_response_channel(response_channel, payload)
|
||||
|
||||
def _write_response_key(self, response_key, ttl_seconds, payload):
|
||||
first_error = None
|
||||
wrote = 0
|
||||
for client in self._response_clients():
|
||||
try:
|
||||
if ttl_seconds > 0:
|
||||
client.setex(response_key, ttl_seconds, payload)
|
||||
else:
|
||||
client.set(response_key, payload)
|
||||
wrote += 1
|
||||
except Exception as exc:
|
||||
if first_error is None:
|
||||
first_error = exc
|
||||
if wrote <= 0 and first_error is not None:
|
||||
raise first_error
|
||||
return wrote
|
||||
|
||||
def _push_response_list(self, response_list, ttl_seconds, payload):
|
||||
first_error = None
|
||||
pushed = 0
|
||||
for client in self._response_clients():
|
||||
try:
|
||||
client.rpush(response_list, payload)
|
||||
if ttl_seconds > 0:
|
||||
client.expire(response_list, ttl_seconds)
|
||||
pushed += 1
|
||||
except Exception as exc:
|
||||
if first_error is None:
|
||||
first_error = exc
|
||||
if pushed <= 0 and first_error is not None:
|
||||
raise first_error
|
||||
return pushed
|
||||
|
||||
def _publish_response_channel(self, response_channel, payload):
|
||||
first_error = None
|
||||
receivers = 0
|
||||
published = 0
|
||||
for client in self._response_clients():
|
||||
try:
|
||||
receivers += int(client.publish(response_channel, payload) or 0)
|
||||
published += 1
|
||||
except Exception as exc:
|
||||
if first_error is None:
|
||||
first_error = exc
|
||||
if published <= 0 and first_error is not None:
|
||||
raise first_error
|
||||
self._published_count += 1
|
||||
if self._published_count <= self.debug_log_limit:
|
||||
print("%s published response receivers=%s" % (self.print_prefix, receivers))
|
||||
return receivers
|
||||
|
||||
# -- non-background drain helpers (used by the strategy adjust thread) -
|
||||
def drain_request_queue(self, max_items=20):
|
||||
processed = 0
|
||||
for _ in range(int(max_items)):
|
||||
try:
|
||||
item = self.listen_redis.lpop(self.request_queue)
|
||||
except Exception as exc:
|
||||
if _is_redis_timeout(exc):
|
||||
print("%s ERROR drain timeout on LPOP queue=%s; skip this tick" % (self.print_prefix, self.request_queue))
|
||||
break
|
||||
raise
|
||||
if not item:
|
||||
break
|
||||
if self.on_raw_payload is not None:
|
||||
self.on_raw_payload(item, "queue-drain")
|
||||
else:
|
||||
self.deliver(_loads(item))
|
||||
processed += 1
|
||||
return processed
|
||||
|
||||
def stop(self):
|
||||
super(RedisTransport, self).stop()
|
||||
pubsub = self._pubsub
|
||||
if pubsub is not None:
|
||||
try:
|
||||
pubsub.close()
|
||||
except Exception:
|
||||
pass
|
||||
thread = self._thread
|
||||
if thread is not None and thread.is_alive():
|
||||
thread.join(1.0)
|
||||
queue_thread = self._queue_thread
|
||||
if queue_thread is not None and queue_thread.is_alive():
|
||||
queue_thread.join(1.0)
|
||||
self._thread = None
|
||||
self._queue_thread = None
|
||||
self._pubsub = None
|
||||
|
||||
|
||||
def _call_redis_rpc(
|
||||
redis_client,
|
||||
account_id,
|
||||
request,
|
||||
timeout_seconds=3.0,
|
||||
transport="queue",
|
||||
request_channel_template=REQUEST_CHANNEL_TEMPLATE,
|
||||
request_queue_template=REQUEST_QUEUE_TEMPLATE,
|
||||
response_channel_template=RESPONSE_CHANNEL_TEMPLATE,
|
||||
response_list_template=RESPONSE_LIST_TEMPLATE,
|
||||
response_key_template=RESPONSE_KEY_TEMPLATE,
|
||||
ttl_seconds=60,
|
||||
):
|
||||
"""Client-side round trip. Accepts a pre-built request envelope."""
|
||||
request_id = request.get("request_id") or uuid.uuid4().hex
|
||||
request_channel = request_channel_template.format(account_id=account_id)
|
||||
request_queue = request_queue_template.format(account_id=account_id)
|
||||
response_channel = response_channel_template.format(
|
||||
account_id=account_id, request_id=request_id
|
||||
)
|
||||
response_list = response_list_template.format(account_id=account_id, request_id=request_id)
|
||||
response_key = response_key_template.format(account_id=account_id, request_id=request_id)
|
||||
# Ensure reply routing is present (the original helper filled these in).
|
||||
request = dict(request)
|
||||
request.setdefault("request_id", request_id)
|
||||
request.setdefault("reply_channel", response_channel)
|
||||
request.setdefault("reply_list", response_list)
|
||||
request.setdefault("reply_key", response_key)
|
||||
request.setdefault("ttl_seconds", ttl_seconds)
|
||||
request["request_id"] = request_id
|
||||
payload = encode_rpc_request_payload(request)
|
||||
|
||||
if str(transport or "queue").lower() in ("queue", "list", "blpop"):
|
||||
redis_client.rpush(request_queue, payload)
|
||||
redis_client.expire(request_queue, max(60, int(ttl_seconds)))
|
||||
wait_timeout = max(1, int(float(timeout_seconds) + 0.999))
|
||||
item = redis_client.blpop(response_list, timeout=wait_timeout)
|
||||
if item:
|
||||
raw_response = (
|
||||
item[1] if isinstance(item, (list, tuple)) and len(item) >= 2 else item
|
||||
)
|
||||
try:
|
||||
redis_client.delete(response_list)
|
||||
except Exception:
|
||||
pass
|
||||
return json.loads(decode_text(raw_response))
|
||||
raw_response = redis_client.get(response_key)
|
||||
if raw_response:
|
||||
return json.loads(decode_text(raw_response))
|
||||
raise TransportTimeout("redis rpc timeout: %s" % request.get("method"))
|
||||
|
||||
pubsub = redis_client.pubsub(ignore_subscribe_messages=True)
|
||||
try:
|
||||
pubsub.subscribe(response_channel)
|
||||
redis_client.publish(request_channel, payload)
|
||||
deadline = time.time() + float(timeout_seconds)
|
||||
while True:
|
||||
remaining = deadline - time.time()
|
||||
if remaining <= 0:
|
||||
break
|
||||
message = pubsub.get_message(timeout=remaining)
|
||||
if not message or message.get("type") != "message":
|
||||
continue
|
||||
response = json.loads(decode_text(message.get("data")))
|
||||
if response.get("request_id") == request_id:
|
||||
return response
|
||||
raw_response = redis_client.get(response_key)
|
||||
if raw_response:
|
||||
return json.loads(decode_text(raw_response))
|
||||
raise TransportTimeout("redis rpc timeout: %s" % request.get("method"))
|
||||
finally:
|
||||
try:
|
||||
pubsub.close()
|
||||
except Exception:
|
||||
pass
|
||||
@@ -1,34 +0,0 @@
|
||||
"""Shared-memory transport stub.
|
||||
|
||||
Reserved for a future low-latency same-host backend. Not implemented because
|
||||
the QMT runtime ships Python 3.6, where ``multiprocessing.shared_memory`` is
|
||||
unavailable (added in 3.8). A ``mmap``-plus-named-mutex implementation is
|
||||
possible but non-trivial; until it lands, selecting this transport raises a
|
||||
clear error so misconfiguration fails fast.
|
||||
"""
|
||||
|
||||
from .base import RpcTransport, TransportError
|
||||
|
||||
|
||||
class SharedMemoryTransport(RpcTransport):
|
||||
name = "shm"
|
||||
|
||||
def __init__(self, account_id="", print_prefix="[bigqmt_rpc]", **kwargs):
|
||||
super(SharedMemoryTransport, self).__init__(
|
||||
account_id=account_id, print_prefix=print_prefix
|
||||
)
|
||||
|
||||
def _unsupported(self):
|
||||
raise TransportError(
|
||||
"shared-memory transport is not implemented yet "
|
||||
"(requires Python 3.8+ shared_memory or a custom mmap ring buffer)"
|
||||
)
|
||||
|
||||
def send_request(self, request, timeout_seconds):
|
||||
self._unsupported()
|
||||
|
||||
def send_response(self, request, response):
|
||||
self._unsupported()
|
||||
|
||||
def start_receiving(self, on_request, **kwargs):
|
||||
self._unsupported()
|
||||
@@ -1,459 +0,0 @@
|
||||
"""ZeroMQ transport for the BigQMT RPC bridge.
|
||||
|
||||
Designed for same-host low latency. Topology:
|
||||
|
||||
* **Server** binds a ``ROUTER`` socket. Each inbound message arrives as
|
||||
``[identity, payload]``; the server remembers ``identity`` keyed by
|
||||
``request_id`` and replies with ``[identity, payload]`` so ZMQ routes the
|
||||
response back to the originating client automatically.
|
||||
* **Client** connects a ``DEALER`` socket (with a unique random identity), sends
|
||||
``[payload]``, then ``poll``/``recv`` for the response. DEALER gives each
|
||||
client an asymmetric async path that pairs naturally with ROUTER.
|
||||
|
||||
Wire framing is a single JSON payload per message. The original b64 stock-code
|
||||
obfuscation (``encode_rpc_request_payload``) is applied too, so payloads stay
|
||||
opaque even though ZMQ does not need it — keeps the wire uniform with Redis.
|
||||
|
||||
Two threads on the server: the ROUTER recv loop, and a per-client is implicit
|
||||
(ZMQ handles multiplexing). One thread on the client for recv is avoided by
|
||||
using DEALER + ``poll`` (synchronous request/response fits the RPC model).
|
||||
"""
|
||||
|
||||
import json
|
||||
import queue
|
||||
import threading
|
||||
import time
|
||||
import uuid
|
||||
|
||||
from ..adapters.redis_common import decode_text
|
||||
from ..redis_rpc import (
|
||||
decode_rpc_request_payload,
|
||||
encode_rpc_request_payload,
|
||||
)
|
||||
from .base import RpcTransport, TransportError, TransportTimeout
|
||||
|
||||
|
||||
# ZMQ does not support ipc:// on Windows (it trips a signaler abort), so the
|
||||
# default endpoint is tcp loopback. The port is derived from the account_id so
|
||||
# distinct accounts don't collide on the same port; override via config when
|
||||
# needed. Base 15560 keeps it clear of common dev ports.
|
||||
DEFAULT_ZMQ_HOST = "127.0.0.1"
|
||||
DEFAULT_ZMQ_BASE_PORT = 15560
|
||||
DEFAULT_ZMQ_PORT_RANGE = 100 # derived port = base + (account_id_int mod range)
|
||||
|
||||
|
||||
def _default_zmq_port(account_id):
|
||||
"""Derive a stable port from account_id so each account gets its own socket."""
|
||||
text = str(account_id or "")
|
||||
digits = "".join(ch for ch in text if ch.isdigit())
|
||||
try:
|
||||
offset = int(digits) % DEFAULT_ZMQ_PORT_RANGE if digits else 0
|
||||
except ValueError:
|
||||
offset = 0
|
||||
return DEFAULT_ZMQ_BASE_PORT + offset
|
||||
|
||||
|
||||
def _default_zmq_address(account_id, host=None):
|
||||
host = host or DEFAULT_ZMQ_HOST
|
||||
return "tcp://%s:%d" % (host, _default_zmq_port(account_id))
|
||||
|
||||
|
||||
def _loads(raw):
|
||||
if isinstance(raw, dict):
|
||||
return dict(raw)
|
||||
text = decode_text(raw)
|
||||
text = decode_rpc_request_payload(text)
|
||||
return json.loads(text)
|
||||
|
||||
|
||||
class ZmqTransport(RpcTransport):
|
||||
"""ZMQ ROUTER/DEALER transport.
|
||||
|
||||
The same instance plays both roles depending on method called:
|
||||
``send_request`` acts as a client (DEALER connect), ``start_receiving`` +
|
||||
``send_response`` act as a server (ROUTER bind). A deployment normally uses
|
||||
one instance per role (the QMT process is the server; the external client
|
||||
is the client).
|
||||
"""
|
||||
|
||||
name = "zmq"
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
bind_address=None,
|
||||
connect_address=None,
|
||||
host=None,
|
||||
port=None,
|
||||
account_id="",
|
||||
print_prefix="[bigqmt_rpc]",
|
||||
io_threads=1,
|
||||
recv_timeout_seconds=1.0,
|
||||
server_hwm=10000,
|
||||
client_linger_ms=0,
|
||||
discovery_redis_client=None,
|
||||
discovery_key_template="bigqmt:zmq:addr:{account_id}",
|
||||
discovery_ttl_seconds=300,
|
||||
port_scan_range=50,
|
||||
):
|
||||
super(ZmqTransport, self).__init__(account_id=account_id, print_prefix=print_prefix)
|
||||
# Address resolution order: explicit bind_address/connect_address win;
|
||||
# otherwise build tcp://host:port from host/port (port defaults to a
|
||||
# value derived from account_id so distinct accounts don't collide).
|
||||
resolved_host = host or DEFAULT_ZMQ_HOST
|
||||
if port is not None:
|
||||
resolved_port = int(port)
|
||||
else:
|
||||
resolved_port = _default_zmq_port(account_id)
|
||||
default_addr = "tcp://%s:%d" % (resolved_host, resolved_port)
|
||||
self.bind_address = bind_address or default_addr
|
||||
self.connect_address = connect_address
|
||||
self.bind_host = resolved_host
|
||||
self.base_port = resolved_port
|
||||
self.io_threads = int(io_threads)
|
||||
self.recv_timeout_seconds = float(recv_timeout_seconds)
|
||||
self.server_hwm = int(server_hwm)
|
||||
self.client_linger_ms = int(client_linger_ms)
|
||||
# Discovery remains available for clients, but a server must bind the
|
||||
# configured address exactly. ``port_scan_range`` is retained only for
|
||||
# backward-compatible config loading and is intentionally not used.
|
||||
self.discovery_redis_client = discovery_redis_client
|
||||
self.discovery_key_template = discovery_key_template
|
||||
self.discovery_ttl_seconds = int(discovery_ttl_seconds)
|
||||
self.port_scan_range = int(port_scan_range)
|
||||
|
||||
self._zmq = None # imported lazily
|
||||
self._ctx = None
|
||||
# server state
|
||||
self._router = None
|
||||
self._router_thread = None
|
||||
self._actual_bind_address = None # set after start_receiving()
|
||||
self._pending_identities = {} # request_id -> client identity bytes
|
||||
self._identity_lock = threading.Lock()
|
||||
self._response_queue = queue.Queue()
|
||||
self._queued_response_count = 0
|
||||
self._sent_response_count = 0
|
||||
# client state
|
||||
self._dealer = None
|
||||
self._client_lock = threading.Lock()
|
||||
|
||||
# -- construction helper ----------------------------------------------
|
||||
@classmethod
|
||||
def from_config(cls, config, account_id="", print_prefix="[bigqmt_rpc]"):
|
||||
config = dict(config or {})
|
||||
return cls(
|
||||
bind_address=config.get("bind_address"),
|
||||
connect_address=config.get("connect_address"),
|
||||
host=config.get("host"),
|
||||
port=config.get("port"),
|
||||
account_id=config.get("account_id", account_id),
|
||||
print_prefix=print_prefix,
|
||||
io_threads=int(config.get("io_threads", 1)),
|
||||
recv_timeout_seconds=float(config.get("recv_timeout_seconds", 1.0)),
|
||||
server_hwm=int(config.get("server_hwm", 10000)),
|
||||
client_linger_ms=int(config.get("client_linger_ms", 0)),
|
||||
discovery_redis_client=config.get("discovery_redis_client"),
|
||||
discovery_key_template=config.get(
|
||||
"discovery_key_template", "bigqmt:zmq:addr:{account_id}"
|
||||
),
|
||||
discovery_ttl_seconds=int(config.get("discovery_ttl_seconds", 300)),
|
||||
port_scan_range=int(config.get("port_scan_range", 50)),
|
||||
)
|
||||
|
||||
# -- shared zmq context -----------------------------------------------
|
||||
def _ensure_zmq(self):
|
||||
if self._zmq is None:
|
||||
try:
|
||||
import zmq # noqa: F401
|
||||
except ImportError as exc: # pragma: no cover - depends on env
|
||||
raise TransportError(
|
||||
"pyzmq is required for the zmq transport: %s" % exc
|
||||
)
|
||||
self._zmq = zmq
|
||||
if self._ctx is None:
|
||||
self._ctx = self._zmq.Context.instance(self.io_threads)
|
||||
return self._zmq, self._ctx
|
||||
|
||||
# -- server side ------------------------------------------------------
|
||||
def _bind_configured_address(self):
|
||||
"""Bind exactly one configured address and reject duplicate servers."""
|
||||
zmq, ctx = self._ensure_zmq()
|
||||
sock = ctx.socket(zmq.ROUTER)
|
||||
sock.setsockopt(zmq.RCVHWM, self.server_hwm)
|
||||
sock.setsockopt(zmq.SNDHWM, self.server_hwm)
|
||||
sock.setsockopt(zmq.RCVTIMEO, int(self.recv_timeout_seconds * 1000))
|
||||
try:
|
||||
sock.bind(self.bind_address)
|
||||
except self._zmq.ZMQError as exc:
|
||||
try:
|
||||
sock.close(linger=0)
|
||||
except Exception:
|
||||
pass
|
||||
if getattr(exc, "errno", None) == zmq.EADDRINUSE:
|
||||
# 端口被占——通常是之前策略实例没正常停止。给出友好提示和解决步骤。
|
||||
print(
|
||||
"%s ZMQ_BIND_CONFLICT: 端口 %s 被占用!"
|
||||
% (self.print_prefix, self.bind_address)
|
||||
)
|
||||
print(
|
||||
"%s 原因:之前的 QMT 策略实例没正常停止,仍占着这个端口。"
|
||||
% self.print_prefix
|
||||
)
|
||||
print(
|
||||
"%s 解决:1) 在 QMT 里停止旧策略再运行;2) 或等 60s 让系统释放端口;"
|
||||
% self.print_prefix
|
||||
)
|
||||
print(
|
||||
"%s 3) 或改配置用别的端口(BIGQMT_REDIS_CONFIG.zmq.port)"
|
||||
% self.print_prefix
|
||||
)
|
||||
raise TransportError(
|
||||
"ZMQ_BIND_CONFLICT address=%s; another bridge instance "
|
||||
"already owns the configured endpoint" % self.bind_address
|
||||
)
|
||||
raise
|
||||
self._router = sock
|
||||
self._actual_bind_address = self.bind_address
|
||||
self._publish_discovery(self.bind_address)
|
||||
|
||||
def _publish_discovery(self, address):
|
||||
if self.discovery_redis_client is None:
|
||||
return
|
||||
key = self.discovery_key_template.format(account_id=self.account_id)
|
||||
try:
|
||||
self.discovery_redis_client.setex(
|
||||
key, self.discovery_ttl_seconds, address
|
||||
)
|
||||
except Exception as exc:
|
||||
print("%s zmq discovery publish failed: %s" % (self.print_prefix, exc))
|
||||
|
||||
def _clear_discovery(self):
|
||||
if self.discovery_redis_client is None:
|
||||
return
|
||||
key = self.discovery_key_template.format(account_id=self.account_id)
|
||||
try:
|
||||
self.discovery_redis_client.delete(key)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def start_receiving(self, on_request, background_threads=True):
|
||||
super(ZmqTransport, self).start_receiving(on_request)
|
||||
zmq, ctx = self._ensure_zmq()
|
||||
self._bind_configured_address()
|
||||
bound = self._actual_bind_address or self.bind_address
|
||||
if not background_threads:
|
||||
print(
|
||||
"%s zmq bound=%s background_threads=False"
|
||||
% (self.print_prefix, bound)
|
||||
)
|
||||
return
|
||||
self._router_thread = threading.Thread(
|
||||
target=self._router_loop, name="bigqmt-zmq-rpc", daemon=True
|
||||
)
|
||||
self._router_thread.start()
|
||||
print(
|
||||
"%s zmq started bound=%s" % (self.print_prefix, self.bind_address)
|
||||
)
|
||||
|
||||
def _router_loop(self):
|
||||
try:
|
||||
while self._running:
|
||||
self._drain_response_queue()
|
||||
request = self._receive_request()
|
||||
if request is not None:
|
||||
self._deliver_request(request)
|
||||
finally:
|
||||
# Close the ROUTER socket on the thread that owns it. On Windows,
|
||||
# closing a ZMQ socket from a different thread trips a signaler
|
||||
# assertion (abort); closing it here is safe because this thread
|
||||
# created and exclusively used it.
|
||||
try:
|
||||
self._router.close(linger=0)
|
||||
except Exception:
|
||||
pass
|
||||
self._router = None
|
||||
|
||||
def _receive_request(self, flags=0):
|
||||
try:
|
||||
frames = self._router.recv_multipart(flags=flags)
|
||||
except self._zmq.Again:
|
||||
return None
|
||||
except Exception as exc:
|
||||
if self._running:
|
||||
print("%s zmq recv failed: %s" % (self.print_prefix, exc))
|
||||
if not flags:
|
||||
time.sleep(0.5)
|
||||
return None
|
||||
if len(frames) < 2:
|
||||
return None
|
||||
identity, payload = frames[0], frames[-1]
|
||||
try:
|
||||
request = _loads(payload)
|
||||
except Exception as exc:
|
||||
print("%s zmq decode failed: %s" % (self.print_prefix, exc))
|
||||
return None
|
||||
request_id = str(request.get("request_id") or uuid.uuid4().hex)
|
||||
with self._identity_lock:
|
||||
self._pending_identities[request_id] = identity
|
||||
return request
|
||||
|
||||
def _deliver_request(self, request):
|
||||
started = time.perf_counter()
|
||||
try:
|
||||
self.deliver(request)
|
||||
except Exception as exc:
|
||||
print("%s zmq deliver failed: %s" % (self.print_prefix, exc))
|
||||
elapsed_ms = (time.perf_counter() - started) * 1000.0
|
||||
if elapsed_ms > 50.0:
|
||||
print("%s zmq slow handler method=%s %.0fms"
|
||||
% (self.print_prefix, request.get("method"), elapsed_ms))
|
||||
|
||||
def _drain_response_queue(self):
|
||||
while True:
|
||||
try:
|
||||
identity, payload = self._response_queue.get_nowait()
|
||||
except queue.Empty:
|
||||
return
|
||||
try:
|
||||
self._router.send_multipart([identity, payload])
|
||||
self._sent_response_count += 1
|
||||
if self._sent_response_count <= 5:
|
||||
print("%s zmq queued response sent" % self.print_prefix)
|
||||
except Exception as exc:
|
||||
print("%s zmq send failed: %s" % (self.print_prefix, exc))
|
||||
|
||||
def send_response(self, request, response):
|
||||
if self._router is None:
|
||||
raise TransportError("zmq server socket is not bound")
|
||||
request_id = str(
|
||||
response.get("request_id") or request.get("request_id") or ""
|
||||
)
|
||||
with self._identity_lock:
|
||||
identity = self._pending_identities.pop(request_id, None)
|
||||
if identity is None:
|
||||
# No matching peer — drop silently (client may have gone away).
|
||||
return
|
||||
payload = encode_rpc_request_payload(response).encode("utf-8")
|
||||
if self._router_thread is not None and threading.current_thread() is not self._router_thread:
|
||||
self._queued_response_count += 1
|
||||
if self._queued_response_count <= 5:
|
||||
print("%s zmq response queued for router thread" % self.print_prefix)
|
||||
self._response_queue.put((identity, payload))
|
||||
return
|
||||
try:
|
||||
self._router.send_multipart([identity, payload])
|
||||
except Exception as exc:
|
||||
print("%s zmq send failed: %s" % (self.print_prefix, exc))
|
||||
|
||||
def drain_request_queue(self, max_items=20):
|
||||
"""Drain requests from the scheduled QMT thread when no receiver thread exists."""
|
||||
if self._router_thread is not None or self._router is None:
|
||||
return 0
|
||||
processed = 0
|
||||
for _index in range(max(int(max_items), 0)):
|
||||
request = self._receive_request(flags=self._zmq.NOBLOCK)
|
||||
if request is None:
|
||||
break
|
||||
self._deliver_request(request)
|
||||
processed += 1
|
||||
return processed
|
||||
|
||||
# -- client side ------------------------------------------------------
|
||||
def _resolve_connect_address(self):
|
||||
"""Resolve the address to connect to.
|
||||
|
||||
Order: explicit connect_address > discovery lookup > default derived.
|
||||
Discovery lets the client find a server that had to move off the
|
||||
default port because of a collision.
|
||||
"""
|
||||
if self.connect_address:
|
||||
return self.connect_address
|
||||
discovered = self._lookup_discovery()
|
||||
if discovered:
|
||||
return discovered
|
||||
return _default_zmq_address(self.account_id)
|
||||
|
||||
def _lookup_discovery(self):
|
||||
if self.discovery_redis_client is None:
|
||||
return None
|
||||
key = self.discovery_key_template.format(account_id=self.account_id)
|
||||
try:
|
||||
raw = self.discovery_redis_client.get(key)
|
||||
except Exception:
|
||||
return None
|
||||
if not raw:
|
||||
return None
|
||||
try:
|
||||
text = raw.decode("utf-8") if isinstance(raw, (bytes, bytearray)) else str(raw)
|
||||
except Exception:
|
||||
return None
|
||||
return text or None
|
||||
|
||||
def _ensure_dealer(self):
|
||||
zmq, ctx = self._ensure_zmq()
|
||||
if self._dealer is None:
|
||||
address = self._resolve_connect_address()
|
||||
sock = ctx.socket(zmq.DEALER)
|
||||
# Unique identity so ROUTER can route replies back to us.
|
||||
sock.setsockopt(zmq.IDENTITY, uuid.uuid4().hex.encode("utf-8")[:16])
|
||||
sock.setsockopt(zmq.LINGER, self.client_linger_ms)
|
||||
sock.connect(address)
|
||||
self._dealer = sock
|
||||
self.connect_address = address
|
||||
return self._dealer
|
||||
|
||||
def send_request(self, request, timeout_seconds, **_kwargs):
|
||||
zmq = self._zmq or self._ensure_zmq()[0]
|
||||
with self._client_lock:
|
||||
dealer = self._ensure_dealer()
|
||||
request = dict(request)
|
||||
request.setdefault("request_id", uuid.uuid4().hex)
|
||||
request_id = request["request_id"]
|
||||
payload = encode_rpc_request_payload(request)
|
||||
try:
|
||||
dealer.send(payload.encode("utf-8"))
|
||||
except Exception as exc:
|
||||
raise TransportError("zmq send failed: %s" % exc)
|
||||
deadline = time.time() + float(timeout_seconds)
|
||||
poller = self._zmq.Poller()
|
||||
poller.register(dealer, self._zmq.POLLIN)
|
||||
while True:
|
||||
remaining = deadline - time.time()
|
||||
if remaining <= 0:
|
||||
break
|
||||
events = dict(poller.poll(timeout=int(remaining * 1000)))
|
||||
if dealer in events:
|
||||
frames = dealer.recv_multipart()
|
||||
raw = frames[-1]
|
||||
response = _loads(raw)
|
||||
if response.get("request_id") == request_id:
|
||||
return response
|
||||
raise TransportTimeout("zmq rpc timeout: %s" % request.get("method"))
|
||||
|
||||
# -- lifecycle --------------------------------------------------------
|
||||
def stop(self):
|
||||
super(ZmqTransport, self).stop()
|
||||
# Clear _running so the router loop exits; the loop closes its own
|
||||
# socket (closing cross-thread trips a Windows signaler abort).
|
||||
thread = self._router_thread
|
||||
if thread is not None and thread.is_alive():
|
||||
thread.join(2.0)
|
||||
if thread is None and self._router is not None:
|
||||
try:
|
||||
self._router.close(linger=0)
|
||||
except Exception:
|
||||
pass
|
||||
self._router = None
|
||||
self._router_thread = None
|
||||
# If we were a server that published a discovery address, clear it so
|
||||
# clients don't keep hitting a dead endpoint.
|
||||
if self._actual_bind_address is not None:
|
||||
self._clear_discovery()
|
||||
self._actual_bind_address = None
|
||||
with self._client_lock:
|
||||
if self._dealer is not None:
|
||||
try:
|
||||
self._dealer.close(linger=self.client_linger_ms)
|
||||
except Exception:
|
||||
pass
|
||||
self._dealer = None
|
||||
# Do NOT terminate the shared context — other sockets/users may rely on it.
|
||||
@@ -1,204 +0,0 @@
|
||||
"""Client-side whole-quote subscription session.
|
||||
|
||||
Owns the per-process state for ``subscribe_whole_quote``: the local
|
||||
subscription table, the shared push-channel subscriber thread, and the
|
||||
keepalive heartbeat thread. One session is shared by every ``subscribe_whole_quote``
|
||||
call in the process (``BigQmtXtData`` delegates here), so all subscriptions ride
|
||||
a single push-channel connection and a single heartbeat loop.
|
||||
|
||||
The big-QMT whole-quote callback is INCREMENTAL (only changed symbols), so a
|
||||
subscription does not by itself deliver an initial full snapshot — callers layer
|
||||
a ``get_full_tick`` prime on top (done in ``BigQmtXtData.subscribe_whole_quote``).
|
||||
"""
|
||||
|
||||
import threading
|
||||
|
||||
|
||||
def _norm_topic(code_list):
|
||||
return ",".join(sorted({str(c).strip().upper() for c in (code_list or []) if str(c or "").strip()}))
|
||||
|
||||
|
||||
class WholeQuoteClientSession(object):
|
||||
def __init__(self, rpc_call, push_channel, client_id, heartbeat_interval_seconds=3.0, sub_id_func=None,
|
||||
push_silence_replay_heartbeats=10):
|
||||
"""``rpc_call`` is ``client.call``-shaped: fn(method, params) -> dict.
|
||||
``push_channel`` is a QuotePushChannel used purely as a subscriber.
|
||||
``sub_id_func`` (optional) mints subscription ids; defaults to a counter.
|
||||
``push_silence_replay_heartbeats``: after this many heartbeat rounds
|
||||
without any push, replay subscriptions (covers server restarts where
|
||||
keepalive keeps succeeding because the redis request queue buffers
|
||||
during the restart window but the subscription table was reset)."""
|
||||
self._rpc = rpc_call
|
||||
self._channel = push_channel
|
||||
self.client_id = str(client_id or "")
|
||||
self._heartbeat_interval = float(heartbeat_interval_seconds)
|
||||
self._push_silence_replay_heartbeats = int(push_silence_replay_heartbeats)
|
||||
self._sub_id_func = sub_id_func
|
||||
self._seq = 0
|
||||
self._lock = threading.RLock()
|
||||
self._subscriptions = {} # sub_id -> {"topic": str, "callback": fn, "codes": [...]}
|
||||
self._started = False
|
||||
self._subscriber_active = False
|
||||
self._subscribed_topics = frozenset() # topic set the subscriber covers now
|
||||
self._heartbeat_thread = None
|
||||
self._last_push_time = None # monotonic time of last incoming push
|
||||
|
||||
# -- subscription lifecycle ---------------------------------------------
|
||||
def subscribe_whole_quote(self, code_list, callback=None):
|
||||
codes = [str(c) for c in (code_list or []) if str(c or "").strip()]
|
||||
if not codes:
|
||||
raise ValueError("code_list is required")
|
||||
with self._lock:
|
||||
sub_id = self._next_sub_id()
|
||||
result = self._rpc(
|
||||
"subscribe_whole_quote",
|
||||
{"client_id": self.client_id, "sub_id": sub_id, "codes": codes},
|
||||
) or {}
|
||||
topic = str(result.get("topic") or result.get("combo_key") or _norm_topic(codes))
|
||||
with self._lock:
|
||||
self._subscriptions[sub_id] = {"topic": topic, "callback": callback, "codes": codes}
|
||||
self._sync_subscriber_locked()
|
||||
return sub_id
|
||||
|
||||
def unsubscribe_quote(self, sub_id):
|
||||
with self._lock:
|
||||
entry = self._subscriptions.pop(sub_id, None)
|
||||
if entry is None:
|
||||
return 0
|
||||
try:
|
||||
self._rpc("unsubscribe_whole_quote", {"client_id": self.client_id, "sub_id": sub_id})
|
||||
finally:
|
||||
with self._lock:
|
||||
self._sync_subscriber_locked()
|
||||
return 0
|
||||
|
||||
def has_subscription(self, sub_id):
|
||||
with self._lock:
|
||||
return sub_id in self._subscriptions
|
||||
|
||||
def replay_subscriptions(self):
|
||||
"""Re-send subscribe for every active sub_id (server restart recovery).
|
||||
Idempotent on the server (keyed by client_id+combo), so replays are safe."""
|
||||
with self._lock:
|
||||
items = [(sid, dict(entry)) for sid, entry in self._subscriptions.items()]
|
||||
for sub_id, entry in items:
|
||||
self._rpc(
|
||||
"subscribe_whole_quote",
|
||||
{"client_id": self.client_id, "sub_id": sub_id, "codes": entry["codes"]},
|
||||
)
|
||||
|
||||
# -- heartbeat -------------------------------------------------------------
|
||||
def start(self):
|
||||
with self._lock:
|
||||
if self._started:
|
||||
return
|
||||
self._started = True
|
||||
self._heartbeat_thread = threading.Thread(
|
||||
target=self._heartbeat_loop, name="bigqmt-quote-keepalive", daemon=True
|
||||
)
|
||||
self._heartbeat_thread.start()
|
||||
|
||||
def stop(self):
|
||||
with self._lock:
|
||||
self._started = False
|
||||
thread = self._heartbeat_thread
|
||||
if thread is not None:
|
||||
thread.join(timeout=1.0)
|
||||
self._heartbeat_thread = None
|
||||
|
||||
def _heartbeat_loop(self):
|
||||
import time
|
||||
|
||||
consecutive_failures = 0
|
||||
silence_rounds = 0
|
||||
prev_last_push = None
|
||||
while True:
|
||||
with self._lock:
|
||||
if not self._started:
|
||||
return
|
||||
sub_ids = list(self._subscriptions.keys())
|
||||
last_push = self._last_push_time
|
||||
if not sub_ids:
|
||||
time.sleep(self._heartbeat_interval)
|
||||
continue
|
||||
failures = 0
|
||||
for sub_id in sub_ids:
|
||||
try:
|
||||
self._rpc("quote_keepalive", {"client_id": self.client_id, "sub_id": sub_id})
|
||||
except Exception:
|
||||
failures += 1
|
||||
if failures:
|
||||
consecutive_failures += 1
|
||||
elif consecutive_failures >= 3:
|
||||
# Server is back after a restart window: replay subscriptions so
|
||||
# the restarted server re-creates the big-QMT subscriptions (its
|
||||
# state is gone). Idempotent on the server, so replays are safe.
|
||||
self.replay_subscriptions()
|
||||
consecutive_failures = 0
|
||||
else:
|
||||
consecutive_failures = 0
|
||||
# Push-silence detection: a server restart can survive with keepalive
|
||||
# succeeding (the redis request queue buffers during the restart
|
||||
# window) while the subscription table was reset, so pushes stop.
|
||||
# Replay when no push arrived for several heartbeat rounds (also
|
||||
# covers the case where the very first prime push never arrived).
|
||||
if last_push != prev_last_push:
|
||||
silence_rounds = 0 # a push arrived since the last round
|
||||
else:
|
||||
silence_rounds += 1
|
||||
prev_last_push = last_push
|
||||
if silence_rounds >= self._push_silence_replay_heartbeats:
|
||||
self.replay_subscriptions()
|
||||
silence_rounds = 0
|
||||
time.sleep(self._heartbeat_interval)
|
||||
|
||||
# -- push routing ------------------------------------------------------------
|
||||
def _on_push(self, topic, data):
|
||||
import time
|
||||
|
||||
now = time.monotonic()
|
||||
with self._lock:
|
||||
self._last_push_time = now
|
||||
callbacks = [
|
||||
entry["callback"]
|
||||
for entry in self._subscriptions.values()
|
||||
if entry["topic"] == topic and entry["callback"] is not None
|
||||
]
|
||||
for callback in callbacks:
|
||||
try:
|
||||
callback(data)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
def _sync_subscriber_locked(self):
|
||||
"""(Re)start the push-channel subscriber to cover exactly the active
|
||||
topics. Reuses an existing subscriber when the topic set is unchanged;
|
||||
stops it before restarting when the set changed. No-op when nothing is
|
||||
subscribed (and stops the running subscriber in that case)."""
|
||||
topics = sorted({entry["topic"] for entry in self._subscriptions.values()})
|
||||
active = frozenset(topics)
|
||||
if active == self._subscribed_topics:
|
||||
return
|
||||
if not active:
|
||||
if self._subscriber_active:
|
||||
try:
|
||||
self._channel.stop()
|
||||
except Exception:
|
||||
pass
|
||||
self._subscriber_active = False
|
||||
self._subscribed_topics = active
|
||||
return
|
||||
if self._subscriber_active:
|
||||
try:
|
||||
self._channel.stop()
|
||||
except Exception:
|
||||
pass
|
||||
self._channel.start_subscriber(topics, self._on_push)
|
||||
self._subscriber_active = True
|
||||
self._subscribed_topics = active
|
||||
|
||||
def _next_sub_id(self):
|
||||
if self._sub_id_func is not None:
|
||||
return self._sub_id_func()
|
||||
self._seq += 1
|
||||
return self._seq
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,102 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Client-side private config example for MiniQMT-compatible replacement.
|
||||
|
||||
Copy this file to:
|
||||
|
||||
src/bigqmt_signal_trader_client_config.py
|
||||
|
||||
Do not commit the real file. It may contain account ids and Redis credentials.
|
||||
"""
|
||||
|
||||
BIGQMT_ACCOUNT_ID = "YOUR_ACCOUNT_ID"
|
||||
BIGQMT_RPC_TIMEOUT_SECONDS = 6.0
|
||||
BIGQMT_DOWNLOAD_WAIT_SECONDS = 1800
|
||||
BIGQMT_DOWNLOAD_POLL_INTERVAL_SECONDS = 0.5
|
||||
|
||||
BIGQMT_REDIS_CONFIG = {
|
||||
"host": "YOUR_REDIS_HOST",
|
||||
"port": 6379,
|
||||
"db": 5,
|
||||
"username": "",
|
||||
"password": "",
|
||||
# Transport selection. Must match the QMT-side server config. Default
|
||||
# "redis" works with the standard DRYRUN; use "zmq" when the server runs
|
||||
# with transport=zmq (e.g. the no-redis version or explicit zmq mode).
|
||||
"transport": "redis",
|
||||
# ZMQ-specific settings (only used when transport=zmq):
|
||||
# "zmq": {
|
||||
# # Explicit connect address. The QMT-side server binds a port derived
|
||||
# # from account_id (default 15563 for account 8886800503). If you know
|
||||
# # the exact address, set it here to skip service discovery.
|
||||
# "connect_address": "tcp://127.0.0.1:15563",
|
||||
# # "host": "127.0.0.1",
|
||||
# # "port": 15563,
|
||||
# },
|
||||
}
|
||||
|
||||
# Default direct mode calls get_full_tick through RPC. Set enabled=True only when
|
||||
# you want client-side get_full_tick to read demand-driven Redis snapshots.
|
||||
BIGQMT_FULL_TICK_CACHE_CONFIG = {
|
||||
"enabled": False,
|
||||
"demand_ttl_seconds": 10,
|
||||
"cache_ttl_seconds": 10,
|
||||
"wait_seconds": 3.5,
|
||||
"poll_interval_seconds": 0.2,
|
||||
}
|
||||
|
||||
# Client-side LOCAL market-data cache.
|
||||
# get_market_data_ex(...) writes returned bars under `dir`; get_local_data(...)
|
||||
# then reads them locally with NO RPC to Big QMT (for offline / repeated local
|
||||
# analysis). download_history_data* submits a server-side Big QMT download job.
|
||||
# - dir: cache folder (default ~/.bigqmt_cache), one pickle per (period, code).
|
||||
# - fallback_rpc: if True, get_local_data auto-fetches+caches a cache miss;
|
||||
# if False (default), a cache-missed code is simply omitted (download first).
|
||||
BIGQMT_LOCAL_CACHE_CONFIG = {
|
||||
"enabled": True,
|
||||
"dir": None, # None -> ~/.bigqmt_cache
|
||||
"fallback_rpc": False,
|
||||
# Storage format: "auto" (parquet if pyarrow installed, else pickle),
|
||||
# "parquet" (columnar/compressed/cross-language — recommended), or "pkl".
|
||||
# One file per (period, dividend_type, code); switching format auto-migrates.
|
||||
"format": "auto",
|
||||
}
|
||||
|
||||
# FormulaServer direct read fast-path (port 58600).
|
||||
# Big QMT's built-in C++ quote/reference service. Routing reads straight to it
|
||||
# bypasses the RPC bridge AND the QMT python thread's GIL: ~0.07ms vs ~13ms
|
||||
# over redis. Enabled by default; you normally do not need this block.
|
||||
#
|
||||
# Covers reference/history reads only. Account, position, order, trade and
|
||||
# 五档 (get_full_tick) calls are NOT served by FormulaServer and always go over
|
||||
# RPC. Every miss — unmapped method, untranslatable params, server down —
|
||||
# falls back to RPC automatically, so an unreachable 58600 changes nothing.
|
||||
BIGQMT_FORMULA_SERVER_CONFIG = {
|
||||
"enabled": True, # or set BIGQMT_FORMULA_ENABLED=0 in the environment
|
||||
# "host": "127.0.0.1", # FormulaServer binds 0.0.0.0, so cross-machine works
|
||||
# # if the firewall allows it
|
||||
# "port": 58600, # unset -> read from qmt_root's formulaserver.ini,
|
||||
# # then fall back to 58600
|
||||
# "qmt_root": r"D:\国金证券QMT交易端",
|
||||
# "timeout_seconds": 3.0,
|
||||
# "methods": [...], # restrict routing to a subset (default: all mapped)
|
||||
# "failure_cooldown_seconds": 30.0, # pause routing this long after a failure
|
||||
}
|
||||
|
||||
# Whole-quote PUSH subscription (xtdata.subscribe_whole_quote, aligned with MiniQMT).
|
||||
# Server pushes each incremental tick batch to every client subscribed to the
|
||||
# same combination; the RPC methods above only manage the subscription lifecycle.
|
||||
# Data flows over a separate push channel matching `transport` above
|
||||
# (redis pub/sub, or zmq PUB/SUB when transport="zmq"), msgpack-encoded
|
||||
# (install the `msgpack` extra; falls back to json if absent).
|
||||
#
|
||||
# quote_client_id: process-stable subscriber id. The server counts references
|
||||
# per (client_id, sub_id) and only tears down the shared big-QMT subscription
|
||||
# after EVERY client of a combination unsubscribes or times out. Unset -> a
|
||||
# persisted id is created at ~/.cache/bigqmt/quote_client_id so a restarted
|
||||
# client is recognised as the same subscriber (needed for replay recovery).
|
||||
# Heartbeat: client sends quote_keepalive every BIGQMT_QUOTE_HEARTBEAT_SECONDS
|
||||
# (default 3.0). Server reaps a client after heartbeat_timeout_seconds
|
||||
# (default 30s = 10 periods, configured server-side).
|
||||
BIGQMT_QUOTE_CLIENT_ID = None # e.g. "my-strategy-1"; None -> persisted auto id
|
||||
# BIGQMT_QUOTE_HEARTBEAT_SECONDS = 3.0 # env var; must be < server timeout/periods
|
||||
|
||||
@@ -1,131 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Big QMT diagnostics for market data and positions.
|
||||
|
||||
This strategy entry never submits orders. It only probes QMT runtime APIs.
|
||||
"""
|
||||
|
||||
|
||||
_ACCOUNT_ID = ""
|
||||
_PROBED = False
|
||||
|
||||
|
||||
def _resolve_runtime_name(name):
|
||||
if name in globals():
|
||||
return globals()[name]
|
||||
try:
|
||||
import builtins
|
||||
return getattr(builtins, name)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def _safe_attr(obj, name, default=None):
|
||||
return getattr(obj, name, default)
|
||||
|
||||
|
||||
def _detect_account():
|
||||
account_value = _resolve_runtime_name("account")
|
||||
return str(account_value or "")
|
||||
|
||||
|
||||
def _probe_market(ContextInfo):
|
||||
code = "000300.SH"
|
||||
try:
|
||||
ticks = ContextInfo.get_full_tick([code])
|
||||
tick = (ticks or {}).get(code)
|
||||
if not tick:
|
||||
print("[bigqmt_diagnostic] market tick missing code=%s raw=%s" % (code, ticks))
|
||||
else:
|
||||
print(
|
||||
"[bigqmt_diagnostic] market tick ok code=%s lastPrice=%s bid1=%s ask1=%s"
|
||||
% (
|
||||
code,
|
||||
tick.get("lastPrice"),
|
||||
(tick.get("bidPrice") or [None])[0],
|
||||
(tick.get("askPrice") or [None])[0],
|
||||
)
|
||||
)
|
||||
except Exception as exc:
|
||||
print("[bigqmt_diagnostic] market tick failed: %s" % exc)
|
||||
|
||||
try:
|
||||
detail = ContextInfo.get_instrumentdetail(code)
|
||||
if not detail:
|
||||
print("[bigqmt_diagnostic] instrument missing code=%s" % code)
|
||||
else:
|
||||
print(
|
||||
"[bigqmt_diagnostic] instrument ok code=%s status=%s up=%s down=%s"
|
||||
% (
|
||||
code,
|
||||
detail.get("InstrumentStatus"),
|
||||
detail.get("UpStopPrice"),
|
||||
detail.get("DownStopPrice"),
|
||||
)
|
||||
)
|
||||
except Exception as exc:
|
||||
print("[bigqmt_diagnostic] instrument failed: %s" % exc)
|
||||
|
||||
|
||||
def _probe_positions(account_id):
|
||||
query = _resolve_runtime_name("get_trade_detail_data")
|
||||
if query is None:
|
||||
print("[bigqmt_diagnostic] position failed: get_trade_detail_data missing")
|
||||
return
|
||||
if not account_id:
|
||||
print("[bigqmt_diagnostic] position skipped: account is empty")
|
||||
return
|
||||
|
||||
try:
|
||||
positions = query(account_id, "STOCK", "POSITION") or []
|
||||
print("[bigqmt_diagnostic] position ok account=%s count=%s" % (account_id, len(positions)))
|
||||
for pos in positions[:8]:
|
||||
print(
|
||||
"[bigqmt_diagnostic] position item code=%s.%s name=%s volume=%s available=%s"
|
||||
% (
|
||||
_safe_attr(pos, "m_strInstrumentID", ""),
|
||||
_safe_attr(pos, "m_strExchangeID", ""),
|
||||
_safe_attr(pos, "m_strInstrumentName", ""),
|
||||
_safe_attr(pos, "m_nVolume", ""),
|
||||
_safe_attr(pos, "m_nCanUseVolume", ""),
|
||||
)
|
||||
)
|
||||
except Exception as exc:
|
||||
print("[bigqmt_diagnostic] position failed account=%s error=%s" % (account_id, exc))
|
||||
|
||||
|
||||
def _probe(ContextInfo, reason):
|
||||
global _PROBED
|
||||
if _PROBED:
|
||||
return
|
||||
_PROBED = True
|
||||
print("[bigqmt_diagnostic] probe start reason=%s account=%s" % (reason, _ACCOUNT_ID))
|
||||
_probe_market(ContextInfo)
|
||||
_probe_positions(_ACCOUNT_ID)
|
||||
print("[bigqmt_diagnostic] probe end")
|
||||
|
||||
|
||||
def init(ContextInfo):
|
||||
global _ACCOUNT_ID
|
||||
_ACCOUNT_ID = _detect_account()
|
||||
if _ACCOUNT_ID and hasattr(ContextInfo, "set_account"):
|
||||
ContextInfo.set_account(_ACCOUNT_ID)
|
||||
print("[bigqmt_diagnostic] init ok account=%s" % _ACCOUNT_ID)
|
||||
_probe(ContextInfo, "init")
|
||||
|
||||
|
||||
def handlebar(ContextInfo):
|
||||
if hasattr(ContextInfo, "is_last_bar") and not ContextInfo.is_last_bar():
|
||||
return None
|
||||
return _probe(ContextInfo, "handlebar")
|
||||
|
||||
|
||||
def adjust(ContextInfo):
|
||||
return handlebar(ContextInfo)
|
||||
|
||||
|
||||
def order_callback(ContextInfo, orderInfo):
|
||||
return None
|
||||
|
||||
|
||||
def deal_callback(ContextInfo, dealInfo):
|
||||
return None
|
||||
@@ -1,28 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Big QMT signal trader dry-run entry.
|
||||
|
||||
Put this file into QMT's python strategy directory and load it from QMT.
|
||||
Current default uses empty signal source and DryRunOrderGateway, so it will not
|
||||
submit real orders.
|
||||
"""
|
||||
|
||||
from bigqmt_signal_trader_strategy import ( # noqa: E402
|
||||
adjust,
|
||||
configure,
|
||||
deal_callback,
|
||||
handlebar,
|
||||
init,
|
||||
order_callback,
|
||||
set_account_id,
|
||||
sync_positions,
|
||||
)
|
||||
|
||||
|
||||
# Fill this before real account testing. Leave empty for dry-run loading tests.
|
||||
ACCOUNT_ID = ""
|
||||
|
||||
|
||||
if ACCOUNT_ID:
|
||||
set_account_id(ACCOUNT_ID)
|
||||
|
||||
configure(mode="dryrun", account_id=ACCOUNT_ID or "dryrun")
|
||||
@@ -1,64 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Local private config example for the QMT python directory.
|
||||
|
||||
Copy this file to the QMT python directory as:
|
||||
|
||||
bigqmt_signal_trader_local_config.py
|
||||
|
||||
Do not commit the real file. It may contain account ids and Redis credentials.
|
||||
"""
|
||||
|
||||
BIGQMT_ACCOUNT_ID = "YOUR_ACCOUNT_ID"
|
||||
|
||||
BIGQMT_REDIS_CONFIG = {
|
||||
"host": "127.0.0.1",
|
||||
"port": 6379,
|
||||
"db": 5,
|
||||
"username": "",
|
||||
"password": "",
|
||||
# Keep order RPC disabled unless you explicitly want remote order/cancel.
|
||||
"rpc_allow_order_methods": False,
|
||||
# Redis and ZMQ can both drain requests through QMT's official
|
||||
# run_time("adjust", ...) callback. This avoids GIL stalls in QMT's process.
|
||||
"rpc_process_in_listener": True,
|
||||
"rpc_listener_methods": ("*",),
|
||||
"rpc_background_threads": False,
|
||||
"schedule_adjust": True,
|
||||
"schedule_adjust_interval": "100nMilliSecond",
|
||||
# The default mode calls get_full_tick through RPC. Enable this cache only
|
||||
# if full-market payloads are too large for your latency/CPU budget.
|
||||
# When a client calls get_full_tick, it renews demand for 10 seconds.
|
||||
# Symbol-list demands refresh every full_tick_refresh_interval_seconds; whole-market
|
||||
# (SH/SZ/BJ/HK) demands refresh on the slower market interval so a ~50k row snapshot
|
||||
# is not pulled every fast tick.
|
||||
"full_tick_cache_enabled": False,
|
||||
"full_tick_demand_ttl_seconds": 10,
|
||||
"full_tick_cache_ttl_seconds": 10,
|
||||
"full_tick_refresh_interval_seconds": 0.5,
|
||||
"full_tick_market_refresh_interval_seconds": 3,
|
||||
# Wall-clock budget for one refresh round; keeps a slow round from stalling the
|
||||
# strategy thread (the in-flight demand always completes).
|
||||
"full_tick_refresh_max_wall_seconds": 0.3,
|
||||
"full_tick_max_requests": 8,
|
||||
# Async download jobs: clients submit download_history_data(2) as a job; the
|
||||
# strategy thread downloads download_job_chunk_size symbols per tick (capped by
|
||||
# download_job_max_wall_seconds), so a long download never blocks the RPC pump.
|
||||
# chunk_size is the smallest per-tick block — keep it modest if downloads are slow.
|
||||
# Disabled: the full terminal's xtdata SDK can't reach a data service to
|
||||
# download. Supplement history via the terminal's 数据管理/补充数据 UI, then read
|
||||
# it over RPC (get_market_data_ex/get_local_data). Enable only where a
|
||||
# MiniQMT/xtdata data service is connectable.
|
||||
"download_jobs_enabled": False,
|
||||
"download_job_chunk_size": 10,
|
||||
"download_job_max_wall_seconds": 0.5,
|
||||
"download_job_ttl_seconds": 3600,
|
||||
# Push order_callback/deal_callback details to Redis so clients get real-time
|
||||
# on_stock_order / on_stock_trade callbacks (MiniQMT style) instead of polling.
|
||||
"exec_events_enabled": True,
|
||||
# Dump the raw order_callback/deal_callback object fields to the QMT output
|
||||
# panel, and attach them to the published event as "raw_fields". Prints on
|
||||
# every callback, so keep it off outside a diagnosis window. Turn it on to
|
||||
# observe what m_nDirection / m_nOffsetFlag actually carry in live callbacks
|
||||
# — the buy/sell mapping in exec_events.py currently assumes 48/49 there.
|
||||
"exec_events_debug_raw_fields": False,
|
||||
}
|
||||
@@ -1,63 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Big QMT Redis dry-run strategy entry.
|
||||
|
||||
This entry reads Redis db5 test signals and writes Redis state, but orders are
|
||||
DryRunOrderGateway orders only. It does not submit real QMT orders.
|
||||
"""
|
||||
|
||||
from bigqmt_signal_trader_strategy import ( # noqa: E402
|
||||
adjust,
|
||||
configure,
|
||||
deal_callback,
|
||||
handlebar,
|
||||
init,
|
||||
order_callback,
|
||||
set_account_id,
|
||||
sync_positions,
|
||||
)
|
||||
|
||||
|
||||
ACCOUNT_ID = "bigqmt_probe"
|
||||
REDIS_HOST = "127.0.0.1"
|
||||
REDIS_PORT = 6379
|
||||
REDIS_DB = 5
|
||||
REDIS_USERNAME = ""
|
||||
REDIS_PASSWORD = ""
|
||||
|
||||
try:
|
||||
from bigqmt_signal_trader_local_config import BIGQMT_REDIS_CONFIG
|
||||
except Exception:
|
||||
BIGQMT_REDIS_CONFIG = {}
|
||||
|
||||
REDIS_HOST = BIGQMT_REDIS_CONFIG.get("host", REDIS_HOST)
|
||||
REDIS_PORT = int(BIGQMT_REDIS_CONFIG.get("port", REDIS_PORT))
|
||||
REDIS_DB = int(BIGQMT_REDIS_CONFIG.get("db", REDIS_DB))
|
||||
REDIS_USERNAME = BIGQMT_REDIS_CONFIG.get("username", REDIS_USERNAME)
|
||||
REDIS_PASSWORD = BIGQMT_REDIS_CONFIG.get("password", REDIS_PASSWORD)
|
||||
|
||||
|
||||
if ACCOUNT_ID:
|
||||
set_account_id(ACCOUNT_ID)
|
||||
|
||||
configure(
|
||||
mode="dryrun",
|
||||
account_id=ACCOUNT_ID,
|
||||
signal_source_type="redis",
|
||||
state_store_type="redis",
|
||||
position_sync_type="redis",
|
||||
redis={
|
||||
"host": REDIS_HOST,
|
||||
"port": REDIS_PORT,
|
||||
"db": REDIS_DB,
|
||||
"username": REDIS_USERNAME,
|
||||
"password": REDIS_PASSWORD,
|
||||
"stream_key_template": "bigqmt:signals:{account_id}",
|
||||
"group_name": "bigqmt-signal-trader",
|
||||
"consumer_name": "bigqmt-probe",
|
||||
"block_ms": 0,
|
||||
"claim_key_template": "bigqmt:signal_claim:{account_id}:{signal_id}",
|
||||
"status_key_template": "bigqmt:signal_status:{account_id}:{signal_id}",
|
||||
"position_key_template": "bigqmt:positions:{account_id}",
|
||||
"position_event_stream_template": "bigqmt:position_events:{account_id}",
|
||||
},
|
||||
)
|
||||
@@ -1,301 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Big QMT Redis Pub/Sub RPC strategy entry.
|
||||
|
||||
This entry does not consume trade signals. RPC order methods are disabled by
|
||||
default; read-only methods and position sync are enabled.
|
||||
"""
|
||||
|
||||
import os
|
||||
import sys
|
||||
|
||||
|
||||
# QMT loads strategy scripts via exec, so __file__ may be undefined. Build a
|
||||
# list of candidate directories (script dir guesses + cwd) and put any that
|
||||
# holds bigqmt_signal_trader_strategy.py on sys.path[0]. This keeps the package
|
||||
# and bigqmt_signal_trader_local_config importable regardless of how QMT
|
||||
# invokes the script.
|
||||
_CANDIDATE_DIRS = []
|
||||
try:
|
||||
_CANDIDATE_DIRS.append(os.path.dirname(os.path.abspath(__file__)))
|
||||
except Exception:
|
||||
pass
|
||||
_CANDIDATE_DIRS.append(os.getcwd())
|
||||
for _up in (".", ".."):
|
||||
_CANDIDATE_DIRS.append(os.path.abspath(os.path.join(os.getcwd(), _up)))
|
||||
for _dir in _CANDIDATE_DIRS:
|
||||
if os.path.exists(os.path.join(_dir, "bigqmt_signal_trader_strategy.py")):
|
||||
if _dir not in sys.path:
|
||||
sys.path.insert(0, _dir)
|
||||
break
|
||||
|
||||
|
||||
try:
|
||||
_load_bridge_module = __bigqmt_load_local_module
|
||||
except NameError:
|
||||
_load_bridge_module = None
|
||||
|
||||
if _load_bridge_module is not None:
|
||||
_strategy_module = _load_bridge_module("bigqmt_signal_trader_strategy")
|
||||
adjust = _strategy_module.adjust
|
||||
bind_qmt_api = _strategy_module.bind_qmt_api
|
||||
configure = _strategy_module.configure
|
||||
deal_callback = _strategy_module.deal_callback
|
||||
handlebar = _strategy_module.handlebar
|
||||
init = _strategy_module.init
|
||||
order_callback = _strategy_module.order_callback
|
||||
set_account_id = _strategy_module.set_account_id
|
||||
sync_positions = _strategy_module.sync_positions
|
||||
else:
|
||||
from bigqmt_signal_trader_strategy import ( # noqa: E402
|
||||
adjust,
|
||||
bind_qmt_api,
|
||||
configure,
|
||||
deal_callback,
|
||||
handlebar,
|
||||
init,
|
||||
order_callback,
|
||||
set_account_id,
|
||||
sync_positions,
|
||||
)
|
||||
|
||||
|
||||
ACCOUNT_ID = ""
|
||||
REDIS_HOST = "127.0.0.1"
|
||||
REDIS_PORT = 6379
|
||||
REDIS_DB = 5
|
||||
REDIS_USERNAME = ""
|
||||
REDIS_PASSWORD = ""
|
||||
RPC_ALLOW_ORDER_METHODS = False
|
||||
RPC_PROCESS_IN_LISTENER = True
|
||||
RPC_BACKGROUND_THREADS = False
|
||||
# "*" expands to read-only RPC methods only. Order/cancel/sync methods still go
|
||||
# through the queue fallback and require schedule_adjust=True when enabled.
|
||||
RPC_LISTENER_METHODS = ("*",)
|
||||
# Transport selection. Default "redis" (zero behavior change). Set to "zmq" /
|
||||
# "mysql" / "shm" in the local config to switch the wire. zmq/mysql sub-config
|
||||
# (bind/connect address, pool sizing, ...) is forwarded verbatim to the factory.
|
||||
RPC_TRANSPORT = "redis"
|
||||
RPC_ZMQ_CONFIG = {}
|
||||
RPC_MYSQL_CONFIG = {}
|
||||
SCHEDULE_ADJUST_ENABLED = True
|
||||
# How often the strategy thread drains the RPC queue (via adjust). Lower = less
|
||||
# queue wait for read RPCs. Verify on the live box that run_time honors sub-3s
|
||||
# intervals (see the adjust cadence log) before trusting a low value.
|
||||
SCHEDULE_ADJUST_INTERVAL = "500nMilliSecond"
|
||||
FULL_TICK_CACHE_ENABLED = False
|
||||
FULL_TICK_DEMAND_TTL_SECONDS = 10
|
||||
FULL_TICK_CACHE_TTL_SECONDS = 10
|
||||
# Symbol-list demands refresh fast; whole-market (SH/SZ/BJ/HK) demands refresh on
|
||||
# a slower cadence so a ~50k row snapshot is not pulled on every fast tick.
|
||||
FULL_TICK_REFRESH_INTERVAL_SECONDS = 0.5
|
||||
FULL_TICK_MARKET_REFRESH_INTERVAL_SECONDS = 3.0
|
||||
# Wall-clock budget for one refresh round to avoid stalling the strategy thread.
|
||||
FULL_TICK_REFRESH_MAX_WALL_SECONDS = 0.3
|
||||
FULL_TICK_MAX_REQUESTS = 8
|
||||
# Async download jobs: the strategy thread drains one queued job at a time,
|
||||
# downloading DOWNLOAD_JOB_CHUNK_SIZE symbols per tick (capped by the wall-clock
|
||||
# budget), so a long download never blocks the RPC pump. chunk_size is the
|
||||
# smallest per-tick block, so keep it modest if per-symbol downloads are slow.
|
||||
# DISABLED by default: the full Big QMT terminal's embedded xtdata SDK has no
|
||||
# reachable data service to download through (raises "无法连接行情服务"). Supplement
|
||||
# history via the terminal's 数据管理/补充数据 UI, then read it over RPC with
|
||||
# get_market_data_ex / get_local_data. Re-enable only where a MiniQMT/xtdata data
|
||||
# service is connectable (set download_jobs_enabled=True in the local config).
|
||||
DOWNLOAD_JOBS_ENABLED = False
|
||||
DOWNLOAD_JOB_CHUNK_SIZE = 10
|
||||
DOWNLOAD_JOB_MAX_WALL_SECONDS = 0.5
|
||||
DOWNLOAD_JOB_TTL_SECONDS = 3600
|
||||
# Push order_callback/deal_callback details to Redis so clients get real-time
|
||||
# on_stock_order/on_stock_trade callbacks (MiniQMT style) instead of polling.
|
||||
EXEC_EVENTS_ENABLED = True
|
||||
# Dump the raw order_callback/deal_callback object fields to the QMT output panel
|
||||
# (and into the published event as "raw_fields"). Off by default — it prints on
|
||||
# every callback. Turn on to settle what m_nDirection/m_nOffsetFlag actually
|
||||
# carry live, which the buy/sell direction mapping currently assumes.
|
||||
EXEC_EVENTS_DEBUG_RAW_FIELDS = False
|
||||
|
||||
try:
|
||||
from bigqmt_signal_trader_local_config import BIGQMT_ACCOUNT_ID, BIGQMT_REDIS_CONFIG
|
||||
except Exception:
|
||||
BIGQMT_ACCOUNT_ID = ""
|
||||
BIGQMT_REDIS_CONFIG = {}
|
||||
|
||||
ACCOUNT_ID = str(BIGQMT_ACCOUNT_ID or ACCOUNT_ID or "")
|
||||
REDIS_HOST = BIGQMT_REDIS_CONFIG.get("host", REDIS_HOST)
|
||||
REDIS_PORT = int(BIGQMT_REDIS_CONFIG.get("port", REDIS_PORT))
|
||||
REDIS_DB = int(BIGQMT_REDIS_CONFIG.get("db", REDIS_DB))
|
||||
REDIS_USERNAME = BIGQMT_REDIS_CONFIG.get("username", REDIS_USERNAME)
|
||||
REDIS_PASSWORD = BIGQMT_REDIS_CONFIG.get("password", REDIS_PASSWORD)
|
||||
RPC_ALLOW_ORDER_METHODS = bool(BIGQMT_REDIS_CONFIG.get("rpc_allow_order_methods", RPC_ALLOW_ORDER_METHODS))
|
||||
RPC_PROCESS_IN_LISTENER = bool(
|
||||
BIGQMT_REDIS_CONFIG.get("rpc_process_in_listener", RPC_PROCESS_IN_LISTENER and not RPC_ALLOW_ORDER_METHODS)
|
||||
)
|
||||
RPC_BACKGROUND_THREADS = bool(BIGQMT_REDIS_CONFIG.get("rpc_background_threads", RPC_BACKGROUND_THREADS))
|
||||
RPC_LISTENER_METHODS = tuple(BIGQMT_REDIS_CONFIG.get("rpc_listener_methods", RPC_LISTENER_METHODS))
|
||||
SCHEDULE_ADJUST_ENABLED = bool(BIGQMT_REDIS_CONFIG.get("schedule_adjust", SCHEDULE_ADJUST_ENABLED))
|
||||
if not RPC_BACKGROUND_THREADS:
|
||||
SCHEDULE_ADJUST_ENABLED = True
|
||||
SCHEDULE_ADJUST_INTERVAL = str(BIGQMT_REDIS_CONFIG.get("schedule_adjust_interval", SCHEDULE_ADJUST_INTERVAL))
|
||||
FULL_TICK_CACHE_ENABLED = bool(BIGQMT_REDIS_CONFIG.get("full_tick_cache_enabled", FULL_TICK_CACHE_ENABLED))
|
||||
FULL_TICK_DEMAND_TTL_SECONDS = float(
|
||||
BIGQMT_REDIS_CONFIG.get("full_tick_demand_ttl_seconds", FULL_TICK_DEMAND_TTL_SECONDS)
|
||||
)
|
||||
FULL_TICK_CACHE_TTL_SECONDS = float(
|
||||
BIGQMT_REDIS_CONFIG.get("full_tick_cache_ttl_seconds", FULL_TICK_CACHE_TTL_SECONDS)
|
||||
)
|
||||
FULL_TICK_REFRESH_INTERVAL_SECONDS = float(
|
||||
BIGQMT_REDIS_CONFIG.get("full_tick_refresh_interval_seconds", FULL_TICK_REFRESH_INTERVAL_SECONDS)
|
||||
)
|
||||
FULL_TICK_MARKET_REFRESH_INTERVAL_SECONDS = float(
|
||||
BIGQMT_REDIS_CONFIG.get("full_tick_market_refresh_interval_seconds", FULL_TICK_MARKET_REFRESH_INTERVAL_SECONDS)
|
||||
)
|
||||
FULL_TICK_REFRESH_MAX_WALL_SECONDS = float(
|
||||
BIGQMT_REDIS_CONFIG.get("full_tick_refresh_max_wall_seconds", FULL_TICK_REFRESH_MAX_WALL_SECONDS)
|
||||
)
|
||||
FULL_TICK_MAX_REQUESTS = int(BIGQMT_REDIS_CONFIG.get("full_tick_max_requests", FULL_TICK_MAX_REQUESTS))
|
||||
DOWNLOAD_JOBS_ENABLED = bool(BIGQMT_REDIS_CONFIG.get("download_jobs_enabled", DOWNLOAD_JOBS_ENABLED))
|
||||
DOWNLOAD_JOB_CHUNK_SIZE = int(BIGQMT_REDIS_CONFIG.get("download_job_chunk_size", DOWNLOAD_JOB_CHUNK_SIZE))
|
||||
DOWNLOAD_JOB_MAX_WALL_SECONDS = float(
|
||||
BIGQMT_REDIS_CONFIG.get("download_job_max_wall_seconds", DOWNLOAD_JOB_MAX_WALL_SECONDS)
|
||||
)
|
||||
DOWNLOAD_JOB_TTL_SECONDS = int(BIGQMT_REDIS_CONFIG.get("download_job_ttl_seconds", DOWNLOAD_JOB_TTL_SECONDS))
|
||||
EXEC_EVENTS_ENABLED = bool(BIGQMT_REDIS_CONFIG.get("exec_events_enabled", EXEC_EVENTS_ENABLED))
|
||||
EXEC_EVENTS_DEBUG_RAW_FIELDS = bool(
|
||||
BIGQMT_REDIS_CONFIG.get("exec_events_debug_raw_fields", EXEC_EVENTS_DEBUG_RAW_FIELDS)
|
||||
)
|
||||
|
||||
|
||||
def _apply_config(account_id):
|
||||
account_id = str(account_id or "")
|
||||
if account_id:
|
||||
set_account_id(account_id)
|
||||
configure(
|
||||
mode="bigqmt",
|
||||
account_id=account_id,
|
||||
position_sync_type="redis" if RPC_TRANSPORT in ("redis", "", "default") else "",
|
||||
enable_rpc=True,
|
||||
schedule_adjust=SCHEDULE_ADJUST_ENABLED,
|
||||
schedule_adjust_interval=SCHEDULE_ADJUST_INTERVAL,
|
||||
redis={
|
||||
"host": REDIS_HOST,
|
||||
"port": REDIS_PORT,
|
||||
"db": REDIS_DB,
|
||||
"username": REDIS_USERNAME,
|
||||
"password": REDIS_PASSWORD,
|
||||
"position_key_template": "bigqmt:positions:{account_id}",
|
||||
"position_event_stream_template": "bigqmt:position_events:{account_id}",
|
||||
},
|
||||
rpc={
|
||||
"enabled": True,
|
||||
"account_id": account_id,
|
||||
"allow_order_methods": RPC_ALLOW_ORDER_METHODS,
|
||||
"request_channel_template": "bigqmt:rpc:req:{account_id}",
|
||||
"response_channel_template": "bigqmt:rpc:resp:{account_id}:{request_id}",
|
||||
"response_key_template": "bigqmt:rpc:resp:{account_id}:{request_id}",
|
||||
"response_ttl_seconds": 60,
|
||||
"drain_max_items": 20,
|
||||
"process_in_listener": RPC_PROCESS_IN_LISTENER,
|
||||
"listener_methods": RPC_LISTENER_METHODS,
|
||||
"background_threads": RPC_BACKGROUND_THREADS,
|
||||
# Transport selection (default redis). Forwarded from the local
|
||||
# config so the factory can pick zmq/mysql/shm.
|
||||
"transport": RPC_TRANSPORT,
|
||||
"zmq": RPC_ZMQ_CONFIG,
|
||||
"mysql": RPC_MYSQL_CONFIG,
|
||||
},
|
||||
full_tick_cache={
|
||||
"enabled": FULL_TICK_CACHE_ENABLED,
|
||||
"account_id": account_id,
|
||||
"demand_ttl_seconds": FULL_TICK_DEMAND_TTL_SECONDS,
|
||||
"cache_ttl_seconds": FULL_TICK_CACHE_TTL_SECONDS,
|
||||
"refresh_interval_seconds": FULL_TICK_REFRESH_INTERVAL_SECONDS,
|
||||
"market_refresh_interval_seconds": FULL_TICK_MARKET_REFRESH_INTERVAL_SECONDS,
|
||||
"refresh_max_wall_seconds": FULL_TICK_REFRESH_MAX_WALL_SECONDS,
|
||||
"max_requests": FULL_TICK_MAX_REQUESTS,
|
||||
},
|
||||
download_jobs={
|
||||
"enabled": DOWNLOAD_JOBS_ENABLED,
|
||||
"account_id": account_id,
|
||||
"chunk_size": DOWNLOAD_JOB_CHUNK_SIZE,
|
||||
"max_wall_seconds": DOWNLOAD_JOB_MAX_WALL_SECONDS,
|
||||
"job_ttl_seconds": DOWNLOAD_JOB_TTL_SECONDS,
|
||||
},
|
||||
exec_events={
|
||||
"enabled": EXEC_EVENTS_ENABLED,
|
||||
"account_id": account_id,
|
||||
"debug_raw_fields": EXEC_EVENTS_DEBUG_RAW_FIELDS,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
def configure_runtime_account(account_id):
|
||||
_apply_config(account_id)
|
||||
|
||||
|
||||
def configure_runtime_redis(redis_config):
|
||||
global REDIS_HOST, REDIS_PORT, REDIS_DB, REDIS_USERNAME, REDIS_PASSWORD, RPC_ALLOW_ORDER_METHODS, RPC_PROCESS_IN_LISTENER, RPC_BACKGROUND_THREADS, RPC_LISTENER_METHODS, SCHEDULE_ADJUST_ENABLED, SCHEDULE_ADJUST_INTERVAL, FULL_TICK_CACHE_ENABLED, FULL_TICK_DEMAND_TTL_SECONDS, FULL_TICK_CACHE_TTL_SECONDS, FULL_TICK_REFRESH_INTERVAL_SECONDS, FULL_TICK_MARKET_REFRESH_INTERVAL_SECONDS, FULL_TICK_REFRESH_MAX_WALL_SECONDS, FULL_TICK_MAX_REQUESTS, RPC_TRANSPORT, RPC_ZMQ_CONFIG, RPC_MYSQL_CONFIG, DOWNLOAD_JOBS_ENABLED, DOWNLOAD_JOB_CHUNK_SIZE, DOWNLOAD_JOB_MAX_WALL_SECONDS, DOWNLOAD_JOB_TTL_SECONDS, EXEC_EVENTS_ENABLED, EXEC_EVENTS_DEBUG_RAW_FIELDS
|
||||
redis_config = dict(redis_config or {})
|
||||
REDIS_HOST = redis_config.get("host", REDIS_HOST)
|
||||
REDIS_PORT = int(redis_config.get("port", REDIS_PORT))
|
||||
REDIS_DB = int(redis_config.get("db", REDIS_DB))
|
||||
REDIS_USERNAME = redis_config.get("username", REDIS_USERNAME)
|
||||
REDIS_PASSWORD = redis_config.get("password", REDIS_PASSWORD)
|
||||
RPC_ALLOW_ORDER_METHODS = bool(redis_config.get("rpc_allow_order_methods", RPC_ALLOW_ORDER_METHODS))
|
||||
RPC_PROCESS_IN_LISTENER = bool(
|
||||
redis_config.get("rpc_process_in_listener", RPC_PROCESS_IN_LISTENER and not RPC_ALLOW_ORDER_METHODS)
|
||||
)
|
||||
RPC_BACKGROUND_THREADS = bool(redis_config.get("rpc_background_threads", RPC_BACKGROUND_THREADS))
|
||||
RPC_LISTENER_METHODS = tuple(redis_config.get("rpc_listener_methods", RPC_LISTENER_METHODS))
|
||||
RPC_TRANSPORT = str(redis_config.get("transport", RPC_TRANSPORT)).lower()
|
||||
RPC_ZMQ_CONFIG = dict(redis_config.get("zmq", RPC_ZMQ_CONFIG))
|
||||
RPC_MYSQL_CONFIG = dict(redis_config.get("mysql", RPC_MYSQL_CONFIG))
|
||||
# schedule_adjust must stay ON for ALL transports — including zmq.
|
||||
# run_time("adjust", interval) is what THROTTLES QMT's strategy callback: with
|
||||
# it, adjust fires on the configured cadence (e.g. 500ms); WITHOUT it QMT calls
|
||||
# adjust in a hot loop (~2500/s) that pegs the GIL and starves the zmq ROUTER
|
||||
# background thread (RPC then times out entirely). It also sets the GIL-release
|
||||
# rhythm the background transport threads rely on. So keep the original rule:
|
||||
# honor an explicit value, default on, and force on when not background-threaded.
|
||||
SCHEDULE_ADJUST_ENABLED = bool(redis_config.get("schedule_adjust", SCHEDULE_ADJUST_ENABLED))
|
||||
if not RPC_BACKGROUND_THREADS:
|
||||
SCHEDULE_ADJUST_ENABLED = True
|
||||
SCHEDULE_ADJUST_INTERVAL = str(redis_config.get("schedule_adjust_interval", SCHEDULE_ADJUST_INTERVAL))
|
||||
FULL_TICK_CACHE_ENABLED = bool(redis_config.get("full_tick_cache_enabled", FULL_TICK_CACHE_ENABLED))
|
||||
FULL_TICK_DEMAND_TTL_SECONDS = float(
|
||||
redis_config.get("full_tick_demand_ttl_seconds", FULL_TICK_DEMAND_TTL_SECONDS)
|
||||
)
|
||||
FULL_TICK_CACHE_TTL_SECONDS = float(redis_config.get("full_tick_cache_ttl_seconds", FULL_TICK_CACHE_TTL_SECONDS))
|
||||
FULL_TICK_REFRESH_INTERVAL_SECONDS = float(
|
||||
redis_config.get("full_tick_refresh_interval_seconds", FULL_TICK_REFRESH_INTERVAL_SECONDS)
|
||||
)
|
||||
FULL_TICK_MARKET_REFRESH_INTERVAL_SECONDS = float(
|
||||
redis_config.get("full_tick_market_refresh_interval_seconds", FULL_TICK_MARKET_REFRESH_INTERVAL_SECONDS)
|
||||
)
|
||||
FULL_TICK_REFRESH_MAX_WALL_SECONDS = float(
|
||||
redis_config.get("full_tick_refresh_max_wall_seconds", FULL_TICK_REFRESH_MAX_WALL_SECONDS)
|
||||
)
|
||||
FULL_TICK_MAX_REQUESTS = int(redis_config.get("full_tick_max_requests", FULL_TICK_MAX_REQUESTS))
|
||||
DOWNLOAD_JOBS_ENABLED = bool(redis_config.get("download_jobs_enabled", DOWNLOAD_JOBS_ENABLED))
|
||||
DOWNLOAD_JOB_CHUNK_SIZE = int(redis_config.get("download_job_chunk_size", DOWNLOAD_JOB_CHUNK_SIZE))
|
||||
DOWNLOAD_JOB_MAX_WALL_SECONDS = float(
|
||||
redis_config.get("download_job_max_wall_seconds", DOWNLOAD_JOB_MAX_WALL_SECONDS)
|
||||
)
|
||||
DOWNLOAD_JOB_TTL_SECONDS = int(redis_config.get("download_job_ttl_seconds", DOWNLOAD_JOB_TTL_SECONDS))
|
||||
EXEC_EVENTS_ENABLED = bool(redis_config.get("exec_events_enabled", EXEC_EVENTS_ENABLED))
|
||||
EXEC_EVENTS_DEBUG_RAW_FIELDS = bool(
|
||||
redis_config.get("exec_events_debug_raw_fields", EXEC_EVENTS_DEBUG_RAW_FIELDS)
|
||||
)
|
||||
_apply_config(ACCOUNT_ID)
|
||||
|
||||
|
||||
def bind_runtime_api(passorder_func=None, cancel_func=None, get_trade_detail_data_func=None,
|
||||
extra_funcs=None):
|
||||
bind_qmt_api(
|
||||
passorder_func=passorder_func,
|
||||
cancel_func=cancel_func,
|
||||
get_trade_detail_data_func=get_trade_detail_data_func,
|
||||
extra_funcs=extra_funcs,
|
||||
)
|
||||
|
||||
|
||||
_apply_config(ACCOUNT_ID)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,9 +0,0 @@
|
||||
"""Optional xtquant import shim backed by Big QMT Redis RPC.
|
||||
|
||||
Put this package before the real xtquant package on PYTHONPATH only when the
|
||||
caller intentionally wants Big QMT RPC compatibility.
|
||||
"""
|
||||
|
||||
from . import xtconstant, xtdata, xttrader, xttype
|
||||
|
||||
__all__ = ["xtconstant", "xtdata", "xttrader", "xttype"]
|
||||
@@ -1,128 +0,0 @@
|
||||
"""MiniQMT-compatible constant definitions (xtconstant).
|
||||
|
||||
Mirrors the native ``xtquant/xtconstant.py`` so code that does
|
||||
``from xtquant.xtconstant import STOCK_BUY`` keeps working against the
|
||||
Big QMT bridge. Values are defined once in
|
||||
``bigqmt_signal_trader.xtquant_compat`` and re-exported here.
|
||||
"""
|
||||
|
||||
from bigqmt_signal_trader.xtquant_compat import (
|
||||
# 账号类型
|
||||
CREDIT_ACCOUNT,
|
||||
FUTURE_ACCOUNT,
|
||||
FUTURE_OPTION_ACCOUNT,
|
||||
HUGANGTONG_ACCOUNT,
|
||||
SECURITY_ACCOUNT,
|
||||
SHENGANGTONG_ACCOUNT,
|
||||
STOCK_OPTION_ACCOUNT,
|
||||
# 委托类型 - 期货
|
||||
FUTURE_ARBITRAGE_CLOSE_HISTORY_FIRST,
|
||||
FUTURE_ARBITRAGE_CLOSE_TODAY_FIRST,
|
||||
FUTURE_ARBITRAGE_OPEN,
|
||||
FUTURE_CLOSE,
|
||||
FUTURE_CLOSE_LONG_HISTORY,
|
||||
FUTURE_CLOSE_LONG_HISTORY_FIRST,
|
||||
FUTURE_CLOSE_LONG_HISTORY_TODAY_THEN_OPEN_SHORT,
|
||||
FUTURE_CLOSE_LONG_TODAY,
|
||||
FUTURE_CLOSE_LONG_TODAY_FIRST,
|
||||
FUTURE_CLOSE_LONG_TODAY_HISTORY_THEN_OPEN_SHORT,
|
||||
FUTURE_CLOSE_SHORT_HISTORY,
|
||||
FUTURE_CLOSE_SHORT_HISTORY_FIRST,
|
||||
FUTURE_CLOSE_SHORT_HISTORY_TODAY_THEN_OPEN_LONG,
|
||||
FUTURE_CLOSE_SHORT_TODAY,
|
||||
FUTURE_CLOSE_SHORT_TODAY_FIRST,
|
||||
FUTURE_CLOSE_SHORT_TODAY_HISTORY_THEN_OPEN_LONG,
|
||||
FUTURE_OPEN,
|
||||
FUTURE_OPEN_LONG,
|
||||
FUTURE_OPEN_SHORT,
|
||||
FUTURE_RENEW_LONG_CLOSE_HISTORY_FIRST,
|
||||
FUTURE_RENEW_LONG_CLOSE_TODAY_FIRST,
|
||||
FUTURE_RENEW_SHORT_CLOSE_HISTORY_FIRST,
|
||||
FUTURE_RENEW_SHORT_CLOSE_TODAY_FIRST,
|
||||
# 委托类型 - 股票 / 信用
|
||||
CREDIT_BUY,
|
||||
CREDIT_BUY_SECU_REPAY,
|
||||
CREDIT_BUY_SECU_REPAY_SPECIAL,
|
||||
CREDIT_DIRECT_CASH_REPAY,
|
||||
CREDIT_DIRECT_CASH_REPAY_SPECIAL,
|
||||
CREDIT_DIRECT_SECU_REPAY,
|
||||
CREDIT_DIRECT_SECU_REPAY_SPECIAL,
|
||||
CREDIT_FIN_BUY,
|
||||
CREDIT_FIN_BUY_SPECIAL,
|
||||
CREDIT_SELL,
|
||||
CREDIT_SELL_SECU_REPAY,
|
||||
CREDIT_SELL_SECU_REPAY_SPECIAL,
|
||||
CREDIT_SLO_SELL,
|
||||
CREDIT_SLO_SELL_SPECIAL,
|
||||
STOCK_BUY,
|
||||
STOCK_SELL,
|
||||
# 委托类型 - 股票期权 / 期货期权
|
||||
OPTION_FUTURE_OPTION_EXERCISE,
|
||||
STOCK_OPTION_BUY_CLOSE,
|
||||
STOCK_OPTION_BUY_OPEN,
|
||||
STOCK_OPTION_CALL_EXERCISE,
|
||||
STOCK_OPTION_COVERED_CLOSE,
|
||||
STOCK_OPTION_COVERED_OPEN,
|
||||
STOCK_OPTION_PUT_EXERCISE,
|
||||
STOCK_OPTION_SECU_LOCK,
|
||||
STOCK_OPTION_SECU_UNLOCK,
|
||||
STOCK_OPTION_SELL_CLOSE,
|
||||
STOCK_OPTION_SELL_OPEN,
|
||||
# 报价类型(市价)
|
||||
FIX_PRICE,
|
||||
LATEST_PRICE,
|
||||
MARKET_MINE_PRICE_FIRST,
|
||||
MARKET_PEER_PRICE_FIRST,
|
||||
MARKET_SH_CONVERT_5_CANCEL,
|
||||
MARKET_SH_CONVERT_5_LIMIT,
|
||||
MARKET_SZ_CONVERT_5_CANCEL,
|
||||
MARKET_SZ_FULL_OR_CANCEL,
|
||||
MARKET_SZ_INSTBUSI_RESTCANCEL,
|
||||
# 市场代码
|
||||
SH_MARKET,
|
||||
SZ_MARKET,
|
||||
# 委托状态
|
||||
ORDER_CANCELED,
|
||||
ORDER_JUNK,
|
||||
ORDER_PARTSUCC_CANCEL,
|
||||
ORDER_PART_CANCEL,
|
||||
ORDER_PART_SUCC,
|
||||
ORDER_REPORTED,
|
||||
ORDER_REPORTED_CANCEL,
|
||||
ORDER_SUCCEEDED,
|
||||
ORDER_UNKNOWN,
|
||||
ORDER_UNREPORTED,
|
||||
ORDER_WAIT_REPORTING,
|
||||
# 账号状态
|
||||
ACCOUNT_STATUS_ASSIS_FAIL,
|
||||
ACCOUNT_STATUS_CLOSED,
|
||||
ACCOUNT_STATUS_CORRECTING,
|
||||
ACCOUNT_STATUS_DISABLEBYSYS,
|
||||
ACCOUNT_STATUS_DISABLEBYUSER,
|
||||
ACCOUNT_STATUS_FAIL,
|
||||
ACCOUNT_STATUS_INITING,
|
||||
ACCOUNT_STATUS_INVALID,
|
||||
ACCOUNT_STATUS_OK,
|
||||
ACCOUNT_STATUS_WAITING_LOGIN,
|
||||
ACCOUNT_STATUSING,
|
||||
)
|
||||
|
||||
# 合法委托类型集合(对齐原生 ORDER_TYPE_SET)
|
||||
ORDER_TYPE_SET = {
|
||||
STOCK_BUY,
|
||||
STOCK_SELL,
|
||||
CREDIT_BUY,
|
||||
CREDIT_SELL,
|
||||
CREDIT_FIN_BUY,
|
||||
CREDIT_SLO_SELL,
|
||||
CREDIT_BUY_SECU_REPAY,
|
||||
CREDIT_DIRECT_SECU_REPAY,
|
||||
CREDIT_SELL_SECU_REPAY,
|
||||
CREDIT_DIRECT_CASH_REPAY,
|
||||
CREDIT_FIN_BUY_SPECIAL,
|
||||
CREDIT_SLO_SELL_SPECIAL,
|
||||
CREDIT_BUY_SECU_REPAY_SPECIAL,
|
||||
CREDIT_DIRECT_SECU_REPAY_SPECIAL,
|
||||
CREDIT_SELL_SECU_REPAY_SPECIAL,
|
||||
CREDIT_DIRECT_CASH_REPAY_SPECIAL,
|
||||
}
|
||||
@@ -1,165 +0,0 @@
|
||||
import bigqmt_signal_trader.xtquant_compat as _compat
|
||||
|
||||
|
||||
def __getattr__(name):
|
||||
return getattr(_compat.xtdata, name)
|
||||
|
||||
|
||||
def get_full_tick(code_list):
|
||||
return _compat.xtdata.get_full_tick(code_list)
|
||||
|
||||
|
||||
def get_market_data(field_list=[], stock_list=[], period="1d", start_time="", end_time="", count=-1, dividend_type="none", fill_data=True):
|
||||
return _compat.xtdata.get_market_data(field_list, stock_list, period, start_time, end_time, count, dividend_type, fill_data)
|
||||
|
||||
|
||||
def get_market_data_ex(field_list=[], stock_list=[], period="1d", start_time="", end_time="", count=-1, dividend_type="none", fill_data=True):
|
||||
return _compat.xtdata.get_market_data_ex(field_list, stock_list, period, start_time, end_time, count, dividend_type, fill_data)
|
||||
|
||||
|
||||
def get_local_data(field_list=[], stock_list=[], period="1d", start_time="", end_time="", count=-1, dividend_type="none", fill_data=True, data_dir=None):
|
||||
return _compat.xtdata.get_local_data(field_list, stock_list, period, start_time, end_time, count, dividend_type, fill_data, data_dir)
|
||||
|
||||
|
||||
def get_instrument_detail(stock_code):
|
||||
return _compat.xtdata.get_instrument_detail(stock_code)
|
||||
|
||||
|
||||
def get_instrumentdetail(stock_code):
|
||||
return _compat.xtdata.get_instrumentdetail(stock_code)
|
||||
|
||||
|
||||
def get_instrument_type(stock_code, variety_list=None):
|
||||
return _compat.xtdata.get_instrument_type(stock_code, variety_list)
|
||||
|
||||
|
||||
def get_stock_list_in_sector(sector_name, real_timetag=-1):
|
||||
return _compat.xtdata.get_stock_list_in_sector(sector_name, real_timetag=real_timetag)
|
||||
|
||||
|
||||
def get_sector_list():
|
||||
return _compat.xtdata.get_sector_list()
|
||||
|
||||
|
||||
def get_sector_info(sector_name=""):
|
||||
return _compat.xtdata.get_sector_info(sector_name)
|
||||
|
||||
|
||||
def subscribe_quote(stock_code, period="1d", start_time="", end_time="", count=0, callback=None):
|
||||
return _compat.xtdata.subscribe_quote(stock_code, period, start_time, end_time, count, callback)
|
||||
|
||||
|
||||
def subscribe_quote2(stock_code, period="1d", start_time="", end_time="", count=0, dividend_type=None, callback=None):
|
||||
return _compat.xtdata.subscribe_quote2(stock_code, period, start_time, end_time, count, dividend_type, callback)
|
||||
|
||||
|
||||
def subscribe_whole_quote(code_list, callback=None):
|
||||
return _compat.xtdata.subscribe_whole_quote(code_list, callback=callback)
|
||||
|
||||
|
||||
def unsubscribe_quote(seq):
|
||||
return _compat.xtdata.unsubscribe_quote(seq)
|
||||
|
||||
|
||||
def run():
|
||||
return _compat.xtdata.run()
|
||||
|
||||
|
||||
def get_divid_factors(stock_code, start_time="", end_time=""):
|
||||
return _compat.xtdata.get_divid_factors(stock_code, start_time, end_time)
|
||||
|
||||
|
||||
def getDividFactors(*args, **kwargs):
|
||||
return _compat.xtdata.get_divid_factors(*args, **kwargs)
|
||||
|
||||
|
||||
def submit_download_history_data(stock_code, period, start_time="", end_time="", incrementally=None):
|
||||
return _compat.xtdata.submit_download_history_data(stock_code, period, start_time, end_time, incrementally)
|
||||
|
||||
|
||||
def submit_download_history_data2(stock_list, period, start_time="", end_time="", incrementally=None):
|
||||
return _compat.xtdata.submit_download_history_data2(stock_list, period, start_time, end_time, incrementally)
|
||||
|
||||
|
||||
def get_download_status(job_id):
|
||||
return _compat.xtdata.get_download_status(job_id)
|
||||
|
||||
|
||||
def wait_download(job_id, timeout=None, poll_interval=None, callback=None):
|
||||
return _compat.xtdata.wait_download(job_id, timeout, poll_interval, callback)
|
||||
|
||||
|
||||
def download_history_data(stock_code, period, start_time="", end_time="", incrementally=None):
|
||||
return _compat.xtdata.download_history_data(stock_code, period, start_time, end_time, incrementally)
|
||||
|
||||
|
||||
def download_history_data2(stock_list, period, start_time="", end_time="", callback=None, incrementally=None):
|
||||
return _compat.xtdata.download_history_data2(stock_list, period, start_time, end_time, callback, incrementally)
|
||||
|
||||
|
||||
def get_trading_dates(market, start_time="", end_time="", count=-1):
|
||||
return _compat.xtdata.get_trading_dates(market, start_time, end_time, count)
|
||||
|
||||
|
||||
def get_holidays():
|
||||
return _compat.xtdata.get_holidays()
|
||||
|
||||
|
||||
def download_holiday_data(incrementally=True):
|
||||
return _compat.xtdata.download_holiday_data(incrementally)
|
||||
|
||||
|
||||
def get_ipo_info(start_time="", end_time=""):
|
||||
return _compat.xtdata.get_ipo_info(start_time, end_time)
|
||||
|
||||
|
||||
def get_etf_info():
|
||||
return _compat.xtdata.get_etf_info()
|
||||
|
||||
|
||||
def download_etf_info():
|
||||
return _compat.xtdata.download_etf_info()
|
||||
|
||||
|
||||
def get_option_list(undl_code, dedate, opttype="", isavailavle=False):
|
||||
return _compat.xtdata.get_option_list(undl_code, dedate, opttype, isavailavle)
|
||||
|
||||
|
||||
def get_his_option_list(undl_code, dedate):
|
||||
return _compat.xtdata.get_his_option_list(undl_code, dedate)
|
||||
|
||||
|
||||
def get_his_option_list_batch(undl_code, start_time="", end_time=""):
|
||||
return _compat.xtdata.get_his_option_list_batch(undl_code, start_time, end_time)
|
||||
|
||||
|
||||
def get_financial_data(stock_list, table_list=[], start_time="", end_time="", report_type="report_time"):
|
||||
return _compat.xtdata.get_financial_data(stock_list, table_list, start_time, end_time, report_type)
|
||||
|
||||
|
||||
def download_financial_data(stock_list, table_list=[], start_time="", end_time="", incrementally=None):
|
||||
return _compat.xtdata.download_financial_data(stock_list, table_list, start_time, end_time, incrementally)
|
||||
|
||||
|
||||
def download_financial_data2(stock_list, table_list=[], start_time="", end_time="", callback=None):
|
||||
return _compat.xtdata.download_financial_data2(stock_list, table_list, start_time, end_time, callback)
|
||||
|
||||
|
||||
def call_formula(formula_name, stock_code, period, start_time="", end_time="", count=-1, dividend_type=None, extend_param={}):
|
||||
return _compat.xtdata.call_formula(formula_name, stock_code, period, start_time, end_time, count, dividend_type, extend_param)
|
||||
|
||||
|
||||
def subscribe_formula(formula_name, stock_code, period, start_time="", end_time="", count=-1, dividend_type=None, extend_param={}, callback=None):
|
||||
return _compat.xtdata.subscribe_formula(formula_name, stock_code, period, start_time, end_time, count, dividend_type, extend_param, callback)
|
||||
|
||||
|
||||
def unsubscribe_formula(request_id):
|
||||
return _compat.xtdata.unsubscribe_formula(request_id)
|
||||
|
||||
|
||||
def get_formula_result(request_id, start_time="", end_time="", count=-1, timeout_second=-1):
|
||||
return _compat.xtdata.get_formula_result(request_id, start_time, end_time, count, timeout_second)
|
||||
|
||||
|
||||
def gen_factor_index(data_name, formula_name, vars, sector_list, start_time="", end_time="", period="1d", dividend_type="none"):
|
||||
return _compat.xtdata.gen_factor_index(data_name, formula_name, vars, sector_list, start_time, end_time, period, dividend_type)
|
||||
@@ -1,7 +0,0 @@
|
||||
from bigqmt_signal_trader.xtquant_compat import (
|
||||
BigQmtXtTrader,
|
||||
XtQuantTrader,
|
||||
XtQuantTraderCallback,
|
||||
)
|
||||
|
||||
__all__ = ["BigQmtXtTrader", "XtQuantTrader", "XtQuantTraderCallback"]
|
||||
@@ -1,3 +0,0 @@
|
||||
from bigqmt_signal_trader.xtquant_compat import StockAccount
|
||||
|
||||
__all__ = ["StockAccount"]
|
||||
@@ -1,300 +0,0 @@
|
||||
# coding: utf-8
|
||||
"""Systematically test all RPC APIs and MiniQMT alias mapping (end-to-end validation).
|
||||
|
||||
For each method: call it with sensible params, report ok/error, and VALIDATE the
|
||||
result is actually correct (not just "call succeeded"). Catches silent failures
|
||||
like:
|
||||
- get_positions returns {} when the account HAS positions
|
||||
- submit_order returns SUBMITTED but the order never entered the system
|
||||
- query_orders returns [] because strategy_name didn't match
|
||||
- client transport (redis) doesn't match server (zmq) → timeout
|
||||
|
||||
Config is read from bigqmt_signal_trader_local_config (gitignored) or env vars;
|
||||
no credentials are hard-coded here. Run from a dir where that config module
|
||||
resolves, e.g.:
|
||||
|
||||
PYTHONPATH="src;D:\\国金证券QMT交易端\\python" python test_all_apis.py
|
||||
|
||||
or set BIGQMT_ACCOUNT_ID / BIGQMT_REDIS_HOST / BIGQMT_REDIS_PORT /
|
||||
BIGQMT_REDIS_DB / BIGQMT_REDIS_PASSWORD env vars.
|
||||
"""
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
|
||||
# Add src to path so bigqmt_signal_trader resolves when run from repo root.
|
||||
sys.path.insert(0, os.path.join(os.path.dirname(os.path.abspath(__file__)), "src"))
|
||||
|
||||
import redis
|
||||
|
||||
from bigqmt_signal_trader.redis_rpc import call_redis_rpc
|
||||
|
||||
|
||||
def _load_account_and_redis():
|
||||
cfg = {}
|
||||
try:
|
||||
import bigqmt_signal_trader_local_config as _c # noqa
|
||||
cfg = getattr(_c, "BIGQMT_REDIS_CONFIG", {}) or {}
|
||||
account = getattr(_c, "BIGQMT_ACCOUNT_ID", None) or cfg.get("account_id")
|
||||
except Exception:
|
||||
account = None
|
||||
account = account or os.environ.get("BIGQMT_ACCOUNT_ID", "")
|
||||
redis_cfg = dict(
|
||||
host=cfg.get("host") or os.environ.get("BIGQMT_REDIS_HOST", "127.0.0.1"),
|
||||
port=int(cfg.get("port") or os.environ.get("BIGQMT_REDIS_PORT", 6379)),
|
||||
db=int(cfg.get("db") or os.environ.get("BIGQMT_REDIS_DB", 5)),
|
||||
password=cfg.get("password", os.environ.get("BIGQMT_REDIS_PASSWORD", "")),
|
||||
socket_timeout=15,
|
||||
)
|
||||
if not redis_cfg["password"]:
|
||||
redis_cfg.pop("password")
|
||||
return str(account), redis_cfg
|
||||
|
||||
|
||||
ACCOUNT, REDIS = _load_account_and_redis()
|
||||
# account_id placeholder filled in main() once ACCOUNT is confirmed.
|
||||
_ACCT_PARAM = {"account_id": None}
|
||||
|
||||
# (method, params, label) — params chosen to be valid during/after market hours
|
||||
TESTS = [
|
||||
# --- 行情快照 ---
|
||||
("get_full_tick", {"codes": ["000001.SZ"]}, "tick"),
|
||||
("get_ticks", {"codes": ["000001.SZ"]}, "ticks-alias"),
|
||||
# --- 合约/品种 ---
|
||||
("get_instrument", {"code": "000001.SZ"}, "instrument"),
|
||||
("get_instrument_detail", {"code": "000001.SZ"}, "instrument-alias"),
|
||||
("get_instrumentdetail", {"code": "000001.SZ"}, "instrument-alias2"),
|
||||
("get_instrument_type", {"code": "000001.SZ", "variety_list": ["stock", "fund"]}, "inst-type"),
|
||||
# --- K线/历史 ---
|
||||
("get_market_data_ex", {"field_list": ["close"], "stock_list": ["000001.SZ"], "period": "1d", "count": 3}, "md-ex"),
|
||||
("get_market_data", {"field_list": ["close"], "stock_list": ["000001.SZ"], "period": "1d", "count": 3}, "md"),
|
||||
("get_local_data", {"field_list": ["close"], "stock_list": ["000001.SZ"], "period": "1d", "count": 3}, "local-data"),
|
||||
# --- 板块 ---
|
||||
("get_sector_list", {}, "sector-list"),
|
||||
("get_stock_list_in_sector", {"sector_name": "沪深A股"}, "sector-stocks"),
|
||||
# --- 交易日历 ---
|
||||
("get_trading_dates", {"market": "SH", "count": 3}, "trade-dates"),
|
||||
("get_holidays", {}, "holidays"),
|
||||
("get_markets", {}, "markets"),
|
||||
("get_market_last_trade_date", {"market": "SH"}, "last-trade-date"),
|
||||
# --- 账户 ---
|
||||
("get_asset", {}, "asset"),
|
||||
("get_positions", {}, "positions"),
|
||||
("query_stock_asset", dict(_ACCT_PARAM), "asset-alias"),
|
||||
("query_stock_positions", dict(_ACCT_PARAM), "positions-alias"),
|
||||
("query_stock_position", dict(stock_code="000001.SZ", **_ACCT_PARAM), "position-single"),
|
||||
]
|
||||
|
||||
|
||||
def data_summary(data):
|
||||
"""One-line summary of returned data for readability."""
|
||||
if data is None:
|
||||
return "None"
|
||||
if isinstance(data, dict):
|
||||
if not data:
|
||||
return "{}"
|
||||
if "__bigqmt_type__" in data:
|
||||
return "[%s cols=%s records=%d]" % (
|
||||
data.get("__bigqmt_type__"),
|
||||
data.get("columns"),
|
||||
len(data.get("records") or []),
|
||||
)
|
||||
keys = list(data.keys())[:3]
|
||||
return "{%s%s: ...}(%d keys)" % (keys, "" if len(keys) < 3 else ", ...", len(data))
|
||||
if isinstance(data, list):
|
||||
return "[list len=%d]" % len(data)
|
||||
return repr(data)[:60]
|
||||
|
||||
|
||||
def _is_empty(data):
|
||||
return data is None or data == {} or data == [] or data == ""
|
||||
|
||||
|
||||
def _call(r, method, params, timeout=12):
|
||||
"""Call and return (response, latency_ms, error_str)."""
|
||||
t0 = time.time()
|
||||
try:
|
||||
resp = call_redis_rpc(r, ACCOUNT, method, params, timeout_seconds=timeout)
|
||||
return resp, (time.time() - t0) * 1000, None
|
||||
except Exception as e:
|
||||
return None, (time.time() - t0) * 1000, str(e)
|
||||
|
||||
|
||||
def main():
|
||||
if not ACCOUNT:
|
||||
raise SystemExit("ACCOUNT is empty: set BIGQMT_ACCOUNT_ID or configure bigqmt_signal_trader_local_config")
|
||||
# Fill the account_id into the account-query test params now that we know it.
|
||||
for i, (method, params, label) in enumerate(TESTS):
|
||||
if "account_id" in params and params["account_id"] is None:
|
||||
params["account_id"] = ACCOUNT
|
||||
|
||||
r = redis.Redis(**REDIS)
|
||||
|
||||
print("=" * 90)
|
||||
print("全量 API 测试 (account=%s) — 端到端验证" % ACCOUNT)
|
||||
print("=" * 90)
|
||||
|
||||
# === 端到端验证 0: 客户端/服务端 transport 一致性 ===
|
||||
print("\n--- 端到端验证: 客户端/服务端一致性 ---")
|
||||
# 检测客户端配置里的 transport
|
||||
client_transport = "redis" # 默认
|
||||
try:
|
||||
import bigqmt_signal_trader_local_config as _c
|
||||
client_transport = str(getattr(_c, "BIGQMT_REDIS_CONFIG", {}).get("transport", "redis")).lower()
|
||||
except Exception:
|
||||
pass
|
||||
print("客户端配置 transport: %s" % client_transport)
|
||||
|
||||
# 如果客户端是 zmq 但服务端不是, ping 会超时
|
||||
ping_resp, ping_ms, ping_err = _call(r, "ping", {}, timeout=8)
|
||||
if ping_err:
|
||||
print("❌ ping 失败: %s" % ping_err)
|
||||
if "timeout" in ping_err.lower():
|
||||
print(" 可能原因: 客户端 transport 和服务端不匹配")
|
||||
print(" - 客户端配置 transport=%s" % client_transport)
|
||||
print(" - 如果服务端是 zmq, 客户端也要设 transport=zmq")
|
||||
print(" - 如果服务端是 redis, 客户端保持 redis 即可")
|
||||
return
|
||||
print("✅ ping OK (%.0fms) — 客户端/服务端连通" % ping_ms)
|
||||
|
||||
# === 端到端验证 2: 账户有持仓时 get_positions 必须返回非空 ===
|
||||
print("\n--- 端到端验证: 持仓查询 ---")
|
||||
pos_resp, pos_ms, pos_err = _call(r, "get_positions", {}, timeout=12)
|
||||
if pos_err:
|
||||
print("❌ get_positions 失败: %s" % pos_err)
|
||||
elif not pos_resp.get("ok"):
|
||||
print("❌ get_positions 返回错误: %s" % pos_resp.get("error"))
|
||||
else:
|
||||
positions = pos_resp.get("data") or {}
|
||||
if len(positions) > 0:
|
||||
print("✅ get_positions OK (%.0fms) — 返回 %d 只持仓" % (pos_ms, len(positions)))
|
||||
else:
|
||||
print("⚠️ get_positions 返回空 — 账户可能真的没持仓, 或查询失败 (检查 QMT 上下文)")
|
||||
|
||||
# === 端到端验证 3: query_orders 验证 (strategy_name 陷阱) ===
|
||||
print("\n--- 端到端验证: 委托查询 ---")
|
||||
ord_resp, ord_ms, ord_err = _call(r, "query_orders", {}, timeout=12)
|
||||
if ord_err:
|
||||
print("❌ query_orders 失败: %s" % ord_err)
|
||||
elif not ord_resp.get("ok"):
|
||||
print("❌ query_orders 返回错误: %s" % ord_resp.get("error"))
|
||||
else:
|
||||
orders = ord_resp.get("data") or []
|
||||
if len(orders) > 0:
|
||||
print("✅ query_orders OK (%.0fms) — 返回 %d 条委托" % (ord_ms, len(orders)))
|
||||
else:
|
||||
print("⚠️ query_orders 返回空 — 可能 strategy_name 不匹配 (默认应为 '' 返回全部)")
|
||||
|
||||
# === 端到端验证 4: 买入/卖出后委托必须进系统 ===
|
||||
print("\n--- 端到端验证: 买入/卖出 (仅交易时段) ---")
|
||||
# 用极低价格买入 (确保不成交), 然后查委托确认进了系统
|
||||
# 先拿一只股票的现价
|
||||
tick_resp, _, tick_err = _call(r, "get_full_tick", {"codes": ["600654.SH"]}, timeout=12)
|
||||
if tick_err or not tick_resp.get("ok"):
|
||||
print("⚠️ 跳过买入测试 (get_full_tick 失败: %s)" % (tick_err or tick_resp.get("error")))
|
||||
else:
|
||||
d = (tick_resp.get("data") or {}).get("600654.SH", {})
|
||||
last_close = float(d.get("lastClose") or d.get("lastPrice") or 3.0)
|
||||
buy_price = round(last_close * 0.8, 2) # 跌停价, 确保不成交
|
||||
print(" 用 600654.SH @%.2f 买入 100 股 (跌停价, 不成交)" % buy_price)
|
||||
|
||||
# 下单前委托数
|
||||
ord_before, _, _ = _call(r, "query_orders", {}, timeout=12)
|
||||
before_count = len((ord_before or {}).get("data") or []) if ord_before else 0
|
||||
|
||||
# 下单
|
||||
sub_resp, sub_ms, sub_err = _call(r, "submit_order", {
|
||||
"stock_code": "600654.SH", "action": "BUY", "volume": 100,
|
||||
"price": buy_price, "price_type": "LIMIT", "strategy_name": "rpc_test",
|
||||
"signal_id": "e2e-test-%d" % int(time.time()),
|
||||
}, timeout=15)
|
||||
if sub_err:
|
||||
print("❌ submit_order 失败: %s" % sub_err)
|
||||
elif not sub_resp.get("ok"):
|
||||
print("❌ submit_order 返回错误: %s" % sub_resp.get("error"))
|
||||
else:
|
||||
server_err = sub_resp.get("server_error") or ""
|
||||
print("✅ submit_order OK (%.0fms)" % sub_ms)
|
||||
if server_err:
|
||||
print(" ⚠️ server_error: %s" % server_err)
|
||||
|
||||
# 等 1s 让 QMT 处理, 然后查委托确认进了系统
|
||||
time.sleep(1)
|
||||
ord_after, _, _ = _call(r, "query_orders", {}, timeout=12)
|
||||
after_orders = (ord_after or {}).get("data") or [] if ord_after else []
|
||||
found = any(
|
||||
str(o.get("stock_code") or "").upper() == "600654.SH"
|
||||
and str(o.get("action") or "").upper() == "BUY"
|
||||
and abs(float(o.get("price") or 0) - buy_price) < 0.01
|
||||
for o in after_orders
|
||||
)
|
||||
if found:
|
||||
print("✅ 委托已进系统 (query_orders 确认)")
|
||||
# 尝试撤单
|
||||
oid = None
|
||||
for o in after_orders:
|
||||
if (str(o.get("stock_code") or "").upper() == "600654.SH"
|
||||
and str(o.get("action") or "").upper() == "BUY"
|
||||
and abs(float(o.get("price") or 0) - buy_price) < 0.01):
|
||||
oid = str(o.get("order_sys_id") or "")
|
||||
break
|
||||
if oid:
|
||||
cancel_resp, cancel_ms, cancel_err = _call(r, "cancel_order", {
|
||||
"order_sys_id": oid, "market": "SH"
|
||||
}, timeout=15)
|
||||
if cancel_err:
|
||||
print("⚠️ cancel_order 失败: %s" % cancel_err)
|
||||
elif cancel_resp and cancel_resp.get("ok"):
|
||||
print("✅ cancel_order OK (%.0fms) — 已撤单" % cancel_ms)
|
||||
else:
|
||||
print("⚠️ cancel_order 返回: %s" % (cancel_resp or {}))
|
||||
else:
|
||||
print("❌ 委托没进系统 — submit_order 成功但 query_orders 找不到")
|
||||
print(" 这是静默失败 (passorder 被 QMT 拒绝但没报错)")
|
||||
print(" 检查: 1) 价格是否超出范围 2) 账户权限 3) QMT 风控")
|
||||
|
||||
# === 全量 API 测试 ===
|
||||
print("\n" + "=" * 90)
|
||||
print("全量 API 测试")
|
||||
print("=" * 90)
|
||||
print("%-22s %-8s %-8s %s" % ("method", "ok", "ms", "data summary"))
|
||||
print("-" * 90)
|
||||
|
||||
results = {"ok": [], "ok_empty": [], "fail": [], "timeout": []}
|
||||
for method, params, label in TESTS:
|
||||
resp, dt, err = _call(r, method, params, timeout=12)
|
||||
if err:
|
||||
is_timeout = "timeout" in err.lower()
|
||||
bucket = "timeout" if is_timeout else "fail"
|
||||
results[bucket].append((method, err[:60]))
|
||||
print("%-22s %-8s %6.0f %s" % (method, "TIMEOUT" if is_timeout else "ERROR", dt, err[:50]))
|
||||
continue
|
||||
ok = resp.get("ok")
|
||||
data = resp.get("data")
|
||||
error = resp.get("error", "")
|
||||
server_err = resp.get("server_error", "")
|
||||
empty = _is_empty(data)
|
||||
if ok and not empty:
|
||||
results["ok"].append(method)
|
||||
status = "OK"
|
||||
elif ok and empty:
|
||||
results["ok_empty"].append(method)
|
||||
status = "EMPTY"
|
||||
else:
|
||||
results["fail"].append((method, error))
|
||||
status = "FAIL"
|
||||
summary = data_summary(data) if ok else error[:50]
|
||||
if server_err:
|
||||
summary += " [server_error: %s]" % server_err[:40]
|
||||
print("%-22s %-8s %6.0f %s" % (method, status, dt, summary))
|
||||
|
||||
print("-" * 90)
|
||||
print("\n=== 汇总 ===")
|
||||
print("有数据 (OK): %d 个" % len(results["ok"]))
|
||||
print("成功但空 (EMPTY): %d 个 %s" % (len(results["ok_empty"]), results["ok_empty"]))
|
||||
print("失败 (FAIL): %d 个 %s" % (len(results["fail"]), [m for m, _ in results["fail"]]))
|
||||
print("超时 (TIMEOUT): %d 个 %s" % (len(results["timeout"]), [m for m, _ in results["timeout"]]))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,286 +0,0 @@
|
||||
import csv
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_backtest.data_feed import CsvBarFeed
|
||||
from bigqmt_backtest.engine import BacktestConfig, BacktestEngine
|
||||
|
||||
|
||||
FIELDS = (
|
||||
"datetime",
|
||||
"symbol",
|
||||
"open",
|
||||
"high",
|
||||
"low",
|
||||
"close",
|
||||
"volume",
|
||||
"prev_close",
|
||||
)
|
||||
|
||||
|
||||
def _write_bars(path, rows):
|
||||
with open(path, "w", encoding="utf-8", newline="") as handle:
|
||||
writer = csv.DictWriter(handle, fieldnames=FIELDS)
|
||||
writer.writeheader()
|
||||
writer.writerows(rows)
|
||||
|
||||
|
||||
def _rows():
|
||||
return [
|
||||
{
|
||||
"datetime": "2026-01-05 09:30:00",
|
||||
"symbol": "600000.SH",
|
||||
"open": 10.00,
|
||||
"high": 10.10,
|
||||
"low": 9.95,
|
||||
"close": 10.05,
|
||||
"volume": 100000,
|
||||
"prev_close": 9.90,
|
||||
},
|
||||
{
|
||||
"datetime": "2026-01-05 09:31:00",
|
||||
"symbol": "600000.SH",
|
||||
"open": 10.10,
|
||||
"high": 10.20,
|
||||
"low": 10.00,
|
||||
"close": 10.15,
|
||||
"volume": 100000,
|
||||
"prev_close": 9.90,
|
||||
},
|
||||
{
|
||||
"datetime": "2026-01-05 09:32:00",
|
||||
"symbol": "600000.SH",
|
||||
"open": 10.20,
|
||||
"high": 10.25,
|
||||
"low": 10.10,
|
||||
"close": 10.18,
|
||||
"volume": 100000,
|
||||
"prev_close": 9.90,
|
||||
},
|
||||
{
|
||||
"datetime": "2026-01-06 09:30:00",
|
||||
"symbol": "600000.SH",
|
||||
"open": 10.30,
|
||||
"high": 10.35,
|
||||
"low": 10.20,
|
||||
"close": 10.25,
|
||||
"volume": 100000,
|
||||
"prev_close": 10.18,
|
||||
},
|
||||
{
|
||||
"datetime": "2026-01-06 09:31:00",
|
||||
"symbol": "600000.SH",
|
||||
"open": 10.25,
|
||||
"high": 10.30,
|
||||
"low": 10.15,
|
||||
"close": 10.20,
|
||||
"volume": 100000,
|
||||
"prev_close": 10.18,
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def _config(run_id, output_dir):
|
||||
return BacktestConfig(
|
||||
run_id=run_id,
|
||||
output_dir=output_dir,
|
||||
initial_cash=100000,
|
||||
buy_commission_rate=0.0003,
|
||||
sell_commission_rate=0.0003,
|
||||
min_commission=5,
|
||||
stamp_tax_rate=0.0005,
|
||||
transfer_fee_rate=0.00001,
|
||||
max_volume_participation=1.0,
|
||||
slippage_bps=0,
|
||||
)
|
||||
|
||||
|
||||
class CsvBarFeedTest(unittest.TestCase):
|
||||
def test_loads_chronologically_and_never_exposes_future_rows(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
path = os.path.join(tmp, "bars.csv")
|
||||
rows = list(reversed(_rows()))
|
||||
_write_bars(path, rows)
|
||||
|
||||
feed = CsvBarFeed(path)
|
||||
|
||||
self.assertEqual(len(feed), 5)
|
||||
self.assertEqual(feed.frame(0)["datetime"], "2026-01-05 09:30:00")
|
||||
self.assertEqual(
|
||||
list(feed.history("600000.SH", end_index=1, count=10, fields=["close"]))[-1]["close"],
|
||||
10.15,
|
||||
)
|
||||
self.assertNotIn(10.18, [item["close"] for item in feed.history("600000.SH", 1, 10)])
|
||||
self.assertEqual(len(feed.data_hash), 64)
|
||||
|
||||
def test_duplicate_symbol_timestamp_is_rejected(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
path = os.path.join(tmp, "bars.csv")
|
||||
rows = _rows()
|
||||
_write_bars(path, rows + [dict(rows[0])])
|
||||
|
||||
with self.assertRaisesRegex(ValueError, "duplicate bar"):
|
||||
CsvBarFeed(path)
|
||||
|
||||
|
||||
class BacktestEngineTest(unittest.TestCase):
|
||||
def _engine(self, tmp, run_id="run-a"):
|
||||
path = os.path.join(tmp, "bars.csv")
|
||||
_write_bars(path, _rows())
|
||||
return BacktestEngine(CsvBarFeed(path), _config(run_id, os.path.join(tmp, run_id)))
|
||||
|
||||
def test_order_is_filled_at_next_bar_open_not_current_close(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
engine = self._engine(tmp)
|
||||
started = engine.start()
|
||||
order = engine.submit_order(
|
||||
{"symbol": "600000.SH", "side": "BUY", "quantity": 100, "order_type": "MARKET"}
|
||||
)
|
||||
|
||||
self.assertEqual(started["frame_index"], 0)
|
||||
self.assertEqual(order["status"], "PENDING")
|
||||
self.assertEqual(engine.state()["positions"], {})
|
||||
|
||||
advanced = engine.next_bar()
|
||||
|
||||
self.assertEqual(advanced["frame_index"], 1)
|
||||
self.assertEqual(advanced["fills"][0]["price"], 10.10)
|
||||
self.assertEqual(advanced["fills"][0]["commission"], 5.0)
|
||||
self.assertEqual(advanced["fills"][0]["transfer_fee"], 0.01)
|
||||
self.assertEqual(advanced["cash"], 98984.99)
|
||||
self.assertEqual(advanced["positions"]["600000.SH"]["quantity"], 100)
|
||||
self.assertEqual(advanced["positions"]["600000.SH"]["available"], 0)
|
||||
|
||||
def test_t_plus_one_rejects_same_day_sell_and_allows_next_day(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
engine = self._engine(tmp)
|
||||
engine.start()
|
||||
engine.submit_order(
|
||||
{"symbol": "600000.SH", "side": "BUY", "quantity": 100, "order_type": "MARKET"}
|
||||
)
|
||||
engine.next_bar()
|
||||
|
||||
rejected = engine.submit_order(
|
||||
{"symbol": "600000.SH", "side": "SELL", "quantity": 100, "order_type": "MARKET"}
|
||||
)
|
||||
self.assertEqual(rejected["status"], "REJECTED")
|
||||
self.assertEqual(rejected["reject_reason"], "t_plus_one_unavailable")
|
||||
|
||||
engine.next_bar()
|
||||
next_day = engine.next_bar()
|
||||
self.assertEqual(next_day["positions"]["600000.SH"]["available"], 100)
|
||||
accepted = engine.submit_order(
|
||||
{"symbol": "600000.SH", "side": "SELL", "quantity": 100, "order_type": "MARKET"}
|
||||
)
|
||||
self.assertEqual(accepted["status"], "PENDING")
|
||||
filled = engine.next_bar()
|
||||
self.assertEqual(filled["fills"][0]["side"], "SELL")
|
||||
self.assertEqual(filled["fills"][0]["stamp_tax"], 0.51)
|
||||
self.assertEqual(filled["total_fees"], 10.53)
|
||||
self.assertEqual(filled["cash"], 100004.47)
|
||||
self.assertEqual(filled["positions"], {})
|
||||
|
||||
def test_limit_locked_bar_does_not_assume_fill(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
path = os.path.join(tmp, "locked.csv")
|
||||
rows = _rows()[:2]
|
||||
rows[1].update({"open": 10.89, "high": 10.89, "low": 10.89, "close": 10.89, "prev_close": 9.90})
|
||||
_write_bars(path, rows)
|
||||
engine = BacktestEngine(CsvBarFeed(path), _config("locked", os.path.join(tmp, "locked")))
|
||||
engine.start()
|
||||
engine.submit_order(
|
||||
{"symbol": "600000.SH", "side": "BUY", "quantity": 100, "order_type": "MARKET"}
|
||||
)
|
||||
|
||||
state = engine.next_bar()
|
||||
|
||||
self.assertEqual(state["fills"], [])
|
||||
self.assertEqual(engine.orders()[0]["status"], "EXPIRED")
|
||||
self.assertEqual(engine.orders()[0]["reject_reason"], "limit_up_locked")
|
||||
|
||||
def test_client_order_id_is_idempotent_but_cannot_change_payload(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
engine = self._engine(tmp)
|
||||
engine.start()
|
||||
payload = {
|
||||
"client_order_id": "stable-1",
|
||||
"symbol": "600000.SH",
|
||||
"side": "BUY",
|
||||
"quantity": 100,
|
||||
"order_type": "MARKET",
|
||||
}
|
||||
first = engine.submit_order(payload)
|
||||
repeated = engine.submit_order(dict(payload))
|
||||
|
||||
self.assertEqual(first["order_id"], repeated["order_id"])
|
||||
self.assertEqual(len(engine.orders()), 1)
|
||||
changed = dict(payload, quantity=200)
|
||||
with self.assertRaisesRegex(ValueError, "different order payload"):
|
||||
engine.submit_order(changed)
|
||||
|
||||
def test_volume_participation_is_shared_across_orders_on_the_same_bar(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
path = os.path.join(tmp, "volume.csv")
|
||||
rows = _rows()[:2]
|
||||
rows[1]["volume"] = 1000
|
||||
_write_bars(path, rows)
|
||||
config = _config("volume", os.path.join(tmp, "volume"))
|
||||
config.max_volume_participation = 0.1
|
||||
engine = BacktestEngine(CsvBarFeed(path), config)
|
||||
engine.start()
|
||||
for order_id in ("first", "second"):
|
||||
engine.submit_order(
|
||||
{
|
||||
"client_order_id": order_id,
|
||||
"symbol": "600000.SH",
|
||||
"side": "BUY",
|
||||
"quantity": 100,
|
||||
"order_type": "MARKET",
|
||||
}
|
||||
)
|
||||
|
||||
advanced = engine.next_bar()
|
||||
|
||||
self.assertEqual(sum(fill["quantity"] for fill in advanced["fills"]), 100)
|
||||
self.assertEqual(engine.orders()[1]["status"], "EXPIRED")
|
||||
self.assertEqual(engine.orders()[1]["reject_reason"], "volume_participation_exhausted")
|
||||
|
||||
def test_finish_writes_evidence_and_is_deterministic(self):
|
||||
signatures = []
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
for run_id in ("det-a", "det-b"):
|
||||
engine = self._engine(tmp, run_id=run_id)
|
||||
engine.start()
|
||||
engine.submit_order(
|
||||
{
|
||||
"client_order_id": "buy-1",
|
||||
"symbol": "600000.SH",
|
||||
"side": "BUY",
|
||||
"quantity": 100,
|
||||
"order_type": "MARKET",
|
||||
}
|
||||
)
|
||||
while not engine.next_bar()["done"]:
|
||||
pass
|
||||
result = engine.finish()
|
||||
signatures.append(result["deterministic_signature"])
|
||||
output_dir = os.path.join(tmp, run_id)
|
||||
for name in ("meta.json", "result.json", "orders.csv", "fills.csv", "equity.csv", "positions.csv"):
|
||||
self.assertTrue(os.path.isfile(os.path.join(output_dir, name)), name)
|
||||
with open(os.path.join(output_dir, "meta.json"), encoding="utf-8") as handle:
|
||||
meta = json.load(handle)
|
||||
self.assertEqual(meta["data_hash"], engine.feed.data_hash)
|
||||
self.assertFalse(meta["live_ready"])
|
||||
|
||||
self.assertEqual(signatures[0], signatures[1])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,162 +0,0 @@
|
||||
import datetime as dt
|
||||
import os
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
SRC = os.path.join(ROOT, "src")
|
||||
sys.path.insert(0, SRC)
|
||||
|
||||
from bigqmt_backtest.data_feed import StreamingBarFeed
|
||||
from bigqmt_backtest.engine import BacktestConfig, StreamingBacktestEngine
|
||||
from bigqmt_backtest.qmt_runtime import QmtBarExtractor, QmtNativeBacktestSession
|
||||
|
||||
|
||||
class FakeQmtContext(object):
|
||||
stockcode = "600000"
|
||||
market = "SH"
|
||||
period = "1m"
|
||||
barpos = 0
|
||||
do_back_test = True
|
||||
|
||||
values = {
|
||||
"open": [10.0],
|
||||
"high": [10.2],
|
||||
"low": [9.9],
|
||||
"close": [10.1],
|
||||
"volume": [10000],
|
||||
"amount": [101000],
|
||||
"preClose": [9.8],
|
||||
}
|
||||
|
||||
def get_bar_timetag(self, barpos):
|
||||
value = dt.datetime(2026, 1, 5, 9, 30)
|
||||
return int(value.timestamp() * 1000)
|
||||
|
||||
def get_history_data(self, count, period, field):
|
||||
return {"600000.SH": self.values.get(field, [])}
|
||||
|
||||
def set_account(self, account_id):
|
||||
self.account_id = account_id
|
||||
|
||||
|
||||
class QmtBarExtractorTest(unittest.TestCase):
|
||||
def test_extracts_qmt_bar_without_live_account_or_order_api(self):
|
||||
row = QmtBarExtractor().extract(FakeQmtContext())
|
||||
|
||||
self.assertEqual(row["symbol"], "600000.SH")
|
||||
self.assertEqual(row["datetime"], "2026-01-05 09:30:00")
|
||||
self.assertEqual(row["open"], 10.0)
|
||||
self.assertEqual(row["prev_close"], 9.8)
|
||||
|
||||
def test_qmt_entry_is_isolated_and_binds_native_order_callbacks(self):
|
||||
path = os.path.join(SRC, "BIGQMT_ZMQ_BACKTEST.py")
|
||||
with open(path, "r", encoding="gbk") as handle:
|
||||
source = handle.read()
|
||||
|
||||
self.assertNotIn("bigqmt_signal_trader", source)
|
||||
self.assertIn("passorder", source)
|
||||
self.assertIn("get_trade_detail_data", source)
|
||||
self.assertIn("order_callback", source)
|
||||
self.assertIn("deal_callback", source)
|
||||
self.assertNotIn("_importlib.import_module =", source)
|
||||
|
||||
runtime_path = os.path.join(SRC, "bigqmt_backtest", "qmt_runtime.py")
|
||||
with open(runtime_path, "r", encoding="utf-8") as handle:
|
||||
runtime_source = handle.read()
|
||||
self.assertNotIn("StreamingBacktestEngine", runtime_source)
|
||||
self.assertNotIn("SimulatedBroker", runtime_source)
|
||||
|
||||
def test_native_session_executes_passorder_on_qmt_callback_thread(self):
|
||||
calls = []
|
||||
|
||||
def fake_passorder(*args):
|
||||
calls.append((threading.get_ident(), args))
|
||||
return "qmt-order-1"
|
||||
|
||||
session = QmtNativeBacktestSession(
|
||||
config={
|
||||
"run_id": "qmt-native-test",
|
||||
"account_id": "test-account",
|
||||
"bar_wait_timeout_seconds": 1,
|
||||
},
|
||||
qmt_api={"passorder": fake_passorder},
|
||||
)
|
||||
qmt_thread = threading.Thread(target=session.on_bar, args=(FakeQmtContext(),))
|
||||
qmt_thread.start()
|
||||
|
||||
state = session.start()
|
||||
queued = session.submit_order({
|
||||
"symbol": "600000.SH",
|
||||
"side": "BUY",
|
||||
"quantity": 100,
|
||||
"order_type": "MARKET",
|
||||
"client_order_id": "external-1",
|
||||
})
|
||||
session.finish()
|
||||
qmt_thread.join(timeout=2)
|
||||
|
||||
self.assertFalse(qmt_thread.is_alive())
|
||||
self.assertEqual(state["execution_backend"], "QMT_NATIVE")
|
||||
self.assertEqual(queued["status"], "QUEUED")
|
||||
self.assertEqual(len(calls), 1)
|
||||
self.assertEqual(calls[0][0], qmt_thread.ident)
|
||||
self.assertEqual(calls[0][1][0], 23)
|
||||
self.assertEqual(calls[0][1][2], "test-account")
|
||||
self.assertEqual(calls[0][1][-2], "external-1")
|
||||
self.assertEqual(session.orders()[0]["status"], "SUBMITTED")
|
||||
|
||||
def test_native_session_rejects_non_backtest_qmt_context(self):
|
||||
context = FakeQmtContext()
|
||||
context.do_back_test = False
|
||||
session = QmtNativeBacktestSession(config={"run_id": "guard"})
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "outside QMT backtest mode"):
|
||||
session.bind_context(context)
|
||||
|
||||
|
||||
class StreamingEngineTest(unittest.TestCase):
|
||||
def test_stream_does_not_report_done_until_qmt_closes_feed(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
feed = StreamingBarFeed()
|
||||
feed.append(
|
||||
{
|
||||
"datetime": "2026-01-05 09:30:00",
|
||||
"symbol": "600000.SH",
|
||||
"open": 10,
|
||||
"high": 10.1,
|
||||
"low": 9.9,
|
||||
"close": 10,
|
||||
"volume": 10000,
|
||||
"prev_close": 9.9,
|
||||
}
|
||||
)
|
||||
engine = StreamingBacktestEngine(
|
||||
feed,
|
||||
BacktestConfig(run_id="stream", output_dir=os.path.join(tmp, "out")),
|
||||
bar_wait_timeout_seconds=0.1,
|
||||
)
|
||||
|
||||
self.assertFalse(engine.start()["done"])
|
||||
feed.append(
|
||||
{
|
||||
"datetime": "2026-01-05 09:31:00",
|
||||
"symbol": "600000.SH",
|
||||
"open": 10.1,
|
||||
"high": 10.2,
|
||||
"low": 10,
|
||||
"close": 10.15,
|
||||
"volume": 10000,
|
||||
"prev_close": 9.9,
|
||||
}
|
||||
)
|
||||
self.assertFalse(engine.next_bar()["done"])
|
||||
feed.close()
|
||||
self.assertTrue(engine.state()["done"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,250 +0,0 @@
|
||||
import csv
|
||||
import os
|
||||
import socket
|
||||
import sys
|
||||
import tempfile
|
||||
import threading
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_backtest.client import BacktestZmqClient
|
||||
from bigqmt_backtest.data_feed import CsvBarFeed
|
||||
from bigqmt_backtest.engine import BacktestConfig, BacktestEngine
|
||||
from bigqmt_backtest.protocol import BacktestBridgeProtocol
|
||||
from bigqmt_backtest.qmt_runtime import QmtNativeBacktestSession
|
||||
from bigqmt_backtest.zmq_server import ZmqBacktestServer
|
||||
|
||||
|
||||
def _free_port():
|
||||
sock = socket.socket()
|
||||
sock.bind(("127.0.0.1", 0))
|
||||
port = sock.getsockname()[1]
|
||||
sock.close()
|
||||
return port
|
||||
|
||||
|
||||
def _feed(tmp):
|
||||
path = os.path.join(tmp, "bars.csv")
|
||||
with open(path, "w", encoding="utf-8", newline="") as handle:
|
||||
writer = csv.DictWriter(
|
||||
handle,
|
||||
fieldnames=("datetime", "symbol", "open", "high", "low", "close", "volume", "prev_close"),
|
||||
)
|
||||
writer.writeheader()
|
||||
writer.writerows(
|
||||
[
|
||||
{
|
||||
"datetime": "2026-01-05 09:30:00",
|
||||
"symbol": "000001.SZ",
|
||||
"open": 10,
|
||||
"high": 10.1,
|
||||
"low": 9.9,
|
||||
"close": 10,
|
||||
"volume": 10000,
|
||||
"prev_close": 9.9,
|
||||
},
|
||||
{
|
||||
"datetime": "2026-01-05 09:31:00",
|
||||
"symbol": "000001.SZ",
|
||||
"open": 10.1,
|
||||
"high": 10.2,
|
||||
"low": 10,
|
||||
"close": 10.15,
|
||||
"volume": 10000,
|
||||
"prev_close": 9.9,
|
||||
},
|
||||
]
|
||||
)
|
||||
return CsvBarFeed(path)
|
||||
|
||||
|
||||
class BacktestProtocolTest(unittest.TestCase):
|
||||
def test_run_and_client_identity_are_enforced_and_requests_are_idempotent(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
engine = BacktestEngine(
|
||||
_feed(tmp),
|
||||
BacktestConfig(run_id="identity", output_dir=os.path.join(tmp, "out")),
|
||||
)
|
||||
protocol = BacktestBridgeProtocol(engine)
|
||||
start = {
|
||||
"schema_version": 1,
|
||||
"request_id": "req-start",
|
||||
"run_id": "identity",
|
||||
"client_id": "client-a",
|
||||
"method": "start",
|
||||
"params": {},
|
||||
}
|
||||
first = protocol.handle(start)
|
||||
repeated = protocol.handle(dict(start))
|
||||
|
||||
self.assertTrue(first["ok"])
|
||||
self.assertEqual(first, repeated)
|
||||
|
||||
reused = dict(start, method="state")
|
||||
reused_response = protocol.handle(reused)
|
||||
self.assertFalse(reused_response["ok"])
|
||||
self.assertIn("different payload", reused_response["error"])
|
||||
|
||||
wrong_run = dict(start, request_id="wrong-run", run_id="other", method="state")
|
||||
self.assertFalse(protocol.handle(wrong_run)["ok"])
|
||||
wrong_client = dict(start, request_id="wrong-client", client_id="client-b", method="next_bar")
|
||||
self.assertFalse(protocol.handle(wrong_client)["ok"])
|
||||
|
||||
|
||||
class ZmqRoundTripTest(unittest.TestCase):
|
||||
def test_external_client_can_complete_a_backtest_over_zmq(self):
|
||||
with tempfile.TemporaryDirectory() as tmp:
|
||||
engine = BacktestEngine(
|
||||
_feed(tmp),
|
||||
BacktestConfig(
|
||||
run_id="zmq-run",
|
||||
output_dir=os.path.join(tmp, "out"),
|
||||
max_volume_participation=1.0,
|
||||
),
|
||||
)
|
||||
endpoint = "tcp://127.0.0.1:%d" % _free_port()
|
||||
server = ZmqBacktestServer(
|
||||
BacktestBridgeProtocol(engine), endpoint=endpoint, exit_on_finish=True
|
||||
)
|
||||
thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
thread.start()
|
||||
self.assertTrue(server.wait_until_ready(3.0))
|
||||
|
||||
client = BacktestZmqClient(
|
||||
endpoint=endpoint,
|
||||
run_id="zmq-run",
|
||||
client_id="external-strategy",
|
||||
timeout_seconds=3.0,
|
||||
)
|
||||
started = client.start()
|
||||
order = client.submit_order(
|
||||
symbol="000001.SZ", side="BUY", quantity=100, order_type="MARKET"
|
||||
)
|
||||
advanced = client.next_bar()
|
||||
history = client.history("000001.SZ", count=10, fields=["close"])
|
||||
result = client.finish()
|
||||
client.close()
|
||||
thread.join(timeout=3.0)
|
||||
|
||||
self.assertEqual(started["frame_index"], 0)
|
||||
self.assertEqual(order["status"], "PENDING")
|
||||
self.assertEqual(advanced["fills"][0]["price"], 10.1)
|
||||
self.assertEqual([row["close"] for row in history], [10.0, 10.15])
|
||||
self.assertEqual(result["run_id"], "zmq-run")
|
||||
self.assertFalse(thread.is_alive())
|
||||
|
||||
def test_qmt_native_service_bridges_orders_to_qmt_matching(self):
|
||||
endpoint = "tcp://127.0.0.1:%d" % _free_port()
|
||||
calls = []
|
||||
holder = {}
|
||||
|
||||
class Context(object):
|
||||
do_back_test = True
|
||||
stockcode = "600000"
|
||||
market = "SH"
|
||||
period = "1m"
|
||||
|
||||
def __init__(self, barpos, close):
|
||||
self.barpos = barpos
|
||||
self.close_value = close
|
||||
|
||||
def set_account(self, account_id):
|
||||
self.account_id = account_id
|
||||
|
||||
def get_bar_timetag(self, barpos):
|
||||
return int((1704072600 + barpos * 60) * 1000)
|
||||
|
||||
def get_history_data(self, count, period, field):
|
||||
values = {
|
||||
"open": self.close_value,
|
||||
"high": self.close_value + 0.1,
|
||||
"low": self.close_value - 0.1,
|
||||
"close": self.close_value,
|
||||
"volume": 10000,
|
||||
"amount": self.close_value * 10000,
|
||||
"preClose": self.close_value - 0.1,
|
||||
}
|
||||
return {"600000.SH": [values[field]]}
|
||||
|
||||
def fake_passorder(*args):
|
||||
calls.append((threading.get_ident(), args))
|
||||
session = holder["session"]
|
||||
session.on_order({
|
||||
"m_strOrderSysID": "qmt-order-1",
|
||||
"m_strRemark": args[-2],
|
||||
"m_strInstrumentID": "600000",
|
||||
"m_strExchangeID": "SH",
|
||||
"m_nOffsetFlag": 48,
|
||||
"m_nVolumeTotalOriginal": args[6],
|
||||
"m_nVolumeTraded": args[6],
|
||||
"m_nOrderStatus": "FILLED",
|
||||
})
|
||||
session.on_trade({
|
||||
"m_strTradeID": "qmt-fill-1",
|
||||
"m_strOrderSysID": "qmt-order-1",
|
||||
"m_strRemark": args[-2],
|
||||
"m_strInstrumentID": "600000",
|
||||
"m_strExchangeID": "SH",
|
||||
"m_nOffsetFlag": 48,
|
||||
"m_nVolume": args[6],
|
||||
"m_dPrice": 10.1,
|
||||
"m_strTradeTime": "09:31:00",
|
||||
})
|
||||
return "qmt-order-1"
|
||||
|
||||
session = QmtNativeBacktestSession(
|
||||
config={
|
||||
"run_id": "qmt-zmq-native",
|
||||
"account_id": "backtest-account",
|
||||
"bar_wait_timeout_seconds": 2,
|
||||
},
|
||||
qmt_api={"passorder": fake_passorder},
|
||||
)
|
||||
holder["session"] = session
|
||||
server = ZmqBacktestServer(
|
||||
BacktestBridgeProtocol(session), endpoint=endpoint, exit_on_finish=True
|
||||
)
|
||||
server_thread = threading.Thread(target=server.serve_forever, daemon=True)
|
||||
server_thread.start()
|
||||
self.assertTrue(server.wait_until_ready(2))
|
||||
|
||||
def qmt_backtest_loop():
|
||||
session.bind_context(Context(0, 10.0))
|
||||
session.on_bar(Context(0, 10.0))
|
||||
session.on_bar(Context(1, 10.1))
|
||||
session.on_qmt_stop()
|
||||
|
||||
qmt_thread = threading.Thread(target=qmt_backtest_loop)
|
||||
qmt_thread.start()
|
||||
|
||||
client = BacktestZmqClient(endpoint, run_id="", timeout_seconds=2)
|
||||
description = client.describe()
|
||||
first = client.start()
|
||||
queued = client.submit_order("600000.SH", "BUY", 100, client_order_id="native-buy-1")
|
||||
second = client.next_bar()
|
||||
done = client.next_bar()
|
||||
result = client.finish()
|
||||
client.close()
|
||||
qmt_thread.join(timeout=2)
|
||||
server_thread.join(timeout=2)
|
||||
|
||||
self.assertEqual(description["engine_owner"], "QMT")
|
||||
self.assertEqual(description["matching_owner"], "QMT")
|
||||
self.assertEqual(client.run_id, "qmt-zmq-native")
|
||||
self.assertEqual(first["frame_index"], 0)
|
||||
self.assertEqual(queued["status"], "QUEUED")
|
||||
self.assertEqual(second["frame_index"], 1)
|
||||
self.assertEqual(second["fills"][0]["fill_id"], "qmt-fill-1")
|
||||
self.assertTrue(done["done"])
|
||||
self.assertEqual(result["result_owner"], "QMT")
|
||||
self.assertEqual(len(calls), 1)
|
||||
self.assertEqual(calls[0][0], qmt_thread.ident)
|
||||
self.assertFalse(qmt_thread.is_alive())
|
||||
self.assertFalse(server_thread.is_alive())
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,176 +0,0 @@
|
||||
import datetime
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.app import SignalTradingApp
|
||||
from bigqmt_signal_trader.models import AssetSnapshot, PositionSnapshot, TradeSignal
|
||||
|
||||
|
||||
def _signal(**kwargs):
|
||||
payload = {
|
||||
"signal_id": "sig-buy-001",
|
||||
"account_id": "test",
|
||||
"action": "BUY",
|
||||
"stock_code": "000001.SZ",
|
||||
"amount": 100,
|
||||
"price_type": "AUTO_LIMIT",
|
||||
"remark": "web_buy_command",
|
||||
"created_at": "2026-06-30 09:31:00",
|
||||
"expire_at": "2026-06-30 09:36:00",
|
||||
"schema_version": 1,
|
||||
}
|
||||
payload.update(kwargs)
|
||||
return TradeSignal.from_dict(payload)
|
||||
|
||||
|
||||
class FakeSignalSource:
|
||||
def __init__(self, items):
|
||||
self.items = items
|
||||
self.acked = []
|
||||
|
||||
def fetch(self, account_id, limit):
|
||||
return self.items[:limit]
|
||||
|
||||
def ack(self, signal):
|
||||
self.acked.append(signal.signal_id)
|
||||
|
||||
|
||||
class FakeMarketDataProvider:
|
||||
def get_ticks(self, codes):
|
||||
return {
|
||||
code: {
|
||||
"lastPrice": 10.0,
|
||||
"askPrice": [10.01, 10.02],
|
||||
"bidPrice": [9.99, 9.98],
|
||||
}
|
||||
for code in codes
|
||||
}
|
||||
|
||||
def get_instrument(self, code):
|
||||
return {"InstrumentStatus": 0, "UpStopPrice": 11.0, "DownStopPrice": 9.0}
|
||||
|
||||
|
||||
class FakePositionProvider:
|
||||
def __init__(self):
|
||||
self.positions = {
|
||||
"000001.SZ": PositionSnapshot(
|
||||
stock_code="000001.SZ",
|
||||
volume=1000,
|
||||
available=500,
|
||||
cost=9.5,
|
||||
)
|
||||
}
|
||||
|
||||
def get_positions(self, account_id):
|
||||
return self.positions
|
||||
|
||||
def get_asset(self, account_id):
|
||||
return AssetSnapshot(account_id=account_id, cash=100000.0, total_asset=200000.0)
|
||||
|
||||
|
||||
class FakeOrderGateway:
|
||||
def __init__(self):
|
||||
self.submitted = []
|
||||
|
||||
def submit(self, request):
|
||||
self.submitted.append(request)
|
||||
from bigqmt_signal_trader.models import OrderSubmitResult
|
||||
|
||||
return OrderSubmitResult(status="SUBMITTED", user_order_id="bq:sig:1")
|
||||
|
||||
def cancel(self, order_ref):
|
||||
return None
|
||||
|
||||
def query_orders(self, account_id, strategy_name):
|
||||
return []
|
||||
|
||||
def query_trades(self, account_id, strategy_name):
|
||||
return []
|
||||
|
||||
|
||||
class FakePositionSyncSink:
|
||||
def __init__(self):
|
||||
self.snapshots = []
|
||||
|
||||
def publish(self, snapshot):
|
||||
self.snapshots.append(snapshot)
|
||||
|
||||
|
||||
class FakeStateStore:
|
||||
def __init__(self, claim_result=True):
|
||||
self.claim_result = claim_result
|
||||
self.claimed = []
|
||||
self.submitted = []
|
||||
self.finished = []
|
||||
|
||||
def claim(self, signal, consumer_id):
|
||||
self.claimed.append((signal.signal_id, consumer_id))
|
||||
return self.claim_result
|
||||
|
||||
def mark_submitted(self, signal_id, result):
|
||||
self.submitted.append((signal_id, result.status))
|
||||
|
||||
def mark_finished(self, signal_id, status, message=""):
|
||||
self.finished.append((signal_id, status, message))
|
||||
|
||||
|
||||
class SignalTradingAppTest(unittest.TestCase):
|
||||
def test_tick_submits_buy_signal_with_replaceable_adapters(self):
|
||||
source = FakeSignalSource([_signal()])
|
||||
state = FakeStateStore()
|
||||
orders = FakeOrderGateway()
|
||||
sync = FakePositionSyncSink()
|
||||
app = SignalTradingApp(
|
||||
account_id="test",
|
||||
signal_source=source,
|
||||
market_data=FakeMarketDataProvider(),
|
||||
position_provider=FakePositionProvider(),
|
||||
order_gateway=orders,
|
||||
position_sync_sink=sync,
|
||||
state_store=state,
|
||||
consumer_id="consumer-a",
|
||||
)
|
||||
|
||||
app.tick(datetime.datetime(2026, 6, 30, 9, 31))
|
||||
|
||||
self.assertEqual(orders.submitted[0].stock_code, "000001.SZ")
|
||||
self.assertEqual(orders.submitted[0].volume, 100)
|
||||
self.assertEqual(state.submitted, [("sig-buy-001", "SUBMITTED")])
|
||||
self.assertEqual(source.acked, ["sig-buy-001"])
|
||||
self.assertEqual(sync.snapshots[0].account_id, "test")
|
||||
|
||||
def test_tick_sells_by_percentage_using_available_position(self):
|
||||
source = FakeSignalSource([
|
||||
_signal(
|
||||
signal_id="sig-sell-001",
|
||||
action="SELL",
|
||||
amount=None,
|
||||
percentage=50,
|
||||
remark="web_sell_command",
|
||||
)
|
||||
])
|
||||
orders = FakeOrderGateway()
|
||||
app = SignalTradingApp(
|
||||
account_id="test",
|
||||
signal_source=source,
|
||||
market_data=FakeMarketDataProvider(),
|
||||
position_provider=FakePositionProvider(),
|
||||
order_gateway=orders,
|
||||
position_sync_sink=FakePositionSyncSink(),
|
||||
state_store=FakeStateStore(),
|
||||
consumer_id="consumer-a",
|
||||
)
|
||||
|
||||
app.tick(datetime.datetime(2026, 6, 30, 9, 31))
|
||||
|
||||
self.assertEqual(orders.submitted[0].action, "SELL")
|
||||
self.assertEqual(orders.submitted[0].volume, 200)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,226 +0,0 @@
|
||||
"""frozen_cash must survive the whole chain: QMT row -> AssetSnapshot -> RPC
|
||||
serialization -> client CompatObject, plus the Redis cached-asset fallback.
|
||||
|
||||
Field names follow MiniQMT's XtAsset(account_id, cash, frozen_cash,
|
||||
market_value, total_asset), where total_asset = cash + frozen_cash + market_value.
|
||||
"""
|
||||
|
||||
import datetime as _dt
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.adapters import position_bigqmt
|
||||
from bigqmt_signal_trader.adapters.position_bigqmt import BigQmtPositionProvider
|
||||
from bigqmt_signal_trader.adapters.position_sync_redis import RedisPositionSyncSink
|
||||
from bigqmt_signal_trader.models import AccountSnapshot, AssetSnapshot, PositionSnapshot
|
||||
from bigqmt_signal_trader.redis_rpc import to_jsonable
|
||||
from bigqmt_signal_trader.xtquant_compat import BigQmtXtTrader, StockAccount
|
||||
|
||||
|
||||
class Row:
|
||||
def __init__(self, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
|
||||
def _provider(row):
|
||||
def query(account_id, account_type, detail_type):
|
||||
return [row] if detail_type in ("ACCOUNT", "ASSET") else []
|
||||
|
||||
return BigQmtPositionProvider(query)
|
||||
|
||||
|
||||
class AssetSnapshotModelTest(unittest.TestCase):
|
||||
def test_defaults_keep_existing_positional_callers_working(self):
|
||||
snapshot = AssetSnapshot("acct", 100.0, 1000.0)
|
||||
|
||||
self.assertEqual(snapshot.cash, 100.0)
|
||||
self.assertEqual(snapshot.total_asset, 1000.0)
|
||||
self.assertIsNone(snapshot.frozen_cash)
|
||||
self.assertIsNone(snapshot.market_value)
|
||||
|
||||
def test_carries_the_full_xtasset_field_set(self):
|
||||
snapshot = AssetSnapshot("acct", 100.0, 1000.0, frozen_cash=50.0, market_value=850.0)
|
||||
|
||||
self.assertEqual(snapshot.frozen_cash, 50.0)
|
||||
self.assertEqual(snapshot.market_value, 850.0)
|
||||
|
||||
|
||||
class QmtCollectionTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
position_bigqmt._missing_field_reported.clear()
|
||||
|
||||
def test_reads_frozen_cash_from_the_account_row(self):
|
||||
provider = _provider(
|
||||
Row(m_dAvailable=100.0, m_dBalance=1000.0, m_dFrozenCash=50.0, m_dInstrumentValue=850.0)
|
||||
)
|
||||
|
||||
asset = provider.get_asset("acct")
|
||||
|
||||
self.assertEqual(asset.cash, 100.0)
|
||||
self.assertEqual(asset.frozen_cash, 50.0)
|
||||
self.assertEqual(asset.market_value, 850.0)
|
||||
self.assertEqual(asset.total_asset, 1000.0)
|
||||
|
||||
def test_accepts_alternate_broker_spellings(self):
|
||||
for field in ("m_dFrozenCash", "m_dFrozen", "m_dFrozenBalance", "frozen_cash"):
|
||||
provider = _provider(Row(m_dAvailable=100.0, m_dBalance=1000.0, **{field: 50.0}))
|
||||
|
||||
self.assertEqual(provider.get_asset("acct").frozen_cash, 50.0, field)
|
||||
|
||||
def test_derived_market_value_excludes_frozen_cash(self):
|
||||
"""Without this, market value is overstated by the frozen amount."""
|
||||
provider = _provider(Row(m_dAvailable=100.0, m_dBalance=1000.0, m_dFrozenCash=50.0))
|
||||
|
||||
self.assertEqual(provider.get_asset("acct").market_value, 850.0)
|
||||
|
||||
def test_derivation_falls_back_when_frozen_is_absent(self):
|
||||
provider = _provider(Row(m_dAvailable=100.0, m_dBalance=1000.0))
|
||||
asset = provider.get_asset("acct")
|
||||
|
||||
self.assertIsNone(asset.frozen_cash)
|
||||
self.assertEqual(asset.market_value, 900.0) # legacy behaviour preserved
|
||||
|
||||
def test_missing_frozen_field_reports_what_the_row_actually_has(self):
|
||||
"""The ThinkTrader spelling is unverified offline; make it self-reporting
|
||||
instead of silently returning None forever."""
|
||||
import io
|
||||
import contextlib
|
||||
|
||||
provider = _provider(Row(m_dAvailable=100.0, m_dBalance=1000.0, m_dWhateverElse=1.0))
|
||||
buffer = io.StringIO()
|
||||
with contextlib.redirect_stdout(buffer):
|
||||
provider.get_asset("acct")
|
||||
output = buffer.getvalue()
|
||||
|
||||
self.assertIn("frozen_cash not found", output)
|
||||
self.assertIn("m_dWhateverElse", output) # names the real fields
|
||||
|
||||
def test_missing_field_is_reported_once_not_per_call(self):
|
||||
import io
|
||||
import contextlib
|
||||
|
||||
provider = _provider(Row(m_dAvailable=100.0, m_dBalance=1000.0))
|
||||
buffer = io.StringIO()
|
||||
with contextlib.redirect_stdout(buffer):
|
||||
for _ in range(5):
|
||||
provider.get_asset("acct")
|
||||
|
||||
self.assertEqual(buffer.getvalue().count("frozen_cash not found"), 1)
|
||||
|
||||
def test_empty_rows_still_degrade_to_all_none(self):
|
||||
provider = BigQmtPositionProvider(lambda *args: [])
|
||||
|
||||
asset = provider.get_asset("acct")
|
||||
|
||||
self.assertIsNone(asset.cash)
|
||||
self.assertIsNone(asset.frozen_cash)
|
||||
|
||||
|
||||
class RpcSerializationTest(unittest.TestCase):
|
||||
def test_to_jsonable_carries_frozen_cash_over_the_wire(self):
|
||||
snapshot = AssetSnapshot("acct", 100.0, 1000.0, frozen_cash=50.0, market_value=850.0)
|
||||
|
||||
payload = to_jsonable(snapshot)
|
||||
|
||||
self.assertEqual(payload["frozen_cash"], 50.0)
|
||||
self.assertEqual(payload["market_value"], 850.0)
|
||||
|
||||
|
||||
class FakeRedis:
|
||||
def __init__(self):
|
||||
self.kv = {}
|
||||
self.streams = {}
|
||||
|
||||
def set(self, key, value):
|
||||
self.kv[key] = value
|
||||
|
||||
def setex(self, key, ttl_seconds, value):
|
||||
self.kv[key] = value
|
||||
|
||||
def xadd(self, key, fields, maxlen=None, approximate=None):
|
||||
self.streams.setdefault(key, []).append(fields)
|
||||
return b"1-0"
|
||||
|
||||
def publish(self, key, value):
|
||||
return 1
|
||||
|
||||
|
||||
class PositionSyncTest(unittest.TestCase):
|
||||
def test_cached_snapshot_includes_frozen_cash(self):
|
||||
redis_client = FakeRedis()
|
||||
RedisPositionSyncSink(redis_client).publish(
|
||||
AccountSnapshot(
|
||||
account_id="acct",
|
||||
asset=AssetSnapshot("acct", 100.0, 1000.0, frozen_cash=50.0, market_value=850.0),
|
||||
positions={"600000.SH": PositionSnapshot("600000.SH", 100, 100, 10.0, "PF")},
|
||||
reason="test",
|
||||
updated_at=_dt.datetime(2026, 7, 1, 9, 31),
|
||||
)
|
||||
)
|
||||
|
||||
payload = json.loads(redis_client.kv["bigqmt:positions:acct"])
|
||||
|
||||
self.assertEqual(payload["asset"]["frozen_cash"], 50.0)
|
||||
self.assertEqual(payload["asset"]["market_value"], 850.0)
|
||||
|
||||
|
||||
class ClientSurfaceTest(unittest.TestCase):
|
||||
"""The reported AttributeError: asset.frozen_cash must simply exist."""
|
||||
|
||||
def _trader(self, response):
|
||||
trader = BigQmtXtTrader(account_id="acct")
|
||||
trader.client.call = lambda method, params=None, account_id=None, **kw: response
|
||||
return trader
|
||||
|
||||
def test_frozen_cash_is_exposed(self):
|
||||
asset = self._trader(
|
||||
{"cash": 100.0, "total_asset": 1000.0, "frozen_cash": 50.0, "market_value": 850.0}
|
||||
).query_stock_asset(StockAccount("acct"))
|
||||
|
||||
self.assertEqual(asset.frozen_cash, 50.0)
|
||||
self.assertEqual(asset.cash, 100.0)
|
||||
self.assertEqual(asset.market_value, 850.0)
|
||||
self.assertEqual(asset.total_asset, 1000.0)
|
||||
|
||||
def test_frozen_cash_defaults_to_zero_not_missing(self):
|
||||
"""A server that predates this field must not resurrect the
|
||||
AttributeError, and callers do arithmetic on it."""
|
||||
asset = self._trader({"cash": 100.0, "total_asset": 1000.0}).query_stock_asset(
|
||||
StockAccount("acct")
|
||||
)
|
||||
|
||||
self.assertEqual(asset.frozen_cash, 0.0)
|
||||
self.assertEqual(asset.cash + asset.frozen_cash, 100.0)
|
||||
|
||||
def test_derived_market_value_excludes_frozen_cash(self):
|
||||
asset = self._trader(
|
||||
{"cash": 100.0, "total_asset": 1000.0, "frozen_cash": 50.0}
|
||||
).query_stock_asset(StockAccount("acct"))
|
||||
|
||||
self.assertEqual(asset.market_value, 850.0)
|
||||
|
||||
def test_server_market_value_wins_over_derivation(self):
|
||||
asset = self._trader(
|
||||
{"cash": 100.0, "total_asset": 1000.0, "frozen_cash": 50.0, "market_value": 111.0}
|
||||
).query_stock_asset(StockAccount("acct"))
|
||||
|
||||
self.assertEqual(asset.market_value, 111.0)
|
||||
|
||||
def test_components_reconstruct_total_asset(self):
|
||||
asset = self._trader(
|
||||
{"cash": 100.0, "total_asset": 1000.0, "frozen_cash": 50.0, "market_value": 850.0}
|
||||
).query_stock_asset(StockAccount("acct"))
|
||||
|
||||
self.assertAlmostEqual(
|
||||
asset.cash + asset.frozen_cash + asset.market_value, asset.total_asset
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,275 +0,0 @@
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.adapter_factory import build_app
|
||||
from bigqmt_signal_trader.adapters.market_bigqmt import BigQmtMarketDataProvider
|
||||
from bigqmt_signal_trader.adapters.order_bigqmt import BigQmtOrderGateway
|
||||
from bigqmt_signal_trader.adapters.position_bigqmt import BigQmtPositionProvider
|
||||
from bigqmt_signal_trader.models import OrderRef, OrderRequest
|
||||
|
||||
|
||||
class Obj:
|
||||
def __init__(self, **kwargs):
|
||||
self.__dict__.update(kwargs)
|
||||
|
||||
|
||||
class FakeContext:
|
||||
def __init__(self):
|
||||
self.tick_codes = []
|
||||
self.instrument_codes = []
|
||||
|
||||
def get_full_tick(self, codes):
|
||||
self.tick_codes.append(list(codes))
|
||||
return {codes[0]: {"lastPrice": 10.0}}
|
||||
|
||||
def get_instrumentdetail(self, code):
|
||||
self.instrument_codes.append(code)
|
||||
return {"InstrumentStatus": 0}
|
||||
|
||||
|
||||
class FakeMarketDataContext(FakeContext):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
self.market_calls = []
|
||||
|
||||
def get_market_data_ex(
|
||||
self,
|
||||
fields=None,
|
||||
stock_code=None,
|
||||
period="1d",
|
||||
start_time="",
|
||||
end_time="",
|
||||
count=-1,
|
||||
dividend_type="none",
|
||||
):
|
||||
self.market_calls.append(
|
||||
{
|
||||
"method": "get_market_data_ex",
|
||||
"fields": fields,
|
||||
"stock_code": stock_code,
|
||||
"period": period,
|
||||
"start_time": start_time,
|
||||
"end_time": end_time,
|
||||
"count": count,
|
||||
"dividend_type": dividend_type,
|
||||
}
|
||||
)
|
||||
return {"600000.SH": {"close": [10.0]}}
|
||||
|
||||
|
||||
class FakeMarketDataFallbackContext(FakeContext):
|
||||
def get_market_data(self, fields=None, stock_code=None, period="1d", **kwargs):
|
||||
return {
|
||||
"method": "get_market_data",
|
||||
"fields": fields,
|
||||
"stock_code": stock_code,
|
||||
"period": period,
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
|
||||
|
||||
class BigQmtAdaptersTest(unittest.TestCase):
|
||||
def test_market_provider_normalizes_codes_before_context_call(self):
|
||||
context = FakeContext()
|
||||
provider = BigQmtMarketDataProvider(context)
|
||||
|
||||
ticks = provider.get_ticks(["600000"])
|
||||
instrument = provider.get_instrument("sz000001")
|
||||
|
||||
self.assertIn("600000.SH", ticks)
|
||||
self.assertEqual(context.tick_codes, [["600000.SH"]])
|
||||
self.assertEqual(context.instrument_codes, ["000001.SZ"])
|
||||
self.assertEqual(instrument["InstrumentStatus"], 0)
|
||||
|
||||
def test_market_provider_passes_market_codes_to_full_tick(self):
|
||||
context = FakeContext()
|
||||
provider = BigQmtMarketDataProvider(context)
|
||||
|
||||
provider.get_ticks(["SH", "sz"])
|
||||
|
||||
self.assertEqual(context.tick_codes, [["SH", "SZ"]])
|
||||
|
||||
def test_market_provider_supports_bigqmt_market_data_ex_signature(self):
|
||||
context = FakeMarketDataContext()
|
||||
provider = BigQmtMarketDataProvider(context)
|
||||
|
||||
data = provider.get_market_data_ex(field_list=["close"], stock_list=["600000.SH"], count=1)
|
||||
|
||||
self.assertEqual(data["600000.SH"]["close"], [10.0])
|
||||
self.assertEqual(context.market_calls[0]["fields"], ["close"])
|
||||
self.assertEqual(context.market_calls[0]["stock_code"], ["600000.SH"])
|
||||
self.assertEqual(context.market_calls[0]["count"], 1)
|
||||
|
||||
def test_market_provider_falls_back_to_market_data_when_ex_is_missing(self):
|
||||
context = FakeMarketDataFallbackContext()
|
||||
provider = BigQmtMarketDataProvider(context)
|
||||
|
||||
data = provider.get_market_data_ex(field_list=["close"], stock_list=["600000.SH"], period="1m")
|
||||
|
||||
self.assertEqual(data["method"], "get_market_data")
|
||||
self.assertEqual(data["fields"], ["close"])
|
||||
self.assertEqual(data["stock_code"], ["600000.SH"])
|
||||
self.assertEqual(data["period"], "1m")
|
||||
|
||||
def test_position_provider_maps_qmt_position_objects(self):
|
||||
calls = []
|
||||
|
||||
def fake_query(account, account_type, detail_type, *args):
|
||||
calls.append((account, account_type, detail_type, args))
|
||||
if detail_type == "POSITION":
|
||||
return [
|
||||
Obj(
|
||||
m_strInstrumentID="510300",
|
||||
m_strExchangeID="SH",
|
||||
m_nVolume=1000,
|
||||
m_nCanUseVolume=800,
|
||||
m_dOpenPrice=3.456,
|
||||
m_dLastPrice=3.789,
|
||||
m_dMarketValue=3789.0,
|
||||
m_nFrozenVolume=200,
|
||||
m_nOnRoadVolume=10,
|
||||
m_nYesterdayVolume=900,
|
||||
m_strInstrumentName="ETF",
|
||||
)
|
||||
]
|
||||
return []
|
||||
|
||||
provider = BigQmtPositionProvider(fake_query)
|
||||
positions = provider.get_positions("acct")
|
||||
|
||||
self.assertEqual(calls[0], ("acct", "STOCK", "POSITION", ()))
|
||||
self.assertEqual(positions["510300.SH"].volume, 1000)
|
||||
self.assertEqual(positions["510300.SH"].available, 800)
|
||||
self.assertEqual(positions["510300.SH"].cost, 3.456)
|
||||
self.assertEqual(positions["510300.SH"].price, 3.789)
|
||||
self.assertEqual(positions["510300.SH"].market_value, 3789.0)
|
||||
self.assertEqual(positions["510300.SH"].frozen_volume, 200)
|
||||
self.assertEqual(positions["510300.SH"].on_road_volume, 10)
|
||||
self.assertEqual(positions["510300.SH"].yesterday_volume, 900)
|
||||
|
||||
def test_order_gateway_submit_uses_qmt_jq_trade_passorder_shape(self):
|
||||
calls = []
|
||||
|
||||
def fake_passorder(*args):
|
||||
calls.append(args)
|
||||
|
||||
context = object()
|
||||
gateway = BigQmtOrderGateway(context_info=context, passorder_func=fake_passorder)
|
||||
request = OrderRequest(
|
||||
signal_id="sig-001",
|
||||
account_id="acct",
|
||||
action="BUY",
|
||||
stock_code="600000",
|
||||
volume=300,
|
||||
price=10.12,
|
||||
price_type=44,
|
||||
strategy_name="bigqmt_signal_trader",
|
||||
remark="manual",
|
||||
)
|
||||
|
||||
result = gateway.submit(request)
|
||||
|
||||
self.assertEqual(result.status, "SUBMITTED")
|
||||
self.assertEqual(result.user_order_id, "manual")
|
||||
self.assertEqual(calls[0][0:9], (23, 1101, "acct", "600000.SH", 44, 10.12, 300, "bigqmt_signal_trader", 2))
|
||||
self.assertEqual(calls[0][9], result.user_order_id)
|
||||
self.assertIs(calls[0][10], context)
|
||||
|
||||
def test_order_gateway_cancel_and_query_orders(self):
|
||||
cancel_calls = []
|
||||
|
||||
def fake_cancel(*args):
|
||||
cancel_calls.append(args)
|
||||
return True
|
||||
|
||||
def fake_query(account, account_type, detail_type, strategy_name):
|
||||
self.assertEqual((account, account_type, detail_type, strategy_name), ("acct", "STOCK", "ORDER", "s"))
|
||||
return [
|
||||
Obj(
|
||||
m_strOrderSysID="ord1",
|
||||
m_strRemark="remark1",
|
||||
m_strInstrumentID="000001",
|
||||
m_strExchangeID="SZ",
|
||||
m_nOffsetFlag=49,
|
||||
m_nVolumeTotalOriginal=1000,
|
||||
m_nVolumeTraded=200,
|
||||
m_nOrderStatus=50,
|
||||
)
|
||||
]
|
||||
|
||||
context = object()
|
||||
gateway = BigQmtOrderGateway(
|
||||
context_info=context,
|
||||
account_id="acct",
|
||||
cancel_func=fake_cancel,
|
||||
get_trade_detail_data_func=fake_query,
|
||||
)
|
||||
|
||||
cancel_result = gateway.cancel(OrderRef("ord1"))
|
||||
orders = gateway.query_orders("acct", "s")
|
||||
|
||||
self.assertTrue(cancel_result.success)
|
||||
self.assertEqual(cancel_calls, [("ord1", "acct", "STOCK", context)])
|
||||
self.assertEqual(orders[0].stock_code, "000001.SZ")
|
||||
self.assertEqual(orders[0].action, "SELL")
|
||||
self.assertEqual(orders[0].traded_volume, 200)
|
||||
|
||||
def test_query_trades_without_strategy_omits_strategy_filter(self):
|
||||
calls = []
|
||||
|
||||
def fake_query(*args):
|
||||
calls.append(args)
|
||||
return [
|
||||
Obj(
|
||||
m_strTradeID="manual-trade-1",
|
||||
m_strOrderSysID="manual-order-1",
|
||||
m_strInstrumentID="600276",
|
||||
m_strExchangeID="SH",
|
||||
m_nOffsetFlag=48,
|
||||
m_nVolume=100,
|
||||
m_dPrice=54.76,
|
||||
m_strTradeTime="130524",
|
||||
m_strRemark="",
|
||||
)
|
||||
]
|
||||
|
||||
gateway = BigQmtOrderGateway(
|
||||
context_info=object(),
|
||||
get_trade_detail_data_func=fake_query,
|
||||
)
|
||||
|
||||
trades = gateway.query_trades_strict("acct", "")
|
||||
|
||||
self.assertEqual(calls, [("acct", "STOCK", "DEAL")])
|
||||
self.assertEqual(trades[0].trade_id, "manual-trade-1")
|
||||
self.assertEqual(trades[0].stock_code, "600276.SH")
|
||||
self.assertEqual(trades[0].action, "BUY")
|
||||
self.assertEqual(trades[0].volume, 100)
|
||||
self.assertEqual(trades[0].price, 54.76)
|
||||
|
||||
def test_factory_bigqmt_mode_wires_real_adapters(self):
|
||||
app = build_app(
|
||||
FakeContext(),
|
||||
{
|
||||
"mode": "bigqmt",
|
||||
"account_id": "acct",
|
||||
"qmt_api": {
|
||||
"passorder": lambda *args: None,
|
||||
"cancel": lambda *args: True,
|
||||
"get_trade_detail_data": lambda *args: [],
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
self.assertIsInstance(app.market_data, BigQmtMarketDataProvider)
|
||||
self.assertIsInstance(app.position_provider, BigQmtPositionProvider)
|
||||
self.assertIsInstance(app.order_gateway, BigQmtOrderGateway)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
-73
@@ -1,73 +0,0 @@
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.adapters.market_bigqmt import BigQmtMarketDataProvider
|
||||
|
||||
|
||||
class RawMarketContext:
|
||||
def __init__(self, payload=None):
|
||||
self.payload = payload or {}
|
||||
self.calls = []
|
||||
|
||||
def get_market_data_ex_ori(
|
||||
self,
|
||||
fields=None,
|
||||
stock_code=None,
|
||||
period="1d",
|
||||
start_time="",
|
||||
end_time="",
|
||||
count=-1,
|
||||
dividend_type="none",
|
||||
):
|
||||
self.calls.append(
|
||||
{
|
||||
"fields": fields,
|
||||
"stock_code": stock_code,
|
||||
"period": period,
|
||||
"start_time": start_time,
|
||||
"end_time": end_time,
|
||||
"count": count,
|
||||
"dividend_type": dividend_type,
|
||||
}
|
||||
)
|
||||
return self.payload
|
||||
|
||||
def get_market_data_ex(self, *args, **kwargs):
|
||||
raise AssertionError("DataFrame-producing QMT API must not be called")
|
||||
|
||||
|
||||
class BigQmtRawMarketBridgeTest(unittest.TestCase):
|
||||
def test_market_data_ex_uses_raw_context_api(self):
|
||||
rows = [[1784014200000, 55.1], [1784014260000, 55.2]]
|
||||
context = RawMarketContext({"600276.SH": rows})
|
||||
provider = BigQmtMarketDataProvider(context)
|
||||
|
||||
data = provider.get_market_data_ex(
|
||||
field_list=["close"], stock_list=["600276.SH"], period="1m", count=2
|
||||
)
|
||||
|
||||
self.assertEqual("DataFrame", data["600276.SH"]["__bigqmt_type__"])
|
||||
self.assertEqual(["stime", "close"], data["600276.SH"]["columns"])
|
||||
self.assertEqual(rows, data["600276.SH"]["records"])
|
||||
self.assertEqual(["close"], context.calls[0]["fields"])
|
||||
self.assertEqual(["600276.SH"], context.calls[0]["stock_code"])
|
||||
|
||||
def test_market_data_ex_returns_empty_frame_for_requested_symbol(self):
|
||||
context = RawMarketContext({})
|
||||
provider = BigQmtMarketDataProvider(context)
|
||||
|
||||
data = provider.get_market_data_ex(
|
||||
field_list=["close"], stock_list=["600276.SH"], period="1m", count=2
|
||||
)
|
||||
|
||||
self.assertEqual([], data["600276.SH"]["records"])
|
||||
self.assertEqual(["stime", "close"], data["600276.SH"]["columns"])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,38 +0,0 @@
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.code_utils import (
|
||||
normalize_stock_code,
|
||||
round_buy_volume,
|
||||
round_sell_volume,
|
||||
)
|
||||
|
||||
|
||||
class CodeUtilsTest(unittest.TestCase):
|
||||
def test_normalize_stock_code_accepts_common_formats(self):
|
||||
self.assertEqual(normalize_stock_code("600000"), "600000.SH")
|
||||
self.assertEqual(normalize_stock_code("000001"), "000001.SZ")
|
||||
self.assertEqual(normalize_stock_code("SZ000001"), "000001.SZ")
|
||||
self.assertEqual(normalize_stock_code("sh600000"), "600000.SH")
|
||||
self.assertEqual(normalize_stock_code("600000.SH"), "600000.SH")
|
||||
|
||||
def test_normalize_stock_code_keeps_etf_tradable(self):
|
||||
self.assertEqual(normalize_stock_code("510300"), "510300.SH")
|
||||
self.assertEqual(normalize_stock_code("159915"), "159915.SZ")
|
||||
|
||||
def test_round_buy_volume_by_lot(self):
|
||||
self.assertEqual(round_buy_volume("000001.SZ", 1234), 1200)
|
||||
self.assertEqual(round_buy_volume("688001.SH", 234), 200)
|
||||
|
||||
def test_round_sell_volume_keeps_all_when_sell_all(self):
|
||||
self.assertEqual(round_sell_volume("000001.SZ", 1234, sell_all=False), 1200)
|
||||
self.assertEqual(round_sell_volume("000001.SZ", 1234, sell_all=True), 1234)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,181 +0,0 @@
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.download_jobs import (
|
||||
_enc,
|
||||
current_key,
|
||||
job_key,
|
||||
pump_download_jobs,
|
||||
queue_key,
|
||||
read_download_status,
|
||||
submit_download_job,
|
||||
wait_download_job,
|
||||
)
|
||||
|
||||
|
||||
class FakeRedis:
|
||||
def __init__(self):
|
||||
self.kv = {}
|
||||
self.lists = {}
|
||||
self.expired = []
|
||||
|
||||
def setex(self, key, ttl, value):
|
||||
self.kv[key] = value
|
||||
self.expired.append((key, ttl))
|
||||
return True
|
||||
|
||||
def set(self, key, value):
|
||||
self.kv[key] = value
|
||||
return True
|
||||
|
||||
def get(self, key):
|
||||
return self.kv.get(key)
|
||||
|
||||
def delete(self, key):
|
||||
self.kv.pop(key, None)
|
||||
return 1
|
||||
|
||||
def rpush(self, key, value):
|
||||
self.lists.setdefault(key, []).append(value)
|
||||
return len(self.lists[key])
|
||||
|
||||
def lpop(self, key):
|
||||
lst = self.lists.get(key) or []
|
||||
if not lst:
|
||||
return None
|
||||
return lst.pop(0)
|
||||
|
||||
def expire(self, key, ttl):
|
||||
self.expired.append((key, ttl))
|
||||
return True
|
||||
|
||||
|
||||
class FakeMarketData:
|
||||
def __init__(self, fail_on=None, sleep=0.0):
|
||||
self.data2_calls = []
|
||||
self.data_calls = []
|
||||
self.fail_on = fail_on
|
||||
self.sleep = sleep
|
||||
|
||||
def download_history_data2(self, stock_list, period, start_time, end_time, incrementally):
|
||||
if self.sleep:
|
||||
time.sleep(self.sleep)
|
||||
if self.fail_on and self.fail_on in stock_list:
|
||||
raise RuntimeError("boom")
|
||||
self.data2_calls.append(list(stock_list))
|
||||
|
||||
def download_history_data(self, code, period, start_time, end_time, incrementally):
|
||||
if self.sleep:
|
||||
time.sleep(self.sleep)
|
||||
if self.fail_on and code == self.fail_on:
|
||||
raise RuntimeError("boom")
|
||||
self.data_calls.append(code)
|
||||
|
||||
|
||||
class DownloadJobsTest(unittest.TestCase):
|
||||
def test_submit_queues_pending_job(self):
|
||||
r = FakeRedis()
|
||||
job = submit_download_job(r, "acct", ["600000.SH", "000001.SZ"], "1d", chunk_size=1)
|
||||
|
||||
self.assertEqual(job["state"], "pending")
|
||||
self.assertEqual(job["total"], 2)
|
||||
self.assertIn(_enc(job["job_id"]), r.lists[queue_key("acct")])
|
||||
self.assertEqual(read_download_status(r, "acct", job["job_id"])["state"], "pending")
|
||||
|
||||
def test_pump_completes_and_chunks_the_symbol_list(self):
|
||||
r = FakeRedis()
|
||||
md = FakeMarketData()
|
||||
submit_download_job(r, "acct", ["a", "b", "c", "d", "e"], "1d", chunk_size=2)
|
||||
|
||||
# max_wall_seconds=0 disables the budget, so one tick drains the whole job.
|
||||
res = pump_download_jobs(r, md, "acct", chunk_size=2, max_wall_seconds=0)
|
||||
|
||||
self.assertEqual(res["state"], "done")
|
||||
self.assertEqual(res["done"], 5)
|
||||
self.assertEqual(md.data2_calls, [["a", "b"], ["c", "d"], ["e"]])
|
||||
# current pointer cleared when the job finishes.
|
||||
self.assertIsNone(r.get(current_key("acct")))
|
||||
|
||||
def test_pump_spreads_across_ticks_under_wall_budget(self):
|
||||
r = FakeRedis()
|
||||
md = FakeMarketData(sleep=0.02)
|
||||
submit_download_job(r, "acct", ["a", "b", "c"], "1d", chunk_size=1)
|
||||
|
||||
res1 = pump_download_jobs(r, md, "acct", max_wall_seconds=0.005)
|
||||
res2 = pump_download_jobs(r, md, "acct", max_wall_seconds=0.005)
|
||||
res3 = pump_download_jobs(r, md, "acct", max_wall_seconds=0.005)
|
||||
|
||||
# One chunk per tick (each chunk exceeds the tiny budget), progress resumes.
|
||||
self.assertEqual((res1["state"], res1["done"]), ("running", 1))
|
||||
self.assertEqual((res2["state"], res2["done"]), ("running", 2))
|
||||
self.assertEqual((res3["state"], res3["done"]), ("done", 3))
|
||||
self.assertEqual(md.data_calls if md.data_calls else md.data2_calls, [["a"], ["b"], ["c"]])
|
||||
|
||||
def test_pump_marks_failed_and_clears_current(self):
|
||||
r = FakeRedis()
|
||||
md = FakeMarketData(fail_on="b")
|
||||
job = submit_download_job(
|
||||
r, "acct", ["a", "b", "c"], "1d", method="download_history_data", chunk_size=1
|
||||
)
|
||||
|
||||
res = pump_download_jobs(r, md, "acct", max_wall_seconds=0)
|
||||
|
||||
self.assertEqual(res["state"], "failed")
|
||||
self.assertEqual(md.data_calls, ["a"])
|
||||
status = read_download_status(r, "acct", job["job_id"])
|
||||
self.assertEqual(status["state"], "failed")
|
||||
self.assertTrue(status["error"])
|
||||
self.assertIsNone(r.get(current_key("acct")))
|
||||
|
||||
def test_pump_with_no_job_returns_none(self):
|
||||
self.assertIsNone(pump_download_jobs(FakeRedis(), FakeMarketData(), "acct"))
|
||||
|
||||
def test_wait_returns_terminal_status(self):
|
||||
r = FakeRedis()
|
||||
job = submit_download_job(r, "acct", ["a"], "1d", chunk_size=1)
|
||||
status = read_download_status(r, "acct", job["job_id"])
|
||||
status["state"] = "done"
|
||||
status["done"] = 1
|
||||
r.set(job_key("acct", job["job_id"]), _enc(json.dumps(status)))
|
||||
|
||||
res = wait_download_job(r, "acct", job["job_id"], wait_seconds=1, poll_interval_seconds=0.01)
|
||||
|
||||
self.assertEqual(res["state"], "done")
|
||||
|
||||
def test_stored_values_are_digit_free_and_compliance_safe(self):
|
||||
import re
|
||||
|
||||
from bigqmt_signal_trader.download_jobs import _dec
|
||||
|
||||
# The QMT redis compliance filter blocks a response only when it contains a
|
||||
# stock-code pattern (which requires digits). Encoded tokens are all letters.
|
||||
stock_re = re.compile(
|
||||
"(^|[^\\d])+([36]0[\\d]{4}|00(000[1-9]|[1-9][\\d]{3}|[\\d][1-9][\\d]{2}|[\\d]{2}[1-9][\\d]))([^\\d]|$)+"
|
||||
)
|
||||
blob = json.dumps({"stock_list": ["600000.SH", "300750.SZ", "000001.SZ"], "chunk_size": 1})
|
||||
self.assertTrue(stock_re.search(blob)) # plaintext WOULD trip the filter
|
||||
token = _enc(blob)
|
||||
self.assertTrue(all(not c.isdigit() for c in token), "encoded token must be digit-free")
|
||||
self.assertIsNone(stock_re.search(token), "encoded token must not match the stock-code filter")
|
||||
self.assertEqual(_dec(token), blob) # round-trips
|
||||
self.assertIsNone(_dec(None))
|
||||
self.assertIsNone(_dec(""))
|
||||
|
||||
# what actually lands in Redis on submit must also be digit-free
|
||||
r = FakeRedis()
|
||||
job = submit_download_job(r, "acct", ["600000.SH"], "1d", chunk_size=1)
|
||||
stored_blob = r.kv[job_key("acct", job["job_id"])]
|
||||
queued = r.lists[queue_key("acct")][0]
|
||||
self.assertTrue(all(not c.isdigit() for c in stored_blob))
|
||||
self.assertTrue(all(not c.isdigit() for c in queued))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,50 +0,0 @@
|
||||
import glob
|
||||
import os
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
SRC = os.path.join(ROOT, "src")
|
||||
|
||||
|
||||
class EntryEncodingTest(unittest.TestCase):
|
||||
"""Guard against the QMT load crash: a file declaring ``#coding:gbk`` but
|
||||
containing non-GBK (e.g. UTF-8 Chinese) bytes fails to load under QMT's
|
||||
GBK-based Python with 'gbk codec can't decode ...'. Entry files loaded by the
|
||||
QMT editor must stay GBK-decodable (ASCII is the safe subset)."""
|
||||
|
||||
def test_gbk_declared_files_are_gbk_decodable(self):
|
||||
bad = []
|
||||
for path in glob.glob(os.path.join(SRC, "**", "*.py"), recursive=True):
|
||||
if "__pycache__" in path:
|
||||
continue
|
||||
data = open(path, "rb").read()
|
||||
first_line = data.split(b"\n", 1)[0].lower().replace(b" ", b"")
|
||||
if b"coding:gbk" not in first_line and b"coding=gbk" not in first_line:
|
||||
continue
|
||||
try:
|
||||
data.decode("gbk")
|
||||
except UnicodeDecodeError as exc:
|
||||
bad.append("%s (byte %d)" % (os.path.relpath(path, ROOT), exc.start))
|
||||
self.assertEqual(
|
||||
bad,
|
||||
[],
|
||||
"files declare #coding:gbk but are not GBK-decodable; QMT will fail to load them: %s" % bad,
|
||||
)
|
||||
|
||||
def test_qmt_loader_stops_previous_service_before_clearing_modules(self):
|
||||
"""A QMT strategy restart must release the old ZMQ port first."""
|
||||
path = os.path.join(SRC, "BIGQMT_REDIS_DRYRUN.py")
|
||||
with open(path, "r", encoding="gbk") as source_file:
|
||||
source = source_file.read()
|
||||
self.assertIn("def _stop_previous_rpc_service():", source)
|
||||
stop_call = source.index("\n_stop_previous_rpc_service()\n")
|
||||
clear_call = source.index("\n_clear_local_modules()\n")
|
||||
self.assertLess(
|
||||
stop_call,
|
||||
clear_call,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,633 +0,0 @@
|
||||
import json
|
||||
import time
|
||||
import threading
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.exec_events import (
|
||||
enrich_order_identity,
|
||||
format_raw_snapshot,
|
||||
normalize_cancel_error_event,
|
||||
normalize_order_error_event,
|
||||
normalize_order_event,
|
||||
remember_order_identity,
|
||||
normalize_trade_event,
|
||||
order_channel,
|
||||
order_error_channel,
|
||||
cancel_error_channel,
|
||||
publish_order_event,
|
||||
publish_trade_event,
|
||||
raw_field_snapshot,
|
||||
trade_channel,
|
||||
)
|
||||
from bigqmt_signal_trader.xtquant_compat import BigQmtXtTrader, XtQuantTraderCallback
|
||||
|
||||
|
||||
class FakeDeal:
|
||||
m_strAccountID = "acct"
|
||||
m_strInstrumentID = "600000.SH"
|
||||
m_dPrice = 10.5
|
||||
m_nVolume = 100
|
||||
m_strTradeID = "T1"
|
||||
m_strOrderSysID = "O1"
|
||||
m_strTradeTime = "2026-07-02 10:00:00"
|
||||
m_nDirection = 48
|
||||
m_dTradeAmount = 1050.0
|
||||
m_dComssion = 0.5
|
||||
|
||||
|
||||
class FakeOrder:
|
||||
m_strAccountID = "acct"
|
||||
m_strInstrumentID = "000001.SZ"
|
||||
m_nOrderStatus = 50
|
||||
m_nVolumeTotal = 200
|
||||
m_nVolumeTraded = 50
|
||||
m_dLimitPrice = 9.9
|
||||
m_strOrderSysID = "O2"
|
||||
m_nDirection = 49
|
||||
strategyName = "s1"
|
||||
m_strRemark = "remark-1"
|
||||
m_strOptName = "限价买入"
|
||||
|
||||
|
||||
class FakeRedis:
|
||||
def __init__(self):
|
||||
self.xadds = []
|
||||
self.pubs = []
|
||||
self.kv = {}
|
||||
|
||||
def xadd(self, key, fields, maxlen=None, approximate=None):
|
||||
self.xadds.append((key, fields))
|
||||
return b"1-0"
|
||||
|
||||
def publish(self, key, value):
|
||||
self.pubs.append((key, value))
|
||||
return 1
|
||||
|
||||
def setex(self, key, _ttl, value):
|
||||
self.kv[key] = value
|
||||
return True
|
||||
|
||||
def get(self, key):
|
||||
return self.kv.get(key)
|
||||
|
||||
|
||||
class RecordingCallback(XtQuantTraderCallback):
|
||||
def __init__(self):
|
||||
self.orders = []
|
||||
self.trades = []
|
||||
self.order_errors = []
|
||||
self.cancel_errors = []
|
||||
self.async_responses = []
|
||||
self.cancel_async_responses = []
|
||||
self.account_statuses = []
|
||||
|
||||
def on_stock_order(self, order):
|
||||
self.orders.append(order)
|
||||
|
||||
def on_stock_trade(self, trade):
|
||||
self.trades.append(trade)
|
||||
|
||||
def on_order_error(self, order_error):
|
||||
self.order_errors.append(order_error)
|
||||
|
||||
def on_cancel_error(self, cancel_error):
|
||||
self.cancel_errors.append(cancel_error)
|
||||
|
||||
def on_order_stock_async_response(self, response):
|
||||
self.async_responses.append(response)
|
||||
|
||||
def on_cancel_order_stock_async_response(self, response):
|
||||
self.cancel_async_responses.append(response)
|
||||
|
||||
def on_account_status(self, status):
|
||||
self.account_statuses.append(status)
|
||||
|
||||
|
||||
class ExecEventsServerTest(unittest.TestCase):
|
||||
def test_normalize_trade_event_maps_thinktrader_fields(self):
|
||||
ev = normalize_trade_event(FakeDeal(), "acct")
|
||||
|
||||
self.assertEqual(ev["event_type"], "trade")
|
||||
self.assertEqual(ev["stock_code"], "600000.SH")
|
||||
self.assertEqual(ev["trade_id"], "T1")
|
||||
self.assertEqual(ev["order_sys_id"], "O1")
|
||||
self.assertEqual(ev["volume"], 100)
|
||||
self.assertEqual(ev["price"], 10.5)
|
||||
self.assertEqual(ev["action"], "BUY") # m_nDirection 48 -> buy
|
||||
self.assertEqual(ev["traded_at"], "2026-07-02 10:00:00")
|
||||
self.assertEqual(ev["commission"], 0.5)
|
||||
|
||||
def test_normalize_order_event_maps_thinktrader_fields(self):
|
||||
ev = normalize_order_event(FakeOrder(), "acct")
|
||||
|
||||
self.assertEqual(ev["event_type"], "order")
|
||||
self.assertEqual(ev["stock_code"], "000001.SZ")
|
||||
self.assertEqual(ev["order_sys_id"], "O2")
|
||||
self.assertEqual(ev["order_volume"], 200)
|
||||
self.assertEqual(ev["traded_volume"], 50)
|
||||
self.assertEqual(ev["status"], 50)
|
||||
self.assertEqual(ev["action"], "SELL") # m_nDirection 49 -> sell
|
||||
self.assertEqual(ev["strategy_name"], "s1")
|
||||
self.assertEqual(ev["remark"], "remark-1")
|
||||
self.assertEqual(ev["user_order_id"], "remark-1")
|
||||
self.assertEqual(ev["opt_name"], "限价买入")
|
||||
|
||||
def test_order_event_fills_strategy_from_remark_identity(self):
|
||||
class CallbackOrder:
|
||||
m_strAccountID = "acct"
|
||||
m_strInstrumentID = "159518"
|
||||
m_strRemark = "涨停价买入1手"
|
||||
m_strOptName = "限价买入"
|
||||
|
||||
redis_client = FakeRedis()
|
||||
remember_order_identity(redis_client, "acct", "涨停价买入1手", "rpc_test", "159518")
|
||||
ev = enrich_order_identity(redis_client, "acct", normalize_order_event(CallbackOrder(), "acct"))
|
||||
|
||||
self.assertEqual(ev["strategy_name"], "rpc_test")
|
||||
self.assertEqual(ev["remark"], "涨停价买入1手")
|
||||
|
||||
def test_publish_writes_stream_and_channel(self):
|
||||
r = FakeRedis()
|
||||
publish_trade_event(r, "acct", {"event_type": "trade", "trade_id": "T1"})
|
||||
|
||||
self.assertEqual(r.pubs[0][0], trade_channel("acct"))
|
||||
self.assertEqual(r.xadds[0][0], trade_channel("acct"))
|
||||
self.assertIn("T1", r.pubs[0][1])
|
||||
|
||||
publish_order_event(r, "acct", {"event_type": "order"})
|
||||
self.assertEqual(r.pubs[1][0], order_channel("acct"))
|
||||
|
||||
def test_arbitration_resolves_direction_offset_conflict_via_op_type(self):
|
||||
"""When m_nDirection and m_nOffsetFlag disagree (futures: sell+open),
|
||||
arbitration via m_nOpType picks the semantically correct field."""
|
||||
class Deal:
|
||||
m_strInstrumentID = "600000.SH"
|
||||
m_nDirection = 49 # EEntrustBS sell
|
||||
m_nOffsetFlag = 48 # offset 48 = 开仓 (open)
|
||||
m_nOpType = 24 # STOCK_SELL — arbiter confirms sell
|
||||
m_nVolume = 10
|
||||
m_dPrice = 1.0
|
||||
m_strTradeID = "X"
|
||||
|
||||
ev = normalize_trade_event(Deal(), "acct")
|
||||
|
||||
self.assertEqual(ev["action"], "SELL") # from direction via arbitration
|
||||
self.assertEqual(ev["direction"], 49) # direction field = m_nDirection
|
||||
self.assertEqual(ev["offset_flag"], 48) # raw offset preserved, not conflated
|
||||
|
||||
def test_arbitration_stock_sell_wrong_direction_fixed_by_op_type(self):
|
||||
"""Stock sell: m_nDirection=48 (bug: always 48), m_nOffsetFlag=49,
|
||||
m_nOpType=24 → arbitration picks offset (49→SELL)."""
|
||||
class SellOrder:
|
||||
m_strInstrumentID = "601398.SH"
|
||||
m_nDirection = 48 # QMT bug — always 48 in live callbacks
|
||||
m_nOffsetFlag = 49 # 平仓 = sell (correct)
|
||||
m_nOpType = 24 # STOCK_SELL (correct)
|
||||
m_nVolumeTotal = 100
|
||||
m_nVolumeTraded = 0
|
||||
m_dLimitPrice = 6.34
|
||||
m_strOrderSysID = "S123"
|
||||
|
||||
ev = normalize_order_event(SellOrder(), "acct")
|
||||
self.assertEqual(ev["action"], "SELL")
|
||||
self.assertEqual(ev["direction"], 49) # offset_flag wins via arbitration
|
||||
|
||||
def test_arbitration_stock_buy_agree(self):
|
||||
"""Stock buy: m_nDirection=48, m_nOffsetFlag=48 → agree → BUY."""
|
||||
class BuyOrder:
|
||||
m_strInstrumentID = "601398.SH"
|
||||
m_nDirection = 48
|
||||
m_nOffsetFlag = 48
|
||||
m_nOpType = 23
|
||||
m_nVolumeTotal = 100
|
||||
m_nVolumeTraded = 0
|
||||
m_dLimitPrice = 5.0
|
||||
m_strOrderSysID = "B456"
|
||||
|
||||
ev = normalize_order_event(BuyOrder(), "acct")
|
||||
self.assertEqual(ev["action"], "BUY")
|
||||
self.assertEqual(ev["direction"], 48)
|
||||
|
||||
def test_direction_zero_falls_back_to_offset(self):
|
||||
"""m_nDirection=0 is treated as absent; offset determines direction."""
|
||||
class SellOrder:
|
||||
m_strInstrumentID = "601398.SH"
|
||||
m_nDirection = 0
|
||||
m_nOffsetFlag = 49
|
||||
m_nVolumeTotal = 100
|
||||
m_nVolumeTraded = 0
|
||||
m_dLimitPrice = 6.34
|
||||
m_strOrderSysID = "S123"
|
||||
|
||||
ev = normalize_order_event(SellOrder(), "acct")
|
||||
self.assertEqual(ev["action"], "SELL")
|
||||
self.assertEqual(ev["direction"], 49)
|
||||
|
||||
def test_direction_none_falls_back_to_offset(self):
|
||||
"""m_nDirection=None → offset determines direction."""
|
||||
class BuyOrder:
|
||||
m_strInstrumentID = "601398.SH"
|
||||
m_nDirection = None
|
||||
m_nOffsetFlag = 48
|
||||
m_nVolumeTotal = 100
|
||||
m_nVolumeTraded = 0
|
||||
m_dLimitPrice = 5.0
|
||||
m_strOrderSysID = "B456"
|
||||
|
||||
ev = normalize_order_event(BuyOrder(), "acct")
|
||||
self.assertEqual(ev["action"], "BUY")
|
||||
self.assertEqual(ev["direction"], 48)
|
||||
|
||||
def test_pledge_direction_has_no_buy_sell_action(self):
|
||||
class Deal:
|
||||
m_strInstrumentID = "600000.SH"
|
||||
m_nDirection = 81 # 质押入库
|
||||
m_nVolume = 10
|
||||
m_dPrice = 1.0
|
||||
|
||||
ev = normalize_trade_event(Deal(), "acct")
|
||||
|
||||
self.assertEqual(ev["action"], "") # pledge is neither buy nor sell
|
||||
self.assertEqual(ev["direction"], 81) # raw direction preserved
|
||||
|
||||
def test_normalize_order_error_event_maps_fields(self):
|
||||
class OrderError:
|
||||
m_strAccountID = "acct"
|
||||
m_strInstrumentID = "600654.SH"
|
||||
m_strOrderSysID = "sys-err-1"
|
||||
m_nErrorID = 2147483647
|
||||
m_strErrorMsg = "废单"
|
||||
|
||||
ev = normalize_order_error_event(OrderError(), "acct")
|
||||
|
||||
self.assertEqual(ev["event_type"], "order_error")
|
||||
self.assertEqual(ev["account_id"], "acct")
|
||||
self.assertEqual(ev["stock_code"], "600654.SH")
|
||||
self.assertEqual(ev["order_sys_id"], "sys-err-1")
|
||||
self.assertEqual(ev["error_id"], 2147483647)
|
||||
self.assertEqual(ev["error_msg"], "废单")
|
||||
|
||||
def test_normalize_cancel_error_event_maps_fields(self):
|
||||
class CancelError:
|
||||
m_strAccountID = "acct"
|
||||
m_strInstrumentID = "600654.SH"
|
||||
m_strOrderSysID = "sys-cancel-1"
|
||||
m_nErrorID = 99
|
||||
m_strErrorMsg = "撤单失败"
|
||||
|
||||
ev = normalize_cancel_error_event(CancelError(), "acct")
|
||||
|
||||
self.assertEqual(ev["event_type"], "cancel_error")
|
||||
self.assertEqual(ev["account_id"], "acct")
|
||||
self.assertEqual(ev["order_sys_id"], "sys-cancel-1")
|
||||
self.assertEqual(ev["error_id"], 99)
|
||||
self.assertEqual(ev["error_msg"], "撤单失败")
|
||||
|
||||
def test_error_channels_are_account_scoped(self):
|
||||
self.assertTrue(order_error_channel("acct").endswith(":acct"))
|
||||
self.assertTrue(cancel_error_channel("acct").endswith(":acct"))
|
||||
|
||||
|
||||
class RawFieldSnapshotTest(unittest.TestCase):
|
||||
"""The snapshot exists to settle what live callbacks actually carry, so it
|
||||
must capture m_* and MiniQMT fields alike and never raise."""
|
||||
|
||||
def test_captures_thinktrader_and_miniqmt_fields(self):
|
||||
snap = raw_field_snapshot(FakeOrder())
|
||||
|
||||
self.assertIn("m_nDirection", snap)
|
||||
self.assertIn("49", snap["m_nDirection"])
|
||||
self.assertIn("int", snap["m_nDirection"])
|
||||
self.assertIn("m_strInstrumentID", snap)
|
||||
|
||||
def test_captures_miniqmt_style_object(self):
|
||||
class XtOrderLike:
|
||||
stock_code = "601398.SH"
|
||||
order_type = 24
|
||||
order_volume = 100
|
||||
|
||||
snap = raw_field_snapshot(XtOrderLike())
|
||||
|
||||
self.assertIn("24", snap["order_type"])
|
||||
self.assertIn("601398.SH", snap["stock_code"])
|
||||
|
||||
def test_captures_dict_payload(self):
|
||||
snap = raw_field_snapshot({"m_nOffsetFlag": 48, "order_type": 24})
|
||||
|
||||
self.assertIn("48", snap["m_nOffsetFlag"])
|
||||
self.assertIn("24", snap["order_type"])
|
||||
|
||||
def test_skips_callables_and_dunders(self):
|
||||
class WithMethod:
|
||||
m_nDirection = 49
|
||||
|
||||
def m_method(self):
|
||||
return 1
|
||||
|
||||
snap = raw_field_snapshot(WithMethod())
|
||||
|
||||
self.assertIn("m_nDirection", snap)
|
||||
self.assertNotIn("m_method", snap)
|
||||
|
||||
def test_unreadable_attribute_does_not_raise(self):
|
||||
class Exploding:
|
||||
m_nDirection = 49
|
||||
|
||||
@property
|
||||
def m_nOffsetFlag(self):
|
||||
raise RuntimeError("boom")
|
||||
|
||||
snap = raw_field_snapshot(Exploding())
|
||||
|
||||
self.assertIn("m_nDirection", snap)
|
||||
self.assertIn("unreadable", snap["m_nOffsetFlag"])
|
||||
|
||||
def test_format_is_a_single_ascii_safe_line(self):
|
||||
line = format_raw_snapshot("order", FakeOrder())
|
||||
|
||||
self.assertNotIn("\n", line)
|
||||
self.assertTrue(line.startswith("[bigqmt_exec_raw] order"))
|
||||
self.assertIn("m_nDirection", line)
|
||||
|
||||
|
||||
class ExecEventsClientDispatchTest(unittest.TestCase):
|
||||
def _trader(self):
|
||||
trader = BigQmtXtTrader(account_id="acct")
|
||||
cb = RecordingCallback()
|
||||
trader.register_callback(cb)
|
||||
return trader, cb
|
||||
|
||||
def test_dispatch_trade_invokes_on_stock_trade(self):
|
||||
trader, cb = self._trader()
|
||||
event = {
|
||||
"event_type": "trade",
|
||||
"account_id": "acct",
|
||||
"stock_code": "600000.SH",
|
||||
"order_sys_id": "sys-1",
|
||||
"trade_id": "t-1",
|
||||
"volume": 100,
|
||||
"price": 10.5,
|
||||
"action": "BUY",
|
||||
"traded_at": "2026-07-02 10:00:00",
|
||||
}
|
||||
trader._dispatch_event(json.dumps(event).encode("utf-8"))
|
||||
|
||||
self.assertEqual(len(cb.trades), 1)
|
||||
trade = cb.trades[0]
|
||||
self.assertEqual(trade.stock_code, "600000.SH")
|
||||
self.assertEqual(trade.trade_id, "t-1")
|
||||
self.assertEqual(trade.traded_volume, 100)
|
||||
self.assertEqual(trade.traded_price, 10.5)
|
||||
self.assertEqual(trade.order_type, 23) # BUY -> STOCK_BUY
|
||||
|
||||
def test_dispatch_order_invokes_on_stock_order(self):
|
||||
trader, cb = self._trader()
|
||||
event = {
|
||||
"event_type": "order",
|
||||
"account_id": "acct",
|
||||
"stock_code": "000001.SZ",
|
||||
"order_sys_id": "sys-2",
|
||||
"order_volume": 200,
|
||||
"traded_volume": 50,
|
||||
"price": 9.9,
|
||||
"status": 50,
|
||||
"action": "SELL",
|
||||
}
|
||||
trader._dispatch_event(json.dumps(event).encode("utf-8"))
|
||||
|
||||
self.assertEqual(len(cb.orders), 1)
|
||||
order = cb.orders[0]
|
||||
self.assertEqual(order.stock_code, "000001.SZ")
|
||||
self.assertEqual(order.order_volume, 200)
|
||||
self.assertEqual(order.traded_volume, 50)
|
||||
self.assertEqual(order.order_status, 50)
|
||||
self.assertEqual(order.order_type, 24) # SELL -> STOCK_SELL
|
||||
|
||||
def test_dispatch_without_callback_is_noop(self):
|
||||
trader = BigQmtXtTrader(account_id="acct")
|
||||
# No callback registered; must not raise.
|
||||
trader._dispatch_event(json.dumps({"event_type": "trade"}).encode("utf-8"))
|
||||
|
||||
def test_dispatch_order_error_invokes_on_order_error(self):
|
||||
trader, cb = self._trader()
|
||||
event = {
|
||||
"event_type": "order_error",
|
||||
"account_id": "acct",
|
||||
"stock_code": "600654.SH",
|
||||
"order_sys_id": "sys-err-1",
|
||||
"error_id": 2147483647,
|
||||
"error_msg": "废单",
|
||||
}
|
||||
trader._dispatch_event(json.dumps(event).encode("utf-8"))
|
||||
|
||||
self.assertEqual(len(cb.order_errors), 1)
|
||||
err = cb.order_errors[0]
|
||||
self.assertEqual(err.order_id, "sys-err-1")
|
||||
self.assertEqual(err.error_id, 2147483647)
|
||||
self.assertEqual(err.error_msg, "废单")
|
||||
self.assertEqual(err.stock_code, "600654.SH")
|
||||
|
||||
def test_dispatch_cancel_error_invokes_on_cancel_error(self):
|
||||
trader, cb = self._trader()
|
||||
event = {
|
||||
"event_type": "cancel_error",
|
||||
"account_id": "acct",
|
||||
"stock_code": "600654.SH",
|
||||
"order_sys_id": "sys-cancel-1",
|
||||
"error_id": 99,
|
||||
"error_msg": "撤单失败",
|
||||
}
|
||||
trader._dispatch_event(json.dumps(event).encode("utf-8"))
|
||||
|
||||
self.assertEqual(len(cb.cancel_errors), 1)
|
||||
err = cb.cancel_errors[0]
|
||||
self.assertEqual(err.order_id, "sys-cancel-1")
|
||||
self.assertEqual(err.error_id, 99)
|
||||
self.assertEqual(err.error_msg, "撤单失败")
|
||||
|
||||
def _run_async(self, trader, result=None, raises=None):
|
||||
"""Submit one async order with order_stock_result stubbed, and wait.
|
||||
|
||||
order_stock_async is fire-and-forget since issue #50: it returns the seq
|
||||
without touching the network and the submit happens on a worker thread,
|
||||
so the callback assertions need the queue drained first. The stub is
|
||||
restored only after the worker is done with it.
|
||||
"""
|
||||
original = trader.order_stock_result
|
||||
|
||||
def fake(*args, **kwargs):
|
||||
if raises is not None:
|
||||
raise raises
|
||||
return result
|
||||
|
||||
trader.order_stock_result = fake
|
||||
try:
|
||||
seq = trader.order_stock_async("acct", "600654.SH", 23, 100, 11, 10.0, "s", "r")
|
||||
self.assertTrue(trader.wait_async_orders(timeout=5.0), "async order did not finish")
|
||||
finally:
|
||||
trader.order_stock_result = original
|
||||
return seq
|
||||
|
||||
def test_order_stock_async_returns_seq_without_submitting(self):
|
||||
"""issue #50: the seq must come back before any RPC happens."""
|
||||
trader, _cb = self._trader()
|
||||
started = threading.Event()
|
||||
release = threading.Event()
|
||||
|
||||
def blocking(*args, **kwargs):
|
||||
started.set()
|
||||
release.wait(5.0)
|
||||
return {"order_sys_id": "sys-slow"}
|
||||
|
||||
trader.order_stock_result = blocking
|
||||
try:
|
||||
t0 = time.time()
|
||||
seq = trader.order_stock_async("acct", "600654.SH", 23, 100, 11, 10.0, "s", "r")
|
||||
elapsed = time.time() - t0
|
||||
|
||||
self.assertGreater(seq, 0)
|
||||
self.assertLess(elapsed, 0.2, "order_stock_async blocked for %.3fs" % elapsed)
|
||||
self.assertTrue(started.wait(5.0), "submit never ran on the worker")
|
||||
finally:
|
||||
release.set()
|
||||
trader.wait_async_orders(timeout=5.0)
|
||||
|
||||
def test_order_stock_async_fires_response_when_submitted(self):
|
||||
trader, cb = self._trader()
|
||||
seq = self._run_async(trader, result={"order_sys_id": "sys-ok-1", "user_order_id": "u-1"})
|
||||
|
||||
self.assertGreater(seq, 0)
|
||||
self.assertEqual(len(cb.async_responses), 1)
|
||||
resp = cb.async_responses[0]
|
||||
self.assertEqual(resp.order_id, "sys-ok-1")
|
||||
self.assertEqual(resp.account_id, "acct")
|
||||
self.assertEqual(resp.seq, seq)
|
||||
|
||||
def test_order_stock_async_requests_no_settlement_wait(self):
|
||||
"""The server must not hold the reply for the order id on this path."""
|
||||
trader, _cb = self._trader()
|
||||
seen = {}
|
||||
original = trader.order_stock_result
|
||||
|
||||
def fake(*args, **kwargs):
|
||||
seen.update(kwargs)
|
||||
return {"order_sys_id": "sys-1"}
|
||||
|
||||
trader.order_stock_result = fake
|
||||
try:
|
||||
trader.order_stock_async("acct", "600654.SH", 23, 100, 11, 10.0, "s", "r")
|
||||
trader.wait_async_orders(timeout=5.0)
|
||||
finally:
|
||||
trader.order_stock_result = original
|
||||
|
||||
self.assertIs(seen.get("wait_settlement"), False)
|
||||
|
||||
def test_order_stock_async_minus_one_fires_order_error(self):
|
||||
trader, cb = self._trader()
|
||||
seq = self._run_async(trader, result=-1) # MiniQMT: submit failed
|
||||
|
||||
self.assertGreater(seq, 0)
|
||||
self.assertEqual(len(cb.order_errors), 1)
|
||||
err = cb.order_errors[0]
|
||||
self.assertEqual(err.error_id, -1)
|
||||
self.assertEqual(err.stock_code, "600654.SH")
|
||||
self.assertEqual(err.seq, seq) # correlate the failure to the seq
|
||||
# No success response for a failed submit.
|
||||
self.assertEqual(len(cb.async_responses), 0)
|
||||
|
||||
def test_order_stock_async_submitted_without_sysid_fires_response_not_error(self):
|
||||
# issue #38: passorder 已提交但委托号还没分配到(order_sys_id 为空)时,
|
||||
# 必须回调成功响应而不是误报 on_order_error。issue #50 之后这是常态:
|
||||
# 异步路径不再等待委托号,它由 order_callback 推送。
|
||||
trader, cb = self._trader()
|
||||
seq = self._run_async(
|
||||
trader, result={"status": "SUBMITTED", "user_order_id": "u-1", "order_sys_id": ""})
|
||||
|
||||
self.assertGreater(seq, 0)
|
||||
self.assertEqual(len(cb.order_errors), 0)
|
||||
self.assertEqual(len(cb.async_responses), 1)
|
||||
resp = cb.async_responses[0]
|
||||
self.assertEqual(resp.order_id, "u-1") # 委托号未知时回退到 user_order_id
|
||||
self.assertEqual(resp.order_sys_id, "")
|
||||
|
||||
def test_order_stock_async_server_error_fires_order_error_with_reason(self):
|
||||
# server_error(委托没进系统)由 call() 转成异常后,async 必须把真实
|
||||
# 原因回调给 on_order_error(issue #38)。
|
||||
trader, cb = self._trader()
|
||||
seq = self._run_async(trader, raises=RuntimeError(
|
||||
"Big QMT order_stock server_error: passorder submitted but "
|
||||
"order not found in system (stock=600654.SH action=BUY price=10.00 "
|
||||
"volume=100). QMT may have silently rejected it."))
|
||||
|
||||
self.assertGreater(seq, 0)
|
||||
self.assertEqual(len(cb.async_responses), 0)
|
||||
self.assertEqual(len(cb.order_errors), 1)
|
||||
err = cb.order_errors[0]
|
||||
self.assertIn("not found in system", err.error_msg)
|
||||
self.assertEqual(err.stock_code, "600654.SH")
|
||||
|
||||
def test_async_orders_keep_submission_order(self):
|
||||
"""One worker, so responses arrive in the order the calls were made."""
|
||||
trader, cb = self._trader()
|
||||
original = trader.order_stock_result
|
||||
|
||||
def fake(*args, **kwargs):
|
||||
return {"order_sys_id": "sys-%s" % args[1]}
|
||||
|
||||
trader.order_stock_result = fake
|
||||
try:
|
||||
for code in ("A.SH", "B.SH", "C.SH"):
|
||||
trader.order_stock_async("acct", code, 23, 100, 11, 10.0, "s", "r")
|
||||
self.assertTrue(trader.wait_async_orders(timeout=5.0))
|
||||
finally:
|
||||
trader.order_stock_result = original
|
||||
|
||||
self.assertEqual([r.order_id for r in cb.async_responses],
|
||||
["sys-A.SH", "sys-B.SH", "sys-C.SH"])
|
||||
self.assertEqual([r.seq for r in cb.async_responses],
|
||||
sorted(r.seq for r in cb.async_responses))
|
||||
|
||||
def test_cancel_order_stock_async_fires_response(self):
|
||||
trader, cb = self._trader()
|
||||
original = trader.cancel_order_stock_sysid
|
||||
|
||||
def fake_cancel(account, market, sysid):
|
||||
return True
|
||||
|
||||
trader.cancel_order_stock_sysid = fake_cancel
|
||||
try:
|
||||
seq = trader.cancel_order_stock_sysid_async("acct", "SH", "sys-1")
|
||||
finally:
|
||||
trader.cancel_order_stock_sysid = original
|
||||
|
||||
self.assertGreater(seq, 0)
|
||||
self.assertEqual(len(cb.cancel_async_responses), 1)
|
||||
resp = cb.cancel_async_responses[0]
|
||||
self.assertTrue(resp.success)
|
||||
self.assertEqual(resp.order_sys_id, "sys-1")
|
||||
self.assertEqual(resp.account_id, "acct")
|
||||
self.assertEqual(resp.seq, seq)
|
||||
|
||||
def test_connect_and_subscribe_fire_account_status(self):
|
||||
trader, cb = self._trader()
|
||||
trader.client.account_id = "acct"
|
||||
# connect() calls ping via RPC — stub it.
|
||||
trader.client.call = lambda *a, **k: {"ok": True}
|
||||
trader.connect()
|
||||
trader.subscribe("acct")
|
||||
|
||||
self.assertEqual(len(cb.account_statuses), 2)
|
||||
status = cb.account_statuses[0]
|
||||
self.assertEqual(status.account_id, "acct")
|
||||
self.assertEqual(status.account_type, "STOCK")
|
||||
self.assertEqual(status.status, 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,340 +0,0 @@
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader import formula_server as fs
|
||||
|
||||
|
||||
class BsonCodecTest(unittest.TestCase):
|
||||
"""The built-in codec is the no-dependency path; it must round-trip every
|
||||
type this wire carries and stay byte-compatible with pymongo's bson."""
|
||||
|
||||
def _round_trip(self, document):
|
||||
return fs._decode_document(fs._encode_document(document.items()), 0)[0]
|
||||
|
||||
def test_round_trips_scalars(self):
|
||||
doc = {"s": "平安银行", "i": 42, "big": 2 ** 40, "f": 10.29, "t": True, "f2": False, "n": None}
|
||||
|
||||
self.assertEqual(self._round_trip(doc), doc)
|
||||
|
||||
def test_round_trips_nested_containers(self):
|
||||
doc = {"func": "getMarketData", "params": {"fields": ["close", "volume"], "count": -1}}
|
||||
|
||||
self.assertEqual(self._round_trip(doc), doc)
|
||||
|
||||
def test_round_trips_the_actual_request_envelope(self):
|
||||
doc = {
|
||||
"func": "getMarketData",
|
||||
"params": {
|
||||
"fields": ["close"],
|
||||
"stockCodes": ["000001.SZ"],
|
||||
"startTime": "",
|
||||
"endTime": "",
|
||||
"period": "1d",
|
||||
"dividendType": "none",
|
||||
"count": 3,
|
||||
},
|
||||
}
|
||||
|
||||
self.assertEqual(self._round_trip(doc), doc)
|
||||
|
||||
def test_array_order_is_preserved(self):
|
||||
doc = {"codes": ["a", "b", "c", "d", "e", "f", "g", "h", "i", "j", "k"]}
|
||||
|
||||
self.assertEqual(self._round_trip(doc)["codes"], doc["codes"])
|
||||
|
||||
def test_matches_pymongo_bson_when_available(self):
|
||||
try:
|
||||
import bson
|
||||
except ImportError:
|
||||
self.skipTest("pymongo bson not installed")
|
||||
doc = {"func": "getLastVolume", "params": {"stockCode": "000001.SZ", "n": 1.5}}
|
||||
|
||||
self.assertEqual(fs._encode_document(doc.items()), bson.BSON.encode(doc))
|
||||
self.assertEqual(fs._decode_document(bson.BSON.encode(doc), 0)[0], doc)
|
||||
|
||||
def test_unsupported_type_is_rejected(self):
|
||||
with self.assertRaises(TypeError):
|
||||
fs._encode_document({"bad": object()}.items())
|
||||
|
||||
|
||||
class AddressResolutionTest(unittest.TestCase):
|
||||
def test_reads_port_from_formulaserver_ini(self):
|
||||
import tempfile
|
||||
|
||||
root = tempfile.mkdtemp()
|
||||
ini_dir = os.path.join(root, "config", "formulaserver")
|
||||
os.makedirs(ini_dir)
|
||||
with open(os.path.join(ini_dir, "formulaserver.ini"), "w") as handle:
|
||||
handle.write("[server_formula]\naddress = 0.0.0.0:58600\n")
|
||||
|
||||
self.assertEqual(fs.read_formulaserver_port(root), 58600)
|
||||
self.assertEqual(fs.resolve_address({"qmt_root": root}), ("127.0.0.1", 58600))
|
||||
|
||||
def test_missing_ini_falls_back_to_default_port(self):
|
||||
self.assertIsNone(fs.read_formulaserver_port(os.path.join(ROOT, "no-such-dir")))
|
||||
host, port = fs.resolve_address({"qmt_root": os.path.join(ROOT, "no-such-dir")})
|
||||
|
||||
self.assertEqual((host, port), ("127.0.0.1", fs.DEFAULT_PORT))
|
||||
|
||||
def test_explicit_config_wins(self):
|
||||
self.assertEqual(
|
||||
fs.resolve_address({"host": "10.0.0.5", "port": 59999}), ("10.0.0.5", 59999)
|
||||
)
|
||||
|
||||
|
||||
class FakeClient(object):
|
||||
host = "127.0.0.1"
|
||||
port = 58600
|
||||
|
||||
def __init__(self, responses=None, error=None):
|
||||
self.responses = responses or {}
|
||||
self.error = error
|
||||
self.calls = []
|
||||
|
||||
def request(self, func, params=None):
|
||||
self.calls.append((func, dict(params or {})))
|
||||
if self.error is not None:
|
||||
raise self.error
|
||||
return self.responses.get(func, {"result": None})
|
||||
|
||||
def close(self):
|
||||
pass
|
||||
|
||||
|
||||
class ParamTranslationTest(unittest.TestCase):
|
||||
def _router(self, responses=None, error=None):
|
||||
client = FakeClient(responses=responses, error=error)
|
||||
return fs.FormulaServerRouter(client=client), client
|
||||
|
||||
def test_instrument_aliases_the_misspelled_volume_fields(self):
|
||||
"""FormulaServer ships FloatVolumn/TotalVolumn; the xtdata SDK spells
|
||||
them FloatVolume/TotalVolume. Downstream reads the SDK spelling."""
|
||||
router, _ = self._router(
|
||||
{"getInstrumentDetail": {"result": {"FloatVolumn": 1.0, "TotalVolumn": 2.0}}}
|
||||
)
|
||||
|
||||
out = router.call("get_instrument", {"code": "000001.SZ"})
|
||||
|
||||
self.assertEqual(out["FloatVolume"], 1.0)
|
||||
self.assertEqual(out["TotalVolume"], 2.0)
|
||||
self.assertEqual(out["FloatVolumn"], 1.0) # raw key still present
|
||||
|
||||
def test_sector_normalizes_the_minus_one_sentinel(self):
|
||||
router, client = self._router({"getStockListInSector": {"result": ["600000.SH"]}})
|
||||
|
||||
router.call("get_stock_list_in_sector", {"sector_name": "沪深300", "real_timetag": -1})
|
||||
|
||||
self.assertEqual(client.calls[0][1], {"sectorName": "沪深300", "realtime": 0})
|
||||
|
||||
def test_market_data_refuses_adjusted_bars(self):
|
||||
"""dividendType is not honoured by the server; serving an adjusted
|
||||
request from here would hand back unadjusted prices silently."""
|
||||
router, client = self._router({"getMarketData": {"result": []}})
|
||||
|
||||
for dividend_type in ("front", "back", "front_ratio"):
|
||||
with self.assertRaises(fs.Unroutable):
|
||||
router.call(
|
||||
"get_market_data_ex",
|
||||
{
|
||||
"field_list": ["close"],
|
||||
"stock_list": ["000001.SZ"],
|
||||
"dividend_type": dividend_type,
|
||||
},
|
||||
)
|
||||
self.assertEqual(client.calls, [])
|
||||
|
||||
def test_market_data_allows_unadjusted(self):
|
||||
router, client = self._router({"getMarketData": {"result": []}})
|
||||
|
||||
router.call(
|
||||
"get_market_data_ex",
|
||||
{"field_list": ["close"], "stock_list": ["000001.SZ"], "dividend_type": "none"},
|
||||
)
|
||||
|
||||
self.assertEqual(client.calls[0][1]["dividendType"], "none")
|
||||
|
||||
def test_market_data_translates_flat_wire_shape(self):
|
||||
router, _ = self._router(
|
||||
{
|
||||
"getMarketData": {
|
||||
"result": [
|
||||
"000001.SZ",
|
||||
["20260703", ["close", 10.29, "volume", 863327.0]],
|
||||
"600000.SH",
|
||||
["20260703", ["close", 8.69, "volume", 695133.0]],
|
||||
]
|
||||
}
|
||||
}
|
||||
)
|
||||
|
||||
out = router.call(
|
||||
"get_market_data_ex",
|
||||
{"field_list": ["close", "volume"], "stock_list": ["000001.SZ", "600000.SH"]},
|
||||
)
|
||||
|
||||
self.assertEqual(out["000001.SZ"]["columns"], ["stime", "close", "volume"])
|
||||
self.assertEqual(
|
||||
out["000001.SZ"]["records"],
|
||||
[{"stime": "20260703", "close": 10.29, "volume": 863327.0}],
|
||||
)
|
||||
self.assertEqual(out["600000.SH"]["records"][0]["close"], 8.69)
|
||||
|
||||
def test_market_data_keeps_requested_codes_with_no_bars(self):
|
||||
router, _ = self._router({"getMarketData": {"result": []}})
|
||||
|
||||
out = router.call(
|
||||
"get_market_data_ex",
|
||||
{"field_list": ["close"], "stock_list": ["000001.SZ", "600000.SH"]},
|
||||
)
|
||||
|
||||
self.assertEqual(sorted(out), ["000001.SZ", "600000.SH"])
|
||||
self.assertEqual(out["000001.SZ"]["records"], [])
|
||||
|
||||
def test_missing_required_params_is_unroutable_not_a_crash(self):
|
||||
router, client = self._router()
|
||||
|
||||
with self.assertRaises(fs.Unroutable):
|
||||
router.call("get_instrument", {})
|
||||
self.assertEqual(client.calls, [])
|
||||
|
||||
|
||||
class FallbackBehaviourTest(unittest.TestCase):
|
||||
def test_unmapped_method_is_not_supported(self):
|
||||
router = fs.FormulaServerRouter(client=FakeClient())
|
||||
|
||||
self.assertFalse(router.supports("get_asset"))
|
||||
self.assertFalse(router.supports("submit_order"))
|
||||
self.assertFalse(router.supports("get_full_tick"))
|
||||
|
||||
def test_trading_dates_and_dividends_stay_on_rpc(self):
|
||||
"""Their FormulaServer params mean something different from ours."""
|
||||
router = fs.FormulaServerRouter(client=FakeClient())
|
||||
|
||||
self.assertFalse(router.supports("get_trading_dates"))
|
||||
self.assertFalse(router.supports("get_divid_factors"))
|
||||
self.assertFalse(router.supports("get_risk_free_rate"))
|
||||
|
||||
def test_transport_failure_trips_the_cooldown(self):
|
||||
router = fs.FormulaServerRouter(
|
||||
client=FakeClient(error=fs.FormulaServerUnavailable("down")),
|
||||
failure_cooldown_seconds=60,
|
||||
)
|
||||
|
||||
with self.assertRaises(fs.Unroutable):
|
||||
router.call("get_last_volume", {"stock": "000001.SZ"})
|
||||
# Breaker is open: no further attempts until the cooldown expires.
|
||||
self.assertFalse(router.supports("get_last_volume"))
|
||||
|
||||
def test_method_not_found_disables_only_that_method(self):
|
||||
router = fs.FormulaServerRouter(
|
||||
client=FakeClient(
|
||||
error=fs.FormulaServerError("nope", error_id=fs.ERROR_METHOD_NOT_FOUND)
|
||||
)
|
||||
)
|
||||
|
||||
with self.assertRaises(fs.Unroutable):
|
||||
router.call("get_main_contract", {"code_market": "IF00.IF"})
|
||||
|
||||
self.assertFalse(router.supports("get_main_contract"))
|
||||
self.assertTrue(router.supports("get_last_volume")) # breaker not tripped
|
||||
|
||||
def test_disabled_router_supports_nothing(self):
|
||||
router = fs.build_router({"enabled": False})
|
||||
|
||||
for method in fs.SUPPORTED_METHODS:
|
||||
self.assertFalse(router.supports(method))
|
||||
|
||||
def test_enabled_accepts_string_flags(self):
|
||||
self.assertFalse(fs.build_router({"enabled": "false"}).enabled)
|
||||
self.assertFalse(fs.build_router({"enabled": "0"}).enabled)
|
||||
self.assertTrue(fs.build_router({"enabled": "true", "port": 1}).enabled)
|
||||
|
||||
|
||||
class ClientCallIntegrationTest(unittest.TestCase):
|
||||
"""BigQmtRpcClient.call must prefer the router and fall back cleanly."""
|
||||
|
||||
def _client(self, router):
|
||||
from bigqmt_signal_trader.xtquant_compat import BigQmtRpcClient
|
||||
|
||||
client = BigQmtRpcClient(account_id="acct")
|
||||
client._formula_router_instance = router
|
||||
return client
|
||||
|
||||
def test_routed_method_never_touches_rpc(self):
|
||||
router = fs.FormulaServerRouter(
|
||||
client=FakeClient({"getLastVolume": {"result": 123.0}})
|
||||
)
|
||||
client = self._client(router)
|
||||
|
||||
def explode(*args, **kwargs):
|
||||
raise AssertionError("RPC must not be used for a routed method")
|
||||
|
||||
client._transport = explode
|
||||
|
||||
self.assertEqual(client.call("get_last_volume", {"stock": "000001.SZ"}), 123.0)
|
||||
|
||||
def test_unroutable_falls_back_to_rpc(self):
|
||||
router = fs.FormulaServerRouter(
|
||||
client=FakeClient(error=fs.FormulaServerUnavailable("down"))
|
||||
)
|
||||
client = self._client(router)
|
||||
calls = []
|
||||
|
||||
class FakeTransport:
|
||||
def send_request(self, request, timeout):
|
||||
calls.append(request["method"])
|
||||
return {"ok": True, "data": "from-rpc"}
|
||||
|
||||
client._transport = lambda: FakeTransport()
|
||||
|
||||
self.assertEqual(client.call("get_last_volume", {"stock": "000001.SZ"}), "from-rpc")
|
||||
self.assertEqual(calls, ["get_last_volume"])
|
||||
|
||||
def test_unmapped_method_goes_straight_to_rpc(self):
|
||||
router = fs.FormulaServerRouter(client=FakeClient())
|
||||
client = self._client(router)
|
||||
calls = []
|
||||
|
||||
class FakeTransport:
|
||||
def send_request(self, request, timeout):
|
||||
calls.append(request["method"])
|
||||
return {"ok": True, "data": {"cash": 1.0}}
|
||||
|
||||
client._transport = lambda: FakeTransport()
|
||||
|
||||
self.assertEqual(client.call("get_asset", {}), {"cash": 1.0})
|
||||
self.assertEqual(calls, ["get_asset"])
|
||||
|
||||
def test_routed_dataframe_payload_is_restored_like_rpc(self):
|
||||
router = fs.FormulaServerRouter(
|
||||
client=FakeClient(
|
||||
{
|
||||
"getMarketData": {
|
||||
"result": ["000001.SZ", ["20260703", ["close", 10.29]]]
|
||||
}
|
||||
}
|
||||
)
|
||||
)
|
||||
client = self._client(router)
|
||||
|
||||
out = client.call(
|
||||
"get_market_data_ex", {"field_list": ["close"], "stock_list": ["000001.SZ"]}
|
||||
)
|
||||
|
||||
frame = out["000001.SZ"]
|
||||
# _restore_jsonable rebuilds a DataFrame when pandas is present, and
|
||||
# degrades to the record list otherwise — same as the RPC path.
|
||||
if hasattr(frame, "columns"):
|
||||
self.assertEqual(list(frame.columns), ["stime", "close"])
|
||||
self.assertEqual(frame.iloc[0]["close"], 10.29)
|
||||
else:
|
||||
self.assertEqual(frame, [{"stime": "20260703", "close": 10.29}])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,114 +0,0 @@
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.full_tick_cache import (
|
||||
full_tick_demand_key,
|
||||
full_tick_request_id,
|
||||
read_full_tick_cache,
|
||||
refresh_full_tick_cache,
|
||||
request_full_tick_cache,
|
||||
)
|
||||
|
||||
|
||||
class FakeRedis:
|
||||
def __init__(self):
|
||||
self.hashes = {}
|
||||
self.kv = {}
|
||||
self.deleted = []
|
||||
self.expired = []
|
||||
|
||||
def hset(self, key, field, value):
|
||||
self.hashes.setdefault(key, {})[field] = value
|
||||
return 1
|
||||
|
||||
def hgetall(self, key):
|
||||
return self.hashes.get(key, {})
|
||||
|
||||
def hdel(self, key, field):
|
||||
self.deleted.append((key, field))
|
||||
self.hashes.setdefault(key, {}).pop(field, None)
|
||||
return 1
|
||||
|
||||
def expire(self, key, seconds):
|
||||
self.expired.append((key, seconds))
|
||||
return True
|
||||
|
||||
def setex(self, key, seconds, value):
|
||||
self.kv[key] = value
|
||||
self.expired.append((key, seconds))
|
||||
return True
|
||||
|
||||
def get(self, key):
|
||||
return self.kv.get(key)
|
||||
|
||||
|
||||
class FakeContext:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
def get_full_tick(self, codes):
|
||||
self.calls.append(list(codes))
|
||||
return {codes[0]: {"lastPrice": 10.0, "bidPrice": [9.9], "askPrice": [10.1]}}
|
||||
|
||||
|
||||
class FullTickCacheTest(unittest.TestCase):
|
||||
def test_request_then_refresh_writes_fresh_snapshot(self):
|
||||
redis_client = FakeRedis()
|
||||
context = FakeContext()
|
||||
|
||||
demand = request_full_tick_cache(redis_client, "acct", ["600000"], demand_ttl_seconds=10)
|
||||
refreshed = refresh_full_tick_cache(redis_client, context, "acct", cache_ttl_seconds=10)
|
||||
ticks = read_full_tick_cache(redis_client, "acct", ["600000.SH"], max_age_seconds=10)
|
||||
|
||||
self.assertEqual(demand["codes"], ["600000.SH"])
|
||||
self.assertEqual(refreshed, 1)
|
||||
self.assertEqual(context.calls, [["600000.SH"]])
|
||||
self.assertEqual(ticks["600000.SH"]["lastPrice"], 10.0)
|
||||
|
||||
def test_expired_demand_is_removed_without_refreshing(self):
|
||||
redis_client = FakeRedis()
|
||||
context = FakeContext()
|
||||
key = full_tick_demand_key("acct")
|
||||
request_id = full_tick_request_id(["600000.SH"])
|
||||
redis_client.hset(
|
||||
key,
|
||||
request_id,
|
||||
'{"request_id":"%s","codes":["600000.SH"],"requested_at_ts":1,"expires_at_ts":1}' % request_id,
|
||||
)
|
||||
|
||||
refreshed = refresh_full_tick_cache(redis_client, context, "acct", cache_ttl_seconds=10)
|
||||
|
||||
self.assertEqual(refreshed, 0)
|
||||
self.assertEqual(context.calls, [])
|
||||
self.assertIn((key, request_id), redis_client.deleted)
|
||||
|
||||
def test_refresh_kind_symbol_skips_market_demands(self):
|
||||
redis_client = FakeRedis()
|
||||
context = FakeContext()
|
||||
request_full_tick_cache(redis_client, "acct", ["600000"], demand_ttl_seconds=10)
|
||||
request_full_tick_cache(redis_client, "acct", ["SH", "SZ"], demand_ttl_seconds=10)
|
||||
|
||||
refreshed = refresh_full_tick_cache(redis_client, context, "acct", cache_ttl_seconds=10, kind="symbol")
|
||||
|
||||
self.assertEqual(refreshed, 1)
|
||||
self.assertEqual(context.calls, [["600000.SH"]])
|
||||
|
||||
def test_refresh_kind_market_skips_symbol_demands(self):
|
||||
redis_client = FakeRedis()
|
||||
context = FakeContext()
|
||||
request_full_tick_cache(redis_client, "acct", ["600000"], demand_ttl_seconds=10)
|
||||
request_full_tick_cache(redis_client, "acct", ["SH", "SZ"], demand_ttl_seconds=10)
|
||||
|
||||
refreshed = refresh_full_tick_cache(redis_client, context, "acct", cache_ttl_seconds=10, kind="market")
|
||||
|
||||
self.assertEqual(refreshed, 1)
|
||||
self.assertEqual(context.calls, [["SH", "SZ"]])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -1,348 +0,0 @@
|
||||
import os
|
||||
import shutil
|
||||
import sys
|
||||
import tempfile
|
||||
import unittest
|
||||
|
||||
|
||||
ROOT = os.path.dirname(os.path.dirname(os.path.dirname(__file__)))
|
||||
sys.path.insert(0, os.path.join(ROOT, "src"))
|
||||
|
||||
from bigqmt_signal_trader.local_cache import LocalMarketCache
|
||||
|
||||
|
||||
def _has_pyarrow():
|
||||
try:
|
||||
import pyarrow # noqa: F401
|
||||
|
||||
return True
|
||||
except Exception:
|
||||
return False
|
||||
|
||||
|
||||
class LocalMarketCacheTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.dir = tempfile.mkdtemp()
|
||||
|
||||
def tearDown(self):
|
||||
shutil.rmtree(self.dir, ignore_errors=True)
|
||||
|
||||
def test_write_read_merge_dedupe(self):
|
||||
import pandas as pd
|
||||
|
||||
c = LocalMarketCache(self.dir)
|
||||
c.write("600000.SH", "1d", pd.DataFrame({"stime": ["20260101", "20260102"], "close": [1.0, 2.0]}))
|
||||
# overlapping second write: 20260102 should be replaced (keep last), 20260103 appended
|
||||
c.write("600000.SH", "1d", pd.DataFrame({"stime": ["20260102", "20260103"], "close": [2.5, 3.0]}))
|
||||
|
||||
df = c.read("600000.SH", "1d")
|
||||
self.assertEqual(list(df["stime"]), ["20260101", "20260102", "20260103"])
|
||||
self.assertEqual(df[df["stime"] == "20260102"]["close"].iloc[0], 2.5)
|
||||
|
||||
def test_range_and_count_filters(self):
|
||||
import pandas as pd
|
||||
|
||||
c = LocalMarketCache(self.dir)
|
||||
c.write("X", "1d", pd.DataFrame({"stime": ["20260101", "20260102", "20260103"], "close": [1, 2, 3]}))
|
||||
|
||||
self.assertEqual(list(c.read("X", "1d", start_time="20260102")["stime"]), ["20260102", "20260103"])
|
||||
self.assertEqual(list(c.read("X", "1d", end_time="20260102")["stime"]), ["20260101", "20260102"])
|
||||
self.assertEqual(list(c.read("X", "1d", count=1)["stime"]), ["20260103"])
|
||||
self.assertIsNone(c.read("MISSING", "1d"))
|
||||
self.assertEqual(c.covered("X", "1d"), ("20260101", "20260103", 3))
|
||||
|
||||
def test_drops_zero_fill_placeholder_rows(self):
|
||||
import pandas as pd
|
||||
|
||||
c = LocalMarketCache(self.dir)
|
||||
df = pd.DataFrame(
|
||||
{"stime": ["20200101", "20200102", "20260701"], "close": [0.0, 0.0, 8.65], "open": [0.0, 0.0, 8.58]}
|
||||
)
|
||||
c.write("X", "1d", df)
|
||||
self.assertEqual(list(c.read("X", "1d")["stime"]), ["20260701"]) # 0-fill dropped
|
||||
|
||||
# an all-placeholder write must not create/overwrite a cache file
|
||||
self.assertEqual(c.write("Y", "1d", pd.DataFrame({"stime": ["20200101"], "close": [0.0]})), 0)
|
||||
self.assertIsNone(c.read("Y", "1d"))
|
||||
|
||||
def test_dividend_type_keeps_separate_caches(self):
|
||||
import pandas as pd
|
||||
|
||||
c = LocalMarketCache(self.dir)
|
||||
c.write("X", "1d", pd.DataFrame({"stime": ["20260101"], "close": [10.0]}), dividend_type="none")
|
||||
c.write("X", "1d", pd.DataFrame({"stime": ["20260101"], "close": [9.0]}), dividend_type="front")
|
||||
|
||||
self.assertEqual(c.read("X", "1d", dividend_type="none")["close"].iloc[0], 10.0)
|
||||
self.assertEqual(c.read("X", "1d", dividend_type="front")["close"].iloc[0], 9.0)
|
||||
self.assertIsNone(c.read("X", "1d", dividend_type="back"))
|
||||
|
||||
def test_pickle_format_roundtrip(self):
|
||||
import pandas as pd
|
||||
|
||||
c = LocalMarketCache(self.dir, fmt="pkl")
|
||||
c.write("X", "1d", pd.DataFrame({"stime": ["20260101", "20260102"], "close": [1.0, 2.0]}))
|
||||
self.assertTrue(c.path("X", "1d").endswith(".pkl"))
|
||||
self.assertEqual(list(c.read("X", "1d")["close"]), [1.0, 2.0])
|
||||
|
||||
@unittest.skipUnless(_has_pyarrow(), "pyarrow not installed")
|
||||
def test_parquet_format_roundtrip(self):
|
||||
import pandas as pd
|
||||
|
||||
c = LocalMarketCache(self.dir, fmt="parquet")
|
||||
c.write("X", "1d", pd.DataFrame({"stime": ["20260101", "20260102"], "close": [1.0, 2.0]}))
|
||||
self.assertTrue(c.path("X", "1d").endswith(".parquet"))
|
||||
self.assertEqual(list(c.read("X", "1d")["close"]), [1.0, 2.0])
|
||||
|
||||
@unittest.skipUnless(_has_pyarrow(), "pyarrow not installed")
|
||||
def test_migrates_pickle_to_parquet(self):
|
||||
import pandas as pd
|
||||
|
||||
LocalMarketCache(self.dir, fmt="pkl").write("X", "1d", pd.DataFrame({"stime": ["20260101"], "close": [1.0]}))
|
||||
pq = LocalMarketCache(self.dir, fmt="parquet")
|
||||
self.assertEqual(list(pq.read("X", "1d")["close"]), [1.0]) # reads the old pkl
|
||||
pq.write("X", "1d", pd.DataFrame({"stime": ["20260102"], "close": [2.0]}))
|
||||
self.assertTrue(os.path.isfile(pq.path("X", "1d"))) # parquet now exists
|
||||
self.assertFalse(os.path.isfile(pq.path("X", "1d")[:-8] + ".pkl")) # old pkl removed
|
||||
self.assertEqual(list(pq.read("X", "1d")["close"]), [1.0, 2.0]) # merged across formats
|
||||
|
||||
|
||||
class FakeClient:
|
||||
def __init__(self, cache_dir, fallback_rpc=False):
|
||||
self.account_id = "acct"
|
||||
self.calls = []
|
||||
self.call_params = []
|
||||
self.local_cache_config = {"enabled": True, "dir": cache_dir, "fallback_rpc": fallback_rpc}
|
||||
|
||||
def _redis(self):
|
||||
return None
|
||||
|
||||
def call(self, method, params=None, account_id=None, timeout_seconds=None):
|
||||
self.calls.append(method)
|
||||
self.call_params.append((method, params))
|
||||
if method == "get_market_data_ex":
|
||||
import pandas as pd
|
||||
|
||||
codes = (params or {}).get("stock_list") or []
|
||||
return {c: pd.DataFrame({"stime": ["20260626", "20260629"], "close": [8.76, 8.73]}) for c in codes}
|
||||
if method == "download_history_data2":
|
||||
# Server-side raw download (raw bars + dividend factors).
|
||||
return True
|
||||
raise AssertionError("unexpected rpc: %s" % method)
|
||||
|
||||
|
||||
class LocalCacheClientTest(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.dir = tempfile.mkdtemp()
|
||||
|
||||
def tearDown(self):
|
||||
shutil.rmtree(self.dir, ignore_errors=True)
|
||||
|
||||
def _xt(self, fallback_rpc=False):
|
||||
from bigqmt_signal_trader.xtquant_compat import BigQmtXtData
|
||||
|
||||
return BigQmtXtData(FakeClient(self.dir, fallback_rpc=fallback_rpc))
|
||||
|
||||
def test_download_caches_then_get_local_reads_without_rpc(self):
|
||||
xt = self._xt()
|
||||
progress = []
|
||||
res = xt.download_history_data2(["600000.SH", "000001.SZ"], "1d", callback=lambda d: progress.append(d))
|
||||
|
||||
self.assertEqual(res, {"finished": 2, "total": 2})
|
||||
self.assertEqual(len(progress), 2)
|
||||
self.assertEqual(progress[-1]["stockcode"], "000001.SZ")
|
||||
self.assertEqual(progress[-1]["finished"], 2)
|
||||
calls_after_download = list(xt.client.calls)
|
||||
|
||||
data = xt.get_local_data(stock_list=["600000.SH", "000001.SZ"], period="1d")
|
||||
|
||||
self.assertIn("600000.SH", data)
|
||||
self.assertIn("000001.SZ", data)
|
||||
self.assertEqual(list(data["600000.SH"]["close"]), [8.76, 8.73])
|
||||
# get_local_data must NOT issue any further RPC — pure local read.
|
||||
self.assertEqual(xt.client.calls, calls_after_download)
|
||||
|
||||
def test_get_market_data_ex_caches_through(self):
|
||||
xt = self._xt()
|
||||
# a plain live read must also populate the cache (cache-through)
|
||||
xt.get_market_data_ex(field_list=["close"], stock_list=["600000.SH"], period="1d")
|
||||
n = len(xt.client.calls)
|
||||
|
||||
data = xt.get_local_data(stock_list=["600000.SH"], period="1d")
|
||||
self.assertIn("600000.SH", data)
|
||||
self.assertEqual(len(xt.client.calls), n) # served from cache, no extra RPC
|
||||
|
||||
def test_get_local_miss_returns_empty_and_no_rpc(self):
|
||||
xt = self._xt()
|
||||
data = xt.get_local_data(stock_list=["600000.SH"], period="1d")
|
||||
self.assertEqual(data, {})
|
||||
self.assertEqual(xt.client.calls, [])
|
||||
|
||||
def test_get_local_fallback_rpc_fetches_and_caches(self):
|
||||
xt = self._xt(fallback_rpc=True)
|
||||
data = xt.get_local_data(stock_list=["600000.SH"], period="1d")
|
||||
self.assertIn("600000.SH", data)
|
||||
self.assertIn("get_market_data_ex", xt.client.calls) # fetched on miss
|
||||
# second read is served from cache — no new RPC
|
||||
n = len(xt.client.calls)
|
||||
xt.get_local_data(stock_list=["600000.SH"], period="1d")
|
||||
self.assertEqual(len(xt.client.calls), n)
|
||||
|
||||
|
||||
class AdjustedDownloadTest(unittest.TestCase):
|
||||
"""Adjusted (front/back) downloads must trigger the server-side raw
|
||||
download FIRST: Big QMT computes adjusted bars from raw bars + dividend
|
||||
factors, and without the server-side download the adjusted result is
|
||||
all zeros (verified live with 600654.SH)."""
|
||||
|
||||
def setUp(self):
|
||||
self.dir = tempfile.mkdtemp()
|
||||
|
||||
def tearDown(self):
|
||||
shutil.rmtree(self.dir, ignore_errors=True)
|
||||
|
||||
def _xt(self):
|
||||
from bigqmt_signal_trader.xtquant_compat import BigQmtXtData
|
||||
|
||||
return BigQmtXtData(FakeClient(self.dir))
|
||||
|
||||
def test_front_download_triggers_server_side_raw_download_first(self):
|
||||
xt = self._xt()
|
||||
xt.download_history_data2(["600000.SH"], "1d", start_time="20200101", dividend_type="front")
|
||||
|
||||
# The server-side raw download must run BEFORE the adjusted pull.
|
||||
method_calls = [m for m, _ in xt.client.call_params]
|
||||
self.assertIn("download_history_data2", method_calls)
|
||||
self.assertIn("get_market_data_ex", method_calls)
|
||||
self.assertLess(
|
||||
method_calls.index("download_history_data2"),
|
||||
method_calls.index("get_market_data_ex"),
|
||||
"server-side raw download must precede the adjusted pull",
|
||||
)
|
||||
# The raw download carries the same codes/period/window.
|
||||
raw_call = next(p for m, p in xt.client.call_params if m == "download_history_data2")
|
||||
self.assertEqual(raw_call["stock_list"], ["600000.SH"])
|
||||
self.assertEqual(raw_call["period"], "1d")
|
||||
self.assertEqual(raw_call["start_time"], "20200101")
|
||||
|
||||
def test_none_download_still_triggers_server_side_download(self):
|
||||
"""issue #47: an unadjusted download used to skip the server RPC and
|
||||
only read what Big QMT already had -- a no-op that still reported
|
||||
{finished: N}. xtdata semantics are "populate the local QMT store", and
|
||||
callers (FormulaServer, get_local_data) depend on that actually happening."""
|
||||
xt = self._xt()
|
||||
xt.download_history_data2(["600000.SH"], "1d", dividend_type="none")
|
||||
|
||||
method_calls = [m for m, _ in xt.client.call_params]
|
||||
self.assertIn("download_history_data2", method_calls)
|
||||
self.assertIn("get_market_data_ex", method_calls)
|
||||
self.assertLess(
|
||||
method_calls.index("download_history_data2"),
|
||||
method_calls.index("get_market_data_ex"),
|
||||
"the download must precede the pull, or the pull reads stale data",
|
||||
)
|
||||
raw_call = next(p for m, p in xt.client.call_params if m == "download_history_data2")
|
||||
self.assertEqual(raw_call["stock_list"], ["600000.SH"])
|
||||
self.assertEqual(raw_call["period"], "1d")
|
||||
|
||||
def test_none_download_survives_server_download_failure(self):
|
||||
"""Same best-effort contract the adjusted path already had: a deployment
|
||||
without the QMT global must still get its bars."""
|
||||
xt = self._xt()
|
||||
original_call = xt.client.call
|
||||
|
||||
def failing_download(method, params=None, account_id=None, timeout_seconds=None):
|
||||
if method == "download_history_data2":
|
||||
raise RuntimeError("global not available")
|
||||
return original_call(method, params, account_id=account_id, timeout_seconds=timeout_seconds)
|
||||
|
||||
xt.client.call = failing_download
|
||||
result = xt.download_history_data2(["600000.SH"], "1d", dividend_type="none")
|
||||
|
||||
self.assertEqual(result["finished"], 1)
|
||||
self.assertIn("get_market_data_ex", [m for m, _ in xt.client.call_params])
|
||||
|
||||
def test_front_download_survives_server_download_failure(self):
|
||||
# Deployments without the QMT global must still get the adjusted pull
|
||||
# (best-effort raw download, never fatal).
|
||||
xt = self._xt()
|
||||
original_call = xt.client.call
|
||||
|
||||
def failing_download(method, params=None, account_id=None, timeout_seconds=None):
|
||||
if method == "download_history_data2":
|
||||
raise RuntimeError("global not available")
|
||||
return original_call(method, params, account_id=account_id, timeout_seconds=timeout_seconds)
|
||||
|
||||
xt.client.call = failing_download
|
||||
result = xt.download_history_data2(["600000.SH"], "1d", dividend_type="front")
|
||||
self.assertEqual(result, {"finished": 1, "total": 1}) # adjusted pull still ran
|
||||
|
||||
|
||||
class _AllZeroThenRealClient(FakeClient):
|
||||
"""First adjusted get_market_data_ex returns all-zero bars (server lacks
|
||||
raw data); after a server-side raw download, subsequent pulls are real."""
|
||||
|
||||
def __init__(self, cache_dir):
|
||||
super(_AllZeroThenRealClient, self).__init__(cache_dir)
|
||||
self._downloaded = False
|
||||
|
||||
def call(self, method, params=None, account_id=None, timeout_seconds=None):
|
||||
self.calls.append(method)
|
||||
self.call_params.append((method, params))
|
||||
import pandas as pd
|
||||
|
||||
if method == "download_history_data2":
|
||||
self._downloaded = True
|
||||
return True
|
||||
if method == "get_market_data_ex":
|
||||
codes = (params or {}).get("stock_list") or []
|
||||
if self._downloaded:
|
||||
return {c: pd.DataFrame({"stime": ["20260626", "20260629"], "close": [8.76, 8.73]}) for c in codes}
|
||||
# all-zero symptom: head zeros, last bar live
|
||||
return {c: pd.DataFrame({"stime": ["20260626", "20260629"], "close": [0.0, 8.73]}) for c in codes}
|
||||
raise AssertionError("unexpected rpc: %s" % method)
|
||||
|
||||
|
||||
class AdjustedReadSelfHealTest(unittest.TestCase):
|
||||
"""Reading adjusted bars that come back all-zero must self-heal:
|
||||
trigger a server-side raw download, wait, and retry once."""
|
||||
|
||||
def setUp(self):
|
||||
self.dir = tempfile.mkdtemp()
|
||||
|
||||
def tearDown(self):
|
||||
shutil.rmtree(self.dir, ignore_errors=True)
|
||||
|
||||
def _xt(self):
|
||||
from bigqmt_signal_trader.xtquant_compat import BigQmtXtData
|
||||
|
||||
return BigQmtXtData(_AllZeroThenRealClient(self.dir))
|
||||
|
||||
def test_front_read_self_heals_all_zero_to_real(self):
|
||||
xt = self._xt()
|
||||
data = xt.get_market_data_ex(
|
||||
field_list=["close"], stock_list=["600000.SH"], period="1d",
|
||||
dividend_type="front",
|
||||
)
|
||||
# After self-heal the retry returns real (non-zero) bars.
|
||||
self.assertEqual(list(data["600000.SH"]["close"]), [8.76, 8.73])
|
||||
# The heal path must have triggered a server-side raw download.
|
||||
method_calls = [m for m, _ in xt.client.call_params]
|
||||
self.assertIn("download_history_data2", method_calls)
|
||||
# get_market_data_ex called twice: initial all-zero pull + retry.
|
||||
self.assertEqual(method_calls.count("get_market_data_ex"), 2)
|
||||
|
||||
def test_none_read_does_not_self_heal(self):
|
||||
xt = self._xt()
|
||||
data = xt.get_market_data_ex(
|
||||
field_list=["close"], stock_list=["600000.SH"], period="1d",
|
||||
dividend_type="none",
|
||||
)
|
||||
# none returns whatever the server sent (no heal, single pull).
|
||||
self.assertEqual(list(data["600000.SH"]["close"]), [0.0, 8.73])
|
||||
method_calls = [m for m, _ in xt.client.call_params]
|
||||
self.assertNotIn("download_history_data2", method_calls)
|
||||
self.assertEqual(method_calls.count("get_market_data_ex"), 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user