from datetime import datetime, timezone

from fastapi import APIRouter, UploadFile, File, Depends, HTTPException, Response, Query
from sqlalchemy import select, func
from sqlalchemy.ext.asyncio import AsyncSession

from app.api.deps import get_db, get_admin_user
from app.core.models.user import User
from app.core.models.task import AsyncTaskLog, AiTaskItem
from . import services
from .schemas import ParseResult, CommitRequest

router = APIRouter(prefix="/bulk-import", tags=["bulk-import"])

# 轮询 AI 异步任务状态（补全/入库共用），前端 GET /api/admin/ai-tasks/{id}
ai_tasks_router = APIRouter(prefix="/ai-tasks", tags=["ai-tasks"])


async def _dispatch_task(parent_task_id: int, celery_task, **kwargs) -> None:
    from .tasks import _update_task

    try:
        celery_task.delay(**kwargs)
    except Exception as exc:
        error = str(exc)[:300]
        await _update_task(parent_task_id, status="failed", result={"error": error})
        raise HTTPException(503, error) from exc

def _authorize_task(row, current_tenant_id: int) -> bool:
    return row.tenant_id == current_tenant_id



@ai_tasks_router.get("", summary="List AI async tasks")
async def list_ai_tasks(
    task_type: str | None = None,
    limit: int = Query(20, ge=1, le=100),
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    from .tasks import _serialize_task_summary

    q = select(AsyncTaskLog).where(AsyncTaskLog.tenant_id == current_user.tenant_id)
    if task_type:
        q = q.where(AsyncTaskLog.task_type == task_type)
    rows = (await db.execute(q.order_by(AsyncTaskLog.id.desc()).limit(limit))).scalars().all()
    return {"items": [_serialize_task_summary(row) for row in rows]}


@ai_tasks_router.get("/{task_id}/items", summary="List AI task items")
async def list_ai_task_items(
    task_id: int,
    page: int = Query(1, ge=1),
    page_size: int = Query(100, ge=1, le=500),
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    task = await db.get(AsyncTaskLog, task_id)
    if task is None or not _authorize_task(task, current_user.tenant_id):
        raise HTTPException(404, "task not found")

    filters = (AiTaskItem.task_id == task_id, AiTaskItem.tenant_id == current_user.tenant_id)
    total = (await db.execute(select(func.count()).select_from(AiTaskItem).where(*filters))).scalar_one()
    rows = (
        await db.execute(
            select(AiTaskItem)
            .where(*filters)
            .order_by(AiTaskItem.row_index)
            .offset((page - 1) * page_size)
            .limit(page_size)
        )
    ).scalars().all()
    return {
        "total": total,
        "page": page,
        "page_size": page_size,
        "items": [
            {
                "id": item.id,
                "task_id": item.task_id,
                "row_index": item.row_index,
                "product_id": item.product_id,
                "status": item.status,
                "source_row": item.source_row,
                "result_row": item.result_row,
                "error": item.error,
            }
            for item in rows
        ],
    }

@ai_tasks_router.get("/{task_id}", summary="轮询 AI 异步任务状态/结果")
async def get_ai_task(task_id: int, db: AsyncSession = Depends(get_db),
                      current_user: User = Depends(get_admin_user)):
    from .tasks import _serialize_task_row
    row = await db.get(AsyncTaskLog, task_id)
    if row is None or not _authorize_task(row, current_user.tenant_id):
        raise HTTPException(404, "任务不存在")
    return _serialize_task_row(row)


@ai_tasks_router.post("/{task_id}/cancel", summary="Cancel AI async task")
async def cancel_ai_task(task_id: int, db: AsyncSession = Depends(get_db),
                         current_user: User = Depends(get_admin_user)):
    from .tasks import _serialize_task_row

    row = await db.get(AsyncTaskLog, task_id)
    if row is None or not _authorize_task(row, current_user.tenant_id):
        raise HTTPException(404, "task not found")
    if row.status not in ("success", "failed", "cancelled"):
        result = row.result or {}
        if isinstance(result, dict):
            result = {**result, "error": "cancelled"}
        else:
            result = {"error": "cancelled"}
        row.status = "cancelled"
        row.result = result
        row.completed_at = datetime.now(timezone.utc)
        await db.commit()
        await db.refresh(row)
    return _serialize_task_row(row)

@ai_tasks_router.post("/{task_id}/retry", summary="Resume cancelled/failed task")
async def retry_ai_task(task_id: int, db: AsyncSession = Depends(get_db),
                        current_user: User = Depends(get_admin_user)):
    from .tasks import _retry_items_stmt, commit_batch_task, enrich_batch_task

    row = await db.get(AsyncTaskLog, task_id)
    if row is None or not _authorize_task(row, current_user.tenant_id):
        raise HTTPException(404, "task not found")
    if row.status not in ("cancelled", "failed"):
        raise HTTPException(400, "only cancelled/failed tasks can be retried")
    await db.execute(_retry_items_stmt(task_id, current_user.tenant_id))
    row.status = "pending"
    row.completed_at = None
    await db.commit()
    if row.task_type == "bulk_import.commit":
        payload = row.payload or {}
        await _dispatch_task(
            task_id, commit_batch_task, task_id=task_id, tenant_id=current_user.tenant_id,
            user_id=payload.get("user_id") or current_user.id, mode=payload.get("mode", "skip"),
            create_missing_taxonomy=payload.get("create_missing_taxonomy", False),
        )
    else:
        await _dispatch_task(task_id, enrich_batch_task, task_id=task_id, tenant_id=current_user.tenant_id)
    return {"task_id": task_id}


@router.get("/template", summary="下载导入模板(xlsx，SKU列为文本格式)")
async def template(current_user: User = Depends(get_admin_user)):
    data = services.build_template_xlsx()
    return Response(
        content=data,
        media_type="application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
        headers={"Content-Disposition": "attachment; filename=import_template.xlsx"},
    )


@router.post("/parse", response_model=ParseResult, summary="上传表格解析为行+列映射")
async def parse(file: UploadFile = File(...), current_user: User = Depends(get_admin_user)):
    content = await file.read()
    name = (file.filename or "").lower()
    if name.endswith(".csv"):
        headers, rows = services.parse_csv(content)
    elif name.endswith(".xlsx"):
        headers, rows = services.parse_xlsx(content)
    else:
        raise HTTPException(400, "仅支持 .csv 或 .xlsx 文件")
    return ParseResult(headers=headers, rows=rows, column_map=services.guess_column_map(headers))


@router.post("/enrich", summary="提交批量AI补全，返回task_id")
async def enrich(body: CommitRequest, current_user: User = Depends(get_admin_user)):
    from .tasks import _create_task_items, _create_task_row, enrich_batch_task

    creates = _creates_from(body)
    task_id = await _create_task_row(current_user.tenant_id, "ai.enrich_batch", {"count": len(creates)})
    await _create_task_items(task_id, current_user.tenant_id, creates)
    await _dispatch_task(task_id, enrich_batch_task, task_id=task_id, tenant_id=current_user.tenant_id)
    return {"task_id": task_id}


def _creates_from(body: CommitRequest) -> list[dict]:
    # 补全后的行已是成品(目标字段名)，不能再按源表头二次映射，否则字段全丢
    return body.rows if body.pre_mapped else services.rows_to_product_creates(body.rows, body.column_map)


@router.post("/preflight", summary="Preview missing categories and brands")
async def preflight(
    body: CommitRequest,
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    creates = _creates_from(body)
    return await services.preflight_taxonomy(db, current_user.tenant_id, creates)

@router.post("/check-duplicates", summary="检测与库中重复的SKU")
async def check_duplicates(
    body: CommitRequest,
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    creates = _creates_from(body)
    existing = await services.find_existing_skus(db, current_user.tenant_id, [c.get("sku") for c in creates])
    return {"count": len(existing), "skus": list(existing.keys())}


@router.post("/commit", summary="确认入库(异步，返回task_id)")
async def commit(
    body: CommitRequest,
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    from .tasks import _create_task_items, _create_task_row, commit_batch_task

    creates = _creates_from(body)
    report = await services.preflight_taxonomy(db, current_user.tenant_id, creates)
    if report["invalid_paths"]:
        raise HTTPException(400, {"invalid_paths": report["invalid_paths"]})
    if report["requires_confirmation"] and not body.create_missing_taxonomy:
        raise HTTPException(409, report)
    task_id = await _create_task_row(
        current_user.tenant_id, "bulk_import.commit",
        {"count": len(creates), "user_id": current_user.id, "mode": body.mode,
         "create_missing_taxonomy": body.create_missing_taxonomy},
    )
    await _create_task_items(task_id, current_user.tenant_id, creates)
    await _dispatch_task(
        task_id, commit_batch_task,
        task_id=task_id, tenant_id=current_user.tenant_id, user_id=current_user.id,
        mode=body.mode, create_missing_taxonomy=body.create_missing_taxonomy,
    )
    return {"task_id": task_id}
