feat(monitor): 新增板块异动监控

This commit is contained in:
shy3130
2026-08-06 19:28:39 +08:00
parent 99bdec875d
commit ecfddb451e
14 changed files with 1214 additions and 30 deletions
+46 -1
View File
@@ -59,15 +59,34 @@ class ConditionModel(BaseModel):
value: float | None = None # op 非 truth 时必填
class SectorTargetModel(BaseModel):
key: str
kind: str
name: str
symbol: str | None = None
source_id: str | None = None
field: str | None = None
source_field: str | None = None
value: str | None = None
level: int | None = None
available: bool = True
member_count: int = 0
class RuleModel(BaseModel):
id: str
name: str
enabled: bool = True
type: str # strategy | signal | price | market
type: str # strategy | signal | price | market | sector
asset_type: str = "stock" # stock | etf (etf: strategy 型走 ETF 历史加载器)
scope: str = "symbols" # symbols | all | sector
symbols: list[str] = []
sector: str | None = None
sector_kind: str | None = None # index | concept | industry
sector_targets: list[SectorTargetModel] = []
sector_trigger: str = "change_pct" # change_pct | momentum
threshold_pct: float = 1.0
window_minutes: int = 5
strategy_id: str | None = None
direction: str = "entry" # entry | exit | both
notify_events: list[str] | None = None
@@ -119,6 +138,10 @@ def get_options(request: Request):
except Exception:
pass
sector_service = getattr(request.app.state, "sector_monitor_service", None)
sector_targets = sector_service.list_targets() if sector_service is not None else {
"index": [], "concept": [], "industry": [],
}
return {
"threshold_fields": threshold_fields,
"builtin_signals": builtin_signals,
@@ -129,6 +152,7 @@ def get_options(request: Request):
{"key": "price", "label": "价格/涨跌"},
{"key": "market", "label": "市场异动"},
{"key": "strategy", "label": "策略监控"},
{"key": "sector", "label": "板块监控"},
],
"scopes": [
{"key": "symbols", "label": "指定标的"},
@@ -152,6 +176,7 @@ def get_options(request: Request):
"intraday_signal_support": intraday_monitor_support(
getattr(request.app.state, "capabilities", None),
),
"sector_targets": sector_targets,
}
@@ -186,6 +211,17 @@ def list_rules(request: Request):
if runtime_warning:
for rule in intraday_rules:
rule["runtime_warning"] = runtime_warning
sector_service = getattr(request.app.state, "sector_monitor_service", None)
if sector_service is not None:
for rule in rules:
if rule.get("type") != "sector":
continue
missing = sector_service.missing_target_keys(rule.get("sector_targets", []))
unavailable = sector_service.unavailable_target_keys(rule.get("sector_targets", []))
if missing:
rule["runtime_warning"] = "部分板块数据已不存在, 请重新选择监控对象"
elif unavailable:
rule["runtime_warning"] = "所选指数未加入实时指数池, 请先在实时监控设置中启用"
# 按 created_at 倒序
rules.sort(key=lambda r: r.get("created_at", ""), reverse=True)
return {"rules": rules}
@@ -232,6 +268,15 @@ def save_rule(req: RuleModel, request: Request):
monitor_rules.validate(rule)
except ValueError as e:
raise HTTPException(status_code=400, detail=str(e)) from e
if rule.get("type") == "sector":
sector_service = getattr(request.app.state, "sector_monitor_service", None)
if sector_service is None:
raise HTTPException(status_code=503, detail="板块监控服务未初始化")
targets = rule.get("sector_targets", [])
if sector_service.missing_target_keys(targets):
raise HTTPException(status_code=400, detail="所选板块数据已变化, 请重新选择")
if sector_service.unavailable_target_keys(targets):
raise HTTPException(status_code=400, detail="所选指数未加入实时指数池, 请先在实时监控设置中启用")
if rule.get("enabled", True) and uses_intraday_signals(rule):
from app.services.kline_sync import intraday_monitor_support
+4
View File
@@ -214,9 +214,12 @@ async def lifespan(app: FastAPI):
from app.strategy.monitor import MonitorRuleEngine
from app.strategy import monitor_rules as mr_store
from app.services import preferences
from app.services.sector_monitor import SectorMonitorService
monitor_engine = MonitorRuleEngine()
sector_monitor_service = SectorMonitorService(repo)
monitor_engine.set_strategy_engine(strategy_engine)
monitor_engine.set_data_dir(store.data_dir)
monitor_engine.set_sector_monitor_service(sector_monitor_service)
# 复用 ScreenerService 的历史窗口加载器 (三级缓存, 启动预计算命中 ~0ms),
# 让声明 filter_history 的策略 (如反包) 也能在实时监控里跑选股 → 盘中触发通知。
monitor_engine.set_history_loader(_screener_svc._load_enriched_history)
@@ -241,6 +244,7 @@ async def lifespan(app: FastAPI):
except Exception as e: # noqa: BLE001
logger.warning("monitor engine load failed: %s", e)
app.state.monitor_engine = monitor_engine
app.state.sector_monitor_service = sector_monitor_service
yield
+18 -3
View File
@@ -1116,6 +1116,11 @@ class QuoteService:
rule_events = engine.evaluate(eval_df, asset_type="stock")
if engine.consume_strategy_result_updates():
self.notify_strategy_results_updated()
if engine.has_rule_type("sector"):
rule_events += engine.evaluate_sectors(
enriched_today if stock_ready else pl.DataFrame(),
self.get_index_quotes(),
)
# ETF 规则轮: 股票快照不含 ETF, 用 ETF enriched 快照单独评估。
# 独立 try —— ETF 轮任何异常都不得丢弃本轮已算出的股票告警。
# refresh=False —— 不在轮询线程上触发 ETF 冷缓存的同步重算 (缓存由 ETF 实时
@@ -1154,7 +1159,7 @@ class QuoteService:
logger.warning("告警落盘失败: %s", e)
# 转为 SSE 推送格式 (兼容旧 alert schema)
for ev in rule_events:
all_alerts.append({
alert = {
"source": ev["source"],
"type": ev["type"],
"rule_id": ev.get("rule_id"),
@@ -1168,7 +1173,16 @@ class QuoteService:
"severity": ev.get("severity", "info"),
"conditions": ev.get("conditions") or [],
"logic": ev.get("logic") or "and",
})
}
for key in (
"sector_kind", "sector_key", "sector_name",
"sector_source_field", "sector_value", "sector_level",
"window_change_pct", "coverage_ratio", "valid_count",
"total_count", "up_count", "down_count", "leader",
):
if key in ev:
alert[key] = ev[key]
all_alerts.append(alert)
# 策略页实时回显: 不写文件 (实时行情每轮更新 enriched, 写文件会被 read_cache
# 的 mtime 校验判过期, 反复读不到)。监控引擎本轮已算出的结果存在内存
@@ -1339,6 +1353,7 @@ class QuoteService:
source_labels = {
"strategy": "策略", "signal": "信号",
"price": "价格", "market": "异动", "ladder": "连板梯队",
"sector": "板块",
}
rules = engine.rules if engine is not None else {}
enqueued = 0
@@ -1391,7 +1406,7 @@ class QuoteService:
source = ev.get("source", "")
source_label = {
"strategy": "策略", "signal": "信号",
"price": "价格", "market": "异动",
"price": "价格", "market": "异动", "sector": "板块",
}.get(source, source or "通知")
name = ev.get("name") or ""
+361
View File
@@ -0,0 +1,361 @@
"""板块监控目标目录与实时聚合快照。"""
from __future__ import annotations
import hashlib
import math
import re
from collections import deque
from datetime import datetime
from pathlib import Path
from typing import Any
import polars as pl
from app.services import preferences
from app.services.ext_data import ExtConfig, ExtConfigStore
CORE_INDICES = {
"000001.SH": "上证指数",
"399001.SZ": "深证成指",
"399006.SZ": "创业板指",
"000680.SH": "科创综指",
}
SECTOR_KINDS = {"index", "concept", "industry"}
_VALUE_SEP = re.compile(r"[\u3001,\uff0c;\uff1b|]+")
_NULL_VALUES = {"nan", "none", "null", "<na>", "n/a", "-"}
_HISTORY_SECONDS = 31 * 60
_WINDOW_TOLERANCE_SECONDS = 90
def _finite(value: Any) -> float | None:
if value is None:
return None
try:
number = float(value)
except (TypeError, ValueError):
return None
return number if math.isfinite(number) else None
def _dimension_kind(field_name: str, field_label: str) -> str | None:
text = f"{field_name} {field_label}".lower()
if any(word in text for word in ("概念", "题材", "concept", "theme")):
return "concept"
if any(word in text for word in ("行业", "申万", "中信", "industry", "sector")):
return "industry"
return None
def _target_key(kind: str, source_id: str, field: str, value: str, level: int | None) -> str:
raw = f"{kind}\0{source_id}\0{field}\0{level or 0}\0{value}"
digest = hashlib.sha1(raw.encode("utf-8")).hexdigest()[:16]
return f"{kind}:{digest}"
class SectorMonitorService:
"""缓存板块成员关系, 并按启用规则构建轻量实时快照。"""
def __init__(self, repo) -> None:
self._repo = repo
self._data_dir: Path = repo.store.data_dir
self._catalog_signature: tuple[tuple[str, int, int], ...] | None = None
self._catalog: dict[str, list[dict]] = {kind: [] for kind in SECTOR_KINDS}
self._targets_by_key: dict[str, dict] = {}
self._members_by_key: dict[str, set[str]] = {}
self._history: dict[str, deque[tuple[float, float]]] = {}
self._history_day: str | None = None
def list_targets(self) -> dict[str, list[dict]]:
self._ensure_catalog()
return {kind: [dict(item) for item in self._catalog[kind]] for kind in self._catalog}
def missing_target_keys(self, targets: list[dict]) -> list[str]:
self._ensure_catalog()
return [str(target.get("key") or "") for target in targets if target.get("key") not in self._targets_by_key]
def unavailable_target_keys(self, targets: list[dict]) -> list[str]:
self._ensure_catalog()
return [
str(target.get("key") or "")
for target in targets
if target.get("key") in self._targets_by_key
and not self._targets_by_key[target["key"]].get("available", True)
]
def build_snapshots(
self,
stock_df: pl.DataFrame,
index_df: pl.DataFrame,
targets: list[dict],
windows: set[int],
*,
now: float,
) -> dict[str, dict]:
if not targets:
return {}
self._ensure_catalog()
self._reset_history_for_day(now)
stock_rows = self._row_map(stock_df, index_values_are_percent=False)
index_rows = self._row_map(index_df, index_values_are_percent=True)
snapshots: dict[str, dict] = {}
for raw_target in targets:
key = str(raw_target.get("key") or "")
target = self._targets_by_key.get(key)
if not target:
continue
if target["kind"] == "index":
snapshot = self._index_snapshot(target, index_rows)
else:
snapshot = self._dimension_snapshot(target, stock_rows)
if snapshot is None:
continue
change_pct = snapshot.get("change_pct")
history = self._history.setdefault(key, deque())
if snapshot["valid"] and change_pct is not None:
history.append((now, float(change_pct)))
while history and history[0][0] < now - _HISTORY_SECONDS:
history.popleft()
snapshot["window_changes"] = {
window: self._window_change(history, now, window, change_pct)
for window in windows
}
snapshots[key] = snapshot
return snapshots
def _ensure_catalog(self) -> None:
signature = self._data_signature()
if signature == self._catalog_signature:
return
catalog = {kind: [] for kind in SECTOR_KINDS}
targets_by_key: dict[str, dict] = {}
members_by_key: dict[str, set[str]] = {}
index_names = dict(CORE_INDICES)
try:
indices = self._repo.get_index_instruments()
if not indices.is_empty() and "symbol" in indices.columns:
for row in indices.to_dicts():
if row.get("asset_type") == "etf":
continue
symbol = str(row.get("symbol") or "").strip().upper()
if symbol:
index_names[symbol] = str(row.get("name") or symbol)
except Exception:
pass
realtime_index_enabled = preferences.get_realtime_pull_index()
realtime_indices = set(preferences.get_realtime_index_symbols() or CORE_INDICES)
all_indices_enabled = preferences.get_realtime_index_mode() == "all"
for symbol, name in sorted(index_names.items()):
target = {
"key": f"index:{symbol}",
"kind": "index",
"name": name,
"symbol": symbol,
"available": realtime_index_enabled and (all_indices_enabled or symbol in realtime_indices),
"member_count": 1,
}
catalog["index"].append(target)
targets_by_key[target["key"]] = target
catalog["index"].sort(key=lambda item: (not item["available"], item["symbol"]))
for config in ExtConfigStore(self._data_dir).load_all():
df = self._read_ext_dataframe(config)
if df.is_empty():
continue
symbol_col = self._symbol_column(config, df)
if not symbol_col:
continue
for field in config.fields:
kind = _dimension_kind(field.name, field.label)
if kind is None or field.name not in df.columns:
continue
for row in df.select([symbol_col, field.name]).iter_rows(named=True):
symbol = str(row.get(symbol_col) or "").strip().upper()
if not symbol:
continue
for raw_value in self._dimension_values(row.get(field.name)):
paths = self._industry_paths(raw_value) if kind == "industry" else [(raw_value, None, raw_value)]
for value, level, name in paths:
key = _target_key(kind, config.id, field.name, value, level)
members_by_key.setdefault(key, set()).add(symbol)
if key not in targets_by_key:
target = {
"key": key,
"kind": kind,
"name": name,
"source_id": config.id,
"field": field.name,
"source_field": f"{config.id}.{field.name}",
"value": value,
"level": level,
"available": True,
}
targets_by_key[key] = target
catalog[kind].append(target)
for kind in ("concept", "industry"):
for target in catalog[kind]:
target["member_count"] = len(members_by_key.get(target["key"], set()))
catalog[kind].sort(key=lambda item: (item.get("level") or 0, item["name"], item["value"]))
self._catalog_signature = signature
self._catalog = catalog
self._targets_by_key = targets_by_key
self._members_by_key = members_by_key
self._history.clear()
def _data_signature(self) -> tuple[tuple[str, int, int], ...]:
base = self._data_dir / "ext_data"
paths: list[Path] = []
for config in ExtConfigStore(self._data_dir).load_all():
if not any(_dimension_kind(field.name, field.label) for field in config.fields):
continue
config_dir = base / config.id
paths.extend(config_dir.rglob("config.json"))
paths.extend(config_dir.rglob("*.parquet"))
signature = [
(str(path), path.stat().st_mtime_ns, path.stat().st_size)
for path in sorted(paths)
if path.is_file()
]
index_mode = preferences.get_realtime_index_mode()
index_enabled = preferences.get_realtime_pull_index()
index_symbols = sorted(preferences.get_realtime_index_symbols() or CORE_INDICES)
signature.append((f"realtime_indices:{index_enabled}:{index_mode}:{','.join(index_symbols)}", 0, 0))
return tuple(signature)
def _read_ext_dataframe(self, config: ExtConfig) -> pl.DataFrame:
base = self._data_dir / "ext_data" / config.id
if config.mode == "timeseries":
files = sorted((base / "timeseries").rglob("*.parquet"))
files = files[-1:] if files else []
else:
files = sorted(base.glob("*.parquet"))
if not files:
return pl.DataFrame()
try:
return pl.read_parquet(files)
except Exception:
return pl.DataFrame()
@staticmethod
def _symbol_column(config: ExtConfig, df: pl.DataFrame) -> str | None:
candidates = ["symbol", "code", "股票代码", "代码"]
for mapping in (config.symbol_map, config.code_map):
if isinstance(mapping, dict) and mapping.get("type") == "mapped":
candidates.append(str(mapping.get("col") or ""))
return next((column for column in candidates if column in df.columns), None)
@staticmethod
def _dimension_values(raw: Any) -> list[str]:
if raw is None:
return []
return [
value.strip()
for value in _VALUE_SEP.split(str(raw))
if value.strip() and value.strip().casefold() not in _NULL_VALUES
]
@staticmethod
def _industry_paths(raw: str) -> list[tuple[str, int, str]]:
parts = [part.strip() for part in raw.split("-") if part.strip()]
if not parts:
return []
return [
("-".join(parts[:level]), level, " / ".join(parts[:level]))
for level in range(1, len(parts) + 1)
]
@staticmethod
def _row_map(df: pl.DataFrame, *, index_values_are_percent: bool) -> dict[str, dict]:
if df.is_empty() or "symbol" not in df.columns:
return {}
rows: dict[str, dict] = {}
for row in df.to_dicts():
symbol = str(row.get("symbol") or "").strip().upper()
if not symbol:
continue
change_pct = _finite(row.get("change_pct"))
if change_pct is not None and index_values_are_percent:
change_pct /= 100
rows[symbol] = {**row, "change_pct": change_pct}
return rows
@staticmethod
def _index_snapshot(target: dict, index_rows: dict[str, dict]) -> dict | None:
row = index_rows.get(str(target.get("symbol") or "").upper())
if not row or row.get("change_pct") is None:
return None
return {
**target,
"valid": True,
"change_pct": row["change_pct"],
"price": _finite(row.get("close") or row.get("last_price")),
"coverage_ratio": 1.0,
"valid_count": 1,
"total_count": 1,
"up_count": int(row["change_pct"] > 0),
"down_count": int(row["change_pct"] < 0),
"leader": None,
}
def _dimension_snapshot(self, target: dict, stock_rows: dict[str, dict]) -> dict | None:
members = self._members_by_key.get(target["key"], set())
if not members:
return None
valid_rows = [
stock_rows[symbol]
for symbol in members
if symbol in stock_rows and stock_rows[symbol].get("change_pct") is not None
]
total_count = len(members)
valid_count = len(valid_rows)
coverage_ratio = valid_count / total_count if total_count else 0.0
valid = total_count >= 5 and coverage_ratio >= 0.8
changes = [float(row["change_pct"]) for row in valid_rows]
leader = max(valid_rows, key=lambda row: row["change_pct"]) if valid_rows else None
return {
**target,
"valid": valid,
"change_pct": sum(changes) / len(changes) if changes else None,
"price": None,
"coverage_ratio": coverage_ratio,
"valid_count": valid_count,
"total_count": total_count,
"up_count": sum(value > 0 for value in changes),
"down_count": sum(value < 0 for value in changes),
"leader": {
"symbol": leader.get("symbol"),
"name": leader.get("name"),
"change_pct": leader.get("change_pct"),
} if leader else None,
}
@staticmethod
def _window_change(
history: deque[tuple[float, float]],
now: float,
window: int,
current: float | None,
) -> float | None:
if current is None:
return None
cutoff = now - window * 60
for timestamp, previous in reversed(history):
if timestamp <= cutoff:
if timestamp < cutoff - _WINDOW_TOLERANCE_SECONDS:
return None
return current - previous
return None
def _reset_history_for_day(self, now: float) -> None:
day = datetime.fromtimestamp(now).date().isoformat()
if self._history_day == day:
return
self._history_day = day
self._history.clear()
+172
View File
@@ -353,6 +353,8 @@ class MonitorRuleEngine:
self._building_strategy_results: dict[str, dict] = {}
# 本轮成功写入股票策略实时结果的策略 ID, 供 QuoteService 在计算完成后精确通知策略页。
self._latest_strategy_result_ids: set[str] = set()
self._sector_monitor_service = None
self._sector_condition_state: dict[tuple[str, str], bool] = {}
def set_strategy_engine(self, engine) -> None:
"""注入 StrategyEngine, type=strategy 规则据此跑选股。"""
@@ -362,6 +364,9 @@ class MonitorRuleEngine:
"""注入数据目录, 用于加载策略的用户覆盖配置。"""
self._data_dir = data_dir
def set_sector_monitor_service(self, service) -> None:
self._sector_monitor_service = service
def invalidate_strategy_state(self) -> None:
"""策略注册表变更后清除选股池、结果和矩阵快照。"""
self._strategy_pools.clear()
@@ -413,6 +418,12 @@ class MonitorRuleEngine:
rule.get("scope", "symbols"),
tuple(sorted(str(symbol) for symbol in rule.get("symbols", []))),
rule.get("sector"),
rule.get("sector_kind"),
tuple(sorted(str(target.get("key")) for target in rule.get("sector_targets", []))),
rule.get("sector_trigger"),
rule.get("direction"),
rule.get("threshold_pct"),
rule.get("window_minutes"),
)
def set_rules(self, rules: list[dict]) -> None:
@@ -450,6 +461,11 @@ class MonitorRuleEngine:
for key, value in list(self._strategy_signal_seen.items())
if key[0] in active_ids
}
self._sector_condition_state = {
key: value
for key, value in list(self._sector_condition_state.items())
if key[0] in active_ids
}
logger.info("MonitorRuleEngine: 装载 %d 条规则", len(self._rules))
def add_rule(self, rule: dict) -> None:
@@ -470,6 +486,9 @@ class MonitorRuleEngine:
self._strategy_signal_seen = {
k: v for k, v in list(self._strategy_signal_seen.items()) if k[0] != rule_id
}
self._sector_condition_state = {
k: v for k, v in self._sector_condition_state.items() if k[0] != rule_id
}
def clear(self) -> None:
self._rules.clear()
@@ -477,6 +496,7 @@ class MonitorRuleEngine:
self._strategy_pools.clear()
self._strategy_signal_state.clear()
self._strategy_signal_seen.clear()
self._sector_condition_state.clear()
@property
def rules(self) -> dict[str, dict]:
@@ -640,6 +660,8 @@ class MonitorRuleEngine:
for rule_id, rule in list(self._rules.items()):
if rule.get("asset_type", "stock") != asset_type:
continue
if rule.get("type") == "sector":
continue
try:
events.extend(self._evaluate_rule(df, rule, now))
except Exception as e:
@@ -652,6 +674,156 @@ class MonitorRuleEngine:
return events
def evaluate_sectors(
self,
stock_df: pl.DataFrame,
index_df: pl.DataFrame,
*,
now: float | None = None,
) -> list[dict]:
"""按板块聚合快照评估 type=sector 规则。"""
if self._sector_monitor_service is None:
return []
rules = [
rule for rule in list(self._rules.values())
if rule.get("enabled", True) and rule.get("type") == "sector"
]
if not rules:
return []
targets_by_key: dict[str, dict] = {}
windows: set[int] = set()
for rule in rules:
for target in rule.get("sector_targets", []):
if target.get("key"):
targets_by_key[str(target["key"])] = target
if rule.get("sector_trigger") == "momentum":
windows.add(int(rule.get("window_minutes", 5)))
timestamp = time.time() if now is None else now
snapshots = self._sector_monitor_service.build_snapshots(
stock_df,
index_df,
list(targets_by_key.values()),
windows,
now=timestamp,
)
events: list[dict] = []
for rule in rules:
try:
events.extend(self._evaluate_sector_rule(rule, snapshots, timestamp))
except Exception as exc: # noqa: BLE001
logger.warning("板块规则评估失败 %s: %s", rule.get("id"), exc)
return events
def _evaluate_sector_rule(self, rule: dict, snapshots: dict[str, dict], now: float) -> list[dict]:
events: list[dict] = []
direction = rule.get("direction", "up")
trigger = rule.get("sector_trigger", "change_pct")
threshold = float(rule.get("threshold_pct", 1.0)) / 100
window = int(rule.get("window_minutes", 5))
for target in rule.get("sector_targets", []):
target_key = str(target.get("key") or "")
snapshot = snapshots.get(target_key)
if not snapshot or not snapshot.get("valid"):
continue
value = (
snapshot.get("change_pct")
if trigger == "change_pct"
else snapshot.get("window_changes", {}).get(window)
)
condition = value is not None and (
value >= threshold if direction == "up" else value <= -threshold
)
state_key = (rule["id"], target_key)
previous = self._sector_condition_state.get(state_key)
self._sector_condition_state[state_key] = condition
if previous is None or previous or not condition:
continue
event_type = f"sector_{trigger}_{direction}"
cooldown_key = (rule["id"], target_key, event_type)
last = self._last_fire.get(cooldown_key)
cooldown = int(rule.get("cooldown_seconds", 3600))
if last is not None and now - last < cooldown:
continue
self._last_fire[cooldown_key] = now
message = rule.get("message", "") or self._sector_message(
snapshot, trigger, direction, threshold, window, value,
)
event = {
"ts": int(now * 1000),
"rule_id": rule["id"],
"rule_name": rule.get("name", ""),
"strategy_id": None,
"source": "sector",
"type": event_type,
"symbol": snapshot.get("symbol") if snapshot.get("kind") == "index" else "",
"name": snapshot.get("name"),
"message": message,
"price": snapshot.get("price"),
"change_pct": snapshot.get("change_pct"),
"window_change_pct": value if trigger == "momentum" else None,
"signals": [],
"severity": rule.get("severity", "info"),
"conditions": [],
"logic": "and",
"sector_kind": snapshot.get("kind"),
"sector_key": target_key,
"sector_name": snapshot.get("name"),
"sector_source_field": snapshot.get("source_field"),
"sector_value": snapshot.get("value"),
"sector_level": snapshot.get("level"),
"coverage_ratio": snapshot.get("coverage_ratio"),
"valid_count": snapshot.get("valid_count"),
"total_count": snapshot.get("total_count"),
"up_count": snapshot.get("up_count"),
"down_count": snapshot.get("down_count"),
"leader": snapshot.get("leader"),
}
events.append(event)
if self._alert_handler:
try:
self._alert_handler(event)
except Exception as exc: # noqa: BLE001
logger.warning("alert handler failed: %s", exc)
return events
@staticmethod
def _sector_message(
snapshot: dict,
trigger: str,
direction: str,
threshold: float,
window: int,
value: float | None,
) -> str:
kind_label = {
"index": "指数", "concept": "概念", "industry": "行业",
}.get(snapshot.get("kind"), "板块")
current = float(snapshot.get("change_pct") or 0)
if trigger == "momentum":
action = "快速拉升" if direction == "up" else "快速下跌"
head = (
f"{kind_label}{snapshot.get('name')}{window}分钟{action} "
f"{float(value or 0) * 100:+.2f}%"
)
else:
action = "涨幅上穿" if direction == "up" else "跌幅下穿"
head = f"{kind_label}{snapshot.get('name')}{action} {threshold * 100:.2f}%"
parts = [head, f"当前 {current * 100:+.2f}%"]
if snapshot.get("kind") != "index":
parts.append(f"上涨 {snapshot.get('up_count', 0)}/{snapshot.get('valid_count', 0)}")
parts.append(f"覆盖 {float(snapshot.get('coverage_ratio') or 0) * 100:.0f}%")
leader = snapshot.get("leader") or {}
if leader.get("name") or leader.get("symbol"):
parts.append(
f"领涨 {leader.get('name') or leader.get('symbol')} "
f"{float(leader.get('change_pct') or 0) * 100:+.2f}%"
)
return "".join(parts)
def _evaluate_rule(self, df: pl.DataFrame, rule: dict, now: float) -> list[dict]:
"""评估单条规则,返回触发的 events。"""
# 1. 按 scope 过滤作用域
+37 -3
View File
@@ -27,7 +27,7 @@ logger = logging.getLogger(__name__)
# ── 常量 ────────────────────────────────────────────────
ID_RE = re.compile(r"^[a-z0-9_]{1,40}$")
RULE_TYPES = {"strategy", "signal", "price", "market", "ladder"}
RULE_TYPES = {"strategy", "signal", "price", "market", "ladder", "sector"}
SCOPES = {"symbols", "all", "sector"}
LOGICS = {"and", "or"}
DIRECTIONS = {"entry", "exit", "both"}
@@ -38,6 +38,9 @@ OPS = {">", ">=", "<", "<=", "==", "!="}
LADDER_METRICS = {"sealed_vol", "sealed_amount"}
# ladder 规则: 方向 (up=涨停炸板预警, down=跌停翘板预警)
LADDER_DIRECTIONS = {"up", "down"}
SECTOR_KINDS = {"index", "concept", "industry"}
SECTOR_TRIGGERS = {"change_pct", "momentum"}
SECTOR_WINDOWS = {1, 3, 5, 10, 15}
# 布尔信号列前缀 (op=truth 时 field 取这些)
_SIGNAL_PREFIXES = ("signal_", "csg_")
@@ -138,6 +141,29 @@ def validate(rule: dict) -> None:
thr = rule.get("threshold")
if not isinstance(thr, (int, float)) or thr < 0:
raise ValueError("threshold 必须是非负数字 (封单 ≤ 此值时报警)")
elif rule.get("type") == "sector":
kind = rule.get("sector_kind")
if kind not in SECTOR_KINDS:
raise ValueError(f"sector_kind 必须是 {SECTOR_KINDS} 之一")
targets = rule.get("sector_targets")
if not isinstance(targets, list) or not targets:
raise ValueError("板块监控至少选择一个监控对象")
if len(targets) > 20:
raise ValueError("板块监控对象最多 20 个")
for target in targets:
if not isinstance(target, dict) or not target.get("key") or not target.get("name"):
raise ValueError("板块监控对象格式错误")
if target.get("kind") != kind:
raise ValueError("板块监控对象类型必须一致")
if rule.get("sector_trigger") not in SECTOR_TRIGGERS:
raise ValueError(f"sector_trigger 必须是 {SECTOR_TRIGGERS} 之一")
if rule.get("direction") not in LADDER_DIRECTIONS:
raise ValueError("板块监控 direction 必须是 up 或 down")
threshold_pct = rule.get("threshold_pct")
if not isinstance(threshold_pct, (int, float)) or not 0 < threshold_pct <= 20:
raise ValueError("板块监控阈值必须大于 0 且不超过 20%")
if rule.get("sector_trigger") == "momentum" and rule.get("window_minutes") not in SECTOR_WINDOWS:
raise ValueError(f"板块异动窗口必须是 {sorted(SECTOR_WINDOWS)} 分钟之一")
else:
# 信号/价格/市场类型: 需要 conditions
conds = rule.get("conditions")
@@ -196,9 +222,14 @@ def normalize(rule: dict) -> dict:
r.setdefault("scope", "symbols")
r.setdefault("symbols", [])
r.setdefault("sector", None)
r.setdefault("sector_kind", None)
r.setdefault("sector_targets", [])
r.setdefault("sector_trigger", "change_pct")
r.setdefault("threshold_pct", 1.0)
r.setdefault("window_minutes", 5)
r.setdefault("strategy_id", None)
# direction 默认值: ladder 用 "up", 其余用 "entry"
r.setdefault("direction", "up" if r.get("type") == "ladder" else "entry")
# direction 默认值: ladder/sector 用 "up", 其余用 "entry"
r.setdefault("direction", "up" if r.get("type") in {"ladder", "sector"} else "entry")
if r.get("type") == "strategy":
if r.get("notify_events") is None:
# 兼容统一监控上线后的旧规则: 当时实际行为是同时通知进入和移出。
@@ -211,6 +242,9 @@ def normalize(rule: dict) -> dict:
# ladder 专属默认字段
r.setdefault("metric", "sealed_vol")
r.setdefault("threshold", 0)
if r.get("type") == "sector":
r["scope"] = "all"
r["symbols"] = []
r.setdefault("logic", "and")
r.setdefault("cooldown_seconds", 3600)
r.setdefault("severity", "info")
+216
View File
@@ -0,0 +1,216 @@
from __future__ import annotations
from types import SimpleNamespace
import polars as pl
import pytest
from app.services import sector_monitor
from app.services.ext_data import ExtConfig, ExtConfigStore, ExtField
from app.services.sector_monitor import SectorMonitorService
from app.strategy import monitor_rules
from app.strategy.monitor import MonitorRuleEngine
class _Repo:
def __init__(self, data_dir, indices: pl.DataFrame | None = None):
self.store = SimpleNamespace(data_dir=data_dir)
self._indices = indices if indices is not None else pl.DataFrame()
def get_index_instruments(self) -> pl.DataFrame:
return self._indices
def _index_target(symbol: str, name: str) -> dict:
return {
"key": f"index:{symbol}",
"kind": "index",
"name": name,
"symbol": symbol,
}
def _sector_rule(targets: list[dict], **overrides) -> dict:
rule = {
"id": "r_sector",
"name": "板块监控",
"enabled": True,
"type": "sector",
"scope": "all",
"sector_kind": targets[0]["kind"],
"sector_targets": targets,
"sector_trigger": "change_pct",
"direction": "up",
"threshold_pct": 1.0,
"window_minutes": 5,
"cooldown_seconds": 0,
"severity": "info",
}
rule.update(overrides)
return monitor_rules.normalize(rule)
def test_validate_accepts_sector_rule_and_rejects_mixed_target_kinds():
rule = _sector_rule([_index_target("000001.SH", "上证指数")])
monitor_rules.validate(rule)
mixed = _sector_rule([
_index_target("000001.SH", "上证指数"),
{
"key": "concept:test:field:人工智能",
"kind": "concept",
"name": "人工智能",
"source_id": "test",
"field": "field",
"value": "人工智能",
},
])
try:
monitor_rules.validate(mixed)
except ValueError as exc:
assert "类型" in str(exc)
else:
raise AssertionError("混合板块类型必须被拒绝")
def test_dimension_values_preserve_names_with_spaces_and_filter_nulls(tmp_path):
service = SectorMonitorService(_Repo(tmp_path))
assert service._dimension_values("中国AI 50;6G概念") == ["中国AI 50", "6G概念"]
assert service._dimension_values("nan") == []
assert service._dimension_values(float("nan")) == []
assert service._industry_paths("电子-半导体-数字芯片设计")[-1] == (
"电子-半导体-数字芯片设计", 3, "电子 / 半导体 / 数字芯片设计",
)
def test_index_targets_are_evaluated_independently(tmp_path):
repo = _Repo(tmp_path)
service = SectorMonitorService(repo)
engine = MonitorRuleEngine()
engine.set_sector_monitor_service(service)
sh = _index_target("000001.SH", "上证指数")
cyb = _index_target("399006.SZ", "创业板指")
engine.set_rules([_sector_rule([sh, cyb])])
first = pl.DataFrame({
"symbol": ["000001.SH", "399006.SZ"],
"name": ["上证指数", "创业板指"],
"close": [3000.0, 2000.0],
"change_pct": [0.8, 0.9],
})
assert engine.evaluate_sectors(pl.DataFrame(), first, now=1000.0) == []
second = first.with_columns(
pl.Series("change_pct", [1.2, 0.95]),
)
events = engine.evaluate_sectors(pl.DataFrame(), second, now=1006.0)
assert [event["sector_name"] for event in events] == ["上证指数"]
assert events[0]["change_pct"] == 0.012
def test_index_availability_updates_when_realtime_pool_changes(tmp_path, monkeypatch):
selected = ["000001.SH"]
monkeypatch.setattr(sector_monitor.preferences, "get_realtime_pull_index", lambda: True)
monkeypatch.setattr(sector_monitor.preferences, "get_realtime_index_mode", lambda: "core")
monkeypatch.setattr(sector_monitor.preferences, "get_realtime_index_symbols", lambda: selected)
service = SectorMonitorService(_Repo(tmp_path))
first = {target["symbol"]: target for target in service.list_targets()["index"]}
assert first["000001.SH"]["available"] is True
assert first["399006.SZ"]["available"] is False
initial_quote = pl.DataFrame({"symbol": ["000001.SH"], "change_pct": [0.2]})
service.build_snapshots(pl.DataFrame(), initial_quote, [first["000001.SH"]], {5}, now=1000.0)
selected[:] = ["399006.SZ"]
second = {target["symbol"]: target for target in service.list_targets()["index"]}
assert second["000001.SH"]["available"] is False
assert second["399006.SZ"]["available"] is True
changed_quote = pl.DataFrame({"symbol": ["000001.SH"], "change_pct": [1.3]})
snapshot = service.build_snapshots(
pl.DataFrame(), changed_quote, [second["000001.SH"]], {5}, now=1300.0,
)
assert snapshot["index:000001.SH"]["window_changes"][5] is None
def test_concept_snapshot_uses_member_average_and_full_window(tmp_path):
config = ExtConfig(
id="concept_test",
label="概念测试",
mode="snapshot",
fields=[
ExtField("symbol", "string", "标的代码"),
ExtField("concept", "string", "所属概念"),
],
)
ExtConfigStore(tmp_path).upsert(config)
ext_dir = tmp_path / "ext_data" / config.id
pl.DataFrame({
"symbol": ["A", "B", "C", "D", "E"],
"concept": ["人工智能", "人工智能", "人工智能", "人工智能", "人工智能"],
}).write_parquet(ext_dir / "part.parquet")
service = SectorMonitorService(_Repo(tmp_path))
target = next(
target for target in service.list_targets()["concept"]
if target["name"] == "人工智能"
)
first = pl.DataFrame({
"symbol": ["A", "B", "C", "D", "E"],
"name": ["", "", "", "", ""],
"close": [10.0] * 5,
"change_pct": [0.01, 0.02, 0.03, -0.01, 0.0],
})
snapshots = service.build_snapshots(first, pl.DataFrame(), [target], {5}, now=1000.0)
assert snapshots[target["key"]]["change_pct"] == pytest.approx(0.01)
assert snapshots[target["key"]]["coverage_ratio"] == 1.0
assert snapshots[target["key"]]["window_changes"][5] is None
second = first.with_columns((pl.col("change_pct") + 0.01).alias("change_pct"))
too_early = service.build_snapshots(second, pl.DataFrame(), [target], {5}, now=1240.0)
assert too_early[target["key"]]["window_changes"][5] is None
unrelated = ExtConfig(
id="hot_test",
label="热度测试",
mode="snapshot",
fields=[
ExtField("symbol", "string", "标的代码"),
ExtField("heat", "float", "市场热度"),
],
)
ExtConfigStore(tmp_path).upsert(unrelated)
unrelated_dir = tmp_path / "ext_data" / unrelated.id
pl.DataFrame({"symbol": ["A"], "heat": [1.0]}).write_parquet(unrelated_dir / "part.parquet")
complete = service.build_snapshots(second, pl.DataFrame(), [target], {5}, now=1300.0)
assert complete[target["key"]]["window_changes"][5] == pytest.approx(0.01)
def test_momentum_rule_triggers_after_complete_window(tmp_path):
service = SectorMonitorService(_Repo(tmp_path))
engine = MonitorRuleEngine()
engine.set_sector_monitor_service(service)
target = _index_target("000001.SH", "上证指数")
engine.set_rules([_sector_rule(
[target],
sector_trigger="momentum",
threshold_pct=1.0,
window_minutes=5,
)])
start = pl.DataFrame({
"symbol": ["000001.SH"],
"name": ["上证指数"],
"close": [3000.0],
"change_pct": [0.2],
})
assert engine.evaluate_sectors(pl.DataFrame(), start, now=1000.0) == []
early = start.with_columns(pl.lit(1.3).alias("change_pct"))
assert engine.evaluate_sectors(pl.DataFrame(), early, now=1240.0) == []
events = engine.evaluate_sectors(pl.DataFrame(), early, now=1300.0)
assert len(events) == 1
assert events[0]["type"] == "sector_momentum_up"
assert events[0]["window_change_pct"] == pytest.approx(0.011)