Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions backend/app/agents/tools/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -4,6 +4,7 @@
`from app.agents.tools import MarketDataTool`.
"""
import logging
from .plan_tool import PlanTool

# Configure logging for the tools package
logging.basicConfig(
Expand All @@ -16,7 +17,6 @@
from .make_trade_tool import MakeTradeTool
from .tweet_post_tool import TweetPostTool
from .database_tool import DatabaseTool
from .plan_tool import PlanTool

# Deprecated: These tools are kept for backwards compatibility
# Use MakeTradeTool instead for simulated trading
Expand All @@ -28,8 +28,8 @@
"MakeTradeTool",
"TweetPostTool",
"DatabaseTool",
"PlanTool",
# Deprecated
"TradeTool",
"PortfolioTool",
"PlanTool"
]
10 changes: 9 additions & 1 deletion backend/app/db/database.py
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand All @@ -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
Expand All @@ -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
Expand Down
8 changes: 6 additions & 2 deletions backend/app/scripts/reset_db.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand All @@ -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)

Expand Down
72 changes: 55 additions & 17 deletions backend/app/scripts/seed_db.py
Original file line number Diff line number Diff line change
@@ -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)
Expand All @@ -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(
Expand All @@ -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(
Expand All @@ -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)
Expand All @@ -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(),
Expand All @@ -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(
Expand All @@ -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(
Expand All @@ -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)
Expand All @@ -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,
)

Expand All @@ -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,
)

Expand Down
2 changes: 1 addition & 1 deletion backend/requirements.txt
Original file line number Diff line number Diff line change
Expand Up @@ -116,7 +116,7 @@ Pygments==2.19.2
PyJWT==2.10.1
PyNaCl==1.6.0
pyparsing==3.2.5
pytest==9.0.1
pytest>=7,<9
pytest-asyncio==0.23.8
python-dateutil==2.9.0
python-dotenv==1.2.1
Expand Down
30 changes: 30 additions & 0 deletions backend/tests/test_debug_plan.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,30 @@
import sys
from pathlib import Path

ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(ROOT))


import asyncio
from uuid import UUID

from app.db.database import AsyncSessionLocal
from app.agents.executor import TradingAgent
from app.agents.tools.database_tool import DatabaseTool

AGENT_UUID = UUID("c067fa4f-8e85-4981-a9d1-a21895cf2d81")
TOURNAMENT_UUID = UUID("f082eee0-39f8-456d-8520-6d1dd2faec38")

async def run():
async with AsyncSessionLocal() as session:
agent = TradingAgent(
agent_id="Gamma Quant",
personality="aggressive",
risk_score=0.7,
agent_uuid=("c067fa4f-8e85-4981-a9d1-a21895cf2d81"),
tournament_uuid=("f082eee0-39f8-456d-8520-6d1dd2faec38"),
database_tool=DatabaseTool(session),
)
agent.make_decision()

asyncio.run(run())