From 897ca3c8e1d2e1f0666b7c71fe94cf7a21779ba6 Mon Sep 17 00:00:00 2001 From: Kayla Zhang Date: Fri, 6 Feb 2026 18:55:08 -0800 Subject: [PATCH] populate dev dabatase with sample tournament data. --- backend/app/db/database.py | 10 ++++- backend/app/scripts/reset_db.py | 8 +++- backend/app/scripts/seed_db.py | 72 +++++++++++++++++++++++++-------- 3 files changed, 70 insertions(+), 20 deletions(-) diff --git a/backend/app/db/database.py b/backend/app/db/database.py index cd18b3d..dbe6f8f 100644 --- a/backend/app/db/database.py +++ b/backend/app/db/database.py @@ -2,7 +2,6 @@ from sqlalchemy.ext.asyncio import create_async_engine, AsyncSession, async_sessionmaker from sqlalchemy.pool import NullPool from dotenv import load_dotenv - # Load environment variables load_dotenv() @@ -29,6 +28,7 @@ DB_DISABLE_SSL = os.getenv("DB_DISABLE_SSL", "False").lower() == "true" connect_args = {} if DB_DISABLE_SSL else {"ssl": "require"} + if USE_NULL_POOL: # NullPool: No connection pooling - creates fresh connection each time # Use this for Celery workers to avoid asyncpg connection conflicts @@ -52,6 +52,14 @@ connect_args=connect_args, ) +from sqlalchemy import create_engine +#sync engine to populate database +SYNC_DATABASE_URL = DATABASE_URL.replace('+asyncpg','') +sync_engine = create_engine( + SYNC_DATABASE_URL, + echo=DB_ECHO, +) + # Create async session factory AsyncSessionLocal = async_sessionmaker( engine, class_=AsyncSession, expire_on_commit=False diff --git a/backend/app/scripts/reset_db.py b/backend/app/scripts/reset_db.py index bcb095d..66d1e4f 100644 --- a/backend/app/scripts/reset_db.py +++ b/backend/app/scripts/reset_db.py @@ -2,7 +2,7 @@ import os # <--- 1. Was missing from dotenv import load_dotenv from sqlalchemy.ext.asyncio import create_async_engine - +from sqlalchemy import text # 3. Correct Import: 'Bet', not 'Bets' from ..db.models import Base, Tournament, Agent, AgentState, Trade, Bet @@ -27,8 +27,12 @@ async def reset_database(): async with engine.begin() as conn: print("🔥 Dropping all tables...") - await conn.run_sync(Base.metadata.drop_all) + await conn.execute(text("DROP SCHEMA public CASCADE")) + await conn.execute(text("CREATE SCHEMA public")) + await conn.execute(text("GRANT ALL ON SCHEMA public TO neondb_owner")) + await conn.execute(text("GRANT ALL ON SCHEMA public TO public")) + print("🏗️ Creating new tables...") await conn.run_sync(Base.metadata.create_all) diff --git a/backend/app/scripts/seed_db.py b/backend/app/scripts/seed_db.py index 0ed5619..f99d207 100644 --- a/backend/app/scripts/seed_db.py +++ b/backend/app/scripts/seed_db.py @@ -1,36 +1,36 @@ # backend/app/scripts/seed_db.py -from uuid import uuid4 -from datetime import datetime, timedelta +from uuid import UUID, uuid4 +from datetime import datetime, timedelta, timezone from decimal import Decimal from sqlmodel import Session -from ..db.database import engine -from ..db.models import Tournament, Agent, Trade, Bet, StatusEnum, ActionEnum +from ..db.database import sync_engine, engine +from ..db.models import Tournament, Agent, Trade, Bet, StatusEnum, ActionEnum, AgentState def seed_database(): """Seed the database with test data""" - with Session(engine) as session: + with Session(sync_engine) as session: # Create Tournaments tournament1 = Tournament( id=uuid4(), - name="Q4 2024 Championship", + name="Q4 2025 Championship", status=StatusEnum.live, start_date=datetime.utcnow(), end_date=datetime.utcnow() + timedelta(days=30), prize_pool=Decimal("10000.00"), - created_at=datetime.utcnow(), + created_at=datetime.now(timezone.utc), ) tournament2 = Tournament( id=uuid4(), - name="Winter Series", + name="Spring Series", status=StatusEnum.upcoming, start_date=datetime.utcnow() + timedelta(days=7), end_date=datetime.utcnow() + timedelta(days=37), prize_pool=Decimal("5000.00"), - created_at=datetime.utcnow(), + created_at=datetime.now(timezone.utc), ) session.add(tournament1) @@ -45,7 +45,7 @@ def seed_database(): avatar_url="https://example.com/avatar1.png", stats={"win_rate": 0.65, "total_trades": 150}, memory={"last_analysis": "Bullish on tech stocks"}, - created_at=datetime.utcnow(), + created_at=datetime.now(timezone.utc), ) agent2 = Agent( @@ -56,7 +56,7 @@ def seed_database(): avatar_url="https://example.com/avatar2.png", stats={"win_rate": 0.58, "total_trades": 200}, memory={"last_analysis": "Focus on fundamentals"}, - created_at=datetime.utcnow(), + created_at=datetime.now(timezone.utc), ) agent3 = Agent( @@ -67,7 +67,7 @@ def seed_database(): avatar_url="https://example.com/avatar3.png", stats={"win_rate": 0.72, "total_trades": 500}, memory={"last_analysis": "Pattern detected in BTC"}, - created_at=datetime.utcnow(), + created_at=datetime.now(timezone.utc), ) session.add(agent1) @@ -76,6 +76,44 @@ def seed_database(): session.commit() + agent_state1 = AgentState( + agent_id=agent1.id, + tournament_id=tournament1.id, + portfolio={"USD": 10000.0}, # Starting cash + portfolio_value_usd=Decimal("10000.00"), + rank=1, + trades_count=0, + last_decision="Initial state", + updated_at=datetime.now(timezone.utc), + ) + + agent_state2 = AgentState( + agent_id=agent2.id, + tournament_id=tournament1.id, + portfolio={"USD": 10000.0}, + portfolio_value_usd=Decimal("10000.00"), + rank=2, + trades_count=0, + last_decision="Initial state", + updated_at=datetime.now(timezone.utc), + ) + + agent_state3 = AgentState( + agent_id=agent3.id, + tournament_id=tournament1.id, + portfolio={"USD": 10000.0}, + portfolio_value_usd=Decimal("10000.00"), + rank=3, + trades_count=0, + last_decision="Initial state", + updated_at=datetime.now(timezone.utc), + ) + + session.add(agent_state1) + session.add(agent_state2) + session.add(agent_state3) + session.commit() + # Create Trades trade1 = Trade( id=uuid4(), @@ -85,7 +123,7 @@ def seed_database(): asset="BTC", amount=Decimal("0.5"), price=Decimal("45000.00"), - timestamp=datetime.utcnow(), + timestamp=datetime.now(timezone.utc), ) trade2 = Trade( @@ -96,7 +134,7 @@ def seed_database(): asset="ETH", amount=Decimal("5.0"), price=Decimal("3000.00"), - timestamp=datetime.utcnow(), + timestamp=datetime.now(timezone.utc), ) trade3 = Trade( @@ -107,7 +145,7 @@ def seed_database(): asset="BTC", amount=Decimal("0.25"), price=Decimal("46000.00"), - timestamp=datetime.utcnow(), + timestamp=datetime.now(timezone.utc), ) session.add(trade1) @@ -122,7 +160,7 @@ def seed_database(): tournament_id=tournament1.id, amount=Decimal("100.00"), odds=Decimal("2.5"), - placed_at=datetime.utcnow(), + placed_at=datetime.now(timezone.utc), settled=False, ) @@ -133,7 +171,7 @@ def seed_database(): tournament_id=tournament1.id, amount=Decimal("250.00"), odds=Decimal("3.0"), - placed_at=datetime.utcnow(), + placed_at=datetime.now(timezone.utc), settled=False, )