package risk

import (
	"context"
	"errors"
	"fmt"
	"math"
	"sync"
	"time"
)

type Manager struct {
	mu                   sync.RWMutex
	maxPositionSize      float64
	maxDrawdown          float64
	stopLoss             float64
	takeProfit           float64
	maxDailyTrades       int
	maxDailyLoss         float64
	maxLeverage          float64
	minVolatility        float64
	maxVolatility        float64
	correlationThreshold float64
	maxConcentration     float64
	minLiquidity         float64
	maxSlippage          float64
	circuitBreaker       CircuitBreaker

	positions        map[string]Position
	dailyStats       DailyStats
	volatilityWindow []float64
	priceHistory     map[string][]PricePoint

	minTradeInterval   time.Duration
	maxTradesPerHour   int
	maxPriceDeviation  float64
	referenceExchanges []string
	maxMarginUsage     float64
	minFreeMargin      float64
	minDepthRatio      float64
	minBidAskSize      float64
	trendWindow        time.Duration
	maxTrendDeviation  float64

	volatilityControl  VolatilityControl
	correlationControl CorrelationControl
	liquidityStress    LiquidityStress

	lastTradeTime    time.Time
	hourlyTradeCount int
	hourlyTradeReset time.Time
}

type CircuitBreaker struct {
	PriceChangeThreshold float64
	TimeWindow           time.Duration
	CoolDown             time.Duration
	LastTriggered        time.Time
	IsTriggered          bool
}

type DailyStats struct {
	Date         time.Time
	TradeCount   int
	TotalPnL     float64
	HighestPrice float64
	LowestPrice  float64
}

type PricePoint struct {
	Price     float64
	Volume    float64
	Timestamp time.Time
}

type Position struct {
	Symbol        string
	EntryPrice    float64
	Amount        float64
	Leverage      float64
	UnrealizedPnL float64
	OpenTime      time.Time
}

type VolatilityControl struct {
	HistoricalWindow int
	ImpliedWeight    float64
	RealizedWeight   float64
	HistoricalVol    float64
	ImpliedVol       float64
}

type CorrelationControl struct {
	MinPairs         int
	MaxCorrelation   float64
	LookbackPeriod   int
	PairCorrelations map[string]map[string]float64
}

type LiquidityStress struct {
	ConfidenceLevel float64
	StressPeriod    int
	MaxLiquidation  float64
	StressScenarios []StressScenario
}

type StressScenario struct {
	PriceChange  float64
	VolumeChange float64
	SpreadChange float64
	Probability  float64
}

func NewManager(config map[string]float64) *Manager {
	return &Manager{
		maxPositionSize: config["max_position_size"],
		maxDrawdown:     config["max_drawdown"],
		stopLoss:        config["stop_loss"],
		takeProfit:      config["take_profit"],
		positions:       make(map[string]Position),
	}
}

func (m *Manager) CheckAndUpdatePosition(ctx context.Context, order exchange.Order) error {
	m.mu.Lock()
	defer m.mu.Unlock()

	// Check position size
	if order.Amount > m.maxPositionSize*m.currentBalance {
		return errors.New("order exceeds maximum position size")
	}

	// Check drawdown
	if m.currentBalance < m.initialBalance*(1-m.maxDrawdown) {
		return errors.New("maximum drawdown reached")
	}

	return nil
}

func (m *Manager) UpdatePositionPrice(symbol string, currentPrice float64) (bool, error) {
	m.mu.Lock()
	defer m.mu.Unlock()

	pos, exists := m.positions[symbol]
	if !exists {
		return false, nil
	}

	// Calculate unrealized PnL
	pnlPercent := (currentPrice - pos.EntryPrice) / pos.EntryPrice

	// Check stop loss
	if pnlPercent <= -m.stopLoss {
		return true, nil // Should close position
	}

	// Check take profit
	if pnlPercent >= m.takeProfit {
		return true, nil // Should close position
	}

	return false, nil
}

func (m *Manager) CheckRisk(ctx context.Context, order exchange.Order, marketData exchange.MarketData) error {
	m.mu.Lock()
	defer m.mu.Unlock()

	// 1. 检查每日交易次数限制
	if m.dailyStats.TradeCount >= m.maxDailyTrades {
		return errors.New("daily trade limit exceeded")
	}

	// 2. 检查每日亏损限制
	if m.dailyStats.TotalPnL <= -m.maxDailyLoss*m.initialBalance {
		return errors.New("daily loss limit reached")
	}

	// 3. 检查杠杆率
	if order.Leverage > m.maxLeverage {
		return errors.New("leverage exceeds maximum allowed")
	}

	// 4. 检查波动率
	volatility := m.calculateVolatility()
	if volatility < m.minVolatility {
		return errors.New("market volatility too low")
	}
	if volatility > m.maxVolatility {
		return errors.New("market volatility too high")
	}

	// 5. 检查流动性
	if !m.checkLiquidity(marketData) {
		return errors.New("insufficient market liquidity")
	}

	// 6. 检查集中度
	if !m.checkConcentration(order) {
		return errors.New("position concentration too high")
	}

	// 7. 检查熔断机制
	if m.checkCircuitBreaker(marketData) {
		return errors.New("circuit breaker triggered")
	}

	// 8. 检查滑点
	if !m.checkSlippage(order, marketData) {
		return errors.New("expected slippage too high")
	}

	// 交易频率检查
	if err := m.checkTradeFrequency(); err != nil {
		return err
	}

	// 价格偏离度检查
	if err := m.checkPriceDeviation(marketData); err != nil {
		return err
	}

	// 保证金使用率检查
	if err := m.checkMarginUsage(order); err != nil {
		return err
	}

	// 订单簿深度检查
	if err := m.checkOrderBookDepth(marketData); err != nil {
		return err
	}

	// 趋势偏离检查
	if err := m.checkTrendDeviation(marketData); err != nil {
		return err
	}

	// 综合波动率检查
	if err := m.checkCompositeVolatility(marketData); err != nil {
		return err
	}

	// 相关性风险检查
	if err := m.checkCorrelationRisk(marketData); err != nil {
		return err
	}

	// 流动性压力测试
	if err := m.checkLiquidityStress(order, marketData); err != nil {
		return err
	}

	return nil
}

func (m *Manager) calculateVolatility() float64 {
	if len(m.volatilityWindow) < 2 {
		return 0
	}

	// 计算对数收益率
	returns := make([]float64, len(m.volatilityWindow)-1)
	for i := 1; i < len(m.volatilityWindow); i++ {
		returns[i-1] = math.Log(m.volatilityWindow[i] / m.volatilityWindow[i-1])
	}

	// 计算标准差
	mean := 0.0
	for _, r := range returns {
		mean += r
	}
	mean /= float64(len(returns))

	variance := 0.0
	for _, r := range returns {
		variance += math.Pow(r-mean, 2)
	}
	variance /= float64(len(returns))

	return math.Sqrt(variance)
}

func (m *Manager) checkLiquidity(data exchange.MarketData) bool {
	// 检查24小时成交量是否满足最小流动性要求
	return data.Volume*data.Price >= m.minLiquidity
}

func (m *Manager) checkConcentration(order exchange.Order) bool {
	totalPositionValue := 0.0
	for _, pos := range m.positions {
		totalPositionValue += pos.Amount * pos.EntryPrice
	}

	// 添加新订单的价值
	newPositionValue := order.Amount * order.Price
	totalPositionValue += newPositionValue

	// 检查单个仓位是否超过总仓位的最大集中度
	for _, pos := range m.positions {
		positionValue := pos.Amount * pos.EntryPrice
		if positionValue/totalPositionValue > m.maxConcentration {
			return false
		}
	}

	return true
}

func (m *Manager) checkCircuitBreaker(data exchange.MarketData) bool {
	if m.circuitBreaker.IsTriggered {
		if time.Since(m.circuitBreaker.LastTriggered) > m.circuitBreaker.CoolDown {
			m.circuitBreaker.IsTriggered = false
			return false
		}
		return true
	}

	// 检查价格变化
	priceHistory := m.priceHistory[data.Symbol]
	if len(priceHistory) == 0 {
		return false
	}

	timeWindow := time.Now().Add(-m.circuitBreaker.TimeWindow)
	var oldPrice float64
	for i := len(priceHistory) - 1; i >= 0; i-- {
		if priceHistory[i].Timestamp.Before(timeWindow) {
			oldPrice = priceHistory[i].Price
			break
		}
	}

	if oldPrice > 0 {
		priceChange := math.Abs(data.Price-oldPrice) / oldPrice
		if priceChange > m.circuitBreaker.PriceChangeThreshold {
			m.circuitBreaker.IsTriggered = true
			m.circuitBreaker.LastTriggered = time.Now()
			return true
		}
	}

	return false
}

func (m *Manager) checkSlippage(order exchange.Order, data exchange.MarketData) bool {
	expectedSlippage := math.Abs(order.Price-data.Price) / data.Price
	return expectedSlippage <= m.maxSlippage
}

func (m *Manager) UpdateDailyStats(pnl float64) {
	today := time.Now().UTC().Truncate(24 * time.Hour)
	if m.dailyStats.Date != today {
		// Reset daily stats
		m.dailyStats = DailyStats{
			Date:         today,
			TradeCount:   0,
			TotalPnL:     0,
			HighestPrice: 0,
			LowestPrice:  math.MaxFloat64,
		}
	}

	m.dailyStats.TradeCount++
	m.dailyStats.TotalPnL += pnl
}

func (m *Manager) UpdatePriceHistory(data exchange.MarketData) {
	if m.priceHistory == nil {
		m.priceHistory = make(map[string][]PricePoint)
	}

	pricePoint := PricePoint{
		Price:     data.Price,
		Volume:    data.Volume,
		Timestamp: data.Timestamp,
	}

	// 保持价格历史在合理范围内
	history := m.priceHistory[data.Symbol]
	if len(history) > 1000 {
		history = history[1:]
	}
	history = append(history, pricePoint)
	m.priceHistory[data.Symbol] = history

	// 更新波动率窗口
	if len(m.volatilityWindow) > 100 {
		m.volatilityWindow = m.volatilityWindow[1:]
	}
	m.volatilityWindow = append(m.volatilityWindow, data.Price)
}

func (m *Manager) checkTradeFrequency() error {
	now := time.Now()

	// 检查最小交易间隔
	if now.Sub(m.lastTradeTime) < m.minTradeInterval {
		return errors.New("trade frequency too high")
	}

	// 检查每小时交易次数
	if now.Sub(m.hourlyTradeReset) >= time.Hour {
		m.hourlyTradeCount = 0
		m.hourlyTradeReset = now
	}
	if m.hourlyTradeCount >= m.maxTradesPerHour {
		return errors.New("hourly trade limit exceeded")
	}

	return nil
}

func (m *Manager) checkPriceDeviation(data exchange.MarketData) error {
	// 获取参考价格
	var prices []float64
	for _, ex := range m.referenceExchanges {
		price, err := m.getPriceFromExchange(ex, data.Symbol)
		if err != nil {
			continue
		}
		prices = append(prices, price)
	}

	if len(prices) == 0 {
		return nil // 无法获取参考价格时暂时跳过检查
	}

	// 计算平均参考价格
	avgPrice := 0.0
	for _, p := range prices {
		avgPrice += p
	}
	avgPrice /= float64(len(prices))

	// 检查价格偏离度
	deviation := math.Abs(data.Price-avgPrice) / avgPrice
	if deviation > m.maxPriceDeviation {
		return fmt.Errorf("price deviation %.2f%% exceeds limit", deviation*100)
	}

	return nil
}

func (m *Manager) checkCompositeVolatility(data exchange.MarketData) error {
	// 计算历史波动率
	historicalVol := m.calculateHistoricalVolatility()

	// 获取期权隐含波动率（如果可用）
	impliedVol := m.getImpliedVolatility(data.Symbol)

	// 计算综合波动率
	compositeVol := historicalVol*m.volatilityControl.RealizedWeight +
		impliedVol*m.volatilityControl.ImpliedWeight

	if compositeVol > m.maxVolatility {
		return fmt.Errorf("composite volatility %.2f%% too high", compositeVol*100)
	}

	return nil
}

func (m *Manager) checkCorrelationRisk(data exchange.MarketData) error {
	// 更新相关性矩阵
	m.updateCorrelationMatrix(data)

	// 检查是否有足够的对冲币对
	hedgePairs := m.findHedgePairs(data.Symbol)
	if len(hedgePairs) < m.correlationControl.MinPairs {
		return errors.New("insufficient hedge pairs")
	}

	// 检查相关性是否在允许范围内
	for _, pair := range hedgePairs {
		if corr := m.getCorrelation(data.Symbol, pair); corr > m.correlationControl.MaxCorrelation {
			return fmt.Errorf("correlation with %s too high: %.2f", pair, corr)
		}
	}

	return nil
}

func (m *Manager) checkLiquidityStress(order exchange.Order, data exchange.MarketData) error {
	// 运行压力测试场景
	for _, scenario := range m.liquidityStress.StressScenarios {
		// 计算在压力情况下的清算风险
		liquidationRisk := m.calculateLiquidationRisk(order, data, scenario)
		if liquidationRisk > m.liquidityStress.MaxLiquidation {
			return fmt.Errorf("stress test liquidation risk %.2f%% too high", liquidationRisk*100)
		}
	}

	return nil
}
