"""
FastAPI Backend Server with WebSocket Support

This server runs the simulation loop and broadcasts updates to connected clients.
The "Schrödinger's Market" architecture: simulation only runs when observers are connected.
"""

import asyncio
import os
from contextlib import asynccontextmanager
from datetime import datetime
from typing import Set

# Load environment variables from .env file
from dotenv import load_dotenv
load_dotenv()

from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.middleware.cors import CORSMiddleware
from fastapi.responses import JSONResponse

from .world_state import WorldState, SocialMediaPost
from .ai_agents import create_agent_swarm, NarratorAgent, SocialMediaAgent


# Company definitions (matching frontend)
COMPANIES = [
    {"ticker": "OMEGA", "name": "Omegacorp Industries", "basePrice": 100.0, "volatility": 0.25},
    {"ticker": "NEXUS", "name": "Nexus Dynamics", "basePrice": 245.0, "volatility": 0.35},
    {"ticker": "TITAN", "name": "Titan Heavy Industries", "basePrice": 78.0, "volatility": 0.18},
    {"ticker": "PULSE", "name": "Pulse BioTech", "basePrice": 156.0, "volatility": 0.42},
    {"ticker": "VORTX", "name": "Vortex Energy", "basePrice": 52.0, "volatility": 0.28},
    {"ticker": "QUBIT", "name": "Qubit Computing", "basePrice": 320.0, "volatility": 0.55},
]

# Global state - track all companies
worlds = {
    company["ticker"]: WorldState(
        ticker=company["ticker"],
        company_name=company["name"],
        initial_price=company["basePrice"],
        initial_volatility=company["volatility"]
    )
    for company in COMPANIES
}
# Keep OMEGA as primary for backward compatibility
world = worlds["OMEGA"]
agents = create_agent_swarm()
connected_clients: Set[WebSocket] = set()
simulation_task = None

# Configuration
TICK_INTERVAL = 1.0  # Seconds between simulation ticks (much faster for active trading)
NEWS_INTERVAL = 5  # Generate news every N ticks (more frequent news)
TRADE_INTERVAL = 1  # Agents trade every tick

# Rate limiting for API calls
MAX_API_CALLS_PER_MINUTE = 14  # Stay under Gemini's 15 RPM limit
api_call_count = 0
last_api_reset = datetime.now()


async def broadcast(message: dict):
    """Send message to all connected WebSocket clients"""
    if not connected_clients:
        return
    
    disconnected = set()
    for client in connected_clients:
        try:
            await client.send_json(message)
        except Exception:
            disconnected.add(client)
    
    # Clean up disconnected clients
    for client in disconnected:
        connected_clients.discard(client)


async def check_rate_limit() -> bool:
    """Check if we can make an API call (rate limiting)"""
    global api_call_count, last_api_reset
    
    now = datetime.now()
    if (now - last_api_reset).seconds >= 60:
        api_call_count = 0
        last_api_reset = now
    
    if api_call_count >= MAX_API_CALLS_PER_MINUTE:
        return False
    
    api_call_count += 1
    return True


async def simulation_loop():
    """
    The main simulation heartbeat.
    
    This loop:
    1. Advances market time (price movement)
    2. Periodically generates news events
    3. Has agents make trading decisions
    4. Broadcasts state to all clients
    """
    global world
    
    print("🚀 Simulation loop started")
    world.is_running = True
    
    try:
        while connected_clients:  # Run only while clients are connected
            tick_start = datetime.now()
            
            # 1. Advance market time for all companies
            for ticker, w in worlds.items():
                w.tick()
            
            # 2. Generate news periodically (mix of macro and micro)
            # Use OMEGA's tick count as the master counter for news generation
            if world.tick_count % NEWS_INTERVAL == 0:
                if await check_rate_limit():
                    narrator = agents["narrator"]
                    if isinstance(narrator, NarratorAgent):
                        # Generate macro news every other interval, micro otherwise
                        import random
                        # Force macro on even multiples, micro on odd
                        news_cycle = world.tick_count // NEWS_INTERVAL
                        is_macro_tick = (news_cycle % 2) == 0
                        news_type = "macro" if is_macro_tick else "micro"
                        
                        print(f"📰 Generating {news_type.upper()} news at tick {world.tick_count} (cycle {news_cycle})")
                        
                        # Generate event - use fallback if AI fails to ensure we get news
                        try:
                            event = await narrator.generate_event(world, news_type)
                            # FORCE category to match requested type
                            event.category = news_type
                            # FORCE affected_tickers for macro
                            if news_type == "macro":
                                event.affected_tickers = ["ALL"]
                                print(f"✅ MACRO NEWS: {event.headline[:60]}... (impact: {event.impact_score:+.2f})")
                            else:
                                # For micro, ensure we have specific tickers
                                if not event.affected_tickers or "ALL" in event.affected_tickers:
                                    all_tickers = list(worlds.keys())
                                    event.affected_tickers = random.sample(all_tickers, random.randint(1, 2))
                                print(f"✅ MICRO NEWS: {event.headline[:60]}... → {event.affected_tickers} (impact: {event.impact_score:+.2f})")
                        except Exception as e:
                            print(f"❌ Error generating {news_type} news: {e}")
                            # Use fallback - this will have correct category
                            event = narrator._generate_fallback_news(news_type)
                            print(f"✅ FALLBACK {news_type.upper()} NEWS: {event.headline[:60]}... (impact: {event.impact_score:+.2f})")
                        
                        # Apply news impact to affected companies
                        # Macro news with "ALL" affects everyone
                        if event.category == "macro" and "ALL" in event.affected_tickers:
                            for ticker, w in worlds.items():
                                w.apply_news_impact(event)
                        else:
                            # Micro news affects specific tickers
                            for ticker in event.affected_tickers:
                                if ticker in worlds:
                                    worlds[ticker].apply_news_impact(event)
                        
                        # Broadcast news event
                        await broadcast({
                            "type": "news",
                            "data": event.to_dict()
                        })
                        
                        # Generate social media posts reacting to news (1-2 posts per news event)
                        if await check_rate_limit():
                            social_agent = agents.get("social_media")
                            if isinstance(social_agent, SocialMediaAgent):
                                import random
                                num_posts = random.randint(1, 2)
                                for _ in range(num_posts):
                                    post = await social_agent.generate_post(event, world)
                                    await broadcast({
                                        "type": "social_media",
                                        "data": post.to_dict()
                                    })
            
            # 3. Agent trading decisions (cycle through traders)
            if world.tick_count % TRADE_INTERVAL == 0:
                # Get all trading agents (exclude narrator)
                trading_agents = [k for k in agents.keys() if k != "narrator"]
                # Pick 2-3 random agents to trade this tick
                import random
                num_traders = random.randint(2, min(4, len(trading_agents)))
                selected_agents = random.sample(trading_agents, num_traders)
                
                for agent_key in selected_agents:
                    if await check_rate_limit():
                        agent = agents[agent_key]
                        try:
                            # Agents trade RANDOM companies, not just OMEGA
                            import random
                            available_tickers = list(worlds.keys())
                            selected_ticker = random.choice(available_tickers)
                            selected_world = worlds[selected_ticker]
                            
                            trade = await agent.decide(selected_world)
                            if trade:
                                # Update trade with correct ticker
                                trade.ticker = selected_ticker
                                selected_world.apply_trade(trade)
                                
                                # Broadcast trade immediately
                                await broadcast({
                                    "type": "trade",
                                    "data": trade.to_dict()
                                })
                                print(f"📊 {trade.agent_name} ({trade.agent_firm}): {trade.action} {trade.quantity} {selected_ticker} @ ${trade.price:.2f}")
                        except Exception as e:
                            print(f"Agent {agent_key} error: {e}")
            
            # 4. Broadcast full state update for all companies
            await broadcast({
                "type": "state",
                "data": {
                    "companies": {ticker: w.to_dict() for ticker, w in worlds.items()},
                    "tick": world.tick_count,
                    "is_running": world.is_running,
                    "observers": world.active_observers,
                    "tick_interval": TICK_INTERVAL  # Include tick interval
                }
            })
            
            # 5. Calculate sleep time to maintain tick interval
            elapsed = (datetime.now() - tick_start).total_seconds()
            sleep_time = max(0.1, TICK_INTERVAL - elapsed)
            await asyncio.sleep(sleep_time)
            
    except asyncio.CancelledError:
        print("🛑 Simulation loop cancelled")
    finally:
        world.is_running = False
        print("💤 Simulation loop stopped - no observers")


async def start_simulation():
    """Start the simulation if not already running"""
    global simulation_task
    
    if simulation_task is None or simulation_task.done():
        simulation_task = asyncio.create_task(simulation_loop())


async def stop_simulation():
    """Stop the simulation gracefully"""
    global simulation_task
    
    if simulation_task and not simulation_task.done():
        simulation_task.cancel()
        try:
            await simulation_task
        except asyncio.CancelledError:
            pass
        simulation_task = None


# FastAPI App
@asynccontextmanager
async def lifespan(app: FastAPI):
    """Startup and shutdown events"""
    print("🌍 AI Market Simulation Server Starting...")
    yield
    print("🌍 Server shutting down...")
    await stop_simulation()


app = FastAPI(
    title="AI Market Simulation",
    description="An agentic AI-driven fictional market simulation",
    version="1.0.0",
    lifespan=lifespan
)

# CORS for frontend
app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],  # In production, specify your frontend URL
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)


# REST Endpoints
@app.get("/")
async def root():
    """Health check and info"""
    return {
        "status": "running",
        "simulation_active": world.is_running,
        "observers": len(connected_clients),
        "tick": world.tick_count,
        "companies": {ticker: {"price": w.price, "tick": w.tick_count} for ticker, w in worlds.items()}
    }


@app.get("/api/state")
async def get_state():
    """Get current market state (REST fallback)"""
    return {
        "companies": {ticker: w.to_dict() for ticker, w in worlds.items()},
        "tick": world.tick_count,
        "is_running": world.is_running,
        "observers": world.active_observers,
        "tick_interval": TICK_INTERVAL  # Include tick interval in response
    }


@app.get("/api/agents")
async def get_agents():
    """Get info about trading agents"""
    return {
        agent_id: {
            "name": agent.name,
            "personality": agent.personality,
            "cash": agent.cash,
            "trade_count": agent.trade_count
        }
        for agent_id, agent in agents.items()
    }


@app.post("/api/reset")
async def reset_simulation():
    """Reset the simulation to initial state"""
    global worlds, world
    
    await stop_simulation()
    
    # Reset all companies to base prices
    for company in COMPANIES:
        ticker = company["ticker"]
        if ticker in worlds:
            # Reset the world state
            worlds[ticker] = WorldState(
                ticker=company["ticker"],
                company_name=company["name"],
                initial_price=company["basePrice"],
                initial_volatility=company["volatility"]
            )
    
    world = worlds["OMEGA"]
    
    # Reset agents
    for agent in agents.values():
        agent.cash = 100_000.0
        agent.positions = []
        agent.trade_count = 0
    
    # Clear news and trade history
    for w in worlds.values():
        w.news_history = []
        w.trade_history = []
        w.price_history = []
        w.active_news_impacts = []
        w._record_price(0)  # Record initial price
    
    if connected_clients:
        await start_simulation()
    
    return {
        "status": "reset", 
        "tick": world.tick_count,
        "companies": {ticker: {"price": w.price} for ticker, w in worlds.items()}
    }


@app.post("/api/user-trade")
async def user_trade(request: dict):
    """Apply user trade to market - REALISTIC LIMITED MARKET IMPACT"""
    from .world_state import Trade
    from datetime import datetime
    
    ticker = request.get("ticker", "OMEGA")
    action = request.get("action", "BUY")  # "BUY" or "SELL"
    quantity = request.get("quantity", 0)
    price = request.get("price", 0.0)
    
    if ticker not in worlds:
        return {"success": False, "message": f"Ticker {ticker} not found"}
    
    world = worlds[ticker]
    
    # Create trade object with user identifier
    trade = Trade(
        id=f"user-{datetime.now().timestamp()}",
        timestamp=datetime.now(),
        agent_id="user",  # This triggers reduced impact in apply_trade
        agent_name="User",
        agent_firm="User Account",
        agent_avatar="👤",
        action=action,
        quantity=quantity,
        price=price or world.price,
        reasoning="User trade",
        ticker=ticker
    )
    
    # Apply trade impact to market (with limited impact for users)
    old_price = world.price
    world.apply_trade(trade)
    price_change = world.price - old_price
    price_change_pct = (price_change / old_price) * 100 if old_price > 0 else 0
    
    # Broadcast the trade impact
    await broadcast({
        "type": "user_trade",
        "data": {
            "ticker": ticker,
            "action": action,
            "quantity": quantity,
            "price": world.price,
            "price_change": price_change,
            "price_change_pct": price_change_pct,
            "market_impact": True
        }
    })
    
    return {
        "success": True,
        "message": f"Trade executed. {ticker} ${old_price:.2f} → ${world.price:.2f} ({price_change_pct:+.2f}%)",
        "new_price": world.price,
        "price_change": price_change,
        "price_change_pct": price_change_pct
    }


# WebSocket Endpoint
@app.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
    """
    WebSocket connection for real-time updates.
    
    The simulation uses "Schrödinger's Market" architecture:
    - When a client connects, the simulation wakes up
    - When all clients disconnect, simulation pauses (saves API costs)
    """
    await websocket.accept()
    connected_clients.add(websocket)
    world.active_observers = len(connected_clients)
    
    print(f"👁️ Observer connected. Total: {len(connected_clients)}")
    
    # Start simulation if this is the first client
    if len(connected_clients) == 1:
        await start_simulation()
    
    # Send initial state for all companies
    await websocket.send_json({
        "type": "init",
        "data": {
            "companies": {ticker: w.to_dict() for ticker, w in worlds.items()},
            "tick": world.tick_count,
            "is_running": world.is_running,
            "observers": world.active_observers
        }
    })
    
    try:
        while True:
            # Listen for client messages (commands)
            data = await websocket.receive_json()
            
            if data.get("type") == "ping":
                await websocket.send_json({"type": "pong"})
            
            elif data.get("type") == "get_explanation":
                # Educational feature: explain a concept
                topic = data.get("topic", "")
                explanation = get_explanation(topic)
                await websocket.send_json({
                    "type": "explanation",
                    "topic": topic,
                    "data": explanation
                })
                
    except WebSocketDisconnect:
        pass
    finally:
        connected_clients.discard(websocket)
        world.active_observers = len(connected_clients)
        print(f"👁️ Observer disconnected. Total: {len(connected_clients)}")
        
        # Stop simulation if no clients
        if not connected_clients:
            await stop_simulation()


# Educational Content
EXPLANATIONS = {
    "delta": {
        "title": "Delta (Δ)",
        "formula": "Δ = ∂V/∂S",
        "description": "Measures how much an option's price changes when the underlying stock moves $1.",
        "insight": "A Delta of 0.50 means if OMEGA goes up $1, your call option goes up $0.50. It's also approximately the probability of expiring in-the-money.",
        "trader_tip": "Professional traders 'delta hedge' by holding stock to offset their option delta, making their position neutral to small price moves."
    },
    "gamma": {
        "title": "Gamma (Γ)",
        "formula": "Γ = ∂²V/∂S² = ∂Δ/∂S",
        "description": "Measures how fast Delta changes as the stock price moves.",
        "insight": "High Gamma means your Delta is unstable - small price moves cause big changes in your exposure. Gamma is highest for at-the-money options near expiration.",
        "trader_tip": "Gamma is the 'curvature' of your P&L. Long gamma positions profit from big moves; short gamma positions profit from stability."
    },
    "theta": {
        "title": "Theta (Θ)",
        "formula": "Θ = -∂V/∂τ",
        "description": "Measures how much value an option loses each day due to time decay.",
        "insight": "Theta is always working against option buyers. A Theta of -0.05 means you lose $0.05 per day just by holding the option.",
        "trader_tip": "Option sellers love Theta - they collect premium as time passes. Buyers fight against it."
    },
    "vega": {
        "title": "Vega (ν)",
        "formula": "ν = ∂V/∂σ",
        "description": "Measures sensitivity to changes in implied volatility.",
        "insight": "If IV increases by 1%, the option price changes by the Vega amount. High Vega means you're betting on volatility changing.",
        "trader_tip": "Before earnings, IV typically rises (Vega helps you). After earnings, IV crushes (Vega hurts you). This is the 'volatility crush'."
    },
    "implied_volatility": {
        "title": "Implied Volatility (IV)",
        "formula": "σ_implied from option prices",
        "description": "The market's expectation of future price movement, derived from option prices.",
        "insight": "IV is mean-reverting. When IV is historically high, options are 'expensive' - good time to sell. When IV is low, options are 'cheap' - good time to buy.",
        "trader_tip": "IV Rank compares current IV to its historical range. IV Rank > 50 = elevated volatility = consider selling premium."
    },
    "black_scholes": {
        "title": "Black-Scholes Model",
        "formula": "C = S·N(d₁) - K·e^(-rT)·N(d₂)",
        "description": "The foundational options pricing model that assumes stock prices follow geometric Brownian motion with constant volatility.",
        "insight": "Black-Scholes assumes European exercise, constant volatility, and no dividends. Real markets violate all these assumptions, which is why we see volatility smiles and skew.",
        "trader_tip": "Black-Scholes is the 'Model T' of options - foundational but limited. Modern traders use it as a starting point, then adjust for real-world effects."
    },
    "sentiment": {
        "title": "Market Sentiment",
        "formula": "Sentiment Index: 0 (Fear) to 1 (Greed)",
        "description": "A measure of the overall mood and emotion in the market, often derived from options flow, VIX, and price action.",
        "insight": "Extreme fear often marks market bottoms (buying opportunity). Extreme greed often marks tops (time to hedge).",
        "trader_tip": "Warren Buffett: 'Be fearful when others are greedy, and greedy when others are fearful.' Contrarian sentiment trading can be powerful."
    }
}


def get_explanation(topic: str) -> dict:
    """Get educational explanation for a topic"""
    topic_lower = topic.lower().replace(" ", "_")
    return EXPLANATIONS.get(topic_lower, {
        "title": topic,
        "description": "Explanation not available for this topic.",
        "insight": "Try asking about: delta, gamma, theta, vega, implied_volatility, black_scholes, sentiment"
    })


# Run with: uvicorn backend.server:app --reload
if __name__ == "__main__":
    import uvicorn
    port = int(os.getenv("PORT", 8000))
    uvicorn.run(app, host="0.0.0.0", port=port)

