#!/usr/bin/env python3

import json
import sys
import requests
from collections import defaultdict
from concurrent.futures import ThreadPoolExecutor
from influxdb_client import InfluxDBClient

# Add apps directory to path for config manager
sys.path.append('/opt/ngon/apps')
from managers.config_manager import (get_pod_miner_counts, get_miner_models,
                                       get_miner_type_specs, get_worker_names_for_site)

# Get miner type hashrates from config (replaces hardcoded values)
def get_miner_hashrates():
    """Get miner hashrates from master config instead of hardcoded values"""
    miner_specs = get_miner_type_specs()
    return {
        miner_type: specs.get('spec_hashrate', 0) 
        for miner_type, specs in miner_specs.items()
    }

def calculate_max_ph():
    """Calculate max potential hashrate in PH/s from config"""
    pod_miner_counts = get_pod_miner_counts()
    miner_types_mapping = get_miner_models()
    miner_hashrates = get_miner_hashrates()
    
    total_hashrate = 0
    for pod_name, miner_count in pod_miner_counts.items():
        if miner_count > 0:
            miner_type = miner_types_mapping.get(pod_name, "M60")
            hashrate_per_miner = miner_hashrates.get(miner_type, 180)
            total_hashrate += miner_count * hashrate_per_miner
    
    return total_hashrate / 1_000  # Convert to PH/s

# InfluxDB credentials
token = "rehPqmCSlCkQnEAaFQEX75JD7J9lAsflaIS6TTYnIh3SetVyzEumX_dRMizxFEsEISG4VUdTd70rUHcJig7V6w=="
org = "ngonsolutions"
url = "https://us-central1-1.gcp.cloud2.influxdata.com"
hashrate_bucket = "MARAGON_Hashrate"

# Connect to InfluxDB
client = InfluxDBClient(url=url, token=token, org=org)

# Foundry pool accounts. Since the 2026-08-01 worker cutover the fleet reports
# into TWO regional sub-accounts (marangonnd001 = ND, marangontx001 = TX) rather
# than one per site, so pool hashrate can no longer be shown on a site row. The
# retired per-site sub-accounts stay listed because a few stubborn miners never
# took the new pool config and still hash into them; they roll up into one
# "Other Subaccounts" figure that should trend to zero. Same shape as
# status_pool_data_cgi.py, which feeds the status page's hashrate modal.
#
# API keys come from the Report Builder's Foundry block rather than being
# duplicated here — one place to rotate. That file is ngon:ngon 660 and
# www-data is in the ngon group.
REPORTING_CONFIG = "/opt/ngon/config/reporting_config.json"

REGIONAL_ACCOUNTS = [
    {"name": "marangonnd001", "region": "ND"},
    {"name": "marangontx001", "region": "TX"},
]

# "label" is the site each retired sub-account used to serve — it's how an
# operator knows where to go chase a straggler that's still hashing here.
LEGACY_ACCOUNTS = [
    {"name": "maratxngonalpha001", "label": "Butz"},
    {"name": "maratxngonhm002", "label": "Fortson"},
    {"name": "maratxngonjohn001", "label": "John"},
    {"name": "maratxngonosprey001", "label": "Osprey"},
    {"name": "maratxngonew001", "label": "EW"},
    {"name": "marandngondanielle001", "label": "Dan"},
    {"name": "marandngonnate001", "label": "Nate"},
    {"name": "maratxngongn001", "label": "GN"},
    {"name": "maratxngonwill001", "label": "Will"},
]

# Live per-subaccount stats (~400 bytes). hashrate5minAvg is the real-time value
# so the pool figure tracks the live graph instead of lagging (matches
# status_pool_data_cgi.py). Foundry hashrate is GH/s -> /1e6 = PH/s.
STATS_ENDPOINT = "{base}/subaccount_stats/{name}"


def load_foundry_config():
    """Base URL + {subaccount: api_key} from the Report Builder config."""
    with open(REPORTING_CONFIG) as f:
        foundry = json.load(f).get("foundry", {})
    return (foundry.get("base_url") or "https://api.foundryusapool.com",
            foundry.get("api_keys") or {})


def fetch_pool_data(account, base_url, api_keys):
    """Live stats for one sub-account. Returns the account dict enriched with
    ph / workers / offline_15m / ok — ok=False on failure (or a missing key) so
    one dead sub-account doesn't break the rest."""
    out = dict(account, ph=0.0, workers=0, offline_15m=0, ok=False)
    api_key = api_keys.get(account['name'])
    if not api_key:
        return out
    url = STATS_ENDPOINT.format(base=base_url, name=account['name'])
    try:
        response = requests.get(url, headers={"X-API-KEY": api_key,
                                              "User-Agent": "ngon-status"}, timeout=15)
        response.raise_for_status()
        stats = response.json()
        out.update({
            'ph': (stats.get('hashrate5minAvg', 0) or 0) / 1000000.0,
            'workers': stats.get('activeWorkers', 0) or 0,
            'offline_15m': stats.get('offline15MinWorkerCount', 0) or 0,
            'ok': True,
        })
    except requests.exceptions.RequestException:
        pass
    return out


def ours_by_pool_account(containers_hashrate):
    """Our own measured PH/s keyed by the Foundry sub-account each pod's site
    reports into (site.worker_name1), so region membership is a config edit
    rather than a hardcoded site list here."""
    out = defaultdict(float)
    for pod_name, ph in containers_hashrate.items():
        site = pod_name.split()[0] if ' ' in pod_name else pod_name
        account = (get_worker_names_for_site(site) or {}).get('worker_name1')
        if account:
            out[account] += ph
    return out


def run_pool_data(containers_hashrate=None):
    """Current pool hashrate split into regional accounts + a legacy roll-up.
    Sub-accounts are fetched in parallel."""
    base_url, api_keys = load_foundry_config()
    accounts = REGIONAL_ACCOUNTS + LEGACY_ACCOUNTS
    with ThreadPoolExecutor(max_workers=len(accounts)) as executor:
        results = list(executor.map(lambda a: fetch_pool_data(a, base_url, api_keys), accounts))

    ours = ours_by_pool_account(containers_hashrate or {})
    regions = [dict(r, ours_ph=ours.get(r['name'], 0.0)) for r in results if 'region' in r]
    legacy = [r for r in results if 'label' in r]
    # Only surface legacy subs still doing something — the rest are retired and
    # would just be a wall of zeros.
    stragglers = sorted([r for r in legacy if r['ph'] > 0 or r['workers'] > 0],
                        key=lambda r: r['ph'], reverse=True)

    return {
        'regions': regions,
        'legacy': {
            'ph': sum(r['ph'] for r in legacy),
            'workers': sum(r['workers'] for r in legacy),
            'accounts': stragglers,
        },
    }

def to_ph(hashrate):
    """Convert hashrate to Petahashes"""
    return hashrate / 1_000_000_000_000_000

def get_known_pod_names():
    """Return the set of pod names defined in master_config. Used to filter
    out phantom/legacy InfluxDB measurements (e.g. 'Ell', 'Ospre', 'Butz')
    that no longer correspond to real pods."""
    return set(get_pod_miner_counts().keys())

def site_from_pod(pod_name):
    """Map a pod name to its site, applying the Ellyson/Walker -> EW merge."""
    base = pod_name.split()[0] if ' ' in pod_name else pod_name
    if base in ('Ellyson', 'Walker'):
        return 'EW'
    return base

def get_current_hashrate():
    """Get current hashrate from Status API (real-time data)"""
    try:
        response = requests.get('http://localhost:5050/api/status', timeout=10)
        response.raise_for_status()
        data = response.json()

        # Navigate: sites -> generator_groups -> pods -> stats -> hashrate
        containers_hashrate = {}  # pod_name -> hashrate in PH/s
        sites = data.get('sites', {})

        for site_name, site_data in sites.items():
            groups = site_data.get('generator_groups', {})
            for group_name, group_data in groups.items():
                pods = group_data.get('pods', {})
                for pod_name, pod_data in pods.items():
                    stats = pod_data.get('stats', {})
                    hashrate_th = stats.get('hashrate', 0) or 0  # TH/s
                    if hashrate_th > 0:
                        containers_hashrate[pod_name] = hashrate_th / 1000  # Convert to PH/s

        total_hashrate = sum(containers_hashrate.values())
        return total_hashrate, containers_hashrate

    except Exception as e:
        # Fallback to InfluxDB if Status API fails
        return get_current_hashrate_from_influx()


def get_current_hashrate_from_influx():
    """Fallback: Get current hashrate from InfluxDB"""
    query = '''
    from(bucket: "MARAGON_Hashrate")
      |> range(start: -1h)
      |> filter(fn: (r) => r["_field"] == "hashrate")
      |> last()
    '''
    result = client.query_api().query(query, org=org)

    containers_hashrate = {}
    for table in result:
        for record in table.records:
            container_name = record.get_measurement()
            raw_hashrate_value = record.get_value()
            hashrate_value = to_ph(raw_hashrate_value)
            containers_hashrate[container_name] = hashrate_value

    total_hashrate = sum(containers_hashrate.values())
    return total_hashrate, containers_hashrate

def get_all_containers_hashrate(days=7, aggregation="1d"):
    """Get all containers' hashrates for the specified number of days with configurable aggregation"""
    # For daily aggregation, use "last" function like the original
    # For hourly or other aggregation, use "mean" function
    if aggregation == "1d":
        fn_type = "last"
        create_empty = "true"
    else:
        fn_type = "last"
        create_empty = "true"
    
    query = f'''
    from(bucket: "MARAGON_Hashrate")
      |> range(start: -{days}d)
      |> filter(fn: (r) => r["_field"] == "hashrate")
      |> aggregateWindow(every: {aggregation}, fn: {fn_type}, createEmpty: {create_empty})
      |> yield(name: "last")
    '''
    result = client.query_api().query(query, org=org)

    known_pods = get_known_pod_names()
    containers_hashrate = {}
    for table in result:
        for record in table.records:
            container_name = record.get_measurement()
            if container_name not in known_pods:
                continue  # skip phantom/legacy measurements
            raw_hashrate_value = record.get_value()

            if raw_hashrate_value is not None:
                hashrate_value = to_ph(raw_hashrate_value)
            else:
                hashrate_value = 0

            time_value = record.get_time().isoformat()

            if container_name not in containers_hashrate:
                containers_hashrate[container_name] = []
            containers_hashrate[container_name].append({"time": time_value, "value": hashrate_value})

    return containers_hashrate

def get_max_potential_hashrate():
    """Calculate maximum potential hashrate based on installed miners"""
    max_hashrate = defaultdict(float)

    pod_miner_counts = get_pod_miner_counts()
    miner_types_mapping = get_miner_models()
    miner_hashrates = get_miner_hashrates()  # Get from config instead of hardcoded

    for pod_name, miner_count in pod_miner_counts.items():
        if miner_count > 0:
            miner_type = miner_types_mapping.get(pod_name, "M60")
            hashrate_per_miner = miner_hashrates.get(miner_type, 180)  # Default to 180 if not found

            max_hashrate[site_from_pod(pod_name)] += miner_count * hashrate_per_miner / 1_000

    return max_hashrate

# Expected hashrate discount factors by miner type
EXPECTED_FACTORS = {
    'M60': 0.95,
    'JPro': 0.88,
    'XP': 0.60,
}

def get_expected_hashrate():
    """Calculate expected hashrate per site (nameplate * miner-type discount factor)"""
    expected_hashrate = defaultdict(float)

    pod_miner_counts = get_pod_miner_counts()
    miner_types_mapping = get_miner_models()
    miner_hashrates = get_miner_hashrates()

    for pod_name, miner_count in pod_miner_counts.items():
        if miner_count > 0:
            miner_type = miner_types_mapping.get(pod_name, "M60")
            hashrate_per_miner = miner_hashrates.get(miner_type, 180)
            factor = EXPECTED_FACTORS.get(miner_type, 0.95)

            expected_hashrate[site_from_pod(pod_name)] += miner_count * hashrate_per_miner * factor / 1_000

    return expected_hashrate

def get_site_avg_hashrate(days=7):
    """Get average hashrate per site for the specified period"""
    query = f'''
    from(bucket: "MARAGON_Hashrate")
      |> range(start: -{days}d)
      |> filter(fn: (r) => r["_field"] == "hashrate")
      |> mean()
    '''
    result = client.query_api().query(query, org=org)

    known_pods = get_known_pod_names()
    pod_avg = {}
    for table in result:
        for record in table.records:
            pod_name = record.get_measurement()
            if pod_name not in known_pods:
                continue  # skip phantom/legacy measurements
            raw_value = record.get_value()
            if raw_value is not None:
                pod_avg[pod_name] = to_ph(raw_value)

    # Group by site
    site_avg = defaultdict(float)
    for pod_name, avg_ph in pod_avg.items():
        site_avg[site_from_pod(pod_name)] += avg_ph

    return dict(site_avg)

def group_and_sum_by_site(containers_hashrate):
    """Group containers by site name and sum their hashrates"""
    grouped_hashrate = defaultdict(float)
    for container, hashrate in containers_hashrate.items():
        grouped_hashrate[site_from_pod(container)] += hashrate
    return grouped_hashrate

def group_containers_by_site_30d(containers_hashrate_30d):
    """Group containers by site name for historical data"""
    grouped_hashrate_30d = defaultdict(list)
    for container, data in containers_hashrate_30d.items():
        site_name = site_from_pod(container)
        for entry in data:
            if not any(d["time"] == entry["time"] for d in grouped_hashrate_30d[site_name]):
                grouped_hashrate_30d[site_name].append({"time": entry["time"], "value": 0})
            for d in grouped_hashrate_30d[site_name]:
                if d["time"] == entry["time"]:
                    d["value"] += entry["value"]
                    d["value"] = round(d["value"], 0)
    return grouped_hashrate_30d

def calculate_total_hashrate(containers_hashrate):
    """Calculate total hashrate per day by summing up individual container hashrates"""
    total_hashrate = {}

    for container, data in containers_hashrate.items():
        for entry in data:
            time = entry["time"]
            value = entry["value"]

            if time not in total_hashrate:
                total_hashrate[time] = 0
            total_hashrate[time] += value

    total_hashrate_list = [{"time": k, "value": round(v, 0)} for k, v in sorted(total_hashrate.items())]
    return total_hashrate_list

def get_performance_data(days=7, aggregation="1d"):
    """Get all performance data for specified number of days with configurable aggregation"""
    try:
        # Get all the data
        total_hashrate_current, containers_hashrate_current = get_current_hashrate()
        containers_hashrate_period = get_all_containers_hashrate(days, aggregation)
        total_hashrate_period = calculate_total_hashrate(containers_hashrate_period)
        pool_regional = run_pool_data(containers_hashrate_current)
        
        # Process the data
        grouped_current_hashrate = group_and_sum_by_site(containers_hashrate_current)
        grouped_containers_hashrate_period = group_containers_by_site_30d(containers_hashrate_period)
        max_potential_hashrate = get_max_potential_hashrate()
        max_ph = calculate_max_ph()  # Calculate dynamically
        
        # Return all data as JSON
        return {
            'success': True,
            'data': {
                'total_hashrate_current': round(total_hashrate_current),
                'max_ph': max_ph,
                'containers_hashrate_current': dict(grouped_current_hashrate),
                'containers_hashrate_30d': dict(grouped_containers_hashrate_period),
                'total_hashrate_30d': total_hashrate_period,
                'max_potential_hashrate': dict(max_potential_hashrate),
                'pool_regional': pool_regional,
                'days': days
            }
        }
    except Exception as e:
        return {
            'success': False,
            'error': str(e)
        }

def get_summary_data(days=7):
    """Get site hashrate summary: nameplate, expected, and average actual"""
    try:
        max_potential = get_max_potential_hashrate()
        expected = get_expected_hashrate()
        site_avg = get_site_avg_hashrate(days)

        return {
            'success': True,
            'data': {
                'max_potential_hashrate': dict(max_potential),
                'expected_hashrate': dict(expected),
                'site_avg_hashrate': site_avg,
                'days': days
            }
        }
    except Exception as e:
        return {
            'success': False,
            'error': str(e)
        }

if __name__ == '__main__':
    print("Content-Type: application/json\n")
    
    # Handle query parameters
    import cgi
    import cgitb
    cgitb.enable()
    
    form = cgi.FieldStorage()
    mode = form.getvalue("mode", "full")
    days = int(form.getvalue("days", "7"))

    if mode == "summary":
        result = get_summary_data(days)
    else:
        aggregation = form.getvalue("aggregation", "1h")
        result = get_performance_data(days, aggregation)

    print(json.dumps(result))
