from __future__ import annotations import json from datetime import datetime from typing import Any class MarketRepositoryMixin: def upsert_stock_master(self, rows: list[dict[str, Any]]) -> int: now = datetime.now().astimezone().isoformat(timespec="seconds") values = [ ( row.get("ts_code", ""), str(row.get("ts_code", "")).split(".")[0], row.get("name") or "--", row.get("industry") or "", row.get("market") or "", str(row.get("list_date") or ""), now, ) for row in rows if row.get("ts_code") ] with self.connect() as connection: connection.executemany( """ INSERT INTO stock_master (ts_code, code, name, industry, market, list_date, updated_at) VALUES (?, ?, ?, ?, ?, ?, ?) ON CONFLICT(ts_code) DO UPDATE SET code=excluded.code, name=excluded.name, industry=excluded.industry, market=excluded.market, list_date=excluded.list_date, updated_at=excluded.updated_at """, values, ) return len(values) def list_stock_master(self) -> list[dict[str, Any]]: with self.connect() as connection: rows = connection.execute( "SELECT ts_code, code, name, industry, market, list_date FROM stock_master" ).fetchall() return [dict(row) for row in rows] def upsert_daily_bars(self, rows: list[dict[str, Any]]) -> int: values = [ ( str(row.get("trade_date") or ""), row.get("ts_code", ""), float(row.get("open") or 0), float(row.get("high") or 0), float(row.get("low") or 0), float(row.get("close") or 0), float(row.get("pct_chg") or 0), float(row.get("vol") or 0), float(row.get("amount") 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 daily_bars (trade_date, ts_code, open, high, low, close, pct_chg, vol, amount) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) ON CONFLICT(trade_date, ts_code) DO UPDATE SET open=excluded.open, high=excluded.high, low=excluded.low, close=excluded.close, pct_chg=excluded.pct_chg, vol=excluded.vol, amount=excluded.amount """, values, ) return len(values) def daily_bars_for_date(self, trade_date: str) -> list[dict[str, Any]]: with self.connect() as connection: rows = connection.execute( "SELECT * FROM daily_bars WHERE trade_date = ? ORDER BY ts_code", (trade_date,), ).fetchall() return [dict(row) for row in rows] def get_snapshot(self, trade_date: str) -> dict[str, Any] | None: with self.connect() as connection: row = connection.execute( "SELECT payload FROM dashboard_snapshots WHERE trade_date = ?", (trade_date,), ).fetchone() if not row: return None try: return json.loads(row["payload"]) except json.JSONDecodeError: return None def get_latest_real_snapshot( self, trade_date: str, strictly_before: bool = False ) -> dict[str, Any] | None: operator = "<" if strictly_before else "<=" with self.connect() as connection: row = connection.execute( f""" SELECT payload FROM dashboard_snapshots WHERE trade_date {operator} ? AND source != 'demo' ORDER BY trade_date DESC LIMIT 1 """, (trade_date,), ).fetchone() if not row: return None try: return json.loads(row["payload"]) except json.JSONDecodeError: return None def save_snapshot(self, trade_date: str, source: str, payload: dict[str, Any]) -> None: updated_at = datetime.now().astimezone().isoformat(timespec="seconds") record_count = sum( len(payload.get(key) or []) for key in ("limits", "broken", "down_limits", "yesterday_limits") ) content = json.dumps(payload, ensure_ascii=False, separators=(",", ":")) with self.connect() as connection: connection.execute( """ INSERT INTO dashboard_snapshots (trade_date, source, payload, record_count, updated_at) VALUES (?, ?, ?, ?, ?) ON CONFLICT(trade_date) DO UPDATE SET source = excluded.source, payload = excluded.payload, record_count = excluded.record_count, updated_at = excluded.updated_at """, (trade_date, source, content, record_count, updated_at), ) def get_data_snapshot(self, kind: str, cache_key: str) -> dict[str, Any] | None: with self.connect() as connection: row = connection.execute( "SELECT payload FROM data_snapshots WHERE kind = ? AND cache_key = ?", (kind, cache_key), ).fetchone() if not row: return None try: return json.loads(row["payload"]) except json.JSONDecodeError: return None def get_latest_data_snapshot( self, kind: str, cache_key_prefix: str, maximum_cache_key: str, exclude_source: str = "", ) -> dict[str, Any] | None: source_clause = " AND source != ?" if exclude_source else "" parameters: list[Any] = [kind, f"{cache_key_prefix}%", maximum_cache_key] if exclude_source: parameters.append(exclude_source) with self.connect() as connection: row = connection.execute( f""" SELECT payload FROM data_snapshots WHERE kind = ? AND cache_key LIKE ? AND cache_key <= ?{source_clause} ORDER BY cache_key DESC LIMIT 1 """, parameters, ).fetchone() if not row: return None try: return json.loads(row["payload"]) except json.JSONDecodeError: return None def save_data_snapshot( self, kind: str, cache_key: str, source: str, payload: dict[str, Any] ) -> None: updated_at = datetime.now().astimezone().isoformat(timespec="seconds") content = json.dumps(payload, ensure_ascii=False, separators=(",", ":")) with self.connect() as connection: connection.execute( """ INSERT INTO data_snapshots (kind, cache_key, source, payload, updated_at) VALUES (?, ?, ?, ?, ?) ON CONFLICT(kind, cache_key) DO UPDATE SET source = excluded.source, payload = excluded.payload, updated_at = excluded.updated_at """, (kind, cache_key, source, content, updated_at), ) def search_stock_master(self, query: str, limit: int = 12) -> list[dict[str, Any]]: text = str(query or "").strip() if not text: return [] escaped = text.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") with self.connect() as connection: rows = connection.execute( """ SELECT ts_code, code, name, industry, market, list_date FROM stock_master WHERE code = ? OR name = ? OR name LIKE ? ESCAPE '\\' ORDER BY CASE WHEN code = ? THEN 0 WHEN name = ? THEN 1 ELSE 2 END, list_date DESC, code LIMIT ? """, (text, text, f"%{escaped}%", text, text, max(1, min(30, int(limit)))), ).fetchall() return [dict(row) for row in rows] def list_snapshot_payloads(self, end_date: str, limit: int = 260) -> list[dict[str, Any]]: with self.connect() as connection: rows = connection.execute( """ SELECT trade_date, payload FROM dashboard_snapshots WHERE trade_date <= ? ORDER BY trade_date DESC LIMIT ? """, (end_date, limit), ).fetchall() result: list[dict[str, Any]] = [] for row in reversed(rows): try: payload = json.loads(row["payload"]) except json.JSONDecodeError: continue payload["_snapshot_date"] = row["trade_date"] result.append(payload) return result def start_sync(self, trade_date: str, source: str) -> int: started_at = datetime.now().astimezone().isoformat(timespec="seconds") with self.connect() as connection: cursor = connection.execute( """ INSERT INTO sync_runs (trade_date, source, status, started_at) VALUES (?, ?, 'running', ?) """, (trade_date, source, started_at), ) return int(cursor.lastrowid) def finish_sync( self, sync_id: int, status: str, record_count: int = 0, message: str = "", source: str | None = None, ) -> None: finished_at = datetime.now().astimezone().isoformat(timespec="seconds") with self.connect() as connection: connection.execute( """ UPDATE sync_runs SET status = ?, finished_at = ?, record_count = ?, message = ?, source = COALESCE(?, source) WHERE id = ? """, (status, finished_at, record_count, message[:1000], source, sync_id), ) def status(self) -> dict[str, Any]: with self.connect() as connection: last_sync = connection.execute( """ SELECT id, trade_date, source, status, started_at, finished_at, record_count, message FROM sync_runs ORDER BY id DESC LIMIT 1 """ ).fetchone() snapshot_stats = connection.execute( """ SELECT COUNT(*) AS dates, COALESCE(SUM(record_count), 0) AS records, MAX(updated_at) AS updated_at FROM dashboard_snapshots """ ).fetchone() watchlist_count = connection.execute("SELECT COUNT(*) FROM watchlist").fetchone()[0] note_count = connection.execute("SELECT COUNT(*) FROM review_notes").fetchone()[0] return { "database": str(self.path.name), "snapshot_dates": int(snapshot_stats["dates"]), "snapshot_records": int(snapshot_stats["records"]), "updated_at": snapshot_stats["updated_at"], "last_sync": dict(last_sync) if last_sync else None, "watchlist_count": int(watchlist_count), "note_count": int(note_count), }