first commit
This commit is contained in:
0
app/core/__init__.py
Normal file
0
app/core/__init__.py
Normal file
105
app/core/auth.py
Normal file
105
app/core/auth.py
Normal 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
116
app/core/config.py
Normal 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
43
app/core/database.py
Normal 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
27
app/core/logger.py
Normal 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
25
app/core/rate_limiter.py
Normal 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())
|
||||
Reference in New Issue
Block a user