#!/usr/bin/env python3
"""
Step 1: Macro-Economic Backdrop Analysis using FRED Data

This script fetches and analyzes restaurant industry macro data from FRED:
- Retail Sales: Food Services and Drinking Places (MRTSSM7225USN)
- Employees: Food Services and Drinking Places (CES7072200001)

Author: K-Dense Coding Agent
Date: 2026-02-03
"""

import os
import sys
import warnings
import requests
import pandas as pd
import numpy as np
import matplotlib
matplotlib.use('Agg')  # Non-interactive backend
import matplotlib.pyplot as plt
import seaborn as sns
from datetime import datetime, timedelta
from pathlib import Path

warnings.filterwarnings('ignore')

# Set up paths
SESSION_DIR = Path("/app/sandbox/session_20260203_085615_5f10d3d2357f")
RESULTS_DIR = SESSION_DIR / "results"
FIGURES_DIR = SESSION_DIR / "figures"

# Ensure directories exist
RESULTS_DIR.mkdir(exist_ok=True)
FIGURES_DIR.mkdir(exist_ok=True)

# Series definitions for restaurant industry
FRED_SERIES = {
    "retail_sales": {
        "id": "MRTSSM7225USN",  # Retail Sales: Food Services and Drinking Places
        "name": "Retail Sales: Food Services",
        "unit": "Millions of Dollars"
    },
    "employment": {
        "id": "CES7072200001",  # All Employees: Food Services and Drinking Places
        "name": "Employment: Food Services",
        "unit": "Thousands of Persons"
    }
}


def fetch_fred_via_datareader(series_id: str, start_date: str, end_date: str) -> pd.DataFrame:
    """
    Fetch FRED data using pandas_datareader.
    """
    try:
        import pandas_datareader.data as web
        print(f"    Attempting pandas_datareader fetch...")
        df = web.DataReader(series_id, "fred", start_date, end_date)
        df.columns = ["value"]
        print(f"    Success: Retrieved {len(df)} observations via pandas_datareader")
        return df
    except Exception as e:
        print(f"    pandas_datareader failed: {str(e)[:100]}")
        return pd.DataFrame()


def fetch_fred_via_api(series_id: str, start_date: str, end_date: str) -> pd.DataFrame:
    """
    Fetch FRED data via direct API call.
    """
    api_key = os.environ.get("FRED_API_KEY", "")
    if not api_key:
        print(f"    No FRED API key found, skipping direct API method")
        return pd.DataFrame()

    print(f"    Attempting direct FRED API fetch...")
    try:
        params = {
            "series_id": series_id,
            "api_key": api_key,
            "file_type": "json",
            "observation_start": start_date,
            "observation_end": end_date,
        }
        response = requests.get(
            "https://api.stlouisfed.org/fred/series/observations",
            params=params,
            timeout=30
        )
        response.raise_for_status()
        data = response.json()

        if "observations" not in data:
            return pd.DataFrame()

        df = pd.DataFrame(data["observations"])
        df["date"] = pd.to_datetime(df["date"])
        df["value"] = pd.to_numeric(df["value"], errors="coerce")
        df = df[["date", "value"]].dropna().set_index("date")
        print(f"    Success: Retrieved {len(df)} observations via FRED API")
        return df
    except Exception as e:
        print(f"    FRED API failed: {str(e)[:100]}")
        return pd.DataFrame()


def fetch_fred_via_fredapi(series_id: str, start_date: str, end_date: str) -> pd.DataFrame:
    """
    Fetch FRED data using fredapi library.
    """
    api_key = os.environ.get("FRED_API_KEY", "")
    if not api_key:
        return pd.DataFrame()

    try:
        from fredapi import Fred
        fred = Fred(api_key=api_key)
        print(f"    Attempting fredapi fetch...")
        series = fred.get_series(series_id, observation_start=start_date, observation_end=end_date)
        df = pd.DataFrame({"value": series})
        print(f"    Success: Retrieved {len(df)} observations via fredapi")
        return df
    except Exception as e:
        print(f"    fredapi failed: {str(e)[:100]}")
        return pd.DataFrame()


def fetch_fred_series(series_id: str, start_date: str, end_date: str) -> pd.DataFrame:
    """
    Fetch a FRED series using multiple methods with fallbacks.
    """
    print(f"  Fetching FRED series: {series_id}")

    # Method 1: pandas_datareader (often works without explicit API key)
    df = fetch_fred_via_datareader(series_id, start_date, end_date)
    if not df.empty:
        return df

    # Method 2: fredapi library
    df = fetch_fred_via_fredapi(series_id, start_date, end_date)
    if not df.empty:
        return df

    # Method 3: Direct API call
    df = fetch_fred_via_api(series_id, start_date, end_date)
    if not df.empty:
        return df

    print(f"    All methods failed for {series_id}")
    return pd.DataFrame()


def calculate_yoy_growth(series: pd.Series, periods: int = 12) -> pd.Series:
    """
    Calculate Year-over-Year percentage growth rate.
    """
    return series.pct_change(periods=periods) * 100


def create_trend_plot(df: pd.DataFrame, output_path: Path) -> None:
    """
    Create dual-axis plot showing raw trends of Sales and Employment.
    """
    print("  Generating trend visualization...")

    fig, ax1 = plt.subplots(figsize=(12, 6))

    # Style settings
    plt.rcParams['font.family'] = 'sans-serif'
    plt.rcParams['font.size'] = 10

    # Plot retail sales on primary axis
    color1 = '#1f77b4'  # Blue
    ax1.set_xlabel('Date', fontsize=11)
    ax1.set_ylabel('Retail Sales (Millions $)', color=color1, fontsize=11)
    line1 = ax1.plot(df.index, df['retail_sales'], color=color1, linewidth=2,
                      label='Retail Sales: Food Services')
    ax1.tick_params(axis='y', labelcolor=color1)
    ax1.yaxis.set_major_formatter(plt.FuncFormatter(lambda x, p: f'{x/1000:.0f}B' if x >= 1000 else f'{x:.0f}M'))

    # Create secondary axis for employment
    ax2 = ax1.twinx()
    color2 = '#ff7f0e'  # Orange
    ax2.set_ylabel('Employment (Thousands)', color=color2, fontsize=11)
    line2 = ax2.plot(df.index, df['employment'], color=color2, linewidth=2,
                      label='Employment: Food Services')
    ax2.tick_params(axis='y', labelcolor=color2)

    # Title and legend
    plt.title('Restaurant Industry Macro Trends: Sales & Employment\n(Food Services and Drinking Places)',
              fontsize=13, fontweight='bold', pad=15)

    # Combine legends
    lines = line1 + line2
    labels = [l.get_label() for l in lines]
    ax1.legend(lines, labels, loc='upper left', frameon=True, fancybox=True)

    # Add grid
    ax1.grid(True, alpha=0.3, linestyle='--')

    # Highlight COVID period if in range
    covid_start = pd.Timestamp('2020-03-01')
    covid_end = pd.Timestamp('2021-03-01')
    if df.index.min() < covid_end and df.index.max() > covid_start:
        ax1.axvspan(max(df.index.min(), covid_start), min(df.index.max(), covid_end),
                    alpha=0.1, color='red', label='COVID Impact Period')

    plt.tight_layout()
    plt.savefig(output_path, dpi=300, bbox_inches='tight', facecolor='white')
    plt.close()
    print(f"    Saved: {output_path}")


def create_yoy_growth_plot(df: pd.DataFrame, output_path: Path) -> None:
    """
    Create plot showing YoY growth rates to highlight momentum.
    """
    print("  Generating YoY growth visualization...")

    fig, (ax1, ax2) = plt.subplots(2, 1, figsize=(12, 8), sharex=True)

    # Plot retail sales YoY growth
    sales_yoy = df['retail_sales_yoy'].dropna()
    colors_sales = ['#2ecc71' if v >= 0 else '#e74c3c' for v in sales_yoy.values]
    ax1.bar(sales_yoy.index, sales_yoy.values, color=colors_sales, alpha=0.7, width=20)
    ax1.axhline(y=0, color='black', linestyle='-', linewidth=0.5)
    ax1.set_ylabel('YoY Growth (%)', fontsize=11)
    ax1.set_title('Retail Sales: Food Services - Year-over-Year Growth', fontsize=12, fontweight='bold')
    ax1.grid(True, alpha=0.3, linestyle='--', axis='y')

    # Add trend line
    if len(sales_yoy) > 1:
        z = np.polyfit(range(len(sales_yoy)), sales_yoy.values, 1)
        p = np.poly1d(z)
        ax1.plot(sales_yoy.index, p(range(len(sales_yoy))),
                 color='navy', linestyle='--', linewidth=2, alpha=0.8, label='Trend')
        ax1.legend(loc='upper right')

    # Plot employment YoY growth
    emp_yoy = df['employment_yoy'].dropna()
    colors_emp = ['#2ecc71' if v >= 0 else '#e74c3c' for v in emp_yoy.values]
    ax2.bar(emp_yoy.index, emp_yoy.values, color=colors_emp, alpha=0.7, width=20)
    ax2.axhline(y=0, color='black', linestyle='-', linewidth=0.5)
    ax2.set_ylabel('YoY Growth (%)', fontsize=11)
    ax2.set_xlabel('Date', fontsize=11)
    ax2.set_title('Employment: Food Services - Year-over-Year Growth', fontsize=12, fontweight='bold')
    ax2.grid(True, alpha=0.3, linestyle='--', axis='y')

    # Add trend line
    if len(emp_yoy) > 1:
        z2 = np.polyfit(range(len(emp_yoy)), emp_yoy.values, 1)
        p2 = np.poly1d(z2)
        ax2.plot(emp_yoy.index, p2(range(len(emp_yoy))),
                 color='navy', linestyle='--', linewidth=2, alpha=0.8, label='Trend')
        ax2.legend(loc='upper right')

    plt.suptitle('Restaurant Industry Momentum Analysis', fontsize=14, fontweight='bold', y=1.02)
    plt.tight_layout()
    plt.savefig(output_path, dpi=300, bbox_inches='tight', facecolor='white')
    plt.close()
    print(f"    Saved: {output_path}")


def generate_macro_summary(df: pd.DataFrame) -> str:
    """
    Generate a text summary of macro findings.
    """
    # Calculate key metrics
    latest_date = df.index.max().strftime('%B %Y')

    # Recent values
    recent_sales = df['retail_sales'].iloc[-1]
    recent_emp = df['employment'].iloc[-1]

    # YoY growth (most recent)
    recent_sales_yoy = df['retail_sales_yoy'].dropna().iloc[-1]
    recent_emp_yoy = df['employment_yoy'].dropna().iloc[-1]

    # Average YoY growth over last 12 months
    avg_sales_yoy = df['retail_sales_yoy'].tail(12).mean()
    avg_emp_yoy = df['employment_yoy'].tail(12).mean()

    # Pre-COVID comparison (Feb 2020)
    covid_comparison = False
    try:
        pre_covid_idx = df.index[df.index <= '2020-02-01']
        if len(pre_covid_idx) > 0:
            pre_covid_sales = df.loc[pre_covid_idx[-1], 'retail_sales']
            pre_covid_emp = df.loc[pre_covid_idx[-1], 'employment']
            sales_vs_precovid = ((recent_sales / pre_covid_sales) - 1) * 100
            emp_vs_precovid = ((recent_emp / pre_covid_emp) - 1) * 100
            covid_comparison = True
    except (IndexError, KeyError) as e:
        pass

    # Determine industry status
    if avg_sales_yoy > 3 and avg_emp_yoy > 2:
        status = "GROWING"
        assessment = "The restaurant industry shows strong growth momentum."
    elif avg_sales_yoy > 0 and avg_emp_yoy > 0:
        status = "RECOVERING/STABLE"
        assessment = "The restaurant industry demonstrates steady recovery and stable growth."
    elif avg_sales_yoy < 0 or avg_emp_yoy < 0:
        status = "CONTRACTING"
        assessment = "The restaurant industry shows signs of contraction."
    else:
        status = "MIXED"
        assessment = "The restaurant industry shows mixed signals."

    summary = f"""## Macro-Economic Assessment: Restaurant Industry

### Data as of {latest_date}

**Industry Status: {status}**

{assessment}

### Key Metrics

| Metric | Current Value | YoY Growth | 12-Month Avg YoY |
|--------|--------------|------------|------------------|
| Retail Sales (Food Services) | ${recent_sales:,.0f}M | {recent_sales_yoy:+.1f}% | {avg_sales_yoy:+.1f}% |
| Employment (Food Services) | {recent_emp:,.0f}K workers | {recent_emp_yoy:+.1f}% | {avg_emp_yoy:+.1f}% |

"""

    if covid_comparison:
        summary += f"""### Recovery from COVID-19 Baseline (vs. Feb 2020)

- **Retail Sales**: {sales_vs_precovid:+.1f}% vs. pre-COVID levels
- **Employment**: {emp_vs_precovid:+.1f}% vs. pre-COVID levels

"""

    # Investment thesis validation
    summary += """### Investment Thesis Validation

**For a "Long" thesis on the restaurant industry:**

"""

    bullish_points = []
    bearish_points = []

    if avg_sales_yoy > 0:
        bullish_points.append(f"- Sales growth remains positive ({avg_sales_yoy:+.1f}% avg YoY)")
    else:
        bearish_points.append(f"- Sales growth is negative ({avg_sales_yoy:+.1f}% avg YoY)")

    if avg_emp_yoy > 0:
        bullish_points.append(f"- Employment growth indicates industry expansion ({avg_emp_yoy:+.1f}% avg YoY)")
    else:
        bearish_points.append(f"- Employment contraction signals industry challenges ({avg_emp_yoy:+.1f}% avg YoY)")

    if covid_comparison and sales_vs_precovid > 0:
        bullish_points.append(f"- Sales have surpassed pre-COVID levels (+{sales_vs_precovid:.1f}%)")

    if bullish_points:
        summary += "**Supportive Factors:**\n" + "\n".join(bullish_points) + "\n\n"

    if bearish_points:
        summary += "**Risk Factors:**\n" + "\n".join(bearish_points) + "\n\n"

    return summary


def main():
    """Main execution function."""
    print("=" * 60)
    print("Step 1: Macro-Economic Backdrop Analysis (FRED Data)")
    print("=" * 60)

    # Define date range (last 10 years)
    end_date = datetime.now().strftime("%Y-%m-%d")
    start_date = (datetime.now() - timedelta(days=365*10)).strftime("%Y-%m-%d")

    print(f"\nDate range: {start_date} to {end_date}")

    # Fetch data
    print("\n[1/5] Fetching FRED data...")

    data_frames = {}
    for key, series_info in FRED_SERIES.items():
        print(f"\n  Series: {series_info['name']} ({series_info['id']})")
        df = fetch_fred_series(series_info['id'], start_date, end_date)

        if not df.empty:
            data_frames[key] = df
            print(f"    Final: {len(df)} observations from {df.index.min().strftime('%Y-%m')} to {df.index.max().strftime('%Y-%m')}")

    if not data_frames:
        print("\n[ERROR] Could not fetch any FRED data.")
        print("All data fetching methods failed. This may be due to:")
        print("1. Network connectivity issues")
        print("2. FRED API rate limiting")
        print("3. Missing API key for authenticated requests")
        print("\nTo manually set a FRED API key:")
        print("  export FRED_API_KEY='your_key_here'")
        print("  Visit https://fredaccount.stlouisfed.org to obtain one.")
        sys.exit(1)

    # Process and merge data
    print("\n[2/5] Processing and aligning data...")

    # Merge dataframes
    combined_df = pd.DataFrame()
    for key, df in data_frames.items():
        if combined_df.empty:
            combined_df = df.rename(columns={'value': key})
        else:
            combined_df = combined_df.join(df.rename(columns={'value': key}), how='outer')

    # Handle missing values with forward fill then backward fill
    combined_df = combined_df.ffill().bfill()

    print(f"  Combined data shape: {combined_df.shape}")
    print(f"  Date range: {combined_df.index.min()} to {combined_df.index.max()}")
    print(f"  Missing values: {combined_df.isnull().sum().to_dict()}")

    # Calculate YoY growth rates
    print("\n[3/5] Calculating Year-over-Year growth rates...")

    for key in data_frames.keys():
        combined_df[f'{key}_yoy'] = calculate_yoy_growth(combined_df[key])

    print(f"  Added YoY growth columns: {[c for c in combined_df.columns if 'yoy' in c]}")

    # Save processed data
    print("\n[4/5] Saving processed data...")

    output_csv = RESULTS_DIR / "restaurant_macro_data.csv"
    combined_df.to_csv(output_csv)
    print(f"  Saved: {output_csv}")
    print(f"  Columns: {list(combined_df.columns)}")
    print(f"  Rows: {len(combined_df)}")

    # Generate visualizations
    print("\n[5/5] Generating visualizations...")

    create_trend_plot(combined_df, FIGURES_DIR / "macro_trends.png")
    create_yoy_growth_plot(combined_df, FIGURES_DIR / "macro_growth_yoy.png")

    # Generate and display summary
    print("\n" + "=" * 60)
    summary = generate_macro_summary(combined_df)
    print(summary)

    print("=" * 60)
    print("Step 1 Complete: Macro-Economic Backdrop Analysis")
    print("=" * 60)

    return combined_df, summary


if __name__ == "__main__":
    df, summary = main()
