149 lines
5.3 KiB
Python
149 lines
5.3 KiB
Python
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()
|