import asyncio
from typing import List, Dict, Any
from celery import shared_task
from app.core.logger import logger
from app.services.embedding_service import get_embedding_service


@shared_task(bind=True, autoretry_for=(Exception,), max_retries=3)
def embed_chunk_task(self, chunk_id: str, chunk_text: str, metadata: Dict[str, Any] = None) -> Dict[str, Any]:
    """
    Celery task to embed a single chunk.
    
    Args:
        chunk_id: Unique identifier for the chunk
        chunk_text: Text content to embed
        metadata: Additional metadata for the chunk
    
    Returns:
        Dictionary with embedding result
    """
    try:
        logger.info(f"Starting embedding task for chunk {chunk_id}")
        
        # Run async embedding in sync context
        embedding_service = get_embedding_service()
        result = asyncio.run(
            embedding_service.embed_chunk(chunk_text, chunk_id)
        )
        
        if result["success"]:
            logger.info(f"Successfully embedded chunk {chunk_id}")
            return {
                "chunk_id": chunk_id,
                "status": "completed",
                "embedding": result["embedding"],
                "token_count": result["token_count"],
            }
        else:
            logger.error(f"Failed to embed chunk {chunk_id}: {result.get('error')}")
            raise Exception(f"Embedding failed: {result.get('error')}")
            
    except Exception as exc:
        logger.error(f"Task failed for chunk {chunk_id}: {exc}")
        raise self.retry(exc=exc, countdown=2 ** self.request.retries)


@shared_task(bind=True, autoretry_for=(Exception,), max_retries=3)
def embed_chunks_batch_task(
    self, 
    chunk_ids: List[str], 
    chunk_texts: List[str],
    source_type: str = "repo_doc"
) -> Dict[str, Any]:
    """
    Celery task to embed multiple chunks in batch.
    
    Args:
        chunk_ids: List of chunk identifiers
        chunk_texts: List of chunk texts
        source_type: Source type for metadata
    
    Returns:
        Dictionary with batch embedding results
    """
    try:
        logger.info(f"Starting batch embedding task for {len(chunk_ids)} chunks")
        
        if len(chunk_ids) != len(chunk_texts):
            raise ValueError("chunk_ids and chunk_texts must have same length")
        
        # Run async embedding in sync context
        embedding_service = get_embedding_service()
        embeddings = asyncio.run(
            embedding_service.embed_texts(chunk_texts)
        )
        
        results = {
            "batch_id": self.request.id,
            "status": "completed",
            "chunks": [],
            "total_chunks": len(chunk_ids),
            "source_type": source_type,
        }
        
        for chunk_id, text, embedding in zip(chunk_ids, chunk_texts, embeddings):
            results["chunks"].append({
                "chunk_id": chunk_id,
                "embedding": embedding,
                "token_count": len(text.split()),  # Rough estimate
            })
        
        logger.info(f"Successfully embedded batch {self.request.id}")
        return results
        
    except Exception as exc:
        logger.error(f"Batch task failed: {exc}")
        raise self.retry(exc=exc, countdown=2 ** self.request.retries)
