"""Storage orphan and migration APIs."""
import asyncio
import json
import math
import os
import re
from datetime import datetime
from urllib.parse import urlparse

import httpx
from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException
from sqlalchemy import select, text
from sqlalchemy.ext.asyncio import AsyncSession

from app.api.deps import get_db, get_admin_user, get_superadmin_user
from app.core.models.user import User
from app.core.storage import check_storage_deps, LocalStorage

router = APIRouter(prefix="/admin/storage", tags=["storage"])

_HERE = os.path.dirname(os.path.abspath(__file__))
_ROOT = os.path.abspath(os.path.join(_HERE, "../../../../.."))
UPLOAD_DIR = os.path.join(_ROOT, "static", "uploads")

# File extension to MIME type map for migration uploads.
_EXT_CT: dict[str, str] = {
    "jpg": "image/jpeg", "jpeg": "image/jpeg",
    "png": "image/png", "webp": "image/webp",
    "gif": "image/gif", "mp4": "video/mp4",
    "webm": "video/webm", "mov": "video/quicktime",
    "pdf": "application/pdf",
    "csv": "text/csv",
    "txt": "text/plain",
    "doc": "application/msword",
    "docx": "application/vnd.openxmlformats-officedocument.wordprocessingml.document",
    "xls": "application/vnd.ms-excel",
    "xlsx": "application/vnd.openxmlformats-officedocument.spreadsheetml.sheet",
}
_WEBP_MIGRATION_SOURCE_EXTS = {"jpg", "jpeg", "png"}


def _prepare_migration_upload(fname: str, data: bytes) -> tuple[str, bytes, str]:
    ext = fname.rsplit(".", 1)[-1].lower() if "." in fname else ""
    if ext in _WEBP_MIGRATION_SOURCE_EXTS:
        try:
            from app.api.routers.upload import _PIL_OK, _to_webp
            if _PIL_OK:
                stem = fname.rsplit(".", 1)[0]
                return f"{stem}.webp", _to_webp(data), "image/webp"
        except Exception:
            pass
    return fname, data, _EXT_CT.get(ext, "application/octet-stream")


def _storage_migration_result(
    *,
    total: int,
    processed: int,
    migrated: int,
    skipped: int,
    already: int,
    errors: list[dict],
    usage_mb: int | None = None,
    done: bool = False,
) -> dict:
    return {
        "total": total,
        "processed": processed,
        "migrated": migrated,
        "skipped": skipped,
        "already": already,
        "errors": len(errors),
        "error_samples": errors[:5],
        "storage_used_mb": usage_mb,
        "done": done,
    }


async def _list_storage_objects(storage, prefix: str) -> list[dict]:
    if isinstance(storage, LocalStorage):
        base_dir = os.path.join(UPLOAD_DIR, prefix) if prefix else UPLOAD_DIR
        objects: list[dict] = []
        if os.path.isdir(base_dir):
            for root, _, files in os.walk(base_dir):
                for fname in files:
                    path = os.path.join(root, fname)
                    rel = os.path.relpath(path, UPLOAD_DIR).replace(os.sep, "/")
                    objects.append({"name": rel, "size": os.path.getsize(path)})
        return objects

    if hasattr(storage, "_session") and hasattr(storage, "_bucket"):
        client_kwargs: dict = {}
        if getattr(storage, "_endpoint", None):
            client_kwargs["endpoint_url"] = storage._endpoint
        objects: list[dict] = []
        async with storage._session.client("s3", **client_kwargs) as s3:
            paginator = s3.get_paginator("list_objects_v2")
            async for page in paginator.paginate(Bucket=storage._bucket, Prefix=prefix):
                for obj in page.get("Contents", []):
                    objects.append({"name": obj.get("Key"), "size": int(obj.get("Size") or 0)})
        return objects

    if hasattr(storage, "_bucket") and hasattr(storage._bucket, "list_objects"):
        loop = asyncio.get_event_loop()
        objects: list[dict] = []
        marker = ""
        while True:
            result = await loop.run_in_executor(
                None,
                lambda m=marker: storage._bucket.list_objects(prefix=prefix, marker=m, max_keys=1000),
            )
            for obj in getattr(result, "object_list", []) or []:
                objects.append({"name": getattr(obj, "key", ""), "size": int(getattr(obj, "size", 0) or 0)})
            if not getattr(result, "is_truncated", False):
                break
            marker = getattr(result, "next_marker", "") or ""
        return objects

    names = await storage.list_files(prefix=prefix)
    objects = []
    for item in names:
        if isinstance(item, dict):
            objects.append({"name": item.get("name"), "size": int(item.get("size") or 0)})
        else:
            objects.append({"name": item, "size": 0})
    return objects


async def _recalculate_tenant_storage_usage(
    db: AsyncSession,
    tenant_id: int,
    storage=None,
    prefix: str | None = None,
) -> int:
    from app.core.models.tenant import Tenant

    if storage is None or prefix is None:
        storage, prefix = await _get_tenant_storage(db, tenant_id)
    total_bytes = sum(obj.get("size", 0) for obj in await _list_storage_objects(storage, prefix or ""))
    used_mb = math.ceil(total_bytes / (1024 * 1024)) if total_bytes > 0 else 0

    result = await db.execute(select(Tenant).where(Tenant.id == tenant_id))
    tenant = result.scalar_one_or_none()
    if tenant:
        tenant.storage_used_mb = used_mb
        await db.commit()
    return used_mb


def _extract_fname(url_val) -> str | None:
    """Extract the filename from a local or CDN URL."""
    if not url_val:
        return None
    path = urlparse(str(url_val)).path
    name = os.path.basename(path)
    return name if ("." in name) else None


def _normalize_key(url_val, tenant_id) -> str | None:
    """从本地/CDN URL 中提取完整对象 key（自 tenant_{id}/ 起）。

    本地 /api/static/uploads/tenant_2/attachments/a.pdf → tenant_2/attachments/a.pdf
    CDN   https://cdn.xxx/tenant_2/x.webp              → tenant_2/x.webp
    找不到 tenant_{id}/ 标记或无文件名时返回 None。
    """
    if not url_val:
        return None
    path = urlparse(str(url_val)).path
    marker = f"tenant_{tenant_id}/"
    idx = path.find(marker)
    if idx < 0:
        return None
    key = path[idx:]
    return key if ("." in os.path.basename(key)) else None


def _select_orphans(all_files: list[dict], ref_keys: set[str], ref_basenames: set[str]) -> list[dict]:
    """筛出孤儿文件：完整 key 未被引用 **且** basename 未被引用才算孤儿。

    basename 仅作防误删安全网，不作精确判断——宁可漏删，不可误删。
    """
    orphans = []
    for f in all_files:
        key = f["name"]
        if key in ref_keys:
            continue
        if os.path.basename(key) in ref_basenames:
            continue
        orphans.append(f)
    return orphans

# Extract embedded URL tokens from rich text HTML.
_RE_URL_IN_TEXT = re.compile(r'(?:https?://[^\s"\'<>]+|/(?:api/static/uploads|static/uploads)/[^\s"\'<>]+)')


def _is_local_url(url: str | None) -> bool:
    """Return whether the URL points at local uploads storage."""
    if not url:
        return False
    s = str(url)
    return s.startswith("/api/static/uploads/") or s.startswith("/static/uploads/")


async def _collect_tenant_url_map(db: AsyncSession, tenant_id: int, only_local: bool = False) -> dict[str, str]:
    """Return {url: filename} for tenant DB URL references.
    Full URLs are keys so identical basenames under different prefixes do not overwrite each other.
    When only_local is true, keep only local uploads paths.
    """
    result: dict[str, str] = {}  # old_url -> filename

    def _keep(url) -> bool:
        return bool(url) and (not only_local or _is_local_url(url))

    def _check(url):
        if _keep(url):
            fname = os.path.basename(urlparse(str(url)).path)
            if fname and "." in fname:
                result[str(url)] = fname

    def _check_text(text_val):
        if not text_val:
            return
        for m in _RE_URL_IN_TEXT.finditer(str(text_val)):
            url = m.group(0)
            if _keep(url):
                fname = os.path.basename(urlparse(url).path)
                if fname and "." in fname:
                    result[url] = fname

    # product_images (join products for tenant filter)
    rows = (await db.execute(text(
        "SELECT pi.url, pi.webp_url FROM product_images pi "
        "JOIN products p ON pi.product_id = p.id "
        "WHERE p.tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for url, webp_url in rows:
        _check(url)
        _check(webp_url)

    # products direct media fields
    rows = (await db.execute(text(
        "SELECT cover_url, og_image FROM products WHERE tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for cover_url, og_image in rows:
        _check(cover_url)
        _check(og_image)

    # products rich-text
    rows = (await db.execute(text(
        "SELECT description, ai_description FROM products WHERE tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for desc, ai_desc in rows:
        _check_text(desc)
        _check_text(ai_desc)

    # categories
    rows = (await db.execute(text(
        "SELECT image_url FROM categories WHERE tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for (url,) in rows:
        _check(url)

    # brands
    rows = (await db.execute(text(
        "SELECT logo_url FROM brands WHERE tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for (url,) in rows:
        _check(url)

    # store_banners
    try:
        rows = (await db.execute(text(
            "SELECT image_url FROM store_banners WHERE tenant_id = :tid"
        ), {"tid": tenant_id})).fetchall()
        for (url,) in rows:
            _check(url)
    except Exception:
        pass

    # product_variants
    try:
        rows = (await db.execute(text(
            "SELECT image_url FROM product_variants WHERE tenant_id = :tid"
        ), {"tid": tenant_id})).fetchall()
        for (url,) in rows:
            _check(url)
    except Exception:
        pass

    # articles
    try:
        rows = (await db.execute(text(
            "SELECT cover_image, og_image, content FROM articles WHERE tenant_id = :tid"
        ), {"tid": tenant_id})).fetchall()
        for cover, og_image, content in rows:
            _check(cover)
            _check(og_image)
            _check_text(content)
    except Exception:
        pass

    # shipping carrier logos
    try:
        rows = (await db.execute(text(
            "SELECT logo_url FROM shipping_carriers WHERE tenant_id = :tid"
        ), {"tid": tenant_id})).fetchall()
        for (url,) in rows:
            _check(url)
    except Exception:
        pass

    # optional product media plugins
    for q in [
        "SELECT image_url FROM product_media_role_items WHERE tenant_id = :tid",
        "SELECT file_url FROM product_attachments WHERE tenant_id = :tid",
    ]:
        try:
            rows = (await db.execute(text(q), {"tid": tenant_id})).fetchall()
            for (url,) in rows:
                _check(url)
        except Exception:
            pass

    # tenant_settings JSON
    try:
        rows = (await db.execute(text(
            "SELECT extra FROM tenant_settings WHERE tenant_id = :tid"
        ), {"tid": tenant_id})).fetchall()
        for (extra,) in rows:
            if extra:
                try:
                    d = json.loads(extra) if isinstance(extra, str) else extra
                    for k in ("logo_url", "og_image", "favicon_url"):
                        _check(d.get(k))
                except Exception:
                    pass
    except Exception:
        pass

    return result


async def _collect_tenant_local_refs(db: AsyncSession, tenant_id: int) -> dict[str, str]:
    """Return local upload references only."""
    return await _collect_tenant_url_map(db, tenant_id, only_local=True)


async def _fetch_existing_bytes(url: str) -> bytes | None:
    """Read bytes from a local upload path or a full http(s) URL.
    Failed reads return None.
    """
    if _is_local_url(url):
        rel = url.split("/static/uploads/", 1)[-1]
        local_path = os.path.join(UPLOAD_DIR, rel)
        if not os.path.isfile(local_path):
            return None
        loop = asyncio.get_event_loop()
        return await loop.run_in_executor(None, lambda p=local_path: open(p, "rb").read())
    try:
        async with httpx.AsyncClient(timeout=5.0) as client:
            resp = await client.get(url)
        if resp.status_code == 200:
            return resp.content
    except Exception:
        pass
    return None


async def _collect_tenant_refs(db: AsyncSession, tenant_id: int) -> tuple[set[str], set[str]]:
    """收集租户被引用文件，返回 (ref_keys, ref_basenames)。

    ref_keys      完整对象 key（tenant_x/...），用于精确判定；
    ref_basenames basename 集合，仅作防误删安全网（含无法归一出完整 key 的历史 URL）。
    """
    ref_keys: set[str] = set()
    ref_basenames: set[str] = set()

    def _add(url_val):
        if not url_val:
            return
        key = _normalize_key(url_val, tenant_id)
        if key:
            ref_keys.add(key)
        else:
            # 归一失败（如历史遗留无 tenant 前缀的 URL）保留 basename 作安全网
            fb = _extract_fname(url_val)
            if fb:
                ref_basenames.add(fb)

    def _add_text(text_val):
        if not text_val:
            return
        for m in _RE_URL_IN_TEXT.finditer(str(text_val)):
            _add(m.group(0))

    rows = (await db.execute(text(
        "SELECT pi.url, pi.webp_url FROM product_images pi "
        "JOIN products p ON pi.product_id = p.id "
        "WHERE p.tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for url, webp_url in rows:
        _add(url)
        _add(webp_url)

    rows = (await db.execute(text(
        "SELECT cover_url, og_image FROM products WHERE tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for cover_url, og_image in rows:
        _add(cover_url)
        _add(og_image)

    rows = (await db.execute(text(
        "SELECT description, ai_description FROM products WHERE tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for desc, ai_desc in rows:
        _add_text(desc)
        _add_text(ai_desc)

    rows = (await db.execute(text(
        "SELECT image_url FROM categories WHERE tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for (url,) in rows:
        _add(url)

    rows = (await db.execute(text(
        "SELECT logo_url FROM brands WHERE tenant_id = :tid"
    ), {"tid": tenant_id})).fetchall()
    for (url,) in rows:
        _add(url)

    for q in [
        "SELECT image_url FROM store_banners WHERE tenant_id = :tid",
        "SELECT image_url FROM product_variants WHERE tenant_id = :tid",
    ]:
        try:
            rows = (await db.execute(text(q), {"tid": tenant_id})).fetchall()
            for (url,) in rows:
                _add(url)
        except Exception:
            pass

    try:
        rows = (await db.execute(text(
            "SELECT cover_image, og_image, content FROM articles WHERE tenant_id = :tid"
        ), {"tid": tenant_id})).fetchall()
        for cover, og_image, content in rows:
            _add(cover)
            _add(og_image)
            _add_text(content)
    except Exception:
        pass

    try:
        rows = (await db.execute(text(
            "SELECT logo_url FROM shipping_carriers WHERE tenant_id = :tid"
        ), {"tid": tenant_id})).fetchall()
        for (url,) in rows:
            _add(url)
    except Exception:
        pass

    for q in [
        "SELECT image_url FROM product_media_role_items WHERE tenant_id = :tid",
        "SELECT file_url FROM product_attachments WHERE tenant_id = :tid",
    ]:
        try:
            rows = (await db.execute(text(q), {"tid": tenant_id})).fetchall()
            for (url,) in rows:
                _add(url)
        except Exception:
            pass

    try:
        rows = (await db.execute(text(
            "SELECT extra FROM tenant_settings WHERE tenant_id = :tid"
        ), {"tid": tenant_id})).fetchall()
        for (extra,) in rows:
            if extra:
                try:
                    d = json.loads(extra) if isinstance(extra, str) else extra
                    for k in ("logo_url", "og_image", "favicon_url"):
                        _add(d.get(k))
                except Exception:
                    pass
    except Exception:
        pass

    # ref_keys 的 basename 也并入安全网，保证完整 key 命中的文件同样受 basename 保护
    for k in ref_keys:
        ref_basenames.add(os.path.basename(k))
    return ref_keys, ref_basenames


async def _update_tenant_urls(
    db: AsyncSession,
    tenant_id: int,
    url_map: dict[str, str],   # old_local_url -> new_cdn_url
) -> None:
    """Replace tenant DB URLs with migrated storage URLs."""
    if not url_map:
        return

    for old_url, new_url in url_map.items():
        fname = os.path.basename(old_url)

        # product_images
        await db.execute(text(
            "UPDATE product_images pi "
            "JOIN products p ON pi.product_id = p.id "
            "SET pi.url = CASE WHEN pi.url = :old THEN :new ELSE pi.url END, "
            "    pi.webp_url = CASE WHEN pi.webp_url = :old THEN :new ELSE pi.webp_url END "
            "WHERE p.tenant_id = :tid AND (pi.url = :old OR pi.webp_url = :old)"
        ), {"new": new_url, "old": old_url, "tid": tenant_id})

        # products direct media fields
        await db.execute(text(
            "UPDATE products "
            "SET cover_url = CASE WHEN cover_url = :old THEN :new ELSE cover_url END, "
            "    og_image = CASE WHEN og_image = :old THEN :new ELSE og_image END "
            "WHERE tenant_id = :tid AND (cover_url = :old OR og_image = :old)"
        ), {"new": new_url, "old": old_url, "tid": tenant_id})

        # products description / ai_description (REPLACE handles embedded URLs)
        await db.execute(text(
            "UPDATE products "
            "SET description = REPLACE(description, :old, :new), "
            "    ai_description = REPLACE(ai_description, :old, :new) "
            "WHERE tenant_id = :tid AND (description LIKE :pat OR ai_description LIKE :pat)"
        ), {"new": new_url, "old": old_url, "tid": tenant_id, "pat": f"%{fname}%"})

        # categories
        await db.execute(text(
            "UPDATE categories SET image_url = :new "
            "WHERE image_url = :old AND tenant_id = :tid"
        ), {"new": new_url, "old": old_url, "tid": tenant_id})

        # brands
        await db.execute(text(
            "UPDATE brands SET logo_url = :new "
            "WHERE logo_url = :old AND tenant_id = :tid"
        ), {"new": new_url, "old": old_url, "tid": tenant_id})

        # store_banners
        try:
            await db.execute(text(
                "UPDATE store_banners SET image_url = :new "
                "WHERE image_url = :old AND tenant_id = :tid"
            ), {"new": new_url, "old": old_url, "tid": tenant_id})
        except Exception:
            pass

        # product_variants
        try:
            await db.execute(text(
                "UPDATE product_variants SET image_url = :new "
                "WHERE image_url = :old AND tenant_id = :tid"
            ), {"new": new_url, "old": old_url, "tid": tenant_id})
        except Exception:
            pass

        # articles
        try:
            await db.execute(text(
                "UPDATE articles "
                "SET cover_image = CASE WHEN cover_image = :old THEN :new ELSE cover_image END, "
                "    og_image = CASE WHEN og_image = :old THEN :new ELSE og_image END, "
                "    content = REPLACE(content, :old, :new) "
                "WHERE tenant_id = :tid AND (cover_image = :old OR og_image = :old OR content LIKE :pat)"
            ), {"new": new_url, "old": old_url, "tid": tenant_id, "pat": f"%{fname}%"})
        except Exception:
            pass

        # shipping carrier logos
        try:
            await db.execute(text(
                "UPDATE shipping_carriers SET logo_url = :new "
                "WHERE logo_url = :old AND tenant_id = :tid"
            ), {"new": new_url, "old": old_url, "tid": tenant_id})
        except Exception:
            pass

        # optional product media plugins
        for q in [
            "UPDATE product_media_role_items SET image_url = :new WHERE image_url = :old AND tenant_id = :tid",
            "UPDATE product_attachments SET file_url = :new WHERE file_url = :old AND tenant_id = :tid",
        ]:
            try:
                await db.execute(text(q), {"new": new_url, "old": old_url, "tid": tenant_id})
            except Exception:
                pass

    # tenant_settings JSON (Python-side: parse, update, re-serialize)
    try:
        r = await db.execute(text(
            "SELECT id, extra FROM tenant_settings WHERE tenant_id = :tid"
        ), {"tid": tenant_id})
        row = r.fetchone()
        if row and row.extra:
            d = json.loads(row.extra) if isinstance(row.extra, str) else dict(row.extra)
            changed = False
            for k in ("logo_url", "og_image", "favicon_url"):
                if d.get(k) in url_map:
                    d[k] = url_map[d[k]]
                    changed = True
            if changed:
                await db.execute(text(
                    "UPDATE tenant_settings SET extra = :extra WHERE id = :id"
                ), {"extra": json.dumps(d, ensure_ascii=False), "id": row.id})
    except Exception:
        pass


async def _list_storage_files(storage, prefix: str = "") -> list[dict]:
    """列出前缀下所有文件，返回完整对象 key + 真实大小（含嵌套目录）。

    复用 _list_storage_objects（本地 os.walk 递归、S3/OSS 返回完整 Key），
    避免旧实现只列第一层、且把嵌套 key 砍成 basename 的问题。
    """
    return await _list_storage_objects(storage, prefix)


async def _get_tenant_storage(db: AsyncSession, tenant_id: int):
    """Resolve the tenant storage backend and object prefix."""
    from app.core.models.tenant import Tenant
    from app.core.models.storage_profile import StorageProfile
    from app.core.storage import get_storage_for_profile

    tr = await db.execute(select(Tenant).where(Tenant.id == tenant_id))
    tenant = tr.scalar_one_or_none()
    if not tenant:
        raise HTTPException(status_code=404, detail="Tenant not found")

    profile = None
    if tenant.storage_profile_id:
        pr = await db.execute(select(StorageProfile).where(StorageProfile.id == tenant.storage_profile_id))
        profile = pr.scalar_one_or_none()

    storage = get_storage_for_profile(profile)
    prefix = f"tenant_{tenant_id}/"
    return storage, prefix


@router.get("/dep-check", summary="Check CDN storage dependencies")
async def cdn_dep_check(
    provider: str = "local",
    _: User = Depends(get_admin_user),
):
    return check_storage_deps(provider)


@router.get("/orphans", summary="Scan orphan files for current tenant")
async def list_orphans(
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    storage, prefix = await _get_tenant_storage(db, current_user.tenant_id)

    ref_keys, ref_basenames = await _collect_tenant_refs(db, current_user.tenant_id)
    all_files = await _list_storage_files(storage, prefix=prefix)
    total_size = sum(f["size"] for f in all_files)

    orphans = _select_orphans(all_files, ref_keys, ref_basenames)
    orphan_size = sum(f["size"] for f in orphans)

    return {
        "total_files": len(all_files),
        "total_size": total_size,
        "used_files": len(all_files) - len(orphans),
        "orphan_files": len(orphans),
        "orphan_size": orphan_size,
        "orphans": sorted(orphans, key=lambda x: x["size"], reverse=True),
    }


@router.delete("/orphans", summary="Delete orphan files for a tenant")
async def delete_orphans(
    tenant_id: int | None = None,
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_superadmin_user),
):
    target_tenant_id = tenant_id if tenant_id is not None else current_user.tenant_id
    storage, prefix = await _get_tenant_storage(db, target_tenant_id)

    ref_keys, ref_basenames = await _collect_tenant_refs(db, target_tenant_id)
    all_files = await _list_storage_files(storage, prefix=prefix)
    orphans = _select_orphans(all_files, ref_keys, ref_basenames)

    deleted, failed = [], []
    for f in orphans:
        try:
            await storage.delete(f["name"])
            deleted.append(f)
        except Exception as e:
            failed.append({"name": f["name"], "error": str(e)})

    freed = sum(f["size"] for f in deleted)
    return {
        "tenant_id": target_tenant_id,
        "deleted": len(deleted),
        "freed_bytes": freed,
        "failed": failed,
    }


STORAGE_MIGRATION_TASK_TYPE = "storage_migration"
STORAGE_MIGRATION_BATCH_SIZE = 100


def _task_response(task, reused: bool = False) -> dict:
    return {
        "task_id": task.id,
        "tenant_id": task.tenant_id,
        "task_type": task.task_type,
        "status": task.status,
        "payload": task.payload or {},
        "result": task.result or {},
        "reused": reused,
    }


async def _get_running_storage_migration_task(db: AsyncSession, tenant_id: int):
    from app.core.models.task import AsyncTaskLog

    result = await db.execute(
        select(AsyncTaskLog)
        .where(
            AsyncTaskLog.tenant_id == tenant_id,
            AsyncTaskLog.task_type == STORAGE_MIGRATION_TASK_TYPE,
            AsyncTaskLog.status.in_(("pending", "running")),
        )
        .order_by(AsyncTaskLog.id.desc())
        .limit(1)
    )
    return result.scalar_one_or_none()


async def _get_storage_migration_task(db: AsyncSession, tenant_id: int, task_id: int) -> dict:
    from app.core.models.task import AsyncTaskLog

    result = await db.execute(
        select(AsyncTaskLog).where(
            AsyncTaskLog.id == task_id,
            AsyncTaskLog.tenant_id == tenant_id,
            AsyncTaskLog.task_type == STORAGE_MIGRATION_TASK_TYPE,
        )
    )
    task = result.scalar_one_or_none()
    if not task:
        raise HTTPException(status_code=404, detail="Migration task not found")
    return _task_response(task)


async def _update_storage_migration_task(
    db: AsyncSession,
    task_id: int | None,
    result: dict,
    status: str = "running",
) -> None:
    if task_id is None:
        return
    from app.core.models.task import AsyncTaskLog

    task = await db.get(AsyncTaskLog, task_id)
    if not task:
        return
    task.status = status
    task.result = result
    if status in ("success", "failed"):
        task.completed_at = datetime.utcnow()
    await db.commit()


async def _run_storage_migration_task(task_id: int) -> None:
    from app.database import AsyncSessionLocal
    from app.core.models.task import AsyncTaskLog

    async with AsyncSessionLocal() as db:
        task = await db.get(AsyncTaskLog, task_id)
        if not task:
            return
        tenant_id = int((task.payload or {}).get("tenant_id") or task.tenant_id)
        batch_size = int((task.payload or {}).get("batch_size") or STORAGE_MIGRATION_BATCH_SIZE)
        try:
            task.status = "running"
            await db.commit()
            result = await _migrate_tenant_to_current_storage(
                db,
                tenant_id,
                task_id=task_id,
                batch_size=batch_size,
            )
            await _update_storage_migration_task(db, task_id, result, status="success")
        except Exception as exc:
            await _update_storage_migration_task(
                db,
                task_id,
                {"error": str(exc), "done": True},
                status="failed",
            )


async def _start_storage_migration_task(
    db: AsyncSession,
    tenant_id: int,
    background_tasks: BackgroundTasks | None = None,
) -> dict:
    from app.core.models.task import AsyncTaskLog

    existing = await _get_running_storage_migration_task(db, tenant_id)
    if existing:
        age = (datetime.utcnow() - existing.updated_at).total_seconds() if existing.updated_at else 0
        if age <= 7200:
            return _task_response(existing, reused=True)
        existing.status = "failed"
        existing.result = {"error": "stale migration task replaced", "done": True}
        existing.completed_at = datetime.utcnow()
        await db.commit()

    initial = _storage_migration_result(
        total=0,
        processed=0,
        migrated=0,
        skipped=0,
        already=0,
        errors=[],
        usage_mb=None,
        done=False,
    )
    task = AsyncTaskLog(
        tenant_id=tenant_id,
        task_type=STORAGE_MIGRATION_TASK_TYPE,
        status="pending",
        payload={"tenant_id": tenant_id, "batch_size": STORAGE_MIGRATION_BATCH_SIZE},
        result=initial,
    )
    db.add(task)
    await db.commit()
    await db.refresh(task)

    if background_tasks is not None:
        background_tasks.add_task(_run_storage_migration_task, task.id)
    else:
        asyncio.create_task(_run_storage_migration_task(task.id))
    return _task_response(task)


async def _migrate_tenant_to_current_storage(
    db: AsyncSession,
    tenant_id: int,
    task_id: int | None = None,
    batch_size: int = STORAGE_MIGRATION_BATCH_SIZE,
) -> dict:
    """Migrate tenant media/file references to the tenant current storage backend."""
    storage, prefix = await _get_tenant_storage(db, tenant_id)

    if isinstance(storage, LocalStorage):
        raise HTTPException(
            status_code=400,
            detail="Current storage backend is local. Assign cloud storage to this tenant first.",
        )

    url_refs = await _collect_tenant_url_map(db, tenant_id)
    import logging as _logging
    _log = _logging.getLogger("migration")
    _log.warning(f"[migrate] tenant={tenant_id} UPLOAD_DIR={UPLOAD_DIR} url_refs_count={len(url_refs)}")
    if url_refs:
        _log.warning(f"[migrate] sample urls: {list(url_refs.items())[:3]}")
    if not url_refs:
        usage_mb = await _recalculate_tenant_storage_usage(db, tenant_id, storage, prefix)
        return _storage_migration_result(
            total=0,
            processed=0,
            migrated=0,
            skipped=0,
            already=0,
            errors=[],
            usage_mb=usage_mb,
            done=True,
        )

    cdn_prefix = getattr(storage, "_public_base", "") or ""
    skip_hosts: set[str] = set()
    try:
        from app.core.models.tenant import Tenant as _Tenant
        from app.core.models.storage_profile import StorageProfile as _SP
        _tr = await db.execute(select(_Tenant).where(_Tenant.id == tenant_id))
        _tenant = _tr.scalar_one_or_none()
        if _tenant and _tenant.storage_profile_id:
            _pr = await db.execute(select(_SP).where(_SP.id == _tenant.storage_profile_id))
            _profile = _pr.scalar_one_or_none()
            if _profile and _profile.migration_skip_hosts:
                skip_hosts = set(json.loads(_profile.migration_skip_hosts))
    except Exception:
        pass

    to_fetch: list[tuple[str, str]] = []
    skipped_urls: list[str] = []
    already_count = 0
    for old_url, fname in url_refs.items():
        if cdn_prefix and old_url.startswith(cdn_prefix):
            already_count += 1
            continue
        if skip_hosts:
            netloc = urlparse(old_url).netloc
            if any(netloc == h or netloc.endswith("." + h) for h in skip_hosts):
                skipped_urls.append(old_url)
                continue
        to_fetch.append((old_url, fname))

    total = len(url_refs)
    processed = already_count + len(skipped_urls)
    migrated = 0
    errors: list[dict] = []
    batch_size = max(1, int(batch_size or STORAGE_MIGRATION_BATCH_SIZE))

    await _update_storage_migration_task(
        db,
        task_id,
        _storage_migration_result(
            total=total,
            processed=processed,
            migrated=migrated,
            skipped=len(skipped_urls),
            already=already_count,
            errors=errors,
            done=False,
        ),
    )

    sem = asyncio.Semaphore(50)

    async def _fetch_one(url: str) -> bytes | None:
        async with sem:
            return await _fetch_existing_bytes(url)

    for offset in range(0, len(to_fetch), batch_size):
        chunk = to_fetch[offset:offset + batch_size]
        fetch_results = await asyncio.gather(*[_fetch_one(u) for u, _ in chunk])
        url_map: dict[str, str] = {}

        for (old_url, fname), data in zip(chunk, fetch_results):
            processed += 1
            if data is None:
                _log.warning(f"[migrate] SKIP: {old_url}")
                skipped_urls.append(old_url)
                continue
            try:
                upload_fname, upload_data, content_type = _prepare_migration_upload(fname, data)
                new_url = await storage.save(f"{prefix}{upload_fname}", upload_data, content_type)
                _log.warning(f"[migrate] OK: {old_url} -> {new_url}")
                url_map[old_url] = new_url
                migrated += 1
            except Exception as e:
                _log.warning(f"[migrate] ERROR: {old_url} {e}")
                errors.append({"file": fname, "url": old_url, "error": str(e)})

        if url_map:
            await _update_tenant_urls(db, tenant_id, url_map)
            await db.commit()

        await _update_storage_migration_task(
            db,
            task_id,
            _storage_migration_result(
                total=total,
                processed=processed,
                migrated=migrated,
                skipped=len(skipped_urls),
                already=already_count,
                errors=errors,
                done=False,
            ),
        )

    usage_mb = await _recalculate_tenant_storage_usage(db, tenant_id, storage, prefix)
    return _storage_migration_result(
        total=total,
        processed=processed,
        migrated=migrated,
        skipped=len(skipped_urls),
        already=already_count,
        errors=errors,
        usage_mb=usage_mb,
        done=True,
    )


@router.post("/migrate-to-cdn", summary="Start storage migration task")
async def migrate_to_cdn(
    background_tasks: BackgroundTasks,
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    return await _start_storage_migration_task(db, current_user.tenant_id, background_tasks)


@router.get("/migrate-to-cdn/{task_id}", summary="Get storage migration task status")
async def get_migration_task(
    task_id: int,
    db: AsyncSession = Depends(get_db),
    current_user: User = Depends(get_admin_user),
):
    return await _get_storage_migration_task(db, current_user.tenant_id, task_id)
