159 lines
5.4 KiB
Python
159 lines
5.4 KiB
Python
import socket
|
||
from dataclasses import dataclass, field
|
||
from pathlib import Path
|
||
|
||
import yaml
|
||
|
||
|
||
@dataclass(slots=True)
|
||
class SignalConfig:
|
||
"""单个交易信号的数据源及开仓限制配置。"""
|
||
|
||
# 信号接口相对于 api_host 的路径。
|
||
url: str = ""
|
||
|
||
# 允许使用该信号的时间段;"*" 表示不限制时间。
|
||
timezone: str = "*"
|
||
|
||
# 当前价格高于信号昨收价时是否仍允许开仓。
|
||
gt_last_price_is_open: bool = False
|
||
|
||
|
||
@dataclass(slots=True)
|
||
class GlobalConfig:
|
||
"""所有主机共享的系统配置。"""
|
||
|
||
qmt_base_url: str = ""
|
||
qmt_token: str = ""
|
||
api_host: str = ""
|
||
qmt_data_dir: str = ""
|
||
|
||
# Windows 主机名到对应账户配置文件的映射。
|
||
hosts: dict[str, str] = field(default_factory=dict)
|
||
|
||
# 信号名称到信号配置的映射。
|
||
signals: dict[str, SignalConfig] = field(default_factory=dict)
|
||
|
||
|
||
@dataclass(slots=True)
|
||
class AccountConfig:
|
||
"""当前主机所使用的账户及交易策略参数。"""
|
||
|
||
account_id: str = ""
|
||
host_key: str = ""
|
||
buy_value: float = 0
|
||
min_cash_ratio: float = 0
|
||
loss_trigger_pct: float = 0
|
||
grid_step_pct: float = 1
|
||
min_profit_pct: float = 0
|
||
enable_loss_add_position: bool = False
|
||
enable_auto_ipo: bool = True
|
||
signal_allow: list[str] = field(default_factory=list)
|
||
excluded_codes: list[str] = field(default_factory=list)
|
||
# ZT 开仓及每次补仓手数(每手 100 股);0 表示不启动。
|
||
zt_open_hands: int = 0
|
||
zt_sell_ratio: float = 0.5
|
||
zt_buy_fall_pct: float = 1.0
|
||
zt_max_price: float = 200.0
|
||
|
||
# 当前账户启用的策略名称,例如 trend。
|
||
strategy: str = ""
|
||
|
||
|
||
# load() 成功后保存已加载的配置,供策略模块直接读取。
|
||
global_config: GlobalConfig | None = None
|
||
account_config: AccountConfig | None = None
|
||
|
||
# QMT 和外部 HTTP 接口的默认请求超时时间,单位为秒。
|
||
HTTP_TIMEOUT = 5.0
|
||
|
||
|
||
def load(
|
||
etc_dir: str | Path | None = None,
|
||
hostname: str | None = None,
|
||
) -> tuple[GlobalConfig, AccountConfig]:
|
||
"""加载公共配置以及当前主机对应的账户配置。
|
||
|
||
Args:
|
||
etc_dir: 配置文件目录,其中必须包含 ``_global.yaml``;为空时
|
||
默认使用 py-client 下的 ``etc`` 目录。
|
||
hostname: 指定要加载的主机名;为空时使用当前计算机名。
|
||
|
||
Returns:
|
||
由全局配置和账户配置组成的二元组。
|
||
|
||
Raises:
|
||
ValueError: 配置缺失、格式错误或策略参数不合法。
|
||
"""
|
||
global global_config, account_config
|
||
|
||
root = Path(etc_dir) if etc_dir is not None else Path(__file__).parent.parent / "etc"
|
||
raw = _yaml(root / "_global.yaml")
|
||
|
||
# 将原始字典转换为带类型的信号配置,方便业务代码使用属性访问。
|
||
signals = {
|
||
key: SignalConfig(**(value or {}))
|
||
for key, value in (raw.get("signals") or {}).items()
|
||
}
|
||
values = {
|
||
key: raw.get(key, "")
|
||
for key in ("qmt_base_url", "qmt_token", "api_host", "qmt_data_dir")
|
||
}
|
||
|
||
current = hostname or socket.gethostname()
|
||
hosts = raw.get("hosts") or {}
|
||
account_file = next(
|
||
(
|
||
value
|
||
for key, value in hosts.items()
|
||
if key.strip().lower() == current.strip().lower()
|
||
),
|
||
"",
|
||
)
|
||
|
||
# QMT 地址、外部 API 地址和数据目录是启动策略的必要参数。
|
||
if (
|
||
not values["qmt_base_url"]
|
||
or not values["api_host"]
|
||
or values["qmt_data_dir"] == "."
|
||
):
|
||
raise ValueError("Global 配置缺少必要参数")
|
||
|
||
if not account_file:
|
||
raise ValueError(f'_global.yaml 未配置计算机 "{current}"')
|
||
if not Path(account_file).suffix:
|
||
account_file += ".yaml"
|
||
|
||
global_config = GlobalConfig(**values, hosts=hosts, signals=signals)
|
||
|
||
# 策略状态文件写入该目录,启动时提前确保目录存在。
|
||
Path(global_config.qmt_data_dir).mkdir(parents=True, exist_ok=True)
|
||
|
||
account_config = AccountConfig(**_yaml(root / account_file))
|
||
if account_config.buy_value <= 0 or account_config.grid_step_pct <= 0:
|
||
raise ValueError("buy_value、grid_step_pct 必须大于 0")
|
||
if type(account_config.zt_open_hands) is not int or account_config.zt_open_hands < 0:
|
||
raise ValueError("zt_open_hands 必须为非负整数,0 表示不启动 ZT 策略")
|
||
if not 0 < account_config.zt_sell_ratio <= 1:
|
||
raise ValueError("zt_sell_ratio 必须在 (0, 1] 区间")
|
||
if account_config.zt_buy_fall_pct <= 0 or account_config.zt_max_price <= 0:
|
||
raise ValueError("zt_buy_fall_pct、zt_max_price 必须大于 0")
|
||
if not account_config.strategy.strip():
|
||
raise ValueError("strategy 不能为空")
|
||
|
||
# host_key 统一为小写,避免不同模块比较时受大小写影响。
|
||
account_config.host_key = account_config.host_key.lower()
|
||
account_config.strategy = account_config.strategy.lower()
|
||
if account_config.strategy == "zt" and account_config.signal_allow != ["dcm"]:
|
||
raise ValueError("zt 策略的 signal_allow 必须且只能为 [\"dcm\"]")
|
||
return global_config, account_config
|
||
|
||
|
||
def _yaml(path: Path) -> dict:
|
||
"""读取 YAML 文件,并将空文件转换为空字典。"""
|
||
try:
|
||
with path.open(encoding="utf-8") as handle:
|
||
return yaml.safe_load(handle) or {}
|
||
except (OSError, yaml.YAMLError) as exc:
|
||
raise ValueError(f"读取或解析配置 {path} 失败: {exc}") from exc
|