#!/usr/bin/env python3
"""
Fix inventory discrepancies found during migration validation
Adds missing products, stock records, and corrects quantities and prices
"""
import os
import sys
import json
import csv
from decimal import Decimal
from datetime import datetime
import subprocess

def load_env_config():
    """Load database configuration from .env file"""
    config = {}
    env_path = '/home/whgoparts/public_html/whims-dev/.env'

    with open(env_path, 'r') as f:
        for line in f:
            line = line.strip()
            if line and not line.startswith('#') and '=' in line:
                key, value = line.split('=', 1)
                config[key] = value.strip('"').strip("'")

    return config

def execute_mysql_query(query, params=None):
    """Execute MySQL query using mysql command"""
    config = load_env_config()

    db_user = config.get('DB_USERNAME', 'root')
    db_pass = config.get('DB_PASSWORD', '')
    db_name = config.get('DB_DATABASE', 'whims_dev')

    # Build mysql command
    cmd = ['mysql', '-u', db_user, f'-p{db_pass}', db_name, '-e', query]

    try:
        result = subprocess.run(cmd, capture_output=True, text=True, check=True)
        return result.stdout
    except subprocess.CalledProcessError as e:
        print(f"Error executing query: {e.stderr}")
        return None

def get_default_location_info():
    """Get the default location and warehouse IDs (Go-Parts GA warehouse)"""
    query = """
    SELECT l.id, l.warehouse_id
    FROM locations l
    JOIN warehouses w ON l.warehouse_id = w.id
    WHERE w.name = 'Go-Parts GA'
    LIMIT 1
    """
    result = execute_mysql_query(query)
    if result:
        lines = result.strip().split('\n')
        if len(lines) > 1:
            values = lines[1].split('\t')
            return {'location_id': int(values[0]), 'warehouse_id': int(values[1])}
    return None

def load_finale_data():
    """Load all Finale product and stock data"""
    print("Loading Finale data...")

    # Load product prices
    finale_products = {}
    with open('/home/whgoparts/public_html/whims-dev/migration-dev/finale-files/ProductListScreenReport-Oct11.csv', 'r') as f:
        reader = csv.DictReader(f)
        for row in reader:
            product_id = row['Product ID']
            avg_cost = row['Average cost']
            finale_products[product_id] = {
                'price': Decimal(avg_cost) if avg_cost and avg_cost.strip() else Decimal('0')
            }

    # Load stock quantities
    with open('/home/whgoparts/public_html/whims-dev/migration-dev/finale-files/StockQuantityBySublocationInUnitsWDetail-Oct11.json', 'r') as f:
        stock_data = json.load(f)

    finale_stock = {}
    current_product = None

    for record in stock_data:
        product_id = record.get('Product ID')
        if product_id and product_id not in [None, '', ' ', 'TOTAL:'] and 'Units\nQoH' not in record:
            current_product = product_id
            if current_product not in finale_stock:
                finale_stock[current_product] = 0
        elif product_id is None and current_product and 'Stock item description' in record:
            desc = record.get('Stock item description', '')
            if 'TOTAL:' in desc:
                qty = record.get('Units\nQoH', 0)
                if qty:
                    finale_stock[current_product] += int(qty)

    # Combine data
    finale_data = {}
    for pid, qty in finale_stock.items():
        if qty > 0:
            price = finale_products.get(pid, {}).get('price', Decimal('0'))
            finale_data[pid] = {
                'qty': qty,
                'price': price,
                'value': Decimal(qty) * price
            }

    print(f"   ✓ Loaded {len(finale_data):,} products from Finale")
    print(f"   ✓ Total Finale valuation: ${sum(p['value'] for p in finale_data.values()):,.2f}")

    return finale_data

def fix_inventory_discrepancies():
    """Main function to fix all inventory discrepancies"""
    print("="*100)
    print("FIXING INVENTORY DISCREPANCIES")
    print("="*100)

    finale_data = load_finale_data()
    location_info = get_default_location_info()

    if not location_info:
        print("❌ Could not find default location info")
        sys.exit(1)

    default_location_id = location_info['location_id']
    default_warehouse_id = location_info['warehouse_id']

    print(f"\n   Using default warehouse ID: {default_warehouse_id}, location ID: {default_location_id}")

    stats = {
        'products_created': 0,
        'products_updated': 0,
        'stocks_created': 0,
        'stocks_updated': 0,
        'value_added': Decimal('0')
    }

    config = load_env_config()
    db_user = config.get('DB_USERNAME', 'root')
    db_pass = config.get('DB_PASSWORD', '')
    db_name = config.get('DB_DATABASE', 'whims_dev')

    # Build a large SQL transaction file
    sql_file = '/tmp/fix_inventory.sql'

    print("\n[1/3] Generating SQL fix script...")

    with open(sql_file, 'w') as f:
        f.write("START TRANSACTION;\n\n")
        f.write("-- Fix inventory discrepancies\n")
        f.write(f"-- Generated: {datetime.now().isoformat()}\n\n")

        for product_id, info in finale_data.items():
            finale_qty = info['qty']
            finale_price = info['price']
            finale_value = info['value']

            # Escape product_id for SQL
            product_id_escaped = product_id.replace("'", "\\'")

            # Insert or update product
            f.write(f"""
-- Product: {product_id_escaped}
INSERT INTO products (product_id, average_price, status, created_at, updated_at)
VALUES ('{product_id_escaped}', {finale_price}, 'active', NOW(), NOW())
ON DUPLICATE KEY UPDATE
    average_price = IF(average_price = 0 OR average_price IS NULL, {finale_price}, average_price),
    updated_at = NOW();

""")

            # Get product internal ID and check stock
            f.write(f"""
SET @product_internal_id = (SELECT id FROM products WHERE product_id = '{product_id_escaped}');
SET @existing_stock = (SELECT COALESCE(SUM(quantity), 0) FROM stocks WHERE product_id = @product_internal_id AND created_at <= '2025-10-11 23:59:59');
SET @qty_needed = {finale_qty} - @existing_stock;

-- If stock is missing or insufficient, add it
INSERT INTO stocks (product_id, warehouse_id, location_id, quantity, reserved_quantity, average_price, `condition`, is_selling, created_at, updated_at)
SELECT @product_internal_id, {default_warehouse_id}, {default_location_id}, @qty_needed, 0, {finale_price}, 'new', 1, '2025-10-10 00:00:00', '2025-10-10 00:00:00'
WHERE @qty_needed > 0;

""")

        f.write("\nCOMMIT;\n")

    print(f"   ✓ Generated SQL script: {sql_file}")

    print("\n[2/3] Executing SQL fix script...")

    cmd = ['mysql', '-u', db_user, f'-p{db_pass}', db_name]

    try:
        with open(sql_file, 'r') as f:
            result = subprocess.run(cmd, stdin=f, capture_output=True, text=True, check=True)

        print("   ✓ SQL script executed successfully")

    except subprocess.CalledProcessError as e:
        print(f"   ❌ Error executing SQL script: {e.stderr}")
        sys.exit(1)

    print("\n[3/3] Validating fixes...")

    # Query final valuation
    query = """
    SELECT
        COUNT(DISTINCT p.id) as products,
        COALESCE(SUM(s.quantity), 0) as total_qty,
        COALESCE(SUM(s.quantity * p.average_price), 0) as total_value
    FROM stocks s
    JOIN products p ON s.product_id = p.id
    WHERE s.created_at <= '2025-10-11 23:59:59'
    """

    result = execute_mysql_query(query)

    if result:
        lines = result.strip().split('\n')
        if len(lines) > 1:
            values = lines[1].split('\t')
            final_products = int(values[0])
            final_qty = int(float(values[1]))
            final_value = Decimal(values[2])

            print("\n" + "="*100)
            print("FINAL VALIDATION")
            print("="*100)
            print(f"\n   Database (Oct 11) after fix:")
            print(f"   - Products with stock: {final_products:,}")
            print(f"   - Total quantity: {final_qty:,}")
            print(f"   - Total valuation: ${final_value:,.2f}")

            finale_total = sum(p['value'] for p in finale_data.values())
            finale_qty = sum(p['qty'] for p in finale_data.values())

            print(f"\n   Finale (Oct 11) target:")
            print(f"   - Products with stock: {len(finale_data):,}")
            print(f"   - Total quantity: {finale_qty:,}")
            print(f"   - Total valuation: ${finale_total:,.2f}")

            diff = final_value - finale_total

            print(f"\n   Difference: ${diff:,.2f}")

            if abs(diff) < Decimal('1.00'):
                print("\n   ✅ SUCCESS! Valuation matches Finale (within $1)")
            else:
                print(f"\n   ⚠️  Still {abs(diff/finale_total*100):.2f}% off target")

    # Save fix report
    fix_report = {
        'timestamp': datetime.now().isoformat(),
        'finale_target': {
            'products': len(finale_data),
            'quantity': sum(p['qty'] for p in finale_data.values()),
            'valuation': float(sum(p['value'] for p in finale_data.values()))
        },
        'fix_applied': True
    }

    with open('/home/whgoparts/public_html/whims-dev/migration-dev/finale-files/fix_report.json', 'w') as f:
        json.dump(fix_report, f, indent=2)

    print("\n✅ Fix report saved to: finale-files/fix_report.json")
    print("="*100)

if __name__ == "__main__":
    fix_inventory_discrepancies()
