# ============================================================
# MT5 Trade Stream Server - FastAPI Application (v2)
# Supports: trade events + account snapshots + open positions
# ============================================================

from contextlib import asynccontextmanager
from fastapi import FastAPI, HTTPException, Query
from fastapi.middleware.cors import CORSMiddleware
from pydantic import BaseModel
from typing import Optional
import asyncio

from config import HOST, PORT, API_KEY, MAX_RECORDS_PER_REQUEST
from database import (
    init_db, insert_trade, get_trades, get_recent_trades, get_stats,
    cleanup_old_data, upsert_account_snapshot, get_all_accounts,
    replace_open_positions, get_open_positions,
)


# ── Background tasks ────────────────────────────────────────
async def periodic_cleanup():
    while True:
        await asyncio.sleep(6 * 3600)
        try:
            deleted = cleanup_old_data()
            if deleted > 0:
                print(f"[Cleanup] Deleted {deleted} old records")
        except Exception as e:
            print(f"[Cleanup Error] {e}")


@asynccontextmanager
async def lifespan(app: FastAPI):
    init_db()
    print(f"✅ MT5 Stream Server v2 started on {HOST}:{PORT}")
    print(f"📡 API Key: {API_KEY[:8]}{'*' * (len(API_KEY) - 8)}")
    task = asyncio.create_task(periodic_cleanup())
    yield
    task.cancel()


# ── App ──────────────────────────────────────────────────────
app = FastAPI(
    title="MT5 Trade Stream Server",
    version="2.0.0",
    lifespan=lifespan,
)
app.add_middleware(
    CORSMiddleware,
    allow_origins=["*"],
    allow_credentials=True,
    allow_methods=["*"],
    allow_headers=["*"],
)


# ── Models ───────────────────────────────────────────────────
class TradeEvent(BaseModel):
    api_key: str
    account_id: str
    ticket: int
    symbol: str
    type: str
    action: str
    lots: float
    price: float
    sl: float = 0
    tp: float = 0
    profit: float = 0
    comment: str = ""
    trade_time: str


class PositionData(BaseModel):
    ticket: int
    symbol: str
    type: str
    lots: float
    open_price: float
    current_price: float = 0
    sl: float = 0
    tp: float = 0
    profit: float = 0
    swap: float = 0
    comment: str = ""
    open_time: str


class AccountSnapshot(BaseModel):
    api_key: str
    account_id: str
    account_name: str = ""
    broker: str = ""
    balance: float = 0
    equity: float = 0
    margin: float = 0
    free_margin: float = 0
    margin_level: float = 0
    floating_profit: float = 0
    total_profit: float = 0
    net_deposit: float = 0
    positions: list[PositionData] = []


# ── POST: Trade event (from EA OnTradeTransaction) ──────────
@app.post("/api/trade")
async def receive_trade(event: TradeEvent):
    if event.api_key != API_KEY:
        raise HTTPException(status_code=403, detail="Invalid API key")
    try:
        data = event.model_dump()
        data.pop("api_key")
        row_id = insert_trade(data)
        return {"success": True, "id": row_id,
                "message": f"{event.action} {event.type} {event.symbol}"}
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))


# ── POST: Account snapshot (from EA OnTimer every 5s) ───────
@app.post("/api/snapshot")
async def receive_snapshot(snapshot: AccountSnapshot):
    if snapshot.api_key != API_KEY:
        raise HTTPException(status_code=403, detail="Invalid API key")
    try:
        # Upsert account info
        account_data = {
            "account_id": snapshot.account_id,
            "account_name": snapshot.account_name,
            "broker": snapshot.broker,
            "balance": snapshot.balance,
            "equity": snapshot.equity,
            "margin": snapshot.margin,
            "free_margin": snapshot.free_margin,
            "margin_level": snapshot.margin_level,
            "floating_profit": snapshot.floating_profit,
            "total_profit": snapshot.total_profit,
            "net_deposit": snapshot.net_deposit,
            "open_positions": len(snapshot.positions),
        }
        upsert_account_snapshot(account_data)

        # Replace open positions
        positions = [p.model_dump() for p in snapshot.positions]
        replace_open_positions(snapshot.account_id, positions)

        return {"success": True,
                "message": f"Account {snapshot.account_id}: {len(positions)} positions"}
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))


# ── GET: All accounts with latest snapshot ───────────────────
@app.get("/api/accounts")
async def list_accounts():
    try:
        accounts = get_all_accounts()
        return {"success": True, "accounts": accounts}
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))


# ── GET: Open positions ─────────────────────────────────────
@app.get("/api/positions")
async def list_positions(account: Optional[str] = Query(None)):
    try:
        positions = get_open_positions(account)
        return {"success": True, "count": len(positions), "positions": positions}
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))


# ── GET: Full dashboard data (accounts + positions + stats) ──
@app.get("/api/dashboard")
async def dashboard():
    try:
        stats = get_stats()
        positions = get_open_positions()
        return {"success": True, **stats, "all_positions": positions}
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))


# ── GET: Trade history ──────────────────────────────────────
@app.get("/api/trades")
async def fetch_trades(
    since_id: int = Query(0),
    account: Optional[str] = Query(None),
    limit: int = Query(100, le=MAX_RECORDS_PER_REQUEST),
):
    try:
        trades = get_trades(since_id=since_id, account=account, limit=limit)
        return {"success": True, "count": len(trades), "trades": trades}
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))


@app.get("/api/trades/recent")
async def recent_trades(limit: int = Query(50, le=MAX_RECORDS_PER_REQUEST)):
    try:
        return {"success": True, "trades": get_recent_trades(limit)}
    except Exception as e:
        raise HTTPException(status_code=500, detail=str(e))


# ── DELETE: Cleanup ──────────────────────────────────────────
@app.delete("/api/trades/cleanup")
async def manual_cleanup(api_key: str = Query(...)):
    if api_key != API_KEY:
        raise HTTPException(status_code=403, detail="Invalid API key")
    return {"success": True, "deleted": cleanup_old_data()}


# ── Health ───────────────────────────────────────────────────
@app.get("/")
async def health():
    return {"status": "running", "service": "MT5 Trade Stream Server", "version": "2.0.0"}


if __name__ == "__main__":
    import uvicorn
    uvicorn.run("server:app", host=HOST, port=PORT, reload=False)
