mirror of
https://ghfast.top/https://github.com/aeroxw/tick-stock-panel.git
synced 2026-09-12 17:54:15 +08:00
feat(monitor): 新增板块异动监控
This commit is contained in:
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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 ""
|
||||
|
||||
@@ -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()
|
||||
@@ -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 过滤作用域
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
Reference in New Issue
Block a user