import asyncio
import aiohttp
import json
import os
from pathlib import Path
from typing import Dict, List, Set, Optional
import logging
from datetime import datetime
from suno_utils.utils.text import read_jsonl
import time

# Configure logging
logging.basicConfig(level=logging.INFO, format='%(asctime)s - %(levelname)s - %(message)s')
logger = logging.getLogger(__name__)

class ContinuousFileProcessor:
    def __init__(
        self,
        output_dir: str = "processed_results",
        max_concurrent_requests: int = 50,
        request_delay: float = 0.1,  # Delay between starting new requests (seconds)
        save_frequency: int = 100,   # Save results every N completions
        progress_file: str = "processing_progress.json"
    ):
        self.output_dir = Path(output_dir)
        self.output_dir.mkdir(exist_ok=True)
        self.max_concurrent_requests = max_concurrent_requests
        self.request_delay = request_delay
        self.save_frequency = save_frequency
        self.progress_file = progress_file
        
        # State tracking
        self.processed_files = self._load_progress()
        self.pending_results = {}
        self.completed_count = 0
        self.failed_count = 0
        self.last_save_time = time.time()
    
    def _load_progress(self) -> Set[str]:
        """Load previously processed files from progress file."""
        if os.path.exists(self.progress_file):
            try:
                with open(self.progress_file, 'r') as f:
                    data = json.load(f)
                    return set(data.get('processed_files', []))
            except Exception as e:
                logger.warning(f"Could not load progress file: {e}")
        return set()
    
    def _save_progress_and_results(self, force: bool = False):
        """Save current progress and pending results to files."""
        current_time = time.time()
        
        # Save based on frequency or force
        if force or len(self.pending_results) >= self.save_frequency or (current_time - self.last_save_time) > 600:
            if self.pending_results:
                # Save current batch of results
                timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
                results_file = self.output_dir / f"results_{timestamp}_{len(self.pending_results)}.json"
                
                with open(results_file, 'w') as f:
                    json.dump(self.pending_results, f, indent=2)
                
                logger.info(f"Saved {len(self.pending_results)} results to {results_file}")
                self.pending_results.clear()
            
            # Update progress file
            progress_data = {
                'processed_files': list(self.processed_files),
                'last_updated': datetime.now().isoformat(),
                'total_processed': len(self.processed_files),
                'completed_count': self.completed_count,
                'failed_count': self.failed_count
            }
            
            with open(self.progress_file, 'w') as f:
                json.dump(progress_data, f, indent=2)
            
            self.last_save_time = current_time
    
    async def _process_single_file(
        self, 
        session: aiohttp.ClientSession, 
        s3_path: str,
        semaphore: asyncio.Semaphore
    ) -> tuple[str, dict, bool]:  # Returns (path, result, success)
        """Process a single file with the API endpoint."""
        async with semaphore:  # Limit concurrent requests
            params = {
                "gen_id": s3_path,
                "min_loop_length_bars": 2,
                "max_loop_length_bars": 4,
                "output_stems": "false",
            }
            
            max_retries = 3
            for attempt in range(max_retries):
                try:
                    async with session.get(
                        "https://suno-ai--loop-extraction-data-extract-loop-points.modal.run",
                        params=params,
                        timeout=aiohttp.ClientTimeout(total=600)  # 5 minute timeout
                    ) as response:
                        if response.status == 200:
                            response_json = await response.json()
                            return s3_path, response_json, True
                        else:
                            logger.warning(f"HTTP {response.status} for {s3_path}")
                            if attempt == max_retries - 1:
                                return s3_path, {"error": f"HTTP {response.status}"}, False
                
                except asyncio.TimeoutError:
                    logger.warning(f"Timeout for {s3_path} (attempt {attempt + 1})")
                    if attempt == max_retries - 1:
                        return s3_path, {"error": "timeout"}, False
                
                except Exception as e:
                    logger.warning(f"Error processing {s3_path}: {e} (attempt {attempt + 1})")
                    if attempt == max_retries - 1:
                        return s3_path, {"error": str(e)}, False
                
                # Wait before retry
                if attempt < max_retries - 1:
                    await asyncio.sleep(2 ** attempt)  # Exponential backoff
    
    async def _handle_completed_request(self, task: asyncio.Task):
        """Handle a completed request and update tracking."""
        try:
            s3_path, result, success = await task
            
            # Update tracking
            self.processed_files.add(s3_path)
            self.pending_results[s3_path] = result
            
            if success:
                self.completed_count += 1
            else:
                self.failed_count += 1
            
            # Periodic save
            self._save_progress_and_results()
            
            if (self.completed_count + self.failed_count) % 500 == 0:
                logger.info(f"Progress: {self.completed_count} completed, {self.failed_count} failed, "
                          f"{len(self.processed_files)} total processed")
        
        except Exception as e:
            logger.error(f"Error handling completed request: {e}")
    
    async def process_all_files_continuously(self, stems_to_process: List[str]):
        """Process all files with continuous requests and controlled delays."""
        # Filter out already processed files
        remaining_files = [f for f in stems_to_process if f not in self.processed_files]
        
        if not remaining_files:
            logger.info("All files have already been processed!")
            return
        
        logger.info(f"Processing {len(remaining_files)} remaining files out of {len(stems_to_process)} total")
        logger.info(f"Already processed: {len(self.processed_files)} files")
        logger.info(f"Request delay: {self.request_delay}s, Max concurrent: {self.max_concurrent_requests}")
        
        # Create semaphore to limit concurrent requests
        semaphore = asyncio.Semaphore(self.max_concurrent_requests)
        
        # Create aiohttp session
        connector = aiohttp.TCPConnector(
            limit=self.max_concurrent_requests * 2,
            limit_per_host=self.max_concurrent_requests
        )
        
        async with aiohttp.ClientSession(connector=connector) as session:
            active_tasks = set()
            
            # Process files with controlled delays
            for i, s3_path in enumerate(remaining_files):
                # Create and start the task
                task = asyncio.create_task(
                    self._process_single_file(session, s3_path, semaphore)
                )
                active_tasks.add(task)
                
                # Add completion callback
                task.add_done_callback(lambda t: asyncio.create_task(self._handle_completed_request(t)))
                
                # Log progress
                if i % 500 == 0:
                    logger.info(f"Started request {i+1}/{len(remaining_files)}: {s3_path}")
                
                # Clean up completed tasks periodically
                if len(active_tasks) > self.max_concurrent_requests * 2:
                    done_tasks = {task for task in active_tasks if task.done()}
                    active_tasks -= done_tasks
                
                # Delay before next request (except for the last one)
                if i < len(remaining_files) - 1:
                    await asyncio.sleep(self.request_delay)
            
            # Wait for all remaining tasks to complete
            logger.info("All requests started, waiting for completion...")
            
            if active_tasks:
                await asyncio.gather(*active_tasks, return_exceptions=True)
        
        # Final save
        self._save_progress_and_results(force=True)
        
        logger.info(f"Processing complete! Completed: {self.completed_count}, "
                   f"Failed: {self.failed_count}, Total: {len(self.processed_files)}")
    
    def combine_all_results(self, output_file: str = "combined_results.json") -> Dict[str, dict]:
        """Combine all result files into a single result dictionary."""
        combined_results = {}
        
        result_files = sorted(self.output_dir.glob("results_*.json"))
        
        for result_file in result_files:
            try:
                with open(result_file, 'r') as f:
                    file_data = json.load(f)
                    combined_results.update(file_data)
                logger.info(f"Loaded {len(file_data)} results from {result_file}")
            except Exception as e:
                logger.error(f"Failed to load {result_file}: {e}")
        
        # Save combined results
        if combined_results:
            with open(output_file, 'w') as f:
                json.dump(combined_results, f, indent=2)
            
            logger.info(f"Combined {len(combined_results)} results into {output_file}")
        
        return combined_results


# Usage example
async def main():
    extreme_meta = read_jsonl("/home/sara/sfx/extreme_stems_consolidated_analysis.jsonl")

    stems_to_process = []
    for meta in extreme_meta:
        if not meta["is_mostly_silent"] and meta["duration_s"] < 600:
            stems_to_process.append(str(meta["s3_path"]))
    
    # Create processor with custom settings
    processor = ContinuousFileProcessor(
        output_dir="processed_results",
        max_concurrent_requests=300,  # Max concurrent requests
        request_delay=0.15,          # 50ms delay between starting requests
        save_frequency=1000,          # Save every 1000 completed requests
        progress_file="processing_progress.json"
    )
    
    # Process all files continuously
    await processor.process_all_files_continuously(stems_to_process)
    
    # Optionally combine all results into one file
    all_results = processor.combine_all_results("final_results.json")
    
    print(f"Processing complete! Total results: {len(all_results)}")

# Run the async function
if __name__ == "__main__":
    asyncio.run(main())