From 84f861b7998b69b1860bd751d31574a83b1a535a Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EB=B0=95=ED=95=98=EB=AF=BC?= Date: Wed, 8 Jul 2026 10:40:10 +0900 Subject: [PATCH 1/3] =?UTF-8?q?feat=20::=20AI=20DB=20=EC=97=94=ED=8B=B0?= =?UTF-8?q?=ED=8B=B0=EC=99=80=20=EB=A0=88=ED=8F=AC=EC=A7=80=ED=86=A0?= =?UTF-8?q?=EB=A6=AC=20=EC=B6=94=EA=B0=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- alembic.ini | 38 ++++++++++ alembic/env.py | 52 ++++++++++++++ .../20260708_0001_create_ai_tables.py | 72 +++++++++++++++++++ app/__init__.py | 1 + app/main.py | 8 +++ app/mlops/__init__.py | 1 + app/mlops/models.py | 28 ++++++++ app/mlops/repositories.py | 23 ++++++ app/shared/__init__.py | 1 + app/shared/db.py | 35 +++++++++ app/shared/enums.py | 19 +++++ app/trading_ai/__init__.py | 1 + app/trading_ai/models.py | 47 ++++++++++++ app/trading_ai/repositories.py | 40 +++++++++++ compose.yml | 1 + docker/postgres/init-ai-db.sql | 1 + main.py | 9 +-- requirements.txt | 3 + tests/test_models_metadata.py | 25 +++++++ 19 files changed, 397 insertions(+), 8 deletions(-) create mode 100644 alembic.ini create mode 100644 alembic/env.py create mode 100644 alembic/versions/20260708_0001_create_ai_tables.py create mode 100644 app/__init__.py create mode 100644 app/main.py create mode 100644 app/mlops/__init__.py create mode 100644 app/mlops/models.py create mode 100644 app/mlops/repositories.py create mode 100644 app/shared/__init__.py create mode 100644 app/shared/db.py create mode 100644 app/shared/enums.py create mode 100644 app/trading_ai/__init__.py create mode 100644 app/trading_ai/models.py create mode 100644 app/trading_ai/repositories.py create mode 100644 docker/postgres/init-ai-db.sql create mode 100644 tests/test_models_metadata.py 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..d993887 --- /dev/null +++ b/alembic/env.py @@ -0,0 +1,52 @@ +from logging.config import fileConfig + +from alembic import context +from sqlalchemy import engine_from_config, pool + +from app.mlops import models as mlops_models +from app.shared.db import Base, get_database_url +from app.trading_ai import models as trading_ai_models + +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()) + +mlops_models +trading_ai_models + + +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..1a86473 --- /dev/null +++ b/alembic/versions/20260708_0001_create_ai_tables.py @@ -0,0 +1,72 @@ +"""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(), 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_id", "decision_logs", ["challenge_id"]) + op.create_index("idx_decision_logs_decided_at", "decision_logs", ["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(), 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_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_table("model_versions") + op.drop_index("idx_decision_logs_model_version", table_name="decision_logs") + op.drop_index("idx_decision_logs_decided_at", table_name="decision_logs") + op.drop_index("idx_decision_logs_challenge_id", 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..1774e35 --- /dev/null +++ b/app/mlops/models.py @@ -0,0 +1,28 @@ +from datetime import datetime +from typing import Any + +from sqlalchemy import BigInteger, CheckConstraint, DateTime, 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__ = ( + CheckConstraint("status IN ('CHAMPION', 'CHALLENGER', 'RETIRED')", name="ck_model_versions_status"), + UniqueConstraint("version_tag", name="uq_model_versions_version_tag"), + ) + + 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, nullable=False) + performance_metrics: Mapped[dict[str, Any]] = mapped_column(JSONB, nullable=False) + status: Mapped[ModelStatus] = mapped_column( + String(20), + nullable=False, + default=ModelStatus.CHALLENGER.value, + server_default=ModelStatus.CHALLENGER.value, + ) diff --git a/app/mlops/repositories.py b/app/mlops/repositories.py new file mode 100644 index 0000000..7d49fe6 --- /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.value) + 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..5805d3c --- /dev/null +++ b/app/trading_ai/models.py @@ -0,0 +1,47 @@ +from datetime import datetime +from decimal import Decimal +from typing import Any + +from sqlalchemy import BigInteger, CheckConstraint, DateTime, 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__ = ( + CheckConstraint("market IN ('KR', 'US', 'COIN')", name="ck_decision_logs_market"), + CheckConstraint("action IN ('BUY', 'SELL', 'HOLD')", name="ck_decision_logs_action"), + Index("idx_decision_logs_challenge_id", "challenge_id"), + Index("idx_decision_logs_decided_at", "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(String(10), 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(String(10), nullable=False) + model_version: Mapped[str] = mapped_column(String(50), nullable=False) + decided_at: Mapped[datetime] = mapped_column(DateTime, 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..67252c9 --- /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): + assert {"decision_logs", "shap_values", "model_versions"} <= set(Base.metadata.tables) + + def test_spring_boot_ids_are_not_foreign_keys(self): + assert not DecisionLog.__table__.c.challenge_id.foreign_keys + assert not DecisionLog.__table__.c.order_id.foreign_keys + + def test_shap_values_reference_decision_logs(self): + assert ShapValue.__table__.c.decision_id.foreign_keys + + def test_model_version_tag_is_unique(self): + constraints = {constraint.name for constraint in ModelVersion.__table__.constraints} + assert "uq_model_versions_version_tag" in constraints + + +if __name__ == "__main__": + unittest.main() From e1b371ca5cedc96f63215b0b53ca489027b1971b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EB=B0=95=ED=95=98=EB=AF=BC?= Date: Wed, 8 Jul 2026 10:46:30 +0900 Subject: [PATCH 2/3] =?UTF-8?q?docs=20::=20PR=20=EC=9E=91=EC=84=B1=20?= =?UTF-8?q?=EC=8A=A4=ED=82=AC=20=ED=85=9C=ED=94=8C=EB=A6=BF=20=EC=82=AC?= =?UTF-8?q?=EC=9A=A9=20=EB=AA=85=EC=8B=9C?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .agents/skills/write-pr/SKILL.md | 13 ++++++++++--- 1 file changed, 10 insertions(+), 3 deletions(-) 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을 출력한다. From 2e137144ef9377933e77f4598e41249f78765ee4 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=EB=B0=95=ED=95=98=EB=AF=BC?= Date: Wed, 8 Jul 2026 11:14:40 +0900 Subject: [PATCH 3/3] =?UTF-8?q?fix=20::=20PR=20=EB=A6=AC=EB=B7=B0=20?= =?UTF-8?q?=EB=B0=98=EC=98=81=20-=20=EC=B1=94=ED=94=BC=EC=96=B8=20?= =?UTF-8?q?=EC=9C=A0=EC=9D=BC=EC=84=B1/=ED=83=80=EC=9E=84=EC=A1=B4/Enum=20?= =?UTF-8?q?=EC=BB=AC=EB=9F=BC/=EB=B3=B5=ED=95=A9=20=EC=9D=B8=EB=8D=B1?= =?UTF-8?q?=EC=8A=A4/=ED=85=8C=EC=8A=A4=ED=8A=B8=20=EB=8B=A8=EC=96=B8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- alembic/env.py | 7 ++----- .../20260708_0001_create_ai_tables.py | 18 ++++++++++++------ app/mlops/models.py | 15 ++++++++++----- app/mlops/repositories.py | 2 +- app/trading_ai/models.py | 19 +++++++++++-------- tests/test_models_metadata.py | 10 +++++----- 6 files changed, 41 insertions(+), 30 deletions(-) diff --git a/alembic/env.py b/alembic/env.py index d993887..125521f 100644 --- a/alembic/env.py +++ b/alembic/env.py @@ -3,9 +3,9 @@ from alembic import context from sqlalchemy import engine_from_config, pool -from app.mlops import models as mlops_models +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 +from app.trading_ai import models as _trading_ai_models # noqa: F401 config = context.config @@ -16,9 +16,6 @@ config.set_main_option("sqlalchemy.url", get_database_url()) -mlops_models -trading_ai_models - def run_migrations_offline() -> None: context.configure( diff --git a/alembic/versions/20260708_0001_create_ai_tables.py b/alembic/versions/20260708_0001_create_ai_tables.py index 1a86473..d968092 100644 --- a/alembic/versions/20260708_0001_create_ai_tables.py +++ b/alembic/versions/20260708_0001_create_ai_tables.py @@ -29,26 +29,32 @@ def upgrade() -> None: 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(), 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_id", "decision_logs", ["challenge_id"]) - op.create_index("idx_decision_logs_decided_at", "decision_logs", ["decided_at"]) + 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(), 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), @@ -65,8 +71,8 @@ def upgrade() -> None: 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_decided_at", table_name="decision_logs") - op.drop_index("idx_decision_logs_challenge_id", 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/mlops/models.py b/app/mlops/models.py index 1774e35..94a583c 100644 --- a/app/mlops/models.py +++ b/app/mlops/models.py @@ -1,7 +1,7 @@ from datetime import datetime from typing import Any -from sqlalchemy import BigInteger, CheckConstraint, DateTime, String, UniqueConstraint +from sqlalchemy import BigInteger, DateTime, Enum, Index, String, UniqueConstraint from sqlalchemy.dialects.postgresql import JSONB from sqlalchemy.orm import Mapped, mapped_column @@ -12,17 +12,22 @@ class ModelVersion(Base): __tablename__ = "model_versions" __table_args__ = ( - CheckConstraint("status IN ('CHAMPION', 'CHALLENGER', 'RETIRED')", name="ck_model_versions_status"), 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, 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( - String(20), + Enum(ModelStatus, native_enum=False, name="ck_model_versions_status"), nullable=False, - default=ModelStatus.CHALLENGER.value, + default=ModelStatus.CHALLENGER, server_default=ModelStatus.CHALLENGER.value, ) diff --git a/app/mlops/repositories.py b/app/mlops/repositories.py index 7d49fe6..79c2191 100644 --- a/app/mlops/repositories.py +++ b/app/mlops/repositories.py @@ -19,5 +19,5 @@ def get_by_version_tag(self, version_tag: str) -> ModelVersion | None: return self.db.scalar(stmt) def get_champion(self) -> ModelVersion | None: - stmt = select(ModelVersion).where(ModelVersion.status == ModelStatus.CHAMPION.value) + stmt = select(ModelVersion).where(ModelVersion.status == ModelStatus.CHAMPION) return self.db.scalar(stmt) diff --git a/app/trading_ai/models.py b/app/trading_ai/models.py index 5805d3c..f31a526 100644 --- a/app/trading_ai/models.py +++ b/app/trading_ai/models.py @@ -2,7 +2,7 @@ from decimal import Decimal from typing import Any -from sqlalchemy import BigInteger, CheckConstraint, DateTime, ForeignKey, Index, Numeric, String +from sqlalchemy import BigInteger, DateTime, Enum, ForeignKey, Index, Numeric, String from sqlalchemy.dialects.postgresql import JSONB from sqlalchemy.orm import Mapped, mapped_column, relationship @@ -13,10 +13,7 @@ class DecisionLog(Base): __tablename__ = "decision_logs" __table_args__ = ( - CheckConstraint("market IN ('KR', 'US', 'COIN')", name="ck_decision_logs_market"), - CheckConstraint("action IN ('BUY', 'SELL', 'HOLD')", name="ck_decision_logs_action"), - Index("idx_decision_logs_challenge_id", "challenge_id"), - Index("idx_decision_logs_decided_at", "decided_at"), + Index("idx_decision_logs_challenge_decided", "challenge_id", "decided_at"), Index("idx_decision_logs_model_version", "model_version"), ) @@ -24,12 +21,18 @@ class DecisionLog(Base): 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(String(10), 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(String(10), 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, nullable=False) + decided_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False) shap_values: Mapped[list["ShapValue"]] = relationship(back_populates="decision") diff --git a/tests/test_models_metadata.py b/tests/test_models_metadata.py index 67252c9..97eaaad 100644 --- a/tests/test_models_metadata.py +++ b/tests/test_models_metadata.py @@ -7,18 +7,18 @@ class ModelMetadataTest(unittest.TestCase): def test_tables_are_registered(self): - assert {"decision_logs", "shap_values", "model_versions"} <= set(Base.metadata.tables) + self.assertTrue({"decision_logs", "shap_values", "model_versions"} <= set(Base.metadata.tables)) def test_spring_boot_ids_are_not_foreign_keys(self): - assert not DecisionLog.__table__.c.challenge_id.foreign_keys - assert not DecisionLog.__table__.c.order_id.foreign_keys + 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): - assert ShapValue.__table__.c.decision_id.foreign_keys + 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} - assert "uq_model_versions_version_tag" in constraints + self.assertIn("uq_model_versions_version_tag", constraints) if __name__ == "__main__":