from __future__ import annotations import json import sqlite3 from datetime import datetime from typing import Any from backend.features.sentiment.engine import build_sentiment_history def _optional_float(value: Any) -> float | None: if value in (None, ""): return None try: return float(value) except (TypeError, ValueError): return None class ScreenerRepositoryMixin: def upsert_benchmark_bars(self, rows: list[dict[str, Any]]) -> int: values = [ ( str(row.get("trade_date") or ""), str(row.get("ts_code") or ""), float(row.get("close") or 0), float(row.get("pct_chg") or 0), ) for row in rows if row.get("trade_date") and row.get("ts_code") ] with self.connect() as connection: connection.executemany( """ INSERT INTO benchmark_bars (trade_date, ts_code, close, pct_chg) VALUES (?, ?, ?, ?) ON CONFLICT(trade_date, ts_code) DO UPDATE SET close=excluded.close, pct_chg=excluded.pct_chg """, values, ) return len(values) def upsert_daily_indicators(self, rows: list[dict[str, Any]]) -> int: values = [ ( str(row.get("trade_date") or ""), row.get("ts_code", ""), float(row.get("turnover_rate") or 0), float(row.get("volume_ratio") or 0), float(row.get("total_mv") or 0), float(row.get("circ_mv") or 0), _optional_float(row.get("pe_ttm")), _optional_float(row.get("pb")), _optional_float(row.get("ps_ttm")), _optional_float(row.get("dv_ttm")), ) for row in rows if row.get("trade_date") and row.get("ts_code") ] with self.connect() as connection: connection.executemany( """ INSERT INTO daily_indicators (trade_date, ts_code, turnover_rate, volume_ratio, total_mv, circ_mv, pe_ttm, pb, ps_ttm, dv_ttm) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(trade_date, ts_code) DO UPDATE SET turnover_rate=excluded.turnover_rate, volume_ratio=excluded.volume_ratio, total_mv=excluded.total_mv, circ_mv=excluded.circ_mv, pe_ttm=excluded.pe_ttm, pb=excluded.pb, ps_ttm=excluded.ps_ttm, dv_ttm=excluded.dv_ttm """, values, ) return len(values) def upsert_fundamental_indicators(self, rows: list[dict[str, Any]]) -> int: values = [ ( str(row.get("end_date") or ""), str(row.get("ann_date") or ""), str(row.get("ts_code") or ""), _optional_float(row.get("roe")), _optional_float(row.get("roa")), _optional_float(row.get("roic")), _optional_float(row.get("grossprofit_margin")), _optional_float(row.get("netprofit_yoy")), _optional_float(row.get("or_yoy")), _optional_float(row.get("ocf_to_opincome")), ) for row in rows if row.get("end_date") and row.get("ts_code") ] with self.connect() as connection: connection.executemany( """ INSERT INTO fundamental_indicators (end_date, ann_date, ts_code, roe, roa, roic, grossprofit_margin, netprofit_yoy, or_yoy, ocf_to_opincome) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(end_date, ts_code) DO UPDATE SET ann_date=excluded.ann_date, roe=excluded.roe, roa=excluded.roa, roic=excluded.roic, grossprofit_margin=excluded.grossprofit_margin, netprofit_yoy=excluded.netprofit_yoy, or_yoy=excluded.or_yoy, ocf_to_opincome=excluded.ocf_to_opincome """, values, ) return len(values) def upsert_moneyflow(self, rows: list[dict[str, Any]]) -> int: values = [] for row in rows: if not row.get("trade_date") or not row.get("ts_code"): continue large_net = ( float(row.get("buy_lg_amount") or 0) + float(row.get("buy_elg_amount") or 0) - float(row.get("sell_lg_amount") or 0) - float(row.get("sell_elg_amount") or 0) ) medium_net = float(row.get("buy_md_amount") or 0) - float(row.get("sell_md_amount") or 0) small_net = float(row.get("buy_sm_amount") or 0) - float(row.get("sell_sm_amount") or 0) values.append(( str(row["trade_date"]), row["ts_code"], float(row.get("net_mf_amount") or 0), large_net, medium_net, small_net, )) with self.connect() as connection: connection.executemany( """ INSERT INTO moneyflow_daily (trade_date, ts_code, net_mf_amount, large_net_amount, medium_net_amount, small_net_amount) VALUES (?, ?, ?, ?, ?, ?) ON CONFLICT(trade_date, ts_code) DO UPDATE SET net_mf_amount=excluded.net_mf_amount, large_net_amount=excluded.large_net_amount, medium_net_amount=excluded.medium_net_amount, small_net_amount=excluded.small_net_amount """, values, ) return len(values) def upsert_earnings_events(self, rows: list[dict[str, Any]]) -> int: values = [ ( str(row.get("end_date") or ""), str(row.get("ann_date") or ""), str(row.get("ts_code") or ""), _optional_float(row.get("forecast_profit")), _optional_float(row.get("actual_profit")), _optional_float(row.get("surprise_pct")), _optional_float(row.get("revenue_yoy")), _optional_float(row.get("netprofit_yoy")), str(row.get("source") or ""), ) for row in rows if row.get("end_date") and row.get("ann_date") and row.get("ts_code") ] with self.connect() as connection: connection.executemany( """ INSERT INTO earnings_events (end_date, ann_date, ts_code, forecast_profit, actual_profit, surprise_pct, revenue_yoy, netprofit_yoy, source) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(end_date, ann_date, ts_code) DO UPDATE SET forecast_profit=excluded.forecast_profit, actual_profit=excluded.actual_profit, surprise_pct=excluded.surprise_pct, revenue_yoy=excluded.revenue_yoy, netprofit_yoy=excluded.netprofit_yoy, source=excluded.source """, values, ) return len(values) def daily_indicator_dates(self, end_date: str = "", limit: int = 400) -> list[str]: where = "WHERE trade_date <= ?" if end_date else "" parameters: tuple[Any, ...] = (end_date, limit) if end_date else (limit,) with self.connect() as connection: rows = connection.execute( f"SELECT DISTINCT trade_date FROM daily_indicators {where} " "ORDER BY trade_date DESC LIMIT ?", parameters, ).fetchall() return [row["trade_date"] for row in reversed(rows)] def fundamental_periods(self) -> list[str]: with self.connect() as connection: rows = connection.execute( "SELECT DISTINCT end_date FROM fundamental_indicators ORDER BY end_date" ).fetchall() return [str(row["end_date"]) for row in rows] def factor_dates(self, end_date: str = "", limit: int = 80) -> list[str]: where = "WHERE trade_date <= ?" if end_date else "" parameters: tuple[Any, ...] = (end_date, limit) if end_date else (limit,) with self.connect() as connection: rows = connection.execute( f"SELECT DISTINCT trade_date FROM daily_bars {where} ORDER BY trade_date DESC LIMIT ?", parameters, ).fetchall() return [row["trade_date"] for row in reversed(rows)] def factor_health_summary(self, end_date: str) -> dict[str, Any]: dividend_start = f"{max(0, int(end_date[:4] or 0) - 5)}0101" with self.connect() as connection: market = connection.execute( "SELECT EXISTS(SELECT 1 FROM daily_bars WHERE trade_date <= ? LIMIT 1)", (end_date,), ).fetchone()[0] auction = connection.execute( "SELECT EXISTS(SELECT 1 FROM auction_factors WHERE trade_date <= ? LIMIT 1)", (end_date,), ).fetchone()[0] benchmark_rows = connection.execute( "SELECT COUNT(*) FROM benchmark_bars WHERE ts_code = '000300.SH' AND trade_date <= ?", (end_date,), ).fetchone()[0] indicator_date = connection.execute( "SELECT MAX(trade_date) FROM daily_indicators WHERE trade_date <= ?", (end_date,), ).fetchone()[0] if indicator_date: valuation_rows, valuation_available = connection.execute( """ SELECT COUNT(*), COALESCE(MAX(pe_ttm IS NOT NULL), 0) FROM daily_indicators WHERE trade_date = ? """, (indicator_date,), ).fetchone() else: valuation_rows, valuation_available = 0, 0 dividend_years = connection.execute( """ SELECT COUNT(DISTINCT substr(trade_date, 1, 4)) FROM daily_indicators WHERE trade_date <= ? AND trade_date >= ? """, (end_date, dividend_start), ).fetchone()[0] fundamental_rows = connection.execute( """ SELECT COUNT(*) FROM fundamental_indicators fi INNER JOIN ( SELECT ts_code, MAX(ann_date || ':' || end_date) AS latest_key FROM fundamental_indicators WHERE ann_date = '' OR ann_date <= ? GROUP BY ts_code ) latest ON latest.ts_code = fi.ts_code AND latest.latest_key = (fi.ann_date || ':' || fi.end_date) """, (end_date,), ).fetchone()[0] moneyflow_dates = connection.execute( """ SELECT COUNT(DISTINCT trade_date) FROM moneyflow_daily WHERE trade_date IN ( SELECT DISTINCT trade_date FROM daily_bars WHERE trade_date <= ? ORDER BY trade_date DESC LIMIT 5 ) """, (end_date,), ).fetchone()[0] earnings_rows = connection.execute( """ SELECT COUNT(*) FROM earnings_events WHERE ann_date <= ? AND ann_date >= replace(date(?, '-45 day'), '-', '') """, (end_date, f"{end_date[:4]}-{end_date[4:6]}-{end_date[6:8]}"), ).fetchone()[0] popularity_rows = connection.execute( "SELECT COUNT(*) FROM popularity_factors WHERE trade_date = ?", (end_date,), ).fetchone()[0] institution_rows = connection.execute( "SELECT COUNT(*) FROM lhb_institution_daily WHERE trade_date = ?", (end_date,), ).fetchone()[0] return { "market": bool(market), "auction": bool(auction), "benchmark": int(benchmark_rows or 0) >= 60, "benchmark_rows": int(benchmark_rows or 0), "valuation": bool(valuation_available), "fundamental": int(fundamental_rows or 0) >= 100, "dividend_history": int(dividend_years or 0) >= 4, "valuation_rows": int(valuation_rows or 0), "fundamental_rows": int(fundamental_rows or 0), "dividend_years": int(dividend_years or 0), "moneyflow_history": int(moneyflow_dates or 0) >= 5, "moneyflow_dates": int(moneyflow_dates or 0), "earnings_events": int(earnings_rows or 0) > 0, "earnings_event_rows": int(earnings_rows or 0), "popularity": int(popularity_rows or 0) > 0, "popularity_rows": int(popularity_rows or 0), "institutions": int(institution_rows or 0) > 0, "institution_rows": int(institution_rows or 0), } def load_factor_data(self, end_date: str, limit_dates: int = 80) -> dict[str, Any]: dates = self.factor_dates(end_date, limit_dates) if not dates: return { "dates": [], "bars": [], "master": [], "indicators": [], "indicator_history": [], "indicator_series": [], "fundamentals": [], "moneyflow": [], "moneyflow_history": [], "auction": [], "benchmarks": [], "fundamental_history": [], "earnings_events": [], "popularity": [], "institutions": [], } placeholders = ",".join("?" for _ in dates) with self.connect() as connection: bars = connection.execute( f"SELECT * FROM daily_bars WHERE trade_date IN ({placeholders}) ORDER BY trade_date, ts_code", dates, ).fetchall() master = connection.execute("SELECT * FROM stock_master").fetchall() indicators = connection.execute( """ SELECT * FROM daily_indicators WHERE trade_date = ( SELECT MAX(trade_date) FROM daily_indicators WHERE trade_date <= ? ) """, (end_date,), ).fetchall() indicator_history = connection.execute( """ SELECT di.* FROM daily_indicators di INNER JOIN ( SELECT ts_code, substr(trade_date, 1, 4) AS year_key, MAX(trade_date) AS max_date FROM daily_indicators WHERE trade_date <= ? AND trade_date >= ? GROUP BY ts_code, substr(trade_date, 1, 4) ) latest ON latest.ts_code = di.ts_code AND latest.max_date = di.trade_date ORDER BY di.trade_date, di.ts_code """, (end_date, str(max(0, int(end_date[:4] or 0) - 5)) + "0101"), ).fetchall() indicator_series = connection.execute( f""" SELECT trade_date, ts_code, turnover_rate, volume_ratio, total_mv, circ_mv, pe_ttm, pb, ps_ttm, dv_ttm FROM daily_indicators WHERE trade_date IN ({placeholders}) ORDER BY trade_date, ts_code """, dates, ).fetchall() fundamentals = connection.execute( """ SELECT fi.* FROM fundamental_indicators fi INNER JOIN ( SELECT ts_code, MAX(ann_date || ':' || end_date) AS latest_key FROM fundamental_indicators WHERE ann_date = '' OR ann_date <= ? GROUP BY ts_code ) latest ON latest.ts_code = fi.ts_code AND latest.latest_key = (fi.ann_date || ':' || fi.end_date) """, (end_date,), ).fetchall() fundamental_history = connection.execute( """ SELECT * FROM fundamental_indicators WHERE ann_date = '' OR ann_date <= ? ORDER BY ann_date, end_date, ts_code """, (end_date,), ).fetchall() moneyflow = connection.execute( """ SELECT * FROM moneyflow_daily WHERE trade_date = ( SELECT MAX(trade_date) FROM moneyflow_daily WHERE trade_date <= ? ) """, (end_date,), ).fetchall() flow_dates = dates[-min(5, len(dates)):] flow_placeholders = ",".join("?" for _ in flow_dates) moneyflow_history = connection.execute( f""" SELECT * FROM moneyflow_daily WHERE trade_date IN ({flow_placeholders}) ORDER BY trade_date, ts_code """, flow_dates, ).fetchall() auction = connection.execute( """ SELECT * FROM auction_factors WHERE trade_date = ( SELECT MAX(trade_date) FROM auction_factors WHERE trade_date <= ? ) """, (end_date,), ).fetchall() benchmarks = connection.execute( f""" SELECT * FROM benchmark_bars WHERE ts_code = '000300.SH' AND trade_date IN ({placeholders}) ORDER BY trade_date """, dates, ).fetchall() earnings_events = connection.execute( """ SELECT * FROM earnings_events WHERE ann_date <= ? ORDER BY ann_date, end_date, ts_code """, (end_date,), ).fetchall() popularity = connection.execute( "SELECT * FROM popularity_factors WHERE trade_date = ? ORDER BY ts_code", (end_date,), ).fetchall() institutions = connection.execute( "SELECT * FROM lhb_institution_daily WHERE trade_date = ? ORDER BY ts_code", (end_date,), ).fetchall() return { "dates": dates, "bars": [dict(row) for row in bars], "master": [dict(row) for row in master], "indicators": [dict(row) for row in indicators], "indicator_history": [dict(row) for row in indicator_history], "indicator_series": [dict(row) for row in indicator_series], "fundamentals": [dict(row) for row in fundamentals], "fundamental_history": [dict(row) for row in fundamental_history], "moneyflow": [dict(row) for row in moneyflow], "moneyflow_history": [dict(row) for row in moneyflow_history], "auction": [dict(row) for row in auction], "benchmarks": [dict(row) for row in benchmarks], "earnings_events": [dict(row) for row in earnings_events], "popularity": [dict(row) for row in popularity], "institutions": [dict(row) for row in institutions], } def snapshot_summaries(self, end_date: str, limit: int = 10) -> list[dict[str, Any]]: try: from sentiment_engine import build_sentiment_history except ModuleNotFoundError: from .sentiment_engine import build_sentiment_history series = build_sentiment_history(self.list_snapshot_payloads(end_date, 260)) return [ { "trade_date": row["trade_date"], "sentiment_score": row["score"], "seal_rate": row["seal_rate"], "limit_up_count": row["limit_up_count"], "limit_down_count": row["limit_down_count"], "broken_count": row["broken_count"], "up_count": row["up_count"], "down_count": row["down_count"], "amount_billion": row["amount_billion"], } for row in series[-limit:] ] def save_screener_strategy( self, user_id: int | None, name: str, description: str, regimes: list[str], formula: dict[str, Any], builtin: bool = False, strategy_id: int | None = None, ) -> int: now = datetime.now().astimezone().isoformat(timespec="seconds") regimes_json = json.dumps(regimes, ensure_ascii=False) formula_json = json.dumps(formula, ensure_ascii=False, separators=(",", ":")) with self.connect() as connection: if strategy_id: if builtin: cursor = connection.execute( """ UPDATE screener_strategies SET name=?, description=?, regimes=?, formula=?, builtin=1, user_id=NULL, updated_at=? WHERE id=? AND builtin=1 """, (name, description, regimes_json, formula_json, now, strategy_id), ) else: cursor = connection.execute( """ UPDATE screener_strategies SET name=?, description=?, regimes=?, formula=?, updated_at=? WHERE id=? AND builtin=0 AND user_id=? """, (name, description, regimes_json, formula_json, now, strategy_id, int(user_id or 0)), ) if cursor.rowcount == 0: raise ValueError("选股策略不存在。") return strategy_id cursor = connection.execute( """ INSERT INTO screener_strategies (user_id, name, description, regimes, formula, builtin, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?) """, (None if builtin else int(user_id or 0), name, description, regimes_json, formula_json, int(builtin), now, now), ) return int(cursor.lastrowid) def list_screener_strategies(self, user_id: int | None = None) -> list[dict[str, Any]]: with self.connect() as connection: if user_id is None: rows = connection.execute( "SELECT * FROM screener_strategies WHERE builtin = 1 ORDER BY updated_at DESC, id" ).fetchall() else: rows = connection.execute( """ SELECT * FROM screener_strategies WHERE builtin = 1 OR user_id = ? ORDER BY builtin DESC, updated_at DESC, id """, (int(user_id),), ).fetchall() result = [] for row in rows: item = dict(row) item["regimes"] = json.loads(item["regimes"]) item["formula"] = json.loads(item["formula"]) item["builtin"] = bool(item["builtin"]) result.append(item) return result def delete_screener_strategy(self, user_id: int, strategy_id: int) -> bool: with self.connect() as connection: row = connection.execute( "SELECT builtin, user_id FROM screener_strategies WHERE id = ?", (strategy_id,), ).fetchone() if not row: raise ValueError("选股策略不存在。") if bool(row["builtin"]): raise ValueError("内置策略不能删除。") if int(row["user_id"] or 0) != int(user_id): raise ValueError("无权删除其他账号的策略。") cursor = connection.execute( "DELETE FROM screener_strategies WHERE id = ? AND builtin = 0 AND user_id = ?", (strategy_id, int(user_id)), ) return cursor.rowcount > 0 def save_screener_run( self, user_id: int, trade_date: str, regime: str, strategy_name: str, formula: dict[str, Any], result: dict[str, Any], mode: str = "smart", ) -> int: normalized_mode = mode if mode in {"smart", "curated", "quant"} else "smart" now = datetime.now().astimezone().isoformat(timespec="seconds") with self.connect() as connection: cursor = connection.execute( """ INSERT INTO screener_runs (user_id, trade_date, regime, mode, strategy_name, formula, result, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?) """, (None if int(user_id) == 0 else int(user_id), trade_date, regime, normalized_mode, strategy_name, json.dumps(formula, ensure_ascii=False, separators=(",", ":")), json.dumps(result, ensure_ascii=False, separators=(",", ":")), now), ) return int(cursor.lastrowid) @staticmethod def _screener_run_payload(row: sqlite3.Row) -> dict[str, Any] | None: try: result = json.loads(row["result"]) except json.JSONDecodeError: return None result.setdefault("meta", {}).update( { "run_id": int(row["id"]), "trade_date": str(row["trade_date"] or ""), "regime": str(row["regime"] or ""), "mode": str(row["mode"] or "smart"), "strategy_name": str(row["strategy_name"] or ""), "created_at": row["created_at"], } ) return result def latest_screener_run( self, user_id: int, trade_date: str, mode: str = "", ) -> dict[str, Any] | None: owner_clause = "user_id IS NULL" if int(user_id) == 0 else "user_id = ?" parameters: tuple[Any, ...] = () if int(user_id) == 0 else (int(user_id),) parameters += (trade_date,) mode_clause = "" if mode in {"smart", "curated", "quant"}: mode_clause = " AND mode = ?" parameters += (mode,) with self.connect() as connection: row = connection.execute( f""" SELECT id, trade_date, regime, mode, strategy_name, result, created_at FROM screener_runs WHERE {owner_clause} AND trade_date <= ?{mode_clause} ORDER BY id DESC LIMIT 1 """, parameters, ).fetchone() return self._screener_run_payload(row) if row else None def latest_screener_runs(self, user_id: int, trade_date: str) -> dict[str, dict[str, Any]]: owner_clause = "user_id IS NULL" if int(user_id) == 0 else "user_id = ?" parameters: tuple[Any, ...] = () if int(user_id) == 0 else (int(user_id),) parameters += (trade_date,) with self.connect() as connection: rows = connection.execute( f""" SELECT runs.id, runs.trade_date, runs.regime, runs.mode, runs.strategy_name, runs.result, runs.created_at FROM screener_runs runs INNER JOIN ( SELECT mode, MAX(id) AS id FROM screener_runs WHERE {owner_clause} AND trade_date <= ? GROUP BY mode ) latest ON latest.id = runs.id """, parameters, ).fetchall() results: dict[str, dict[str, Any]] = {} for row in rows: mode = str(row["mode"] or "smart") payload = self._screener_run_payload(row) if mode in {"smart", "curated", "quant"} and payload: results[mode] = payload return results def latest_screener_context_runs( self, user_id: int, trade_date: str, limit: int = 60, ) -> list[dict[str, Any]]: safe_limit = max(1, min(120, int(limit))) owner_clause = "user_id IS NULL" if int(user_id) == 0 else "user_id = ?" parameters: tuple[Any, ...] = () if int(user_id) == 0 else (int(user_id),) parameters += (trade_date, safe_limit) with self.connect() as connection: rows = connection.execute( f""" WITH ranked AS ( SELECT id, trade_date, regime, mode, strategy_name, result, created_at, ROW_NUMBER() OVER ( PARTITION BY mode, CASE WHEN mode = 'smart' THEN regime ELSE '' END, CASE WHEN mode IN ('smart', 'curated') THEN strategy_name ELSE '' END ORDER BY id DESC ) AS context_rank FROM screener_runs WHERE {owner_clause} AND trade_date <= ? ) SELECT id, trade_date, regime, mode, strategy_name, result, created_at FROM ranked WHERE context_rank = 1 ORDER BY id DESC LIMIT ? """, parameters, ).fetchall() return [ payload for row in rows if (payload := self._screener_run_payload(row)) is not None ] def screener_runs_for_date( self, user_id: int, trade_date: str, limit: int = 80, ) -> list[dict[str, Any]]: safe_limit = max(1, min(160, int(limit))) owner_clause = "user_id IS NULL" if int(user_id) == 0 else "user_id = ?" parameters: tuple[Any, ...] = () if int(user_id) == 0 else (int(user_id),) parameters += (trade_date, safe_limit) with self.connect() as connection: rows = connection.execute( f""" SELECT id, trade_date, regime, mode, strategy_name, result, created_at FROM screener_runs WHERE {owner_clause} AND trade_date = ? ORDER BY id DESC LIMIT ? """, parameters, ).fetchall() result = [] seen: set[tuple[str, str, str]] = set() for row in rows: key = ( str(row["mode"] or "smart"), str(row["regime"] or ""), str(row["strategy_name"] or ""), ) if key in seen: continue seen.add(key) payload = self._screener_run_payload(row) if payload is not None: result.append(payload) return result def get_screener_run(self, user_id: int, run_id: int) -> dict[str, Any] | None: owner_clause = "user_id IS NULL" if int(user_id) == 0 else "user_id = ?" parameters: tuple[Any, ...] = (int(run_id),) if int(user_id) != 0: parameters += (int(user_id),) with self.connect() as connection: row = connection.execute( f""" SELECT id, trade_date, regime, mode, strategy_name, result, created_at FROM screener_runs WHERE id = ? AND {owner_clause} """, parameters, ).fetchone() if not row: return None result = self._screener_run_payload(row) if result is None: return None result.setdefault("meta", {}).update( { "run_id": int(row["id"]), "trade_date": row["trade_date"], "mode": str(row["mode"] or "smart"), "created_at": row["created_at"], } ) result["strategy_name"] = row["strategy_name"] result["regime"] = row["regime"] return result def save_strategy_tracks( self, user_id: int, run_id: int, selection_date: str, strategy_name: str, candidates: list[dict[str, Any]], ) -> int: now = datetime.now().astimezone().isoformat(timespec="seconds") values = [] for item in candidates: ts_code = str(item.get("ts_code") or "").strip() code = str(item.get("code") or ts_code.split(".")[0]).strip() entry_price = float(item.get("price") or 0) if not ts_code or not code or entry_price <= 0: continue values.append( ( int(user_id), int(run_id), selection_date, strategy_name, ts_code, code, str(item.get("name") or "--"), str(item.get("sector") or "其他"), entry_price, now, now, ) ) with self.connect() as connection: connection.executemany( """ INSERT INTO strategy_tracks (user_id, run_id, selection_date, strategy_name, ts_code, code, name, sector, entry_price, created_at, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(user_id, run_id, ts_code) DO UPDATE SET name=excluded.name, sector=excluded.sector, entry_price=excluded.entry_price, updated_at=excluded.updated_at """, values, ) return len(values) def list_strategy_tracks(self, user_id: int, limit_batches: int = 12) -> list[dict[str, Any]]: limit_batches = max(1, min(50, int(limit_batches))) with self.connect() as connection: rows = connection.execute( """ SELECT * FROM strategy_tracks WHERE user_id = ? AND run_id IN ( SELECT run_id FROM strategy_tracks WHERE user_id = ? GROUP BY run_id ORDER BY run_id DESC LIMIT ? ) ORDER BY run_id DESC, id """, (int(user_id), int(user_id), limit_batches), ).fetchall() return [dict(row) for row in rows] def delete_strategy_track(self, user_id: int, track_id: int) -> bool: with self.connect() as connection: cursor = connection.execute( "DELETE FROM strategy_tracks WHERE id = ? AND user_id = ?", (int(track_id), int(user_id)), ) return cursor.rowcount > 0 def load_tracking_bars( self, targets: list[tuple[str, str]], limit: int = 5 ) -> dict[tuple[str, str], list[dict[str, Any]]]: unique_targets = set(targets) if not unique_targets: return {} codes = sorted({ts_code for ts_code, _ in unique_targets}) earliest_date = min(selection_date for _, selection_date in unique_targets) placeholders = ",".join("?" for _ in codes) with self.connect() as connection: rows = connection.execute( f""" SELECT ts_code, trade_date, open, high, low, close FROM daily_bars WHERE ts_code IN ({placeholders}) AND trade_date > ? ORDER BY ts_code, trade_date """, [*codes, earliest_date], ).fetchall() by_code: dict[str, list[dict[str, Any]]] = {} for row in rows: item = dict(row) by_code.setdefault(str(item["ts_code"]), []).append(item) row_limit = max(1, min(20, int(limit))) return { (ts_code, selection_date): [ row for row in by_code.get(ts_code, []) if row["trade_date"] > selection_date ][:row_limit] for ts_code, selection_date in unique_targets }