first commit

This commit is contained in:
2026-07-16 23:55:16 +09:00
commit 57c07a4e12
40 changed files with 2513 additions and 0 deletions

0
app/core/__init__.py Normal file
View File

105
app/core/auth.py Normal file
View File

@@ -0,0 +1,105 @@
from __future__ import annotations
import asyncio
import hashlib
import hmac
import time
from datetime import datetime, timedelta
import httpx
from app.core.config import get_base_url, settings
from app.core.logger import setup_logger
from app.core.rate_limiter import RateLimiter
logger = setup_logger("auth")
_TOKEN_URL = "/oauth2/tokenP"
_APPROVAL_URL = "/oauth2/Approval"
_HASHKEY_URL = "/uapi/hashkey"
class TokenManager:
def __init__(self) -> None:
self._access_token: str = ""
self._token_expires_at: datetime = datetime.min
self._approval_key: str = ""
self._approval_expires_at: datetime = datetime.min
self._rate_limiter = RateLimiter(requests_per_second=1.0)
self._lock = asyncio.Lock()
self._client = httpx.AsyncClient(timeout=10.0)
@property
def is_vps(self) -> bool:
return settings.kis.server_mode == "vps"
async def _request_token(self) -> None:
url = f"{get_base_url()}{_TOKEN_URL}"
payload = {
"grant_type": "client_credentials",
"appkey": settings.kis.app_key,
"appsecret": settings.kis.app_secret,
}
resp = await self._client.post(url, json=payload)
resp.raise_for_status()
data = resp.json()
self._access_token = data["access_token"]
expires_in = int(data.get("expires_in", 7776000))
self._token_expires_at = datetime.now() + timedelta(seconds=expires_in)
logger.info("REST access_token 발급 완료 (만료: %s)", self._token_expires_at)
async def _request_approval_key(self) -> None:
url = f"{get_base_url()}{_APPROVAL_URL}"
payload = {
"grant_type": "client_credentials",
"appkey": settings.kis.app_key,
"secretkey": settings.kis.app_secret,
}
resp = await self._client.post(url, json=payload)
resp.raise_for_status()
data = resp.json()
self._approval_key = data["approval_key"]
self._approval_expires_at = datetime.now() + timedelta(hours=23, minutes=50)
logger.info("WebSocket approval_key 발급 완료 (만료: %s)", self._approval_expires_at)
async def get_access_token(self) -> str:
async with self._lock:
if datetime.now() >= self._token_expires_at:
await self._request_token()
return self._access_token
async def get_approval_key(self) -> str:
async with self._lock:
if datetime.now() >= self._approval_expires_at:
await self._request_approval_key()
return self._approval_key
async def generate_hashkey(self, data: dict) -> str:
await self._rate_limiter.acquire()
url = f"{get_base_url()}{_HASHKEY_URL}"
headers = {
"content-type": "application/json",
"appkey": settings.kis.app_key,
"appsecret": settings.kis.app_secret,
}
resp = await self._client.post(url, json=data, headers=headers)
resp.raise_for_status()
return resp.json()["HASH"]
def get_auth_headers(self, tr_id: str) -> dict[str, str]:
return {
"Content-Type": "application/json; charset=utf-8",
"authorization": f"Bearer {self._access_token}",
"appKey": settings.kis.app_key,
"appSecret": settings.kis.app_secret,
"tr_id": tr_id,
"custtype": "P",
}
async def close(self) -> None:
await self._client.aclose()
token_manager = TokenManager()

116
app/core/config.py Normal file
View File

@@ -0,0 +1,116 @@
from __future__ import annotations
import os
from pathlib import Path
from typing import Any
import yaml
from pydantic import Field
from pydantic_settings import BaseSettings
class KISConfig(BaseSettings):
app_key: str = Field(default="", alias="KIS_APP_KEY")
app_secret: str = Field(default="", alias="KIS_APP_SECRET")
account_no: str = Field(default="", alias="KIS_ACCOUNT_NO")
account_code: str = Field(default="01", alias="KIS_ACCOUNT_CODE")
hts_id: str = Field(default="", alias="KIS_HTS_ID")
server_mode: str = Field(default="vps", alias="KIS_SERVER_MODE")
model_config = {"env_file": ".env", "extra": "ignore"}
class AppConfig(BaseSettings):
host: str = Field(default="0.0.0.0", alias="APP_HOST")
port: int = Field(default=8000, alias="APP_PORT")
debug: bool = Field(default=False, alias="APP_DEBUG")
db_path: str = Field(default="./data/stock.db", alias="DB_PATH")
log_level: str = Field(default="INFO", alias="LOG_LEVEL")
model_config = {"env_file": ".env", "extra": "ignore"}
class CollectorConfig:
interval_seconds: int = 5
market_open_hour: int = 9
market_close_hour: int = 15
market_close_minute: int = 30
class StrategyConfig:
check_interval_seconds: int = 3
max_daily_trades: int = 50
default_order_type: str = "00"
class TradingConfig:
max_order_amount: int = 10_000_000
min_order_amount: int = 100_000
slippage_percent: float = 0.1
class RateLimitConfig:
requests_per_second: float = 1.0
retry_delay: float = 1.5
class Settings:
def __init__(self) -> None:
self.kis = KISConfig()
self.app = AppConfig()
self.collector = CollectorConfig()
self.strategy = StrategyConfig()
self.trading = TradingConfig()
self.rate_limit = RateLimitConfig()
self._load_yaml()
def _load_yaml(self) -> None:
yaml_path = Path("config.yaml")
if not yaml_path.exists():
return
with open(yaml_path, encoding="utf-8") as f:
data: dict[str, Any] = yaml.safe_load(f) or {}
kis = data.get("kis", {})
if "server_mode" in kis:
self.kis.server_mode = kis["server_mode"]
rl = kis.get("rate_limit", {})
if "requests_per_second" in rl:
self.rate_limit.requests_per_second = rl["requests_per_second"]
if "retry_delay" in rl:
self.rate_limit.retry_delay = rl["retry_delay"]
collector = data.get("collector", {})
for key, val in collector.items():
if hasattr(self.collector, key):
setattr(self.collector, key, val)
strategy = data.get("strategies", {})
for key, val in strategy.items():
if hasattr(self.strategy, key):
setattr(self.strategy, key, val)
trading = data.get("trading", {})
for key, val in trading.items():
if hasattr(self.trading, key):
setattr(self.trading, key, val)
logging_cfg = data.get("logging", {})
if "level" in logging_cfg:
self.app.log_level = logging_cfg["level"]
settings = Settings()
def get_base_url() -> str:
if settings.kis.server_mode == "real":
return "https://openapi.koreainvestment.com:9443"
return "https://openapivts.koreainvestment.com:29443"
def get_ws_url() -> str:
if settings.kis.server_mode == "real":
return "ws://ops.koreainvestment.com:21000"
return "ws://ops.koreainvestment.com:31000"

43
app/core/database.py Normal file
View File

@@ -0,0 +1,43 @@
from __future__ import annotations
from pathlib import Path
from sqlalchemy import create_engine, event
from sqlalchemy.orm import DeclarativeBase, sessionmaker
from app.core.config import settings
class Base(DeclarativeBase):
pass
engine = create_engine(
f"sqlite:///{settings.app.db_path}",
echo=False,
connect_args={"check_same_thread": False},
)
@event.listens_for(engine, "connect")
def _set_sqlite_pragma(dbapi_connection, connection_record) -> None: # noqa: ANN001
cursor = dbapi_connection.cursor()
cursor.execute("PRAGMA journal_mode=WAL")
cursor.execute("PRAGMA foreign_keys=ON")
cursor.close()
SessionLocal = sessionmaker(bind=engine, autoflush=False, autocommit=False)
def init_db() -> None:
Path(settings.app.db_path).parent.mkdir(parents=True, exist_ok=True)
Base.metadata.create_all(bind=engine)
def get_db():
db = SessionLocal()
try:
yield db
finally:
db.close()

27
app/core/logger.py Normal file
View File

@@ -0,0 +1,27 @@
from __future__ import annotations
import logging
from pathlib import Path
LOG_FORMAT = "%(asctime)s | %(levelname)-8s | %(name)s | %(message)s"
def setup_logger(name: str, level: str = "INFO", log_file: str | None = None) -> logging.Logger:
logger = logging.getLogger(name)
logger.setLevel(getattr(logging, level.upper(), logging.INFO))
if not logger.handlers:
formatter = logging.Formatter(LOG_FORMAT, datefmt="%Y-%m-%d %H:%M:%S")
console = logging.StreamHandler()
console.setFormatter(formatter)
logger.addHandler(console)
if log_file:
log_path = Path(log_file)
log_path.parent.mkdir(parents=True, exist_ok=True)
file_handler = logging.FileHandler(log_file, encoding="utf-8")
file_handler.setFormatter(formatter)
logger.addHandler(file_handler)
return logger

25
app/core/rate_limiter.py Normal file
View File

@@ -0,0 +1,25 @@
from __future__ import annotations
import asyncio
import time
from collections import deque
class RateLimiter:
def __init__(self, requests_per_second: float = 1.0) -> None:
self._min_interval = 1.0 / requests_per_second
self._timestamps: deque[float] = deque()
self._lock = asyncio.Lock()
async def acquire(self) -> None:
async with self._lock:
now = time.monotonic()
while self._timestamps and self._timestamps[0] <= now - self._min_interval:
self._timestamps.popleft()
if self._timestamps:
wait_time = self._timestamps[0] + self._min_interval - now
if wait_time > 0:
await asyncio.sleep(wait_time)
self._timestamps.append(time.monotonic())