#!/usr/bin/env python3
"""
Query ROV UHS data and create time series plots.
"""

import psycopg2
import pandas as pd
import matplotlib.pyplot as plt
import matplotlib.dates as mdates
from datetime import datetime, timedelta

# Database connection parameters - adjust as needed
DB_CONFIG = {
    'host': 'coredata-dpkd.dp.mbari.org',
    'database': 'navproc',
    'user': 'navproc_ro',
    'password': 'cant_chang3_data',
    'port': 5432
}

def yearday_to_date(year, yearday):
    """Convert year and yearday to a date."""
    base_date = datetime(year, 1, 1)
    target_date = base_date + timedelta(days=yearday - 1)
    return target_date

def run_query(start_year, start_yearday, end_yearday):
    """Run the SQL query with user-provided parameters."""
    
    # Convert yeardays to dates
    start_date = yearday_to_date(start_year, start_yearday)
    end_date = yearday_to_date(start_year, end_yearday)
    
    print(f"Querying data from {start_date} to {end_date}...")
    
    query = """
    WITH base AS (SELECT TIMESTAMPTZ '2025-01-01 00:00:00+00' AS base_utc),
    bounds AS (
      SELECT base_utc + INTERVAL '%s days' AS start_ts,
             base_utc + INTERVAL '%s days' AS end_ts
      FROM base
    )
    SELECT
      EXTRACT(EPOCH FROM t.pub_timestamp)::double precision  AS epoch_seconds_utc,
      (EXTRACT(EPOCH FROM t.pub_timestamp)*1000)::bigint     AS epoch_millis_utc,
      TO_CHAR(t.pub_timestamp AT TIME ZONE 'America/Los_Angeles', 'YYYY-MM-DD HH24:MI:SS.MS') AS timestamp_pacific,
      t.rov_set_length,
      t.rov_tension,
      t.rov_wire_length,
      t.rov_wire_speed
    FROM public.navproc_rov_uhs_msg t
    CROSS JOIN bounds b
    WHERE t.pub_timestamp >= b.start_ts
      AND t.pub_timestamp <  b.end_ts
    ORDER BY t.pub_timestamp;
    """
    
    # Connect and execute query
    conn = psycopg2.connect(**DB_CONFIG)
    
    # Calculate day offsets from base date (2025-01-01)
    base = datetime(2025, 1, 1)
    start_offset = (start_date - base).days
    end_offset = (end_date - base).days
    
    df = pd.read_sql_query(query % (start_offset, end_offset), conn)
    conn.close()
    
    print(f"Retrieved {len(df)} rows")
    return df

def save_to_csv(df, filename):
    """Save dataframe to CSV."""
    df.to_csv(filename, index=False)
    print(f"Data saved to {filename}")

def create_time_series_plot(df, output_file='rov_timeseries.png'):
    """Create time series plot with each parameter on its own y-axis."""
    
    # Convert timestamp to datetime for plotting
    df['timestamp'] = pd.to_datetime(df['timestamp_pacific'])
    
    # Create figure with 4 subplots
    fig, axes = plt.subplots(4, 1, figsize=(14, 10), sharex=True)
    fig.suptitle('ROV UHS Time Series Data', fontsize=16, fontweight='bold')
    
    # Define the columns to plot and their properties
    plot_configs = [
        {'col': 'rov_set_length', 'color': 'blue', 'label': 'Set Length'},
        {'col': 'rov_tension', 'color': 'red', 'label': 'Tension'},
        {'col': 'rov_wire_length', 'color': 'green', 'label': 'Wire Length'},
        {'col': 'rov_wire_speed', 'color': 'purple', 'label': 'Wire Speed'}
    ]
    
    # Plot each parameter
    for ax, config in zip(axes, plot_configs):
        ax.plot(df['timestamp'], df[config['col']], 
                color=config['color'], linewidth=0.8, alpha=0.8)
        ax.set_ylabel(config['label'], fontweight='bold', color=config['color'])
        ax.tick_params(axis='y', labelcolor=config['color'])
        ax.grid(True, alpha=0.3)
        
        # Add stats to subplot
        mean_val = df[config['col']].mean()
        std_val = df[config['col']].std()
        ax.text(0.02, 0.95, f'μ={mean_val:.2f}, σ={std_val:.2f}', 
                transform=ax.transAxes, fontsize=9,
                verticalalignment='top', bbox=dict(boxstyle='round', 
                facecolor='white', alpha=0.7))
    
    # Format x-axis
    axes[-1].set_xlabel('Time (Pacific)', fontweight='bold')
    axes[-1].xaxis.set_major_formatter(mdates.DateFormatter('%Y-%m-%d %H:%M'))
    plt.setp(axes[-1].xaxis.get_majorticklabels(), rotation=45, ha='right')
    
    plt.tight_layout()
    plt.savefig(output_file, dpi=300, bbox_inches='tight')
    print(f"Plot saved to {output_file}")
    plt.show()

def main():
    """Main function to run the script."""
    print("=" * 60)
    print("ROV UHS Data Query and Visualization")
    print("=" * 60)
    
    # Get user input
    try:
        start_year = int(input("Enter start year (e.g., 2025): "))
        start_yearday = int(input("Enter start yearday (1-366): "))
        end_yearday = int(input("Enter end yearday (1-366): "))
        
        # Validate input
        if not (1 <= start_yearday <= 366) or not (1 <= end_yearday <= 366):
            raise ValueError("Yearday must be between 1 and 366")
        if start_yearday >= end_yearday:
            raise ValueError("Start yearday must be less than end yearday")
            
    except ValueError as e:
        print(f"Error: {e}")
        return
    
    # Run query
    try:
        df = run_query(start_year, start_yearday, end_yearday)
        
        if df.empty:
            print("No data found for the specified date range.")
            return
        
        # Generate output filenames based on date range
        csv_filename = f'rov_data_{start_year}_day{start_yearday:03d}-{end_yearday:03d}.csv'
        plot_filename = f'rov_plot_{start_year}_day{start_yearday:03d}-{end_yearday:03d}.png'
        
        # Save to CSV
        save_to_csv(df, csv_filename)
        
        # Create plot
        create_time_series_plot(df, plot_filename)
        
        print("\nProcessing complete!")
        
    except psycopg2.Error as e:
        print(f"Database error: {e}")
    except Exception as e:
        print(f"Error: {e}")

if __name__ == "__main__":
    main()