from __future__ import annotations

import asyncio
from typing import Any

from celery import shared_task

from app.core.logger import logger
from app.services.pipeline_service import (
    run_commit_analysis_embedding_pipeline_sync,
    run_snapshot_embedding_pipeline_sync,
)


@shared_task(
    bind=True,
    name="embedding.process_snapshot_batch",
    autoretry_for=(Exception,),
    max_retries=3,
    default_retry_delay=60,
)
def process_snapshot_embedding_task(
    self,
    run_id: str,
    repo_id: str,
    snapshot_id: str,
    commit_sha: str,
    artifact_uri: str,
    event_id: str,
) -> dict[str, Any]:
    """Snapshot-scoped embedding job triggered by graph.artifact.ready."""
    logger.info(
        "Processing snapshot embedding job run_id={run_id} repo_id={repo_id} snapshot_id={snapshot_id}",
        run_id=run_id,
        repo_id=repo_id,
        snapshot_id=snapshot_id,
    )
    try:
        return run_snapshot_embedding_pipeline_sync(
            run_id=run_id,
            repo_id=repo_id,
            snapshot_id=snapshot_id,
            commit_sha=commit_sha,
            artifact_uri=artifact_uri,
            event_id=event_id,
        )
    except Exception as exc:
        logger.error(
            "Snapshot embedding task failed run_id={run_id} error={error}",
            run_id=run_id,
            error=str(exc),
        )
        raise self.retry(exc=exc, countdown=60 * (2 ** self.request.retries))


@shared_task(
    bind=True,
    name="embedding.process_commit_analysis",
    autoretry_for=(Exception,),
    max_retries=3,
    default_retry_delay=60,
)
def process_commit_analysis_embedding_task(
    self,
    run_id: str,
    repo_id: str,
    commit_sha: str,
    analysis_id: str,
    event_id: str,
    snapshot_id: str | None = None,
) -> dict[str, Any]:
    """Commit-scoped embedding job triggered by commit.analysis.ready."""
    logger.info(
        "Processing commit analysis embedding job run_id={run_id} analysis_id={analysis_id}",
        run_id=run_id,
        analysis_id=analysis_id,
    )
    try:
        return run_commit_analysis_embedding_pipeline_sync(
            run_id=run_id,
            repo_id=repo_id,
            commit_sha=commit_sha,
            analysis_id=analysis_id,
            event_id=event_id,
            snapshot_id=snapshot_id,
        )
    except Exception as exc:
        logger.error(
            "Commit analysis embedding task failed run_id={run_id} error={error}",
            run_id=run_id,
            error=str(exc),
        )
        raise self.retry(exc=exc, countdown=60 * (2 ** self.request.retries))
