from __future__ import annotations import json from sqlalchemy.orm import Session from app.core.config import settings from app.core.logger import setup_logger from app.models.stock import PriceHistory, Strategy, Trade from app.services.market_data import market_data_service from app.services.trading import trading_service from app.engine.strategies.base import BaseStrategy, Signal from app.engine.strategies.conditional import ConditionalStrategy from app.engine.strategies.technical import TechnicalStrategy from app.engine.strategies.periodic import PeriodicStrategy logger = setup_logger("strategy_engine") STRATEGY_MAP = { "conditional": ConditionalStrategy, "technical": TechnicalStrategy, "periodic": PeriodicStrategy, } class StrategyEngine: def __init__(self) -> None: self._daily_trade_count = 0 def _create_strategy(self, strategy_record: Strategy) -> BaseStrategy | None: cls = STRATEGY_MAP.get(strategy_record.strategy_type) if not cls: logger.warning("알 수 없는 전략 타입: %s", strategy_record.strategy_type) return None params = json.loads(strategy_record.params_json) if strategy_record.params_json else {} return cls( stock_code=strategy_record.stock_code, params=params, strategy_id=strategy_record.id, ) async def evaluate_all(self, db: Session) -> list[Signal]: strategies = db.query(Strategy).filter(Strategy.is_active == True).all() signals: list[Signal] = [] for strat_record in strategies: try: strategy = self._create_strategy(strat_record) if not strategy: continue current_price = await market_data_service.get_current_price(strat_record.stock_code) if not current_price: continue price_history = self._get_price_history(db, strat_record.stock_code) holding = self._get_holding_qty(db, strat_record.stock_code) signal = strategy.evaluate(current_price, price_history, holding) if signal.action != "hold": signal.qty = max(signal.qty, strat_record.qty) signal.strategy_id = strat_record.id signals.append(signal) logger.info( "신호 발생: %s %s %s주 - %s", signal.stock_code, signal.action, signal.qty, signal.reason, ) except Exception as e: logger.error("전략 평가 오류 (ID=%s): %s", strat_record.id, e) return signals async def execute_signals(self, signals: list[Signal], db: Session) -> list[Trade]: if self._daily_trade_count >= settings.strategy.max_daily_trades: logger.warning("일일 최대 매매 횟수 초과 (%d)", settings.strategy.max_daily_trades) return [] trades: list[Trade] = [] for signal in signals: try: order_result = await trading_service.place_order( stock_code=signal.stock_code, side=signal.action, qty=signal.qty, price=signal.price, order_type=settings.strategy.default_order_type, ) stock_name = "" price_data = await market_data_service.get_current_price(signal.stock_code) if price_data: stock_name = price_data.get("stock_name", "") trade = Trade( order_no=order_result.get("order_no", ""), stock_code=signal.stock_code, stock_name=stock_name, side=signal.action, qty=signal.qty, price=signal.price, order_type=settings.strategy.default_order_type, status="filled" if order_result.get("rt_cd") == "0" else "rejected", strategy_id=signal.strategy_id, ) db.add(trade) db.commit() self._daily_trade_count += 1 trades.append(trade) except Exception as e: logger.error("주문 실행 오류: %s", e) return trades def _get_price_history(self, db: Session, stock_code: str, limit: int = 60) -> list[dict]: records = ( db.query(PriceHistory) .filter(PriceHistory.stock_code == stock_code) .order_by(PriceHistory.datetime.desc()) .limit(limit) .all() ) return [ { "date": r.datetime.strftime("%Y%m%d"), "open": r.open, "high": r.high, "low": r.low, "close": r.close, "volume": r.volume, } for r in reversed(records) ] def _get_holding_qty(self, db: Session, stock_code: str) -> int: from app.models.stock import Holding holding = db.query(Holding).filter(Holding.stock_code == stock_code).first() return holding.qty if holding else 0 def reset_daily_count(self) -> None: self._daily_trade_count = 0 strategy_engine = StrategyEngine()