refactor QMT client and optimize API

This commit is contained in:
2026-08-28 22:46:04 +08:00
parent d09f271569
commit 29fee85b3d
18 changed files with 1239 additions and 519 deletions

372
AUDIT_AND_REMEDIATION.md Normal file
View File

@@ -0,0 +1,372 @@
# big-qmt 项目审计与整改建议
- 审计日期2026-08-28
- 审计范围:服务端 `api/`、客户端 `py-client/`
- 审计方式:静态代码检查、调用链核对、只读语法编译
- 当前状态:仅供人工确认,尚未实施代码整改
## 一、总体结论
当前版本不建议直接进入实盘运行。
服务端存在文件编码导致的启动级错误,且所有 QMT 调用和大对象序列化都在 Tornado 主线程同步执行。客户端的持仓管理、止盈、补仓、撤单和下单确认链存在多处必现错误或状态不一致风险。
建议按以下顺序处理:
1. 恢复服务端和客户端的基本可运行性。
2. 修复交易安全相关的订单确认、撤单和持仓数据模型。
3. 为止盈、补仓、订单状态机建立测试。
4. 在确认 QMT 线程约束后优化服务端响应速度。
5. 最后进行结构简化和重复代码清理。
---
## 二、P0启动及交易安全问题
### 2.1 服务端文件编码不一致,程序无法正常编译
位置:`api/QMT_API.py:1`
现状:
```python
# -*- coding: gbk -*-
```
文件实际内容包含 UTF-8 字节,只读编译时报错:
```text
SyntaxError: 'gbk' codec can't decode byte ...
```
影响:服务端可能在载入阶段直接退出,所有 API 不可用。
解决方案:
如果运行环境强制要求 GBK则必须把整个文件真实转换为 GBK不能只修改声明。
验收标准:
- 中文日志和错误响应无乱码。
### 2.2 客户端持仓对象被错误当成字典和二元组使用
位置:`py-client/strategy/trend/positions.py:32-40`
现状:`client.positions()` 返回 `list[Position]`,但代码同时使用:
```python
for idx, pos in positions:
code = pos["stock_code"]
avg_price = pos.get("avg_price", 0)
```
行情结果同样是 `Tick` dataclass却使用字典的 `.get()`
影响:进入持仓管理后必然抛出 `TypeError``AttributeError`,止盈和补仓完全无法执行。
解决方案:
全项目统一使用 SDK dataclass + __slots__,不再混用原始字典。
验收标准:
- 使用真实 `Position``Tick` 对象执行一轮不抛异常。
### 2.3 `handle_profit` 调用参数和函数签名不一致
位置:
- 调用:`py-client/strategy/trend/positions.py:58`
- 定义:`py-client/strategy/trend/positions.py:73`
影响:修复持仓遍历后,下一步仍会立即触发 `TypeError`
解决方案:
删除未使用的 `open_price``strategy_name` 或把它们纳入统一模型。
推荐接口:
```python
def handle_profit(
runtime: Runtime,
position: Position,
tick: Tick,
pnl_rate: float,
) -> ProfitDecision:
...
```
验收标准:
- 静态类型检查能够发现参数数量错误。
- ARMED、RAISED、STEADY、RETREAT 四种状态都有测试。
### 2.4 补仓流程存在多处必现错误
位置:`py-client/strategy/trend/positions.py:118-159`
问题包括:
- `StateItem` 被当成字典调用 `.get()`
- `orders.busy()` 多传入一个 `run` 参数。
- 某些分支只返回 `False`,调用方却解包两个值。
- 使用不存在的 `run.state.STATUS_ING`
- `state.added_num = +1` 每次都赋值为 1并非累加。
- `LOSS_TIERS[added_num]` 可能数组越界。
- 下单后没有扣减本轮剩余预算,多持仓可能超额补仓。
解决方案:
1. 所有 `StateItem` 字段改为属性访问。
2. `orders.busy(code, "BUY")` 使用正确签名。
3. 所有返回分支统一返回结构,推荐使用 dataclass + __slots__
```python
@dataclass(frozen=True)
class TradeDecision:
submitted: bool
message: str = ""
reserved_cash: float = 0.0
```
4. 使用模块常量 `STATUS_ING`,或把状态定义成 `Enum`
5. 补仓次数使用 `state.added_num += 1`
6.`added_num >= len(LOSS_TIERS)` 时明确禁止继续补仓。
7. `RunOnce` 创建本轮 `remaining_cash`,每次成功提交补仓后立即扣减。
验收标准:
- 第 0、1、2 次补仓边界均有测试。
- 超过最大补仓次数不会抛异常或继续下单。
### 2.5 止盈跟踪器每轮重建,无法形成跨轮回撤
位置:`py-client/strategy/trend/positions.py:31`
影响:每轮都会清空最高盈利网格,止盈状态无法从 ARMED/RAISED 演进至 RETREAT。
解决方案:
1. `GridTrailingTracker` 应作为 `Runtime` 字段,在策略启动时只创建一次。
2. 检查是否有定时清理的功能
验收标准:
- 连续输入 2.1%、3.1%、2.9% 能产生 ARMED、RAISED、RETREAT。
- 相同股票不同账户的峰值互不污染。
- 清仓后重新建仓不会继承旧峰值。
### 2.6 “取消过期订单”只查询可撤状态,没有执行撤单
位置:
- 客户端:`py-client/strategy/trend/order.py:74-88`
- 服务端:`api/QMT_API.py:962-969`
影响:过期订单一直保留,订单锁可能长期阻止新交易。
解决方案:
方案 A推荐新增按真实委托号撤单接口。
```text
POST /api/order/cancel_by_id
body: {order_id, account_type}
```
服务端先执行 `can_cancel_order()`,可撤时调用真正的 `cancel()`,并返回撤单请求结果。
验收标准:
- 暂时不做验证,后期验证
### 2.7 低现金资金闸同时跳过卖出管理
位置:`py-client/strategy/trend/boot.py:126-134`
影响:可用资金不足时直接结束整轮流程,持仓止盈和风险退出也被禁止。
解决方案:
把“是否允许新开仓/补仓”和“是否允许卖出”拆成不同条件。
验收标准:
- 可用现金低于阈值时不开仓、可补仓。
- 同一情况下满足止盈条件的持仓仍然能够提交卖单。
### 2.9 客户端订单标签未真正传给 QMT
位置:
- 客户端发送:`py-client/sdk/trade.py:10-18`
- 服务端丢弃:`api/QMT_API.py:663`
现状:客户端发送 `strategyName`,服务端调用 `passorder()` 时却硬编码为 `qmt`
解决方案:
1. 统一strategy_name 为信号的key,m_strRemark为本地业务订单号。
2. 同时修改QMT_API.py
验收标准:
- 下单后在 QMT 委托明细中可以看到客户端标签。
- 能从本地订单 ID 追踪到真实委托号和最终成交。
---
## 三、P1服务端响应速度整改
验收标准:
- 一轮策略账户查询由三次以上 QMT 调用下降为一次快照调用。
- 下单后下一次快照不会返回过期的订单状态。
### 3.3 大量使用 `dir()` 和 `getattr()` 反射序列化
位置:
- `api/QMT_API.py:920-930`
- `api/QMT_API.py:943-950`
- `api/QMT_API.py:979-1031`
- `api/QMT_API.py:1454-1473`
影响对每个对象遍历全部属性、捕获异常并转字符串CPU 开销大,返回字段也不稳定。
解决方案:
1. 为订单、成交、资产、持仓等类型定义固定字段映射。
2. 只返回客户端实际使用的字段。
3. 使用统一的轻量转换函数,不在每个 Handler 复制反射循环。
4. 对未知扩展类型单独保留调试接口,不进入高频生产路径。
验收标准:
- 高频订单查询不再调用 `dir()`
- 返回 JSON 字段固定并有接口契约测试。
- 相同数据量下序列化 CPU 时间明显下降。
### 3.5 回调同步写 JSON 文件
位置:`api/QMT_API.py:1475-1516`
影响:目录创建、反射序列化和格式化写盘可能阻塞 QMT 回调线程。
解决方案:
3. 生产环境关闭 `indent=4`
4. 使用临时文件替换,避免半写文件。
验收标准:
- 回调函数本身在毫秒级返回。
- 磁盘慢或不可写时不会阻塞交易回调。
- 写入失败可监控且不会静默丢失。
### 3.6 客户端 HTTP 没有连接池
位置:`py-client/sdk/client.py:28-47`
影响:每次 `urlopen()` 都可能新建连接,高频轮询产生额外 TCP 开销。
解决方案:
1. 改用支持连接池的 HTTP 客户端,如 `httpx.Client``requests.Session`
2. 整个策略生命周期复用一个 Client。
3. 设置连接、读取和总超时,不只设置单一 timeout。
4. 只对幂等查询配置有限重试;下单和撤单不能自动盲重试。
验收标准:
- 连续请求复用 TCP 连接。
- 查询超时能重试,下单超时进入“结果未知、需对账”状态而不是重复下单。
---
## 四、P1客户端其他逻辑与可靠性问题
### 4.1 新开仓订单锁可能在同一轮失效
位置:
- `py-client/strategy/trend/open.py:24`
- `py-client/strategy/trend/order.py:61-65`
- `py-client/strategy/trend/order.py:99-101`
现状:开仓检查 `busy()`,该方法只查看 `data`;新下单后只把键加入 `index`,没有加入 `data`
解决方案:
1. 统一锁判断,只保留一个权威接口。
2. 下单成功后立即插入本地 pending `OrderItem`
3. 信号进入处理前按证券代码去重。
4. 每轮刷新券商订单后用真实订单覆盖本地 pending 状态。
验收标准:同一轮两个来源返回同一证券信号时最多提交一笔买单。
### 4.2 `Runtime` 文档和字段不一致
位置:`py-client/strategy/trend/runtime.py`
现状:文档描述 `peak_grids`,实际 dataclass 没有该字段;持仓代码仍可能访问它。
解决方案:
2. 如果统一使用 `GridTrailingTracker`,删除 `peak_grids` 及所有引用。
3. 不应同时保留两套止盈峰值实现。
验收标准:项目中只有一种网格峰值状态来源。
### 4.3 `ping_api_host()` 参数无效且吞掉退出信号
位置:`py-client/main.py:59-76`
问题:
- `rpc_host` 参数没有使用。
- `connect_timeout` 参数没有使用。
- 使用裸 `except:`,会捕获 `KeyboardInterrupt``SystemExit`
- 无限重试没有最大日志节流或取消事件。
解决方案:
1. 函数重命名为 `wait_for_qmt_api()`,删除无用参数。
2. 仅捕获网络类异常和 `APIError`
3. 允许 `KeyboardInterrupt` 正常终止。
4. 使用 `threading.Event.wait()` 或可取消等待。
验收标准API 不可用时可以通过 Ctrl+C 立即退出。
### 4.5 状态文件缺少完整对账和生命周期
位置:`py-client/strategy/trend/state.py`
解决方案:
1. 启动时用真实持仓、订单和成交三方对账。
2. `ING` 状态必须根据真实订单结果转为 `OK``FAILED``CANCELED``UNKNOWN`
3. 已清仓证券应从状态中删除,并清除观察器和止盈峰值。
4. 状态文件不增加版本号,不增加新字段。
验收标准:程序在下单后崩溃并重启,能够从券商真实状态恢复,而不会重复下单。
### 4.6 日志调用格式错误且异常上下文不足
位置:`py-client/`
现状:不符合 logging 格式化规则。
解决方案:
1. 统一优化日志打印同时输出至文本文件(每天一个文件)。
验收标准:日志输出期间不出现 logging 自身的格式化异常。

File diff suppressed because it is too large Load Diff

View File

@@ -7,7 +7,7 @@ from pathlib import Path
import yaml import yaml
@dataclass @dataclass(slots=True)
class SignalConfig: class SignalConfig:
"""单个交易信号的数据源及开仓限制配置。""" """单个交易信号的数据源及开仓限制配置。"""
@@ -21,7 +21,7 @@ class SignalConfig:
gt_last_price_is_open: bool = False gt_last_price_is_open: bool = False
@dataclass @dataclass(slots=True)
class GlobalConfig: class GlobalConfig:
"""所有主机共享的系统配置。""" """所有主机共享的系统配置。"""
@@ -37,7 +37,7 @@ class GlobalConfig:
signals: dict[str, SignalConfig] = field(default_factory=dict) signals: dict[str, SignalConfig] = field(default_factory=dict)
@dataclass @dataclass(slots=True)
class AccountConfig: class AccountConfig:
"""当前主机所使用的账户及交易策略参数。""" """当前主机所使用的账户及交易策略参数。"""

View File

@@ -20,7 +20,7 @@ class GridState(str, Enum):
STEADY = "steady" # 仍处于当前峰值网格,继续持有 STEADY = "steady" # 仍处于当前峰值网格,继续持有
@dataclass(frozen=True) @dataclass(frozen=True, slots=True)
class GridObservation: class GridObservation:
"""一次网格观察的不可变结果。""" """一次网格观察的不可变结果。"""

View File

@@ -4,12 +4,12 @@ import secrets
from .http import get_json from .http import get_json
@dataclass @dataclass(slots=True)
class SignalItem: class SignalItem:
signal_key: str = ""; code: str = ""; name: str = ""; desc: str = ""; last_close: float = 0 signal_key: str = ""; code: str = ""; name: str = ""; desc: str = ""; last_close: float = 0
tech_indicator: dict[str, float] = field(default_factory=dict) tech_indicator: dict[str, float] = field(default_factory=dict)
@dataclass @dataclass(slots=True)
class SignalResult: class SignalResult:
code: str = ""; total: int = 0; updated: str = ""; data: dict[str, SignalItem] = field(default_factory=dict); message: str = "" code: str = ""; total: int = 0; updated: str = ""; data: dict[str, SignalItem] = field(default_factory=dict); message: str = ""

View File

@@ -2,23 +2,25 @@
# -*- coding: utf-8 -*- # -*- coding: utf-8 -*-
import logging as log import logging as log
from logging.handlers import TimedRotatingFileHandler
import os import os
import sys import sys
import time import time
import config import config
from dataclasses import dataclass from dataclasses import dataclass
import yaml import yaml
import httpx
PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__)) PROJECT_ROOT = os.path.dirname(os.path.abspath(__file__))
GLOBAL_CONFIG_PATH = os.path.join(PROJECT_ROOT, "etc", "_global.yaml") GLOBAL_CONFIG_PATH = os.path.join(PROJECT_ROOT, "etc", "_global.yaml")
if PROJECT_ROOT not in sys.path: if PROJECT_ROOT not in sys.path:
sys.path.insert(0, PROJECT_ROOT) sys.path.insert(0, PROJECT_ROOT)
from sdk import Client from sdk import APIError, Client
from strategy.trend.boot import StartTrend from strategy.trend.boot import StartTrend
@dataclass(frozen=True) @dataclass(frozen=True, slots=True)
class StrategyDefinition: class StrategyDefinition:
mutex_scope: str mutex_scope: str
start_strategy: object start_strategy: object
@@ -56,24 +58,47 @@ def check_single_instance(project_root: str) -> bool:
return False return False
def ping_api_host( def wait_for_qmt_api(retry_interval: float = 5.0) -> None:
rpc_host: str,
retry_interval: float = 5.0,
connect_timeout: float = 3.0,
) -> None:
"""循环检查 API 地址,连通后才返回。""" """循环检查 API 地址,连通后才返回。"""
client = Client(config.global_config.qmt_base_url, config.global_config.qmt_token, config.HTTP_TIMEOUT) client = Client(config.global_config.qmt_base_url, config.global_config.qmt_token, config.HTTP_TIMEOUT)
while True: retry_event = __import__("threading").Event()
while not retry_event.is_set():
try: try:
assets = client.assets() client.assets()
log.info(f"API 服务已连通:{config.global_config.qmt_base_url}") log.info(f"API 服务已连通:{config.global_config.qmt_base_url}")
client.close()
return return
except: except (APIError, httpx.RequestError) as exc:
log.warning( log.warning(
f"API 服务未就绪:{config.global_config.qmt_base_url}{retry_interval:g} 秒后重试" "API 服务未就绪:%s%g 秒后重试%s",
config.global_config.qmt_base_url,
retry_interval,
exc,
) )
time.sleep(retry_interval) retry_event.wait(retry_interval)
def configure_logging(data_dir: str) -> None:
"""同时输出控制台日志和按天轮转的文本日志。"""
log_dir = os.path.join(data_dir, "logs")
os.makedirs(log_dir, exist_ok=True)
root = log.getLogger()
root.setLevel(log.INFO)
formatter = log.Formatter("%(asctime)s [%(levelname)s] %(message)s")
if not root.handlers:
console = log.StreamHandler()
console.setFormatter(formatter)
root.addHandler(console)
file_handler = TimedRotatingFileHandler(
os.path.join(log_dir, "py-client.log"),
when="midnight",
interval=1,
backupCount=30,
encoding="utf-8",
)
file_handler.setFormatter(formatter)
root.addHandler(file_handler)
def wait_for_any_key() -> None: def wait_for_any_key() -> None:
print("按任意键退出...", flush=True) print("按任意键退出...", flush=True)
@@ -96,12 +121,12 @@ def main() -> int:
config.load() config.load()
if config.global_config is None or config.account_config is None: if config.global_config is None or config.account_config is None:
raise RuntimeError("配置尚未加载,请先调用 config.load()") raise RuntimeError("配置尚未加载,请先调用 config.load()")
configure_logging(config.global_config.qmt_data_dir)
ping_api_host(config.global_config.qmt_base_url) wait_for_qmt_api()
STRATEGIES[config.account_config.strategy].start_strategy() STRATEGIES[config.account_config.strategy].start_strategy()
return 0 return 0
except (OSError, yaml.YAMLError, ValueError) as exc: except (OSError, yaml.YAMLError, ValueError, RuntimeError, KeyError) as exc:
print(f"启动失败: {exc}", file=sys.stderr, flush=True) print(f"启动失败: {exc}", file=sys.stderr, flush=True)
wait_for_any_key() wait_for_any_key()
return 1 return 1

View File

@@ -0,0 +1,2 @@
httpx>=0.27,<1
PyYAML>=6.0

View File

@@ -1,10 +1,9 @@
from __future__ import annotations from __future__ import annotations
import json
from dataclasses import asdict, is_dataclass from dataclasses import asdict, is_dataclass
from typing import Any from typing import Any
from urllib.error import HTTPError
from urllib.request import Request, urlopen import httpx
from .errors import APIError, BusinessError from .errors import APIError, BusinessError
@@ -14,11 +13,28 @@ def csv_join(items: list[str]) -> str:
class Client: class Client:
"""复用连接池的同步 QMT HTTP 客户端。"""
def __init__(self, base_url: str, token: str, timeout: float = 15.0) -> None: def __init__(self, base_url: str, token: str, timeout: float = 15.0) -> None:
self.base_url = base_url.rstrip("/") self.base_url = base_url.rstrip("/")
self.token = token self.token = token
self.timeout = timeout if timeout > 0 else 15.0 self.timeout = timeout if timeout > 0 else 15.0
self.account_type = "stock" self.account_type = "stock"
self.http = httpx.Client(
base_url=self.base_url,
headers={"X-Token": token, "Accept": "application/json"},
timeout=httpx.Timeout(self.timeout),
limits=httpx.Limits(max_connections=20, max_keepalive_connections=10),
)
def close(self) -> None:
self.http.close()
def __enter__(self) -> "Client":
return self
def __exit__(self, *_args: object) -> None:
self.close()
def set_account_type(self, account_type: str) -> "Client": def set_account_type(self, account_type: str) -> "Client":
if account_type.strip(): if account_type.strip():
@@ -26,28 +42,38 @@ class Client:
return self return self
def _request(self, method: str, path: str, body: Any = None) -> Any: def _request(self, method: str, path: str, body: Any = None) -> Any:
data = None if is_dataclass(body):
headers = {"X-Token": self.token, "Accept": "application/json"} body = asdict(body)
if method != "GET": attempts = 2 if _is_idempotent(method, path) else 1
if body is None: body = {} response: httpx.Response | None = None
if is_dataclass(body): body = asdict(body) for attempt in range(attempts):
data = json.dumps(body, ensure_ascii=False).encode() try:
headers["Content-Type"] = "application/json" response = self.http.request(method, path, json=body)
request = Request(self.base_url + path, data=data, headers=headers, method=method) break
except (httpx.ConnectError, httpx.ReadTimeout):
if attempt + 1 == attempts:
raise
assert response is not None
if response.status_code >= 400:
try:
message = response.json().get("error", response.text)
except (ValueError, AttributeError):
message = response.text.strip()
raise APIError(response.status_code, str(message))
if not response.content:
return None
try: try:
with urlopen(request, timeout=self.timeout) as response: return response.json()
raw = response.read() except ValueError as exc:
except HTTPError as exc: raise ValueError(
raw = exc.read() f"invalid JSON from {path}: {response.content[:512]!r}"
try: message = json.loads(raw).get("error", raw.decode(errors="replace")) ) from exc
except (ValueError, AttributeError): message = raw.decode(errors="replace").strip()
raise APIError(exc.code, str(message)) from exc
if not raw: return None
try: return json.loads(raw)
except ValueError as exc: raise ValueError(f"invalid JSON from {path}: {raw[:512]!r}") from exc
def _get(self, path: str) -> Any: return self._request("GET", path) def _get(self, path: str) -> Any:
def _post(self, path: str, body: Any = None) -> Any: return self._request("POST", path, body) return self._request("GET", path)
def _post(self, path: str, body: Any = None) -> Any:
return self._request("POST", path, {} if body is None else body)
def _get_field(self, path: str, key: str) -> Any: def _get_field(self, path: str, key: str) -> Any:
return self._get(path).get(key) return self._get(path).get(key)
@@ -57,3 +83,20 @@ class Client:
if isinstance(result, dict) and result.get("error"): if isinstance(result, dict) and result.get("error"):
raise BusinessError(result["error"]) raise BusinessError(result["error"])
return result.get(key, result) if key and isinstance(result, dict) else result return result.get(key, result) if key and isinstance(result, dict) else result
def _is_idempotent(method: str, path: str) -> bool:
if method == "GET":
return True
prefixes = (
"/api/v2/",
"/api/holding",
"/api/money/",
"/api/context/",
"/api/check/",
"/api/data/",
"/api/trade/trade_detail_data",
"/api/order/deal",
)
unsafe = ("subscribe", "unsubscribe")
return path.startswith(prefixes) and not any(word in path for word in unsafe)

View File

@@ -11,7 +11,7 @@ def _number(value: Any, kind: type = float) -> Any:
return kind() return kind()
@dataclass @dataclass(slots=True)
class Position: class Position:
stock_code: str = "" stock_code: str = ""
stock_name: str = "" stock_name: str = ""
@@ -44,20 +44,20 @@ class Position:
) )
@dataclass @dataclass(slots=True)
class Assets: class Assets:
total: float = 0.0 total: float = 0.0
available: float = 0.0 available: float = 0.0
@dataclass @dataclass(slots=True)
class Tick: class Tick:
last_price: float = 0.0 last_price: float = 0.0
last_close: float = 0.0 last_close: float = 0.0
raw: dict[str, Any] = field(default_factory=dict) raw: dict[str, Any] = field(default_factory=dict)
@dataclass @dataclass(slots=True)
class HistoryDataRequest: class HistoryDataRequest:
length: int = 10 length: int = 10
period: str = "" period: str = ""
@@ -66,7 +66,7 @@ class HistoryDataRequest:
skip_paused: bool = True skip_paused: bool = True
@dataclass @dataclass(slots=True)
class MarketDataRequest: class MarketDataRequest:
fields: list[str] = field(default_factory=list) fields: list[str] = field(default_factory=list)
stocks: list[str] = field(default_factory=list) stocks: list[str] = field(default_factory=list)
@@ -77,7 +77,7 @@ class MarketDataRequest:
count: int = 0 count: int = 0
@dataclass @dataclass(slots=True)
class FinancialDataRequest: class FinancialDataRequest:
tabname: str = ""; colname: str = ""; market: str = ""; code: str = "" tabname: str = ""; colname: str = ""; market: str = ""; code: str = ""
report_type: str = ""; barpos: int = 0 report_type: str = ""; barpos: int = 0
@@ -85,22 +85,22 @@ class FinancialDataRequest:
start_date: str = ""; end_date: str = "" start_date: str = ""; end_date: str = ""
@dataclass @dataclass(slots=True)
class FactorDataRequest: class FactorDataRequest:
field_list: list[str] = field(default_factory=list); stock_list: list[str] = field(default_factory=list) field_list: list[str] = field(default_factory=list); stock_list: list[str] = field(default_factory=list)
stock_code: str = ""; start_date: str = ""; end_date: str = "" stock_code: str = ""; start_date: str = ""; end_date: str = ""
@dataclass @dataclass(slots=True)
class BSMPriceRequest: class BSMPriceRequest:
option_type: str; object_prices: Any; strike_price: float; risk_free: float; sigma: float; days: int; dividend: float option_type: str; object_prices: Any; strike_price: float; risk_free: float; sigma: float; days: int; dividend: float
@dataclass @dataclass(slots=True)
class BSMIVRequest: class BSMIVRequest:
option_type: str; object_prices: float; strike_price: float; option_price: float; risk_free: float; days: int; dividend: float option_type: str; object_prices: float; strike_price: float; option_price: float; risk_free: float; days: int; dividend: float
@dataclass @dataclass(slots=True)
class LocalDataRequest: class LocalDataRequest:
stock_code: str; start_time: str = ""; end_time: str = ""; period: str = ""; divid_type: str = ""; count: int = 0 stock_code: str; start_time: str = ""; end_time: str = ""; period: str = ""; divid_type: str = ""; count: int = 0

View File

@@ -13,9 +13,20 @@ class TradeMixin:
if value: body[key] = value if value: body[key] = value
return self._post("/api/trade/passorder", body) return self._post("/api/trade/passorder", body)
def passorder_latest(self, side, stock, volume): return self.passorder_latest_tagged(side, stock, volume, "") def passorder_latest(self, side, stock, volume): return self.passorder_latest_tagged(side, stock, volume, "", "")
def passorder_latest_tagged(self, side, stock, volume, order_id): def passorder_latest_tagged(self, side, stock, volume, strategy_name, order_id):
return self.passorder(side, stock, volume, ORDER_TYPE_VOLUME, PR_TYPE_LATEST, -1, QUICK_TRADE_NOW, order_id) body = {
"opType": side,
"orderType": ORDER_TYPE_VOLUME,
"stock": stock,
"prType": PR_TYPE_LATEST,
"price": -1,
"volume": volume,
"quickTrade": QUICK_TRADE_NOW,
"strategyName": strategy_name,
"orderId": order_id,
}
return self._post("/api/trade/passorder", body)
def algo_passorder(self, **kwargs): return self._post("/api/trade/algo_passorder", kwargs) def algo_passorder(self, **kwargs): return self._post("/api/trade/algo_passorder", kwargs)
def smart_algo_passorder(self, **kwargs): return self._post("/api/trade/smart_algo_passorder", kwargs) def smart_algo_passorder(self, **kwargs): return self._post("/api/trade/smart_algo_passorder", kwargs)
@@ -46,6 +57,7 @@ class TradeMixin:
def value_by_order_id(self, order_id, datatype): return self._post("/api/trade/value_by_order_id", {"orderId": order_id, "accountType": self.account_type, "datatype": datatype}).get("data") def value_by_order_id(self, order_id, datatype): return self._post("/api/trade/value_by_order_id", {"orderId": order_id, "accountType": self.account_type, "datatype": datatype}).get("data")
def last_order_id(self, datatype): return self._post("/api/trade/last_order_id", {"account": self.account_type, "datatype": datatype}).get("last_order_id") def last_order_id(self, datatype): return self._post("/api/trade/last_order_id", {"account": self.account_type, "datatype": datatype}).get("last_order_id")
def can_cancel_order(self, order_id): return self._post("/api/trade/can_cancel_order", {"orderId": order_id, "accountType": self.account_type}).get("can_cancel") def can_cancel_order(self, order_id): return self._post("/api/trade/can_cancel_order", {"orderId": order_id, "accountType": self.account_type}).get("can_cancel")
def cancel_by_id(self, order_id): return self._post("/api/order/cancel_by_id", {"order_id": order_id, "account_type": self.account_type})
def debt_contract(self): return self._contract("debt_contract") def debt_contract(self): return self._contract("debt_contract")
def assure_contract(self): return self._contract("assure_contract") def assure_contract(self): return self._contract("assure_contract")
def enable_short_contract(self): return self._contract("enable_short_contract") def enable_short_contract(self): return self._contract("enable_short_contract")

View File

@@ -12,6 +12,7 @@ from datetime import datetime
import config import config
from libs import init_signals, market_allow_open, trading_time from libs import init_signals, market_allow_open, trading_time
from sdk import Client from sdk import Client
from libs.grid_take_profit import GridTrailingTracker
from .state import State from .state import State
from .order import OrderBook from .order import OrderBook
from .watch import DipWatch from .watch import DipWatch
@@ -78,7 +79,9 @@ def StartTrend() -> None:
config.account_config.strategy, config.account_config.strategy,
config.account_config.account_id, config.account_config.account_id,
) )
storeState.sync_positions(positions) orders = client.trade_detail_data("order")
deals = client.deals()
storeState.reconcile(positions, orders, deals)
# 获取本策略的信号开仓数据 # 获取本策略的信号开仓数据
signals = init_signals(config.global_config,["morning","tail","arbitrage"]) signals = init_signals(config.global_config,["morning","tail","arbitrage"])
@@ -90,6 +93,7 @@ def StartTrend() -> None:
orders=OrderBook(), orders=OrderBook(),
open_watch=DipWatch(), open_watch=DipWatch(),
add_watch=DipWatch(), add_watch=DipWatch(),
profit_tracker=GridTrailingTracker(config.account_config.grid_step_pct),
) )
logging.info( logging.info(
@@ -129,9 +133,9 @@ def RunOnce(run: Runtime, signals) -> None:
except Exception: except Exception:
logging.exception("获取资产失败") logging.exception("获取资产失败")
return return
if assets.available < assets.total * run.account_cfg.min_cash_ratio: allow_open_by_cash = assets.available >= assets.total * run.account_cfg.min_cash_ratio
if not allow_open_by_cash:
logging.info("资金总闸:可用金额太少,禁止开新仓") logging.info("资金总闸:可用金额太少,禁止开新仓")
return
# 3. 获取大盘状态,只有大盘信号允许时才执行开仓。 # 3. 获取大盘状态,只有大盘信号允许时才执行开仓。
market_ok = market_allow_open(run.global_cfg.api_host) market_ok = market_allow_open(run.global_cfg.api_host)
@@ -143,11 +147,23 @@ def RunOnce(run: Runtime, signals) -> None:
logging.exception("获取持仓失败") logging.exception("获取持仓失败")
return return
active_codes = set(position_codes)
removed_codes = set(run.state.codes) - active_codes
for code in removed_codes:
run.state.delete(code)
run.open_watch.forget(code)
run.add_watch.forget(code)
if removed_codes:
run.state.save()
# 5. 验证有效开仓信号:排除已有持仓,并按 signal_allow 过滤。 # 5. 验证有效开仓信号:排除已有持仓,并按 signal_allow 过滤。
position_code_set = set(position_codes) position_code_set = set(position_codes)
allow_open = [ allow_open = []
signal for signal in signals if signal.code not in position_code_set seen_codes = set(position_code_set)
] for signal in signals:
if signal.code not in seen_codes:
allow_open.append(signal)
seen_codes.add(signal.code)
# 6. 获取持仓和待开仓证券的实时行情 tick。 # 6. 获取持仓和待开仓证券的实时行情 tick。
all_codes = list(position_codes) all_codes = list(position_codes)
@@ -161,7 +177,7 @@ def RunOnce(run: Runtime, signals) -> None:
return return
# 7. 执行开仓:必须同时存在有效信号且大盘允许开仓。 # 7. 执行开仓:必须同时存在有效信号且大盘允许开仓。
if allow_open and market_ok: if allow_open and market_ok and allow_open_by_cash:
open_signal(run, ticks, allow_open) open_signal(run, ticks, allow_open)
# 8. 持仓计算。当前 Go 版本的 managePositions 为空,保留扩展入口。 # 8. 持仓计算。当前 Go 版本的 managePositions 为空,保留扩展入口。

View File

@@ -41,7 +41,14 @@ def open_signal(run, ticks, open_signals) -> None:
# 6. 生成本地订单号并按最新价提交开仓委托。 # 6. 生成本地订单号并按最新价提交开仓委托。
order_id = run.orders.new_order_id("base") order_id = run.orders.new_order_id("base")
request = PlaceOrderRequest(run.client, OP_BUY, item.code, volume, order_id) request = PlaceOrderRequest(
run.client,
OP_BUY,
item.code,
volume,
order_id,
item.signal_key,
)
if not run.orders.place(request): if not run.orders.place(request):
continue continue

View File

@@ -15,7 +15,7 @@ OFFSET_FLAG = {"23": "BUY", "24": "SELL", "48": "BUY", "49": "SELL"}
BUSY_STATUSES = {"48", "49", "50", "51", "52", "55"} BUSY_STATUSES = {"48", "49", "50", "51", "52", "55"}
@dataclass(frozen=True) @dataclass(frozen=True, slots=True)
class PlaceOrderRequest: class PlaceOrderRequest:
"""``OrderBook.place`` 提交委托所需的全部参数。""" """``OrderBook.place`` 提交委托所需的全部参数。"""
@@ -24,9 +24,10 @@ class PlaceOrderRequest:
code: str code: str
volume: int volume: int
order_id: str order_id: str
strategy_name: str
@dataclass @dataclass(slots=True)
class OrderItem: class OrderItem:
"""从 QMT 委托明细转换得到的本地订单记录。""" """从 QMT 委托明细转换得到的本地订单记录。"""
@@ -37,6 +38,7 @@ class OrderItem:
status: str status: str
created_at: datetime | None created_at: datetime | None
volume: int volume: int
local_order_id: str = ""
class OrderBook: class OrderBook:
@@ -50,8 +52,8 @@ class OrderBook:
@staticmethod @staticmethod
def new_order_id(leg: str) -> str: def new_order_id(leg: str) -> str:
"""生成不超过 24 个字符的策略订单号""" """生成短订单号,为 QMT 备注中的信号键预留空间"""
return f"zt-{leg}-{secrets.token_hex(6)}"[:24] return f"zt-{leg[:1]}-{secrets.token_hex(4)}"
def is_lock(self, side: str, code: str) -> bool: def is_lock(self, side: str, code: str) -> bool:
"""判断证券在指定买卖方向上是否已经被委托锁定。""" """判断证券在指定买卖方向上是否已经被委托锁定。"""
@@ -61,8 +63,9 @@ class OrderBook:
def busy(self, code: str, side: str) -> bool: def busy(self, code: str, side: str) -> bool:
"""判断证券是否存在仍在处理中的同方向委托。""" """判断证券是否存在仍在处理中的同方向委托。"""
with self.lock: with self.lock:
order = self.data.get(f"{side}-{code}") key = f"{side}-{code}"
return bool(order and order.status in BUSY_STATUSES) order = self.data.get(key)
return key in self.index or bool(order and order.status in BUSY_STATUSES)
def refresh(self, client: Any) -> None: def refresh(self, client: Any) -> None:
"""从 QMT 刷新当前委托明细和方向索引。""" """从 QMT 刷新当前委托明细和方向索引。"""
@@ -71,7 +74,9 @@ class OrderBook:
] ]
with self.lock: with self.lock:
self.data = {key: item for key, item in parsed_orders} self.data = {key: item for key, item in parsed_orders}
self.index = [key for key, _ in parsed_orders] self.index = [
key for key, item in parsed_orders if item.status in BUSY_STATUSES
]
def cancel_expired(self, client: Any, now: datetime | None = None) -> None: def cancel_expired(self, client: Any, now: datetime | None = None) -> None:
"""尝试撤销超过有效期且具有委托编号的订单。""" """尝试撤销超过有效期且具有委托编号的订单。"""
@@ -85,20 +90,39 @@ class OrderBook:
and current - order.created_at > self.timeout and current - order.created_at > self.timeout
and order.id and order.id
): ):
client.can_cancel_order(order.id) client.cancel_by_id(order.id)
def place(self, request: PlaceOrderRequest) -> bool: def place(self, request: PlaceOrderRequest) -> bool:
"""按最新价提交委托,并立即写入本地方向锁。""" """按最新价提交委托,并立即写入本地方向锁。"""
request.client.passorder_latest_tagged( result = request.client.passorder_latest_tagged(
request.op, request.op,
request.code, request.code,
request.volume, request.volume,
request.strategy_name,
request.order_id, request.order_id,
) )
if not isinstance(result, dict):
return False
order_ref = str(result.get("order_ref") or "").strip().lower()
if result.get("status") != "success" or order_ref in {"", "unknown", "none"}:
return False
side = OFFSET_FLAG.get(str(request.op), "") side = OFFSET_FLAG.get(str(request.op), "")
pending = OrderItem(
id=order_ref,
code=request.code,
side=side,
remark=request.order_id,
status="48",
created_at=datetime.now(),
volume=request.volume,
local_order_id=request.order_id,
)
with self.lock: with self.lock:
self.index.append(f"{side}-{request.code}") key = f"{side}-{request.code}"
self.data[key] = pending
if key not in self.index:
self.index.append(key)
return True return True
@@ -126,6 +150,7 @@ def parse_order(row: dict[str, Any]) -> tuple[str, OrderItem]:
status=str(row.get("m_nOrderStatus") or ""), status=str(row.get("m_nOrderStatus") or ""),
created_at=created_at, created_at=created_at,
volume=volume, volume=volume,
local_order_id=_local_order_id(str(row.get("m_strRemark") or "")),
) )
return f"{item.side}-{item.code}", item return f"{item.side}-{item.code}", item
@@ -146,3 +171,8 @@ def _parse_insert_datetime(row: dict[str, Any]) -> datetime | None:
return datetime.strptime(date + clock, "%Y%m%d%H%M%S") return datetime.strptime(date + clock, "%Y%m%d%H%M%S")
except ValueError: except ValueError:
return None return None
def _local_order_id(remark: str) -> str:
"""兼容 ``local_order_id|signal_key`` 形式的 QMT 备注。"""
return remark.split("|", 1)[0] if remark else ""

View File

@@ -1,167 +1,187 @@
"""趋势策略持仓管理逻辑,对应 Go 版本的 ``logic/positions.go``""" """趋势策略持仓止盈与分级补仓"""
from __future__ import annotations from __future__ import annotations
import logging import logging
from math import floor from dataclasses import dataclass
from libs.calc import calc_buy_volume, calculate_min_profit_rate
from libs.grid_take_profit import GridState
from sdk import OP_BUY, OP_SELL, Position, Tick
from libs.calc import calc_buy_volume,calculate_min_profit_rate
from libs.grid_take_profit import GridState, GridTrailingTracker
from sdk import OP_BUY, OP_SELL
import config
from .order import PlaceOrderRequest from .order import PlaceOrderRequest
from .state import STATUS_ING, STATUS_NONE, STATUS_OK
from .runtime import Runtime from .runtime import Runtime
from .state import STATUS_ING
LEG_BASE = "base" LEG_BASE = "base"
LEG_ADDED = "add" LEG_ADDED = "add"
LOSS_TIERS = (-30.0, -50.0)
# 止盈网格跟踪器延迟初始化,避免导入模块时账户配置尚未加载。
profit_tracker = None
# 分级补仓档位(百分比) @dataclass(frozen=True, slots=True)
LOSS_TIERS = [-30, -50] class TradeDecision:
# 补仓反弹确认阈值(百分比) """一次止盈或补仓判断的统一结果。"""
LOSS_REBOUND_THRESHOLD = 0.5
def manage_positions(run:Runtime, ticks, positions, market_ok: bool,available:float) -> None: submitted: bool
"""执行持仓计算。""" message: str = ""
logging.info(f"持仓:{len(positions)} 支股票,开始处理") reserved_cash: float = 0.0
global profit_tracker
profit_tracker = GridTrailingTracker(step=run.account_cfg.grid_step_pct)
for idx,pos in positions:
code = pos['stock_code']
avg_price = pos.get('avg_price', 0)
volume = pos.get('volume', 0)
can_use_volume = pos.get('can_use_volume', 0)
current_price = ticks.get(code, {}).get('lastPrice', 0)
strategy_name = pos.get('strategy_name', '')
market_value = pos.get('market_value',0)
profit = pos.get('profit_rate', 0)
# 排除指定股票
if code in config.account_config.excluded_codes: def manage_positions(
runtime: Runtime,
ticks: dict[str, Tick],
positions: list[Position],
market_ok: bool,
available: float,
) -> None:
"""处理所有真实持仓,并在本轮内统一控制补仓预算。"""
active_keys = {
_position_key(runtime, position.stock_code)
for position in positions
if position.volume > 0 and position.stock_code
}
runtime.profit_tracker.retain(active_keys)
remaining_cash = max(0.0, available)
logging.info("[持仓] 共 %d 只,开始处理", len(positions))
for position in positions:
code = position.stock_code
tick = ticks.get(code)
if code in runtime.account_cfg.excluded_codes:
continue
if (
not code
or position.open_price <= 0
or position.volume <= 0
or tick is None
or tick.last_price <= 0
):
continue continue
# 过滤无效仓位 pnl_rate = round(
if avg_price == 0 or can_use_volume == 0 or current_price == 0 or volume == 0: (tick.last_price - position.open_price) / position.open_price * 100,
continue 2,
)
minimum_profit = calculate_min_profit_rate(position.open_price, 1)
profit_decision = handle_profit(
runtime=runtime,
position=position,
tick=tick,
pnl_rate=pnl_rate,
minimum_profit=minimum_profit,
)
if profit_decision.message:
logging.info("[止盈] %s %s", code, profit_decision.message)
# 计算盈亏率(百分比) if runtime.account_cfg.enable_loss_add_position and market_ok:
pnl_ratio = (current_price - avg_price) / avg_price * 100 if avg_price != 0 else 0 loss_decision = handle_loss(
pnl_ratio = round(pnl_ratio, 2) runtime=runtime,
position=position,
tick=tick,
pnl_rate=pnl_rate,
available=remaining_cash,
)
remaining_cash -= loss_decision.reserved_cash
if loss_decision.message:
logging.info("[补仓] %s %s", code, loss_decision.message)
# 计算最小利润率1倍
min_profit_rate_val = calculate_min_profit_rate(avg_price, 1)
# 盈利处理 def handle_profit(
is_closed, message = handle_profit(run,code,avg_price, pnl_ratio, min_profit_rate_val, can_use_volume, strategy_name) runtime: Runtime,
if is_closed: position: Position,
logging.info("profit", code, f"止盈执行 | {message}") tick: Tick,
if message != "": pnl_rate: float,
logging.info("profit", code, message) minimum_profit: float,
) -> TradeDecision:
# 补仓处理 """基于跨轮保存的最高盈利网格判断是否提交止盈。"""
if config.account_config.enable_loss_add_position and market_ok: if pnl_rate < minimum_profit:
is_replenished, message = handle_loss(run,code,current_price,pnl_ratio,market_value,market_ok,available) return TradeDecision(False)
if is_replenished:
logging.info("loss", code, f"补仓执行 | {message}")
if message != "":
logging.info("loss", code, message)
# 盈利处理
def handle_profit(run:Runtime, code: str, pnl_rate: float,
min_profit_rate: float, vol: int) -> tuple[bool, str]:
"""
盈利处理 - 基于网格的止盈策略
Args:
code: 股票代码
open_price: 开仓价格
pnl_rate: 当前盈亏率(百分比)
min_profit_rate: 最小利润率阈值
vol: 可用股数
strategy_name: str
Returns:
tuple[bool, str]: (是否执行平仓, 操作说明)
"""
# 预检查:未达到最小利润率
if pnl_rate < min_profit_rate:
return False, ""
position_key = f"{run.account_cfg.account_id}:{code}"
observation = profit_tracker.observe(position_key, pnl_rate)
key = _position_key(runtime, position.stock_code)
observation = runtime.profit_tracker.observe(key, pnl_rate)
if observation.state == GridState.ARMED: if observation.state == GridState.ARMED:
msg = f"首次达到{pnl_rate}%,设置峰值网格{observation.current_grid}" return TradeDecision(
return False, msg False,
f"首次达到 {pnl_rate:.2f}%,峰值网格={observation.current_grid}",
)
if observation.state == GridState.RAISED: if observation.state == GridState.RAISED:
return False, f"上涨至{pnl_rate}%,更新峰值网格{observation.current_grid}" return TradeDecision(
False,
# 执行平仓 f"上涨至 {pnl_rate:.2f}%,峰值网格={observation.current_grid}",
if observation.state == GridState.RETREAT: )
order_id = run.orders.new_order_id(LEG_BASE) if observation.state in {GridState.STEADY}:
request = PlaceOrderRequest(run.client, OP_SELL, code, vol, order_id) return TradeDecision(False)
result = run.orders.place(request) if runtime.orders.busy(position.stock_code, "SELL"):
if result : return TradeDecision(False, "卖出委托处理中")
success_msg = f"✓ 委托成功 | {vol}股 订单号:{result} 等待成交"
logging.info("profit", code, success_msg) volume = position.can_use_volume - position.can_use_volume % 100
return True, success_msg if volume <= 0:
else: return TradeDecision(False, "无可用整手持仓")
fail_msg = f"止盈委托失败: {code}" order_id = runtime.orders.new_order_id(LEG_BASE)
logging.error("profit", code, "✗ 止盈委托失败") request = PlaceOrderRequest(
return False, fail_msg client=runtime.client,
op=OP_SELL,
code=position.stock_code,
volume=volume,
order_id=order_id,
strategy_name=runtime.account_cfg.strategy,
)
if not runtime.orders.place(request):
return TradeDecision(False, "止盈委托失败")
return TradeDecision(True, f"卖出 {volume} 股,订单={order_id}")
def handle_loss(run:Runtime, code: str, current_price,pnl_rate,market_value: float,market_ok: bool, available: float) -> tuple[bool, str]: def handle_loss(
"""满足条件时提交补仓委托,并返回扣减后的剩余预算。""" runtime: Runtime,
state = run.state.get(code) position: Position,
added_num = state.get('added_num',0) tick: Tick,
# 预检查:未达到最低补仓阈值 pnl_rate: float,
if pnl_rate > LOSS_TIERS[added_num]: available: float,
return False, "" ) -> TradeDecision:
"""按亏损档位、反弹确认和本轮剩余预算提交补仓。"""
try:
state = runtime.state.get(position.stock_code)
except KeyError:
return TradeDecision(False, "缺少持仓状态,跳过补仓")
# 强制条件 if state.added_num >= len(LOSS_TIERS):
if current_price>200 or market_value>=60000: return TradeDecision(False, "已达到最大补仓次数")
return False, f"成本价{current_price}>200,仓位价值{market_value}>=60000, 不补仓" if pnl_rate > LOSS_TIERS[state.added_num]:
return TradeDecision(False)
# 1. 大盘必须允许开仓,且价格已从观察低点达到反弹阈值。 if tick.last_price > 200 or position.market_value >= 60_000:
if not market_ok or not run.add_watch.triggered("补仓", code, current_price): return TradeDecision(False, "价格或仓位市值超过补仓限制")
return False if not runtime.add_watch.triggered("补仓", position.stock_code, tick.last_price):
return TradeDecision(False, "等待价格反弹确认")
if runtime.orders.busy(position.stock_code, "BUY"):
return TradeDecision(False, "买入委托处理中")
# 2. 计算补仓数量和预计占用金额。 volume = calc_buy_volume(tick.last_price, runtime.account_cfg.buy_value)
volume = calc_buy_volume(current_price, run.account_cfg.buy_value) amount = tick.last_price * volume
amount = current_price * volume if volume <= 0 or amount > available:
return TradeDecision(False, "本轮可用资金不足")
# 3. 检查预算。 order_id = runtime.orders.new_order_id(LEG_ADDED)
if amount > available: request = PlaceOrderRequest(
return False, f"f{code} f{amount} 仓位资金不够补仓" client=runtime.client,
op=OP_BUY,
code=position.stock_code,
volume=volume,
order_id=order_id,
strategy_name=runtime.account_cfg.strategy,
)
if not runtime.orders.place(request):
return TradeDecision(False, "补仓委托失败")
# 是否已有未完成的买入委托 state.added_num += 1
if run.orders.busy(run, code, "BUY"): state.added_status = STATUS_ING
return False, f"{code}订单锁定中" state.added_order_id = order_id
state.added_qty = volume
# 4. 生成补仓订单号并提交买入委托。 state.added_cost = tick.last_price
order_id = run.orders.new_order_id(LEG_ADDED) runtime.state.set(state)
request = PlaceOrderRequest(run.client, OP_BUY, code, volume, order_id) runtime.state.save()
result = run.orders.place(request) runtime.add_watch.forget(position.stock_code)
if result : return TradeDecision(True, f"买入 {volume} 股,订单={order_id}", amount)
state.added_num = +1
state.added_status = run.state.STATUS_ING
state.added_order_id = order_id
run.state.set(state)
run.state.save()
run.add_watch.forget(code)
return True,f"补仓委托成功: {code} {volume}手, 等待成交确认"
else:
return False,f"补仓失败: {code}"
def forget(run, code: str) -> None: def _position_key(runtime: Runtime, code: str) -> str:
"""持仓退出后清理开仓、补仓观察记录和止盈峰值。""" return f"{runtime.account_cfg.account_id}:{code}"
run.peak_grids.pop(f"{code}|{LEG_ADDED}", None)

View File

@@ -6,6 +6,7 @@ from dataclasses import dataclass, field
from config import AccountConfig, GlobalConfig from config import AccountConfig, GlobalConfig
from sdk import Client from sdk import Client
from libs.grid_take_profit import GridTrailingTracker
from .order import OrderBook from .order import OrderBook
from .state import State from .state import State
@@ -27,7 +28,7 @@ class Runtime:
orders: 当前活动委托和证券方向锁。 orders: 当前活动委托和证券方向锁。
open_watch: 新开仓使用的价格反弹观察器。 open_watch: 新开仓使用的价格反弹观察器。
add_watch: 亏损补仓使用的价格反弹观察器。 add_watch: 亏损补仓使用的价格反弹观察器。
peak_grids: ``证券代码|仓位类型`` 到最高盈利网格的映射 profit_tracker: 跨轮保存的账户持仓最高盈利网格跟踪器
""" """
# 外部服务与账户配置。 # 外部服务与账户配置。
@@ -40,4 +41,4 @@ class Runtime:
orders: OrderBook orders: OrderBook
open_watch: DipWatch open_watch: DipWatch
add_watch: DipWatch add_watch: DipWatch
profit_tracker: GridTrailingTracker

View File

@@ -15,6 +15,9 @@ from sdk import Position
STATUS_NONE = "" STATUS_NONE = ""
STATUS_ING = "ING" STATUS_ING = "ING"
STATUS_OK = "OK" STATUS_OK = "OK"
STATUS_FAILED = "FAILED"
STATUS_CANCELED = "CANCELED"
STATUS_UNKNOWN = "UNKNOWN"
@dataclass(slots=True) @dataclass(slots=True)
@@ -110,6 +113,33 @@ class State:
self.save() self.save()
def reconcile(
self,
positions: Iterable[Position],
orders: list[dict[str, str]],
deals: list[dict[str, str]],
) -> None:
"""用真实持仓、委托和成交恢复本地状态,不增加持久化字段。"""
position_list = list(positions)
self.sync_positions(position_list)
active_codes = {
item.stock_code for item in position_list if item.volume > 0
}
for code in list(self.codes):
if code not in active_codes:
self.delete(code)
for code in list(self.codes):
item = self.get(code)
item.base_status = _reconcile_leg(
item.base_order_id, item.base_status, orders, deals
)
item.added_status = _reconcile_leg(
item.added_order_id, item.added_status, orders, deals
)
self.set(item)
self.save()
def save(self) -> None: def save(self) -> None:
"""将内存状态格式化写入 JSON并原子替换正式文件。""" """将内存状态格式化写入 JSON并原子替换正式文件。"""
with self.lock: with self.lock:
@@ -144,3 +174,40 @@ class State:
} }
except (TypeError, ValueError) as exc: except (TypeError, ValueError) as exc:
raise ValueError(f"[状态] 状态字段无效: {exc}") from exc raise ValueError(f"[状态] 状态字段无效: {exc}") from exc
def _reconcile_leg(
local_order_id: str,
current_status: str,
orders: list[dict[str, str]],
deals: list[dict[str, str]],
) -> str:
if current_status != STATUS_ING or not local_order_id:
return current_status
if any(local_order_id in row.get("m_strRemark", "") for row in deals):
return STATUS_OK
order = next(
(
row for row in orders
if local_order_id in row.get("m_strRemark", "")
),
None,
)
if order is None:
return STATUS_UNKNOWN
traded = _as_int(order.get("m_nVolumeTraded"))
status = str(order.get("m_nOrderStatus", ""))
if traded > 0 and status not in {"48", "49", "50", "51", "52", "55"}:
return STATUS_OK
if status in {"54", "56"}:
return STATUS_CANCELED
if status in {"57", "58"}:
return STATUS_FAILED
return STATUS_ING
def _as_int(value: object) -> int:
try:
return int(value or 0)
except (TypeError, ValueError):
return 0

View File

@@ -4,7 +4,7 @@ from threading import Lock
import logging import logging
@dataclass @dataclass(slots=True)
class _Entry: class _Entry:
last_close: float last_close: float
expires_at: datetime expires_at: datetime

View File

@@ -0,0 +1,137 @@
from __future__ import annotations
import unittest
from tempfile import TemporaryDirectory
from types import SimpleNamespace
from unittest.mock import patch
from libs.grid_take_profit import GridState, GridTrailingTracker
from sdk import Assets, Position, Tick
from strategy.trend.order import OrderBook, PlaceOrderRequest
from strategy.trend.positions import LOSS_TIERS, handle_loss, manage_positions
from strategy.trend.boot import RunOnce
from strategy.trend.state import STATUS_OK, State, StateItem
class FakeClient:
def __init__(self):
self.orders = []
def passorder_latest_tagged(self, op, code, volume, strategy_name, order_id):
self.orders.append((op, code, volume, strategy_name, order_id))
return {"status": "success", "order_ref": f"broker-{len(self.orders)}"}
class TrendTests(unittest.TestCase):
def test_grid_states_and_account_isolation(self):
tracker = GridTrailingTracker(1)
self.assertEqual(tracker.observe("A:code", 2.1).state, GridState.ARMED)
self.assertEqual(tracker.observe("A:code", 3.1).state, GridState.RAISED)
self.assertEqual(tracker.observe("A:code", 2.9).state, GridState.RETREAT)
self.assertEqual(tracker.observe("B:code", 2.9).state, GridState.ARMED)
tracker.retain([])
self.assertEqual(tracker.observe("A:code", 2.9).state, GridState.ARMED)
def test_order_book_locks_duplicate_order(self):
client = FakeClient()
book = OrderBook()
request = PlaceOrderRequest(client, 23, "000001.SZ", 100, "local", "morning")
self.assertTrue(book.place(request))
self.assertTrue(book.busy("000001.SZ", "BUY"))
def test_position_dataclasses_execute_without_type_error(self):
with TemporaryDirectory() as directory:
state = State.for_strategy(directory, "trend", "A")
position = Position(
stock_code="000001.SZ", volume=100, can_use_volume=100,
open_price=10, market_value=1000,
)
state.sync_positions([position])
runtime = SimpleNamespace(
client=FakeClient(), state=state, orders=OrderBook(),
open_watch=SimpleNamespace(forget=lambda _code: None),
add_watch=SimpleNamespace(triggered=lambda *_args: False, forget=lambda _code: None),
profit_tracker=GridTrailingTracker(1),
account_cfg=SimpleNamespace(
account_id="A", excluded_codes=[], grid_step_pct=1,
enable_loss_add_position=False, buy_value=5000,
strategy="trend",
),
)
manage_positions(runtime, {"000001.SZ": Tick(last_price=10.1)}, [position], True, 5000)
def test_loss_tier_boundary_does_not_overflow(self):
self.assertEqual(len(LOSS_TIERS), 2)
with TemporaryDirectory() as directory:
state = State.for_strategy(directory, "trend", "A")
position = Position(stock_code="A", volume=100, open_price=10, market_value=1000)
state.sync_positions([position])
item = state.get("A")
item.added_num = len(LOSS_TIERS)
state.set(item)
runtime = SimpleNamespace(
state=state, account_cfg=SimpleNamespace(buy_value=5000, strategy="trend"),
add_watch=SimpleNamespace(triggered=lambda *_args: True), orders=OrderBook(),
client=FakeClient(),
)
decision = handle_loss(runtime, position, Tick(last_price=5), -60, 5000)
self.assertFalse(decision.submitted)
def test_loss_tiers_zero_and_one(self):
with TemporaryDirectory() as directory:
state = State.for_strategy(directory, "trend", "A")
position = Position(stock_code="A", volume=100, open_price=10, market_value=1000)
state.sync_positions([position])
runtime = SimpleNamespace(
state=state, account_cfg=SimpleNamespace(buy_value=5000, strategy="trend"),
add_watch=SimpleNamespace(triggered=lambda *_args: False),
orders=OrderBook(), client=FakeClient(),
)
first = handle_loss(runtime, position, Tick(last_price=7), -30, 5000)
self.assertIn("等待", first.message)
item = state.get("A")
item.added_num = 1
state.set(item)
before_second_tier = handle_loss(runtime, position, Tick(last_price=6), -40, 5000)
self.assertEqual(before_second_tier.message, "")
second = handle_loss(runtime, position, Tick(last_price=5), -50, 5000)
self.assertIn("等待", second.message)
def test_reconcile_ing_order_from_deal(self):
with TemporaryDirectory() as directory:
state = State.for_strategy(directory, "trend", "A")
position = Position(stock_code="A", volume=100, open_price=10)
state.set(StateItem("A", base_order_id="local-1", base_status="ING"))
state.reconcile(
[position],
[],
[{"m_strRemark": "local-1|morning"}],
)
self.assertEqual(state.get("A").base_status, STATUS_OK)
def test_low_cash_still_runs_position_management(self):
client = SimpleNamespace(
assets=lambda: Assets(total=10000, available=10),
positions=lambda: (["A"], [Position(stock_code="A", volume=100, open_price=10)]),
full_tick=lambda _codes: {"A": Tick(last_price=11)},
)
runtime = SimpleNamespace(
client=client,
account_cfg=SimpleNamespace(min_cash_ratio=0.1),
global_cfg=SimpleNamespace(api_host="http://example"),
orders=SimpleNamespace(cancel_expired=lambda _client: None),
state=SimpleNamespace(codes=["A"]),
)
with (
patch("strategy.trend.boot.trading_time", return_value=True),
patch("strategy.trend.boot.market_allow_open", return_value=True),
patch("strategy.trend.boot.open_signal") as open_mock,
patch("strategy.trend.boot.manage_positions") as manage_mock,
):
RunOnce(runtime, [])
open_mock.assert_not_called()
manage_mock.assert_called_once()
if __name__ == "__main__":
unittest.main()