diff --git a/.agents/skills/write-pr/SKILL.md b/.agents/skills/write-pr/SKILL.md index ece4340..919c6a7 100644 --- a/.agents/skills/write-pr/SKILL.md +++ b/.agents/skills/write-pr/SKILL.md @@ -1,7 +1,6 @@ --- name: write-pr description: Push the current branch, generate a PR title/body from commits since the base branch following this project's convention (no prefix, Korean title), attach best-matching labels if any exist, and open the GitHub PR. -compatibility: Requires git and gh (GitHub CLI) --- ## Step 1 — Determine Base Branch @@ -38,7 +37,15 @@ If there are no commits ahead of base, report that and exit. 예: `JWT 인증 필터 추가` -**PR 본문**: 커밋 목록을 바탕으로 변경 사항을 bullet로 정리하고, 관련 이슈가 있으면 `Closes #<번호>` 형식으로 연결. +**PR 본문**: 저장소에 PR 템플릿이 있으면 반드시 먼저 읽고 그 형식을 채운다. + +```bash +find .github -maxdepth 3 -type f \( -iname '*pull_request_template*' -o -path './.github/PULL_REQUEST_TEMPLATE/*' \) +``` + +- 템플릿이 있으면 섹션명과 체크리스트를 유지하고, 커밋 목록을 바탕으로 내용을 채운다. +- 관련 이슈가 있으면 템플릿의 관련 이슈 섹션에 `Closes #<번호>` 형식으로 연결한다. +- 템플릿이 없을 때만 자유 형식 bullet 본문을 작성한다. **어트리뷰션 주의**: 본문에 "Generated by Codex" 같은 서명이나 트레일러를 추가하지 않는다. @@ -60,4 +67,4 @@ gh pr create --base develop --title "<한글 제목>" --body "<본문>" ## Step 7 — Report -생성된 PR URL을 출력한다. \ No newline at end of file +생성된 PR URL을 출력한다. diff --git a/alembic.ini b/alembic.ini new file mode 100644 index 0000000..0a26948 --- /dev/null +++ b/alembic.ini @@ -0,0 +1,38 @@ +[alembic] +script_location = alembic +prepend_sys_path = . +sqlalchemy.url = postgresql+psycopg://mlflow:mlflow@localhost:5432/arena_ai + +[loggers] +keys = root,sqlalchemy,alembic + +[handlers] +keys = console + +[formatters] +keys = generic + +[logger_root] +level = WARNING +handlers = console +qualname = + +[logger_sqlalchemy] +level = WARNING +handlers = +qualname = sqlalchemy.engine + +[logger_alembic] +level = INFO +handlers = +qualname = alembic + +[handler_console] +class = StreamHandler +args = (sys.stderr,) +level = NOTSET +formatter = generic + +[formatter_generic] +format = %(levelname)-5.5s [%(name)s] %(message)s +datefmt = %H:%M:%S diff --git a/alembic/env.py b/alembic/env.py new file mode 100644 index 0000000..125521f --- /dev/null +++ b/alembic/env.py @@ -0,0 +1,49 @@ +from logging.config import fileConfig + +from alembic import context +from sqlalchemy import engine_from_config, pool + +from app.mlops import models as _mlops_models # noqa: F401 +from app.shared.db import Base, get_database_url +from app.trading_ai import models as _trading_ai_models # noqa: F401 + +config = context.config + +if config.config_file_name is not None: + fileConfig(config.config_file_name) + +target_metadata = Base.metadata + +config.set_main_option("sqlalchemy.url", get_database_url()) + + +def run_migrations_offline() -> None: + context.configure( + url=get_database_url(), + target_metadata=target_metadata, + literal_binds=True, + dialect_opts={"paramstyle": "named"}, + ) + + with context.begin_transaction(): + context.run_migrations() + + +def run_migrations_online() -> None: + connectable = engine_from_config( + config.get_section(config.config_ini_section, {}), + prefix="sqlalchemy.", + poolclass=pool.NullPool, + ) + + with connectable.connect() as connection: + context.configure(connection=connection, target_metadata=target_metadata) + + with context.begin_transaction(): + context.run_migrations() + + +if context.is_offline_mode(): + run_migrations_offline() +else: + run_migrations_online() diff --git a/alembic/versions/20260708_0001_create_ai_tables.py b/alembic/versions/20260708_0001_create_ai_tables.py new file mode 100644 index 0000000..d968092 --- /dev/null +++ b/alembic/versions/20260708_0001_create_ai_tables.py @@ -0,0 +1,78 @@ +"""create ai tables + +Revision ID: 20260708_0001 +Revises: +Create Date: 2026-07-08 +""" + +from typing import Sequence + +from alembic import op +import sqlalchemy as sa +from sqlalchemy.dialects import postgresql + +revision: str = "20260708_0001" +down_revision: str | None = None +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "decision_logs", + sa.Column("id", sa.BigInteger(), autoincrement=True, nullable=False), + sa.Column("challenge_id", sa.BigInteger(), nullable=False), + sa.Column("order_id", sa.BigInteger(), nullable=True), + sa.Column("symbol_code", sa.String(length=20), nullable=False), + sa.Column("market", sa.String(length=10), nullable=False), + sa.Column("feature_snapshot", postgresql.JSONB(astext_type=sa.Text()), nullable=False), + sa.Column("model_output_probability", sa.Numeric(precision=5, scale=4), nullable=False), + sa.Column("action", sa.String(length=10), nullable=False), + sa.Column("model_version", sa.String(length=50), nullable=False), + sa.Column("decided_at", sa.DateTime(timezone=True), nullable=False), + sa.CheckConstraint("action IN ('BUY', 'SELL', 'HOLD')", name="ck_decision_logs_action"), + sa.CheckConstraint("market IN ('KR', 'US', 'COIN')", name="ck_decision_logs_market"), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("idx_decision_logs_challenge_decided", "decision_logs", ["challenge_id", "decided_at"]) + op.create_index("idx_decision_logs_model_version", "decision_logs", ["model_version"]) + + op.create_table( + "model_versions", + sa.Column("id", sa.BigInteger(), autoincrement=True, nullable=False), + sa.Column("version_tag", sa.String(length=50), nullable=False), + sa.Column("trained_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("performance_metrics", postgresql.JSONB(astext_type=sa.Text()), nullable=False), + sa.Column("status", sa.String(length=20), server_default="CHALLENGER", nullable=False), + sa.CheckConstraint("status IN ('CHAMPION', 'CHALLENGER', 'RETIRED')", name="ck_model_versions_status"), + sa.PrimaryKeyConstraint("id"), + sa.UniqueConstraint("version_tag", name="uq_model_versions_version_tag"), + ) + op.create_index( + "uq_model_versions_champion", + "model_versions", + ["status"], + unique=True, + postgresql_where=sa.text("status = 'CHAMPION'"), + ) + op.create_table( + "shap_values", + sa.Column("id", sa.BigInteger(), autoincrement=True, nullable=False), + sa.Column("decision_id", sa.BigInteger(), nullable=False), + sa.Column("feature_name", sa.String(length=50), nullable=False), + sa.Column("contribution", sa.Numeric(precision=8, scale=5), nullable=False), + sa.Column("base_value", sa.Numeric(precision=5, scale=4), nullable=False), + sa.ForeignKeyConstraint(["decision_id"], ["decision_logs.id"]), + sa.PrimaryKeyConstraint("id"), + ) + op.create_index("idx_shap_values_decision_id", "shap_values", ["decision_id"]) + + +def downgrade() -> None: + op.drop_index("idx_shap_values_decision_id", table_name="shap_values") + op.drop_table("shap_values") + op.drop_index("uq_model_versions_champion", table_name="model_versions") + op.drop_table("model_versions") + op.drop_index("idx_decision_logs_model_version", table_name="decision_logs") + op.drop_index("idx_decision_logs_challenge_decided", table_name="decision_logs") + op.drop_table("decision_logs") diff --git a/app/__init__.py b/app/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/app/__init__.py @@ -0,0 +1 @@ + diff --git a/app/main.py b/app/main.py new file mode 100644 index 0000000..e12160d --- /dev/null +++ b/app/main.py @@ -0,0 +1,8 @@ +from fastapi import FastAPI + +app = FastAPI() + + +@app.get("/health") +def health(): + return {"status": "ok"} diff --git a/app/mlops/__init__.py b/app/mlops/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/app/mlops/__init__.py @@ -0,0 +1 @@ + diff --git a/app/mlops/models.py b/app/mlops/models.py new file mode 100644 index 0000000..94a583c --- /dev/null +++ b/app/mlops/models.py @@ -0,0 +1,33 @@ +from datetime import datetime +from typing import Any + +from sqlalchemy import BigInteger, DateTime, Enum, Index, String, UniqueConstraint +from sqlalchemy.dialects.postgresql import JSONB +from sqlalchemy.orm import Mapped, mapped_column + +from app.shared.db import Base +from app.shared.enums import ModelStatus + + +class ModelVersion(Base): + __tablename__ = "model_versions" + __table_args__ = ( + UniqueConstraint("version_tag", name="uq_model_versions_version_tag"), + Index( + "uq_model_versions_champion", + "status", + unique=True, + postgresql_where="status = 'CHAMPION'", + ), + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + version_tag: Mapped[str] = mapped_column(String(50), nullable=False) + trained_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + performance_metrics: Mapped[dict[str, Any]] = mapped_column(JSONB, nullable=False) + status: Mapped[ModelStatus] = mapped_column( + Enum(ModelStatus, native_enum=False, name="ck_model_versions_status"), + nullable=False, + default=ModelStatus.CHALLENGER, + server_default=ModelStatus.CHALLENGER.value, + ) diff --git a/app/mlops/repositories.py b/app/mlops/repositories.py new file mode 100644 index 0000000..79c2191 --- /dev/null +++ b/app/mlops/repositories.py @@ -0,0 +1,23 @@ +from sqlalchemy import select +from sqlalchemy.orm import Session + +from app.mlops.models import ModelVersion +from app.shared.enums import ModelStatus + + +class ModelVersionRepository: + def __init__(self, db: Session): + self.db = db + + def add(self, model_version: ModelVersion) -> ModelVersion: + self.db.add(model_version) + self.db.flush() + return model_version + + def get_by_version_tag(self, version_tag: str) -> ModelVersion | None: + stmt = select(ModelVersion).where(ModelVersion.version_tag == version_tag) + return self.db.scalar(stmt) + + def get_champion(self) -> ModelVersion | None: + stmt = select(ModelVersion).where(ModelVersion.status == ModelStatus.CHAMPION) + return self.db.scalar(stmt) diff --git a/app/shared/__init__.py b/app/shared/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/app/shared/__init__.py @@ -0,0 +1 @@ + diff --git a/app/shared/db.py b/app/shared/db.py new file mode 100644 index 0000000..d389d40 --- /dev/null +++ b/app/shared/db.py @@ -0,0 +1,35 @@ +from functools import lru_cache +import os +from typing import Generator + +from sqlalchemy import create_engine +from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker + + +DEFAULT_DATABASE_URL = "postgresql+psycopg://mlflow:mlflow@localhost:5432/arena_ai" + + +class Base(DeclarativeBase): + pass + + +def get_database_url() -> str: + return os.getenv("DATABASE_URL", DEFAULT_DATABASE_URL) + + +@lru_cache +def get_engine(): + return create_engine(get_database_url(), pool_pre_ping=True) + + +@lru_cache +def get_session_local(): + return sessionmaker(autocommit=False, autoflush=False, bind=get_engine()) + + +def get_db() -> Generator[Session, None, None]: + db = get_session_local()() + try: + yield db + finally: + db.close() diff --git a/app/shared/enums.py b/app/shared/enums.py new file mode 100644 index 0000000..76c7574 --- /dev/null +++ b/app/shared/enums.py @@ -0,0 +1,19 @@ +from enum import Enum + + +class Market(str, Enum): + KR = "KR" + US = "US" + COIN = "COIN" + + +class TradingAction(str, Enum): + BUY = "BUY" + SELL = "SELL" + HOLD = "HOLD" + + +class ModelStatus(str, Enum): + CHAMPION = "CHAMPION" + CHALLENGER = "CHALLENGER" + RETIRED = "RETIRED" diff --git a/app/trading_ai/__init__.py b/app/trading_ai/__init__.py new file mode 100644 index 0000000..8b13789 --- /dev/null +++ b/app/trading_ai/__init__.py @@ -0,0 +1 @@ + diff --git a/app/trading_ai/models.py b/app/trading_ai/models.py new file mode 100644 index 0000000..f31a526 --- /dev/null +++ b/app/trading_ai/models.py @@ -0,0 +1,50 @@ +from datetime import datetime +from decimal import Decimal +from typing import Any + +from sqlalchemy import BigInteger, DateTime, Enum, ForeignKey, Index, Numeric, String +from sqlalchemy.dialects.postgresql import JSONB +from sqlalchemy.orm import Mapped, mapped_column, relationship + +from app.shared.db import Base +from app.shared.enums import Market, TradingAction + + +class DecisionLog(Base): + __tablename__ = "decision_logs" + __table_args__ = ( + Index("idx_decision_logs_challenge_decided", "challenge_id", "decided_at"), + Index("idx_decision_logs_model_version", "model_version"), + ) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + challenge_id: Mapped[int] = mapped_column(BigInteger, nullable=False) + order_id: Mapped[int | None] = mapped_column(BigInteger) + symbol_code: Mapped[str] = mapped_column(String(20), nullable=False) + market: Mapped[Market] = mapped_column( + Enum(Market, native_enum=False, name="ck_decision_logs_market"), + nullable=False, + ) + feature_snapshot: Mapped[dict[str, Any]] = mapped_column(JSONB, nullable=False) + model_output_probability: Mapped[Decimal] = mapped_column(Numeric(5, 4), nullable=False) + action: Mapped[TradingAction] = mapped_column( + Enum(TradingAction, native_enum=False, name="ck_decision_logs_action"), + nullable=False, + ) + model_version: Mapped[str] = mapped_column(String(50), nullable=False) + decided_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) + + shap_values: Mapped[list["ShapValue"]] = relationship(back_populates="decision") + + +class ShapValue(Base): + __tablename__ = "shap_values" + __table_args__ = (Index("idx_shap_values_decision_id", "decision_id"),) + + id: Mapped[int] = mapped_column(BigInteger, primary_key=True, autoincrement=True) + decision_id: Mapped[int] = mapped_column(BigInteger, ForeignKey("decision_logs.id"), nullable=False) + feature_name: Mapped[str] = mapped_column(String(50), nullable=False) + contribution: Mapped[Decimal] = mapped_column(Numeric(8, 5), nullable=False) + base_value: Mapped[Decimal] = mapped_column(Numeric(5, 4), nullable=False) + + decision: Mapped[DecisionLog] = relationship(back_populates="shap_values") diff --git a/app/trading_ai/repositories.py b/app/trading_ai/repositories.py new file mode 100644 index 0000000..b516015 --- /dev/null +++ b/app/trading_ai/repositories.py @@ -0,0 +1,40 @@ +from sqlalchemy import select +from sqlalchemy.orm import Session + +from app.trading_ai.models import DecisionLog, ShapValue + + +class DecisionLogRepository: + def __init__(self, db: Session): + self.db = db + + def add(self, decision_log: DecisionLog) -> DecisionLog: + self.db.add(decision_log) + self.db.flush() + return decision_log + + def get_by_id(self, decision_id: int) -> DecisionLog | None: + return self.db.get(DecisionLog, decision_id) + + def list_by_challenge_id(self, challenge_id: int, limit: int = 100) -> list[DecisionLog]: + stmt = ( + select(DecisionLog) + .where(DecisionLog.challenge_id == challenge_id) + .order_by(DecisionLog.decided_at.desc()) + .limit(limit) + ) + return list(self.db.scalars(stmt)) + + +class ShapValueRepository: + def __init__(self, db: Session): + self.db = db + + def add_many(self, shap_values: list[ShapValue]) -> list[ShapValue]: + self.db.add_all(shap_values) + self.db.flush() + return shap_values + + def list_by_decision_id(self, decision_id: int) -> list[ShapValue]: + stmt = select(ShapValue).where(ShapValue.decision_id == decision_id) + return list(self.db.scalars(stmt)) diff --git a/compose.yml b/compose.yml index 9a6d1a4..d5f8ec6 100644 --- a/compose.yml +++ b/compose.yml @@ -9,6 +9,7 @@ services: - "5432:5432" volumes: - postgres_data:/var/lib/postgresql/data + - ./docker/postgres/init-ai-db.sql:/docker-entrypoint-initdb.d/10-init-ai-db.sql:ro healthcheck: test: ["CMD-SHELL", "pg_isready -U mlflow -d mlflow"] interval: 5s diff --git a/docker/postgres/init-ai-db.sql b/docker/postgres/init-ai-db.sql new file mode 100644 index 0000000..d60a933 --- /dev/null +++ b/docker/postgres/init-ai-db.sql @@ -0,0 +1 @@ +CREATE DATABASE arena_ai; diff --git a/main.py b/main.py index e12160d..7436e63 100644 --- a/main.py +++ b/main.py @@ -1,8 +1 @@ -from fastapi import FastAPI - -app = FastAPI() - - -@app.get("/health") -def health(): - return {"status": "ok"} +from app.main import app diff --git a/requirements.txt b/requirements.txt index 9d6cc3d..3b9a22a 100644 --- a/requirements.txt +++ b/requirements.txt @@ -6,3 +6,6 @@ mlflow==3.14.0 pandas==2.3.3 openai==2.44.0 pydantic==2.13.4 +SQLAlchemy==2.0.51 +alembic==1.18.5 +psycopg[binary]==3.3.4 diff --git a/tests/test_models_metadata.py b/tests/test_models_metadata.py new file mode 100644 index 0000000..97eaaad --- /dev/null +++ b/tests/test_models_metadata.py @@ -0,0 +1,25 @@ +import unittest + +from app.mlops.models import ModelVersion +from app.shared.db import Base +from app.trading_ai.models import DecisionLog, ShapValue + + +class ModelMetadataTest(unittest.TestCase): + def test_tables_are_registered(self): + self.assertTrue({"decision_logs", "shap_values", "model_versions"} <= set(Base.metadata.tables)) + + def test_spring_boot_ids_are_not_foreign_keys(self): + self.assertFalse(DecisionLog.__table__.c.challenge_id.foreign_keys) + self.assertFalse(DecisionLog.__table__.c.order_id.foreign_keys) + + def test_shap_values_reference_decision_logs(self): + self.assertTrue(ShapValue.__table__.c.decision_id.foreign_keys) + + def test_model_version_tag_is_unique(self): + constraints = {constraint.name for constraint in ModelVersion.__table__.constraints} + self.assertIn("uq_model_versions_version_tag", constraints) + + +if __name__ == "__main__": + unittest.main()