419 lines
14 KiB
Python
419 lines
14 KiB
Python
# -*- coding: utf-8 -*-
|
|
import json
|
|
import locale
|
|
import os
|
|
import sys
|
|
from urllib.request import Request, urlopen
|
|
from tornado.web import Application, RequestHandler, HTTPError
|
|
from tornado.ioloop import IOLoop
|
|
import logging
|
|
|
|
# Configuration
|
|
ACCOUNT_ID = os.environ.get('QMT_ACCOUNT_ID', '')
|
|
DATA_DIR = os.environ.get('QMT_DATA_DIR', r'D:\qmt_strategy_data')
|
|
TOKEN="QMTbyYanweidong"
|
|
PORT = 10086
|
|
PASS_CODES_URL = "http://139.224.247.176:13499/a/pass_codes"
|
|
|
|
# ===================================
|
|
logging.basicConfig(level=logging.INFO)
|
|
logger = logging.getLogger(__name__)
|
|
locale.setlocale(locale.LC_CTYPE, 'chinese')
|
|
|
|
|
|
def safe_call(func, *args, **kwargs):
|
|
try:
|
|
return func(*args, **kwargs)
|
|
except HTTPError:
|
|
raise
|
|
except Exception as e:
|
|
raise HTTPError(
|
|
500,
|
|
reason="QMT: %s call failed." % func.__name__,
|
|
) from e
|
|
|
|
|
|
def get_pass_codes(account_id):
|
|
request = Request(
|
|
PASS_CODES_URL,
|
|
headers={"Accept": "application/json", "User-Agent": "big-qmt/1"},
|
|
)
|
|
with urlopen(request, timeout=10) as response:
|
|
payload = json.load(response)
|
|
|
|
remote_codes = payload.get("data")
|
|
if not isinstance(remote_codes, list):
|
|
raise ValueError("pass_codes response data must be an array")
|
|
|
|
positions = safe_call(
|
|
get_trade_detail_data, account_id, 'stock', 'position'
|
|
) or []
|
|
position_codes = [
|
|
position.m_strInstrumentID + '.' + position.m_strExchangeID
|
|
for position in positions
|
|
]
|
|
|
|
codes = []
|
|
seen = set()
|
|
for code in remote_codes + position_codes:
|
|
code = str(code).strip()
|
|
if code and code not in seen:
|
|
seen.add(code)
|
|
codes.append(code)
|
|
return codes
|
|
|
|
|
|
# ============= BaseHandler =============
|
|
|
|
AUTH_EXEMPT = set()
|
|
|
|
|
|
def no_auth(cls):
|
|
AUTH_EXEMPT.add(cls)
|
|
return cls
|
|
|
|
|
|
class BaseHandler(RequestHandler):
|
|
def prepare(self):
|
|
if self.__class__ not in AUTH_EXEMPT:
|
|
token = self.request.headers.get('X-Token')
|
|
if token != TOKEN:
|
|
raise HTTPError(500, "Authentication failed: invalid or missing token")
|
|
|
|
def set_default_headers(self):
|
|
self.set_header("Content-Type", "application/json; charset=utf-8")
|
|
|
|
def write_error(self,status_code, **kwargs):
|
|
self.finish(self._reason)
|
|
|
|
def ctx(self):
|
|
return self.application.ContextInfo
|
|
|
|
def acc(self):
|
|
return self.application.accountID
|
|
|
|
def write_json(self, data, default=None):
|
|
self.write(json.dumps(
|
|
data,
|
|
separators=(',', ':'),
|
|
ensure_ascii=False,
|
|
default=default,
|
|
))
|
|
|
|
|
|
# ============= 1. ContextInfo properties =============
|
|
# "/api/v2/context/info"
|
|
class ContextInfoHandler(BaseHandler):
|
|
def get(self):
|
|
ctx = self.ctx()
|
|
data = {
|
|
"period": ctx.period,
|
|
"barpos": ctx.barpos,
|
|
"time_tick_size": ctx.time_tick_size,
|
|
"stockcode": ctx.stockcode,
|
|
"dividend_type": ctx.dividend_type,
|
|
"market": ctx.market,
|
|
"do_back_test": ctx.do_back_test,
|
|
"benchmark": ctx.benchmark,
|
|
"capital": ctx.capital,
|
|
"timetag":ctx.timetag,
|
|
"universe": ctx.get_universe(),
|
|
}
|
|
self.write_json(data)
|
|
|
|
# ============= 2. Data queries (ContextInfo get_*) =============
|
|
STOCK_HANDLER = {
|
|
# handler_type: (method_name, use_context)
|
|
"stock_name": ("get_stock_name", True),
|
|
"open_date": ("get_open_date", True),
|
|
"last_volume": ("get_last_volume", True),
|
|
"total_share": ("get_total_share", True),
|
|
"svol": ("get_svol", True),
|
|
"bvol": ("get_bvol", True),
|
|
"divid_factors": ("get_divid_factors", True),
|
|
"etf_info": ("get_etf_info", False),
|
|
"etf_iopv": ("get_etf_iopv", False),
|
|
"instrumentdetail": ("get_instrumentdetail", True),
|
|
"his_st_data": ("get_his_st_data", True),
|
|
}
|
|
# "/api/v2/get/*" Stock-related single-symbol queries
|
|
class StockGetHandler(BaseHandler):
|
|
def get(self, handler_type):
|
|
# 快速路径:配置查找
|
|
cfg = STOCK_HANDLER.get(handler_type)
|
|
if not cfg:
|
|
raise HTTPError(500, "Unknown API")
|
|
|
|
# 参数验证
|
|
|
|
query_vals = self.get_query_argument("stock_code", "").strip()
|
|
if not query_vals:
|
|
raise HTTPError(500, "stock_code required")
|
|
|
|
# 方法调用
|
|
method_name, use_context = cfg
|
|
method = getattr(self.ctx(), method_name) if use_context else globals()[method_name]
|
|
result = safe_call(method, query_vals)
|
|
|
|
|
|
# 响应
|
|
self.write_json({
|
|
"stock_code": query_vals,
|
|
"ref": result
|
|
}, default=str)
|
|
|
|
|
|
# Aggregate assets, positions, and orders in one request.
|
|
class PortfolioHandler(BaseHandler):
|
|
def get(self):
|
|
account_id = self.acc()
|
|
account_data = safe_call(get_trade_detail_data, account_id, 'stock', 'account')
|
|
positions = safe_call(get_trade_detail_data, account_id, 'stock', 'position') or []
|
|
orders = safe_call(get_trade_detail_data, account_id, 'stock', 'order') or []
|
|
|
|
result = {
|
|
"assets": format_assets(account_data),
|
|
"positions": format_holding(positions),
|
|
"orders": [fixed_fields(order) for order in orders],
|
|
}
|
|
self.write_json(result)
|
|
|
|
|
|
# get_trade_detail_data('position') - Query positions in the wrapped format
|
|
class HoldingHandler(BaseHandler):
|
|
def get(self):
|
|
positions = safe_call(get_trade_detail_data, self.acc(), 'stock', 'position') or []
|
|
holding = format_holding(positions)
|
|
self.write_json({"data": holding})
|
|
|
|
# get_trade_detail_data('account') - Query account assets
|
|
class AssetsHandler(BaseHandler):
|
|
def get(self):
|
|
_data = safe_call(get_trade_detail_data, self.acc(), 'stock', 'account')
|
|
self.write_json(format_assets(_data))
|
|
|
|
class OrderHandler(BaseHandler):
|
|
def get(self):
|
|
ret = safe_call(get_trade_detail_data, self.acc(), 'stock', 'order')
|
|
if ret is None:
|
|
ret = []
|
|
result = [fixed_fields(obj) for obj in ret]
|
|
self.write_json(result)
|
|
|
|
class DealHandler(BaseHandler):
|
|
def get(self):
|
|
deals = safe_call(get_trade_detail_data, self.acc(), 'stock', 'deal') or []
|
|
rets = [fixed_fields(deal) for deal in deals]
|
|
self.write_json({"deals": rets})
|
|
|
|
# ContextInfo.get_full_tick() - Get full tick data
|
|
class FullTickHandler(BaseHandler):
|
|
def post(self):
|
|
data = json.loads(self.request.body)
|
|
stocks = data.get('stocks', [])
|
|
#if not stocks:
|
|
# raise HTTPError(400, "need args stocks")
|
|
ret = safe_call(self.ctx().get_full_tick, stocks)
|
|
if not ret:
|
|
raise HTTPError(500, "Failed to get tick data")
|
|
self.write_json(ret, default=str)
|
|
|
|
# passorder() - Submit a general trading order
|
|
class PassorderHandler(BaseHandler):
|
|
def post(self):
|
|
try:
|
|
data = json.loads(self.request.body)
|
|
opType = int(data['opType'])
|
|
orderType = int(data.get('orderType', 1101))
|
|
stockCode = data['stockCode']
|
|
pr_type = int(data.get('prType', 11))
|
|
price = float(data['price'])
|
|
volume = int(data['volume'])
|
|
quickTrade = int(data.get('quickTrade', 2))
|
|
strategy_name = str(data.get('strategyName', '')).strip()
|
|
order_id = str(data.get('orderId', '')).strip()
|
|
except (json.JSONDecodeError, KeyError, TypeError, ValueError) as e:
|
|
raise HTTPError(400, reason="Invalid order parameters: %s" % e) from e
|
|
|
|
# QMT stores strategyName in the order remark; preserve the signal key and local order ID.
|
|
try:
|
|
order_ref = passorder(opType, orderType, self.acc(), stockCode, pr_type, price, volume, strategy_name, quickTrade,order_id, self.ctx())
|
|
except HTTPError:
|
|
raise
|
|
except Exception as e:
|
|
logger.exception("passorder failed")
|
|
raise HTTPError(502, reason="QMT order submission failed") from e
|
|
|
|
self.write_json({
|
|
"status": "success",
|
|
"opType": opType,
|
|
"stockCode": stockCode,
|
|
"strategy_name": strategy_name,
|
|
"local_order_id": order_id,
|
|
"order_ref": str(order_ref)
|
|
})
|
|
|
|
|
|
class CancelByIdHandler(BaseHandler):
|
|
"""Cancel an order by its actual system order ID."""
|
|
def post(self):
|
|
data = json.loads(self.request.body)
|
|
order_id = str(data.get('order_id', '')).strip()
|
|
if not order_id:
|
|
raise HTTPError(400, "order_id cannot be empty")
|
|
cancelable = safe_call(can_cancel_order, order_id, self.acc(), 'stock')
|
|
if not cancelable:
|
|
self.write_json({
|
|
"status": "failed", "order_id": order_id,
|
|
"message": "Order does not exist or cannot currently be canceled"
|
|
})
|
|
return
|
|
result = safe_call(cancel, order_id, self.acc(), 'stock', self.ctx())
|
|
self.write_json({
|
|
"status": "success" if result is not False else "failed",
|
|
"order_id": order_id,
|
|
})
|
|
|
|
# get_ipo_data() - Get today's new stock and bond offerings
|
|
class IpoDataHandler(BaseHandler):
|
|
def post(self):
|
|
data = json.loads(self.request.body)
|
|
typ = data.get('type', 'STOCK')
|
|
ret = safe_call(get_ipo_data, typ)
|
|
self.write_json(ret)
|
|
|
|
# sys: Python version information
|
|
class PythonVersionHandler(BaseHandler):
|
|
def get(self):
|
|
version_info = {
|
|
"python_version": sys.version,
|
|
"python_version_info": {
|
|
"major": sys.version_info.major,
|
|
"minor": sys.version_info.minor,
|
|
"micro": sys.version_info.micro,
|
|
"releaselevel": sys.version_info.releaselevel,
|
|
"serial": sys.version_info.serial,
|
|
}
|
|
}
|
|
self.write_json(version_info)
|
|
|
|
|
|
def format_assets(account_data):
|
|
info = account_data[0] if account_data else None
|
|
if not info:
|
|
raise HTTPError(500, "Failed to get account data")
|
|
return {
|
|
"total": round(info.m_dBalance, 2),
|
|
"available": round(info.m_dAvailable, 2),
|
|
}
|
|
|
|
def format_holding(positions):
|
|
holding = {}
|
|
for position in positions:
|
|
stock = position.m_strInstrumentID + '.' + position.m_strExchangeID
|
|
holding[stock] = {
|
|
'StockCode': stock,
|
|
'TradeID':position.m_strTradeID,
|
|
'StockName': position.m_strInstrumentName,
|
|
'Direction': position.m_nDirection,
|
|
'Volume': position.m_nVolume,
|
|
'OpenPrice': position.m_dOpenPrice,
|
|
'FloatProfit': position.m_dFloatProfit,
|
|
'MarketValue': position.m_dMarketValue,
|
|
'StockHolder': position.m_strStockHolder,
|
|
'FrozenVolume': position.m_nFrozenVolume,
|
|
'CanUseVolume': position.m_nCanUseVolume,
|
|
'OnRoadVolume': position.m_nOnRoadVolume,
|
|
'YesterdayVolume': position.m_nYesterdayVolume,
|
|
'LastPrice': position.m_dLastPrice,
|
|
'ProfitRate': position.m_dProfitRate,
|
|
'FutureTradeType': position.m_eFutureTradeType,
|
|
'ExpireDate': position.m_strExpireDate
|
|
}
|
|
return holding
|
|
|
|
TRADE_DETAIL_FIELDS = (
|
|
'm_strOrderSysID', 'm_strInstrumentID', 'm_strExchangeID',
|
|
'm_strInstrumentName', 'm_nOffsetFlag', 'm_nOrderStatus',
|
|
'm_nVolumeTotal', 'm_nVolumeTraded', 'm_nOrderTime',
|
|
'm_strInsertDate', 'm_strInsertTime', 'm_strRemark',
|
|
'm_dPrice', 'm_dTradePrice', 'm_dTradeAmount',
|
|
)
|
|
MISSING = object()
|
|
|
|
|
|
def fixed_fields(obj, fields=TRADE_DETAIL_FIELDS):
|
|
result = {}
|
|
for field in fields:
|
|
try:
|
|
value = getattr(obj, field, MISSING)
|
|
except TypeError:
|
|
continue
|
|
if value is MISSING:
|
|
continue
|
|
if not callable(value):
|
|
result[field] = str(value)
|
|
if not result:
|
|
attrs = getattr(obj, '__dict__', {})
|
|
result = {
|
|
key: str(value) for key, value in attrs.items()
|
|
if not key.startswith('_') and not callable(value)
|
|
}
|
|
return result
|
|
|
|
|
|
# ============= Route registration =============
|
|
def make_app():
|
|
return Application([
|
|
# ContextInfo properties
|
|
(r"/api/context/info", ContextInfoHandler),
|
|
(r"/api/get/(stock_name|open_date|last_volume|total_share|svol|bvol|divid_factors|etf_info|etf_iopv|instrumentdetail|his_st_data)", StockGetHandler),
|
|
|
|
# Portfolio
|
|
(r"/api/portfolio", PortfolioHandler),
|
|
(r"/api/portfolio/positions", HoldingHandler),
|
|
(r"/api/portfolio/assets", AssetsHandler),
|
|
(r"/api/portfolio/order", OrderHandler),
|
|
(r"/api/portfolio/deal", DealHandler),
|
|
|
|
(r"/api/data/full_tick", FullTickHandler),
|
|
(r"/api/trade/ipo_data", IpoDataHandler),
|
|
(r"/api/trade/cancel_by_id", CancelByIdHandler),
|
|
(r"/api/trade/passorder", PassorderHandler),
|
|
|
|
# System
|
|
(r"/api/sys/python_version", PythonVersionHandler),
|
|
|
|
], debug=False)
|
|
|
|
|
|
def init(ContextInfo):
|
|
if not (ACCOUNT_ID or "").strip():
|
|
msg = "ACCOUNT_ID is empty; startup aborted"
|
|
logger.error(msg)
|
|
raise ValueError(msg)
|
|
if not (DATA_DIR or "").strip():
|
|
msg = "DATA_DIR is empty; startup aborted"
|
|
logger.error(msg)
|
|
raise ValueError(msg)
|
|
try:
|
|
ContextInfo.accountID = ACCOUNT_ID
|
|
ContextInfo.set_account(ACCOUNT_ID)
|
|
|
|
codes = get_pass_codes(ContextInfo.accountID)
|
|
ContextInfo.set_universe(list(codes))
|
|
|
|
# Api App
|
|
app = make_app()
|
|
app.ContextInfo = ContextInfo
|
|
app.accountID = ContextInfo.accountID
|
|
app.listen(PORT, address='0.0.0.0')
|
|
logger.info(f"ACCOUNT_ID: {ACCOUNT_ID}")
|
|
logger.info(f"DATA_DIR: {DATA_DIR}")
|
|
logger.info(f"TOKEN: {TOKEN}")
|
|
logger.info(f"Initialized symbol universe with {len(codes)} instruments")
|
|
logger.info(f"QMT HTTP Server started at http://0.0.0.0:{PORT} (all APIs loaded)")
|
|
IOLoop.current().start()
|
|
except Exception as e:
|
|
logger.exception(f"server start failed: {e}")
|