import numpy as np
import yfinance as yf
import torch
import math
import matplotlib.pyplot as plt

# ============================================================
# FIXED CCT-GRADIENT BTC FORECAST
# Key fix: Use gradient descent to find equilibrium, then
# project forward using discovered dynamics
# ============================================================

def find_equilibrium(prices, lr=0.05, num_restarts=10, max_iter=1000):
    """
    Find the equilibrium gap d* using gradient descent on loss.
    
    L(d) = sin²(π·s) where s = √(d² + 4P)
    
    The equilibrium is where s is integer (stable market state).
    """
    P = prices[-1]
    P_t = torch.tensor(float(P), requires_grad=False)
    
    def loss_fn(d):
        s = torch.sqrt(d**2 + 4.0 * P_t)
        loss = torch.sin(torch.pi * s) ** 2
        return loss, s
    
    best_d = None
    best_loss = float('inf')
    
    for restart in range(num_restarts):
        # Initialize near the expected gap (sqrt(P) scale)
        d = torch.randn(1) * np.sqrt(P) * 0.5
        d = d.detach().requires_grad_()
        optimizer = torch.optim.Adam([d], lr=lr)
        
        for step in range(max_iter):
            optimizer.zero_grad()
            loss, s = loss_fn(d)
            loss.backward()
            optimizer.step()
            
            with torch.no_grad():
                d.clamp_(min=1.0)  # Minimum gap of 1
            
            if loss.item() < best_loss:
                best_loss = loss.item()
                best_d = d.item()
            
            if loss.item() < 1e-12:
                break
    
    # Find the integer s closest to equilibrium
    d_star = best_d
    s_sq = d_star**2 + 4 * P
    s_float = math.sqrt(s_sq)
    s_nearest = round(s_float)
    
    # Equilibrium price (fair value) from factorization analogy
    # P ≈ (s - d)/2 × (s + d)/2 when s,d are integers
    # But for BTC, interpret d as "distance from fair value"
    equilibrium_price = P * (s_nearest / s_float) ** 2
    
    return {
        'current_price': P,
        'equilibrium_price': equilibrium_price,
        'gap_d': d_star,
        's_float': s_float,
        's_nearest': s_nearest,
        'loss': best_loss,
        'convergence_ratio': s_float / s_nearest  # How close to integer
    }


def cct_gradient_forecast(prices, forecast_days=40, lookback=100):
    """
    CCT-Gradient BTC Forecast using equilibrium dynamics.
    
    Strategy:
    1. Find current equilibrium via gradient descent
    2. Compute momentum from recent price history
    3. Project forward as damped oscillation toward equilibrium
    """
    
    # Step 1: Analyze recent history for momentum
    recent = prices[-lookback:]
    
    # Compute returns and volatility
    returns = np.diff(np.log(recent))
    mean_return = np.mean(returns)
    volatility = np.std(returns)
    
    # Step 2: Find equilibrium for each historical point
    equilibria = []
    for i in range(0, len(recent), 10):  # Sample every 10 days
        window = recent[:i+1] if i > 0 else recent
        eq = find_equilibrium(window, num_restarts=5, max_iter=500)
        equilibria.append({
            'day': i,
            'price': window[-1],
            'equilibrium': eq['equilibrium_price'],
            'gap': eq['gap_d'],
            'loss': eq['loss']
        })
    
    # Step 3: Compute the "pull" toward equilibrium
    current_eq = equilibria[-1]['equilibrium']
    current_price = prices[-1]
    
    # The gap tells us how much "correction" is needed
    # If equilibrium > current: likely to rise
    # If equilibrium < current: likely to fall
    pull_strength = (current_eq - current_price) / current_price
    
    # Step 4: Project forward using damped oscillation
    # Markets oscillate around equilibrium (like ODE periodicity)
    forecasts = []
    
    # Damping factor (market friction)
    damping = 0.95
    
    for day in range(forecast_days):
        # Damped approach to equilibrium
        pull = pull_strength * (damping ** day)
        predicted_price = current_price * (1 + pull)
        
        # Add stochastic component (market noise)
        noise = volatility * np.random.randn()
        predicted_price *= np.exp(noise * np.sqrt(day + 1) * 0.1)
        
        forecasts.append(predicted_price)
    
    return {
        'forecasts': np.array(forecasts),
        'current_equilibrium': current_eq,
        'pull_strength': pull_strength,
        'volatility': volatility,
        'mean_return': mean_return,
        'equilibria_history': equilibria,
        'recent_returns': returns
    }


def cct_gradient_forecast_v2(prices, forecast_days=40):
    """
    Simpler CCT-Gradient: Direct equilibrium projection.
    
    Uses gradient descent to find the "attractor price"
    and projects forward with damped convergence.
    """
    
    # Find equilibrium for full price history
    result = find_equilibrium(prices, num_restarts=15, max_iter=1000)
    
    P = prices[-1]
    eq_P = result['equilibrium_price']
    
    # The "gap" d tells us about the price manifold
    # Large d → price far from equilibrium
    # Small d → price near equilibrium
    
    # Compute historical gap ratio
    gap_ratios = []
    for i in range(10, len(prices)):
        window = prices[:i]
        eq = find_equilibrium(window, num_restarts=5, max_iter=300)
        ratio = window[-1] / eq['equilibrium_price']
        gap_ratios.append(ratio)
    
    avg_ratio = np.mean(gap_ratios)
    ratio_std = np.std(gap_ratios)
    
    # Project forward
    forecasts = []
    
    # Estimate damping from autocorrelation of gap ratios
    autocorr = np.corrcoef(gap_ratios[:-1], gap_ratios[1:])[0, 1]
    damping = max(0.8, min(0.99, autocorr if not np.isnan(autocorr) else 0.9))
    
    for day in range(forecast_days):
        # Target: equilibrium price
        target = eq_P
        
        # Current position
        if day == 0:
            current = P
        else:
            current = forecasts[day - 1]
        
        # Damped convergence to equilibrium
        correction = (target - current) * 0.1 * (damping ** day)
        
        # Add market oscillation (period ~14 days, common in crypto)
        oscillation = 0.02 * np.sin(2 * np.pi * day / 14)
        
        predicted = current + correction + P * oscillation * np.random.randn() * 0.1
        
        forecasts.append(max(predicted, P * 0.5))  # Floor at 50% of current
    
    return {
        'forecasts': np.array(forecasts),
        'equilibrium_price': eq_P,
        'current_price': P,
        'gap_d': result['gap_d'],
        'damping': damping,
        's_nearest': result['s_nearest']
    }


# ============================================================
# LOAD BTC DATA
# ============================================================
print("Downloading BTC data...")
btc = yf.download("BTC-USD", period="1y", interval="1d")
close = btc['Close'].dropna().values.flatten()
print(f"Loaded {len(close)} days of BTC data")
print(f"Current price: ${close[-1]:,.2f}")
print(f"Price range: ${close.min():,.2f} - ${close.max():,.2f}")

# ============================================================
# RUN CCT-GRADIENT FORECAST
# ============================================================
print("\n" + "="*60)
print("CCT-GRADIENT BTC FORECAST v2")
print("="*60)

# Version 2: Direct equilibrium projection
result = cct_gradient_forecast_v2(close, forecast_days=40)

print(f"\nEquilibrium Analysis:")
print(f"  Current Price: ${result['current_price']:,.2f}")
print(f"  Equilibrium Price: ${result['equilibrium_price']:,.2f}")
print(f"  Gap d: {result['gap_d']:.2f}")
print(f"  Market Invariant s: {result['s_nearest']}")
print(f"  Damping Factor: {result['damping']:.4f}")

# ============================================================
# VISUALIZATION
# ============================================================
fig, axes = plt.subplots(2, 2, figsize=(14, 10))

# 1. Historical + Forecast
ax = axes[0, 0]
days_hist = np.arange(len(close))
days_fc = np.arange(len(close), len(close) + 40)

ax.plot(days_hist[-90:], close[-90:], 'b-', label='Historical', linewidth=1.5)
ax.plot(days_fc, result['forecasts'], 'r-', label='CCT Forecast', linewidth=2)
ax.axhline(y=result['equilibrium_price'], color='green', linestyle='--', 
           label=f'Equilibrium: ${result["equilibrium_price"]:,.0f}')
ax.axvline(x=len(close)-1, color='gray', linestyle=':', label='Today')
ax.fill_between(days_fc, result['forecasts'] * 0.95, result['forecasts'] * 1.05, 
                alpha=0.2, color='red', label='±5% Band')
ax.set_title("BTC-USD: 40-Day CCT-Gradient Forecast")
ax.set_xlabel("Day")
ax.set_ylabel("Price (USD)")
ax.legend()
ax.grid(True, alpha=0.3)

# 2. Forecast details
ax = axes[0, 1]
ax.bar(days_fc, result['forecasts'], color='red', alpha=0.7, width=0.8)
ax.axhline(y=result['current_price'], color='blue', linestyle='-', 
           label=f'Current: ${result["current_price"]:,.0f}')
ax.axhline(y=result['equilibrium_price'], color='green', linestyle='--', 
           label=f'Equilibrium: ${result["equilibrium_price"]:,.0f}')
ax.set_title("Daily Forecast Values")
ax.set_xlabel("Days Ahead")
ax.set_ylabel("Price (USD)")
ax.legend()
ax.grid(True, alpha=0.3)

# 3. Return distribution from recent data
ax = axes[1, 0]
returns = np.diff(np.log(close[-90:])) * 100
ax.hist(returns, bins=30, color='purple', alpha=0.7, edgecolor='black')
ax.axvline(x=0, color='black', linestyle='-', linewidth=1)
ax.set_title(f"Recent Daily Returns (μ={np.mean(returns):.2f}%, σ={np.std(returns):.2f}%)")
ax.set_xlabel("Return (%)")
ax.set_ylabel("Frequency")

# 4. CCT Structure explanation
ax = axes[1, 1]
ax.axis('off')
text = f"""
CCT-GRADIENT STRUCTURE
══════════════════════════════════════

Problem: BTC Price Forecast
Method: Gradient Descent on Loss L(d)

┌────────────────────────────────────────┐
│ L(d) = sin²(π·s)                       │
│ s = √(d² + 4P)                         │
│                                        │
│ d* = equilibrium gap                   │
│ s* = nearest integer to s              │
│                                        │
│ Convergence → Price Manifold           │
└────────────────────────────────────────┘

Results:
  Current: ${result['current_price']:,.2f}
  Equilibrium: ${result['equilibrium_price']:,.2f}
  Gap: {result['gap_d']:.2f}
  S(invariant): {result['s_nearest']}

The gradient descent finds the price
manifold where s → integer (stable).
"""
ax.text(0.1, 0.9, text, transform=ax.transAxes, fontsize=10,
        verticalalignment='top', fontfamily='monospace',
        bbox=dict(boxstyle='round', facecolor='lightblue', alpha=0.5))

plt.tight_layout()
plt.show()

# ============================================================
# PRINT FORECAST TABLE
# ============================================================
print("\n" + "="*60)
print("40-Day BTC Forecast (CCT-Gradient)")
print("="*60)
print(f"{'Day':<6} {'Forecast':<20} {'Change':<15}")
print("-"*60)

baseline = result['current_price']
for d in [1, 5, 10, 20, 40]:
    if d <= len(result['forecasts']):
        fc = result['forecasts'][d-1]
        change = (fc - baseline) / baseline * 100
        arrow = "↑" if change > 0 else "↓"
        print(f"{d:<6} ${fc:>15,.2f}   {arrow}{abs(change):>6.2f}%")

print("="*60)
