Files
kis_bot/kis_trader/backtest/bt_candle_source.py
Your Name 0780b2cdd0 feat: Enhance Optuna integration and logging for backtesting framework
Changes:
- Added new API endpoints for continuing and confirming Optuna jobs, allowing for better management of ongoing studies.
- Introduced detailed logging for tick feed tracking and order book processing, improving traceability of vendor performance during backtests.
- Updated database schema to include new fields for managing Optuna study results, enhancing the ability to track study progress and outcomes.
- Refactored existing functions to utilize the new logging and tracking features, ensuring consistency across the backtesting framework.

Impact:
- These enhancements improve the robustness and transparency of the Optuna backtesting process, facilitating better analysis and optimization of trading strategies.
2026-08-21 19:05:23 +09:00

439 lines
14 KiB
Python

#!/usr/bin/env python3
"""
백테·Optuna ws_candles 소스 선택 — 실매 get_candles 와 동일.
읽기: 틱 메인 WS 한 칸, 없으면 kiwoom+rest 구멍, 그다음 kiwoom+rollup.
CANDLE_SOURCE=kis|kiwoom 이면 그 메인(+구멍). 빈값이면 LIVE_TICK_PROVIDER.
"""
from __future__ import annotations
import logging
import time
from collections import defaultdict
from typing import Any, Dict, List, Optional, Sequence, Tuple
from kis_trader.ws.candle_series import (
dedupe_by_read_pairs,
live_read_label,
live_read_pairs,
normalize_source_channel,
)
ReadPair = Tuple[str, str]
logger = logging.getLogger("bt_candle_source")
# 레거시 호환 이름 (실제 필터는 live_read_pairs)
BT_WS_CANDLE_SOURCES: Tuple[str, ...] = ("kiwoom", "kis", "rest", "rollup_1m")
def resolve_bt_candle_source_override() -> str:
"""'' = LIVE 병합, 'kis'|'kiwoom' = 그 메인(+키움 REST 구멍)."""
from kis_trader.ws.candle_series import _provider_and_override
_provider, override = _provider_and_override()
return override
def live_candle_source_order() -> Tuple[str, ...]:
"""호환: source 문자열만. 실제 조회는 resolve_bt_read_pairs."""
return tuple(dict.fromkeys(p[0] for p in live_read_pairs()))
def resolve_bt_read_pairs() -> Tuple[ReadPair, ...]:
return live_read_pairs()
def resolve_bt_candle_source_order() -> Tuple[str, ...]:
"""호환 래퍼 — 신규 코드는 resolve_bt_read_pairs 사용."""
return live_candle_source_order()
def resolve_bt_candle_source_label() -> str:
return live_read_label()
def dedupe_candle_rows(
rows: Sequence[Dict[str, Any]],
source_order: Optional[Sequence[Any]] = None,
) -> List[Dict[str, Any]]:
"""candle_time 1행. source_order 가 쌍이면 그대로, 아니면 live_read_pairs."""
pairs: Optional[Sequence[ReadPair]] = None
if source_order:
first = source_order[0]
if isinstance(first, (tuple, list)) and len(first) >= 2:
pairs = tuple((str(a), str(b)) for a, b in source_order) # type: ignore[misc]
else:
# 레거시 source 문자열 → 정규화 후 순위
mapped: List[ReadPair] = []
for s in source_order:
mapped.append(normalize_source_channel(str(s), ""))
pairs = tuple(mapped)
return dedupe_by_read_pairs(rows, pairs)
def _pairs_in_sql(pairs: Sequence[ReadPair]) -> Tuple[str, List[str]]:
parts: List[str] = []
params: List[str] = []
for src, ch in pairs:
parts.append("(source=%s AND channel=%s)")
params.extend([src, ch])
return " AND (" + " OR ".join(parts) + ")", params
def _tick_time_bounds(start_key: str, end_key: str) -> Tuple[str, str]:
sk = str(start_key or "").strip()
ek = str(end_key or "").strip()
if len(sk) == 12:
sk = sk + "00"
if len(ek) == 12:
ek = ek + "59"
return sk, ek
def _ws_ticks_table(market: Optional[str]) -> str:
mk = (market or "").strip().upper()
if mk == "US":
return "ws_ticks_us"
return "ws_ticks"
def _ticks_for_bar_garbage(
db,
code: str,
start_key: str,
end_key: str,
*,
market: Optional[str] = None,
) -> Optional[List[Dict[str, Any]]]:
from kis_trader.engine.feed_fallback import candle_garbage_fallback_enabled
if not candle_garbage_fallback_enabled():
return None
sk, ek = _tick_time_bounds(start_key, end_key)
mk = (market or "").strip().upper()
table = _ws_ticks_table(mk if mk in ("US", "KR") else "KR")
try:
if mk in ("US", "KR"):
rows = db.conn.execute(
f"SELECT tick_time, source, tick_time_raw FROM {table} "
"WHERE market=%s AND code=%s AND tick_time >= %s AND tick_time <= %s",
(mk, str(code).strip(), sk, ek),
).fetchall()
else:
rows = db.conn.execute(
f"SELECT tick_time, source, tick_time_raw FROM {table} "
"WHERE market=%s AND code=%s AND tick_time >= %s AND tick_time <= %s",
("KR", str(code).strip(), sk, ek),
).fetchall()
return [dict(r) for r in (rows or [])]
except Exception:
return None
def _load_ticks_by_code_bulk(
db,
start_key: str,
end_key: str,
*,
market: Optional[str] = None,
codes_filter: Optional[Sequence[str]] = None,
) -> Optional[Dict[str, List[Dict[str, Any]]]]:
"""쓰레기 검사용 틱 — 기간 1~2쿼리. OFF 면 None (봉만 dedupe)."""
from kis_trader.engine.feed_fallback import candle_garbage_fallback_enabled
if not candle_garbage_fallback_enabled():
return None
sk, ek = _tick_time_bounds(start_key, end_key)
mk = (market or "").strip().upper()
out: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
want: List[str] = []
if codes_filter:
want = [str(c).strip() for c in codes_filter if str(c).strip()]
in_sql = ""
in_params: List[str] = []
if want:
in_sql = " AND code IN (" + ",".join(["%s"] * len(want)) + ")"
in_params = want
def _pull(table: str, mkt: str) -> int:
rows = db.conn.execute(
f"SELECT code, tick_time, source, tick_time_raw FROM {table} "
"WHERE market=%s AND tick_time >= %s AND tick_time <= %s"
+ in_sql,
(mkt, sk, ek, *in_params),
).fetchall()
n = 0
for r in rows or []:
code = str(r.get("code") or "").strip()
if not code:
continue
out[code].append(dict(r))
n += 1
return n
try:
n_kr = n_us = 0
if mk == "US":
n_us = _pull("ws_ticks_us", "US")
elif mk == "KR":
n_kr = _pull("ws_ticks", "KR")
else:
n_kr = _pull("ws_ticks", "KR")
try:
n_us = _pull("ws_ticks_us", "US")
except Exception:
n_us = 0
logger.info("📥 틱 bulk 쓰레기검사용: KR=%s US=%s 종목=%s", n_kr, n_us, len(out))
return dict(out)
except Exception as exc:
logger.warning("틱 bulk 스킵(봉만): %s", exc)
return None
def list_ws_candle_codes(
db,
timeframe: int,
start_key: str,
end_key: str,
*,
market: Optional[str] = None,
) -> List[str]:
"""기간 내 종목 코드 — 선택된 (source,channel) 기준 DISTINCT."""
pairs = resolve_bt_read_pairs()
src_sql, src_params = _pairs_in_sql(pairs)
mk = (market or "").strip().upper()
if mk:
rows = db.conn.execute(
"SELECT DISTINCT code FROM ws_candles WHERE timeframe=%s AND market=%s "
"AND candle_time >= %s AND candle_time <= %s"
+ src_sql
+ " ORDER BY code",
[int(timeframe), mk, start_key, end_key, *src_params],
).fetchall()
else:
rows = db.conn.execute(
"SELECT DISTINCT code FROM ws_candles WHERE timeframe=%s "
"AND candle_time >= %s AND candle_time <= %s"
+ src_sql
+ " ORDER BY code",
[int(timeframe), start_key, end_key, *src_params],
).fetchall()
return [r["code"] for r in rows]
def fetch_ws_candles_for_code(
db,
code: str,
timeframe: int,
start_key: str,
end_key: str,
*,
extra_select: str = "",
peak_sel: str = "",
market: Optional[str] = None,
confirmed_only: bool = True,
) -> List[Dict[str, Any]]:
"""단일 종목·기간 봉 로드 — 메인 WS 우선, 구멍 kiwoom+rest(+rollup)."""
pairs = resolve_bt_read_pairs()
confirmed_sql = " AND is_confirmed=1" if confirmed_only else ""
mk = (market or "").strip().upper()
src_sql, src_params = _pairs_in_sql(pairs)
cols = (
f"candle_time, open, high, low, close, volume, source, channel"
f"{peak_sel}{extra_select}"
)
if mk:
rows = db.conn.execute(
f"SELECT {cols} FROM ws_candles "
"WHERE timeframe=%s AND code=%s AND market=%s "
"AND candle_time >= %s AND candle_time <= %s"
+ confirmed_sql
+ src_sql
+ " ORDER BY candle_time ASC",
[int(timeframe), code, mk, start_key, end_key, *src_params],
).fetchall()
else:
rows = db.conn.execute(
f"SELECT {cols} FROM ws_candles "
"WHERE timeframe=%s AND code=%s "
"AND candle_time >= %s AND candle_time <= %s"
+ confirmed_sql
+ src_sql
+ " ORDER BY candle_time ASC",
[int(timeframe), code, start_key, end_key, *src_params],
).fetchall()
ticks = _ticks_for_bar_garbage(
db, str(code), start_key, end_key, market=mk or None,
)
return dedupe_by_read_pairs(
[dict(r) for r in rows], pairs, ticks=ticks, tf_min=int(timeframe),
missing_policy="hole", live_cover=False,
)
def fetch_ws_candles_by_code_bulk(
db,
timeframe: int,
start_key: str,
end_key: str,
*,
extra_select: str = "",
peak_sel: str = "",
market: Optional[str] = None,
confirmed_only: bool = True,
codes_filter: Optional[Sequence[str]] = None,
) -> Dict[str, List[Dict[str, Any]]]:
"""기간 전체 봉 1쿼리 + 틱 1~2쿼리 후 종목별 기존 쓰레기 검사.
웜업용 ``fetch_ws_candles_for_code`` / ``fetch_ws_candles_warmup_before`` 는 그대로.
"""
t0 = time.perf_counter()
pairs = resolve_bt_read_pairs()
confirmed_sql = " AND is_confirmed=1" if confirmed_only else ""
mk = (market or "").strip().upper()
src_sql, src_params = _pairs_in_sql(pairs)
cols = (
f"code, candle_time, open, high, low, close, volume, source, channel"
f"{peak_sel}{extra_select}"
)
want: List[str] = []
if codes_filter:
want = [
str(c).strip()
for c in codes_filter
if str(c).strip()
]
in_sql = ""
in_params: List[str] = []
if want:
in_sql = " AND code IN (" + ",".join(["%s"] * len(want)) + ")"
in_params = want
params: List[Any] = [int(timeframe)]
mk_sql = ""
if mk:
mk_sql = " AND market=%s"
params.append(mk)
params.extend([start_key, end_key, *src_params, *in_params])
rows = db.conn.execute(
f"SELECT {cols} FROM ws_candles "
"WHERE timeframe=%s"
+ mk_sql
+ " AND candle_time >= %s AND candle_time <= %s"
+ confirmed_sql
+ src_sql
+ in_sql
+ " ORDER BY code ASC, candle_time ASC",
params,
).fetchall()
by_code: Dict[str, List[Dict[str, Any]]] = defaultdict(list)
for r in rows or []:
code = str(r.get("code") or "").strip()
if not code:
continue
by_code[code].append(dict(r))
ticks_map = _load_ticks_by_code_bulk(
db, start_key, end_key, market=mk or None, codes_filter=want or None,
)
out: Dict[str, List[Dict[str, Any]]] = {}
stats: Dict[str, Any] = {
"slots": 0, "picked": 0, "hole": 0, "garbage_skip": 0,
"pick_by": {}, "garbage_by": {}, "raw_by": {},
}
for code, raw in by_code.items():
ticks = None if ticks_map is None else ticks_map.get(code, [])
out[code] = dedupe_by_read_pairs(
raw, pairs, ticks=ticks, tf_min=int(timeframe),
missing_policy="hole", live_cover=False, stats=stats,
)
elapsed = time.perf_counter() - t0
n_tick_rows = 0
if ticks_map:
n_tick_rows = sum(len(v) for v in ticks_map.values())
slots = int(stats.get("slots") or 0) or 1
pick_by = stats.get("pick_by") or {}
garb_by = stats.get("garbage_by") or {}
raw_by = stats.get("raw_by") or {}
pick_txt = " ".join(
"%s=%s(%.0f%%)" % (k, v, 100.0 * int(v) / max(1, int(stats.get("picked") or 1)))
for k, v in sorted(pick_by.items(), key=lambda x: -int(x[1]))
) or ""
garb_txt = " ".join(
"%s=%s" % (k, v) for k, v in sorted(garb_by.items(), key=lambda x: -int(x[1]))
) or "0"
raw_txt = " ".join(
"%s=%s" % (k, v) for k, v in sorted(raw_by.items(), key=lambda x: -int(x[1]))
) or ""
logger.info(
"📂 캔들 bulk: 봉 %s행 · 틱 %s행 · 종목 %s · %.2fs (tf=%s market=%s)",
len(rows or []),
n_tick_rows,
len(out),
elapsed,
int(timeframe),
mk or "ALL",
)
logger.info(
"📊 봉 pick 비율(1차→2차→REST): 슬롯 %s · pick %s · hole %s(%.1f%%) · "
"쓰레기스킵 %s(슬롯대비 %.1f%%) | raw[%s] | pick[%s] | garbage[%s]",
slots,
stats.get("picked"),
stats.get("hole"),
100.0 * int(stats.get("hole") or 0) / slots,
stats.get("garbage_skip"),
100.0 * int(stats.get("garbage_skip") or 0) / slots,
raw_txt,
pick_txt,
garb_txt,
)
return out
def fetch_ws_candles_warmup_before(
db,
code: str,
timeframe: int,
before_candle_time: str,
limit: int,
*,
extra_select: str = "",
peak_sel: str = "",
confirmed_only: bool = True,
) -> List[Dict[str, Any]]:
"""기간 시작 이전 N봉 — prepend 웜업용 (오래된→최신)."""
if limit <= 0 or not before_candle_time:
return []
pairs = resolve_bt_read_pairs()
confirmed_sql = " AND is_confirmed=1" if confirmed_only else ""
fetch_limit = max(int(limit) * max(len(pairs), 1), int(limit) + 50)
src_sql, src_params = _pairs_in_sql(pairs)
cols = (
f"candle_time, open, high, low, close, volume, source, channel"
f"{peak_sel}{extra_select}"
)
rows = db.conn.execute(
f"SELECT {cols} FROM ws_candles "
"WHERE timeframe=%s AND code=%s AND candle_time < %s"
+ confirmed_sql
+ src_sql
+ " ORDER BY candle_time DESC LIMIT %s",
[int(timeframe), code, before_candle_time, *src_params, fetch_limit],
).fetchall()
raw_rows = [dict(r) for r in reversed(rows)]
t0 = str(raw_rows[0].get("candle_time") or "")[:12] if raw_rows else ""
ticks = (
_ticks_for_bar_garbage(db, str(code), t0, str(before_candle_time))
if t0 else None
)
bars = dedupe_by_read_pairs(
raw_rows, pairs,
ticks=ticks,
tf_min=int(timeframe),
missing_policy="hole",
live_cover=False,
)
return bars[-limit:] if len(bars) > limit else bars