86 lines
3.1 KiB
Python
86 lines
3.1 KiB
Python
from __future__ import annotations
|
|
|
|
import logging
|
|
|
|
from app.clients.modelhub import ModelHubClient
|
|
from app.domain.model_type import engine_for_model
|
|
from app.domain.priorities import PRIORITY_OTHERS_DOWNLOADED
|
|
from app.settings import Settings
|
|
from app.storage.repositories import Repository
|
|
|
|
LOG = logging.getLogger(__name__)
|
|
|
|
SOURCE_MAP = {
|
|
"HuggingFace": "HUGGING_FACE",
|
|
"HUGGING_FACE": "HUGGING_FACE",
|
|
"ModelScope": "MODEL_SCOPE",
|
|
"MODEL_SCOPE": "MODEL_SCOPE",
|
|
}
|
|
|
|
|
|
def normalize_source(raw: str | None) -> str:
|
|
if not raw:
|
|
return "HUGGING_FACE"
|
|
return SOURCE_MAP.get(raw, raw)
|
|
|
|
|
|
class DownloadSuccessCrawler:
|
|
def __init__(self, settings: Settings, repo: Repository, client: ModelHubClient):
|
|
self.settings = settings
|
|
self.repo = repo
|
|
self.client = client
|
|
|
|
def refresh(self) -> int:
|
|
if not self.settings.enable_download_success_crawler:
|
|
return 0
|
|
imported = 0
|
|
seen: set[str] = set()
|
|
page = 1
|
|
while True:
|
|
data = self.client.list_success_download_tasks(page, self.settings.download_success_page_size)
|
|
records = data.get("records") or []
|
|
total = int(data.get("total") or 0)
|
|
if not records:
|
|
break
|
|
|
|
for record in records:
|
|
args = record.get("args") or {}
|
|
model_id = args.get("model_id") or args.get("modelId")
|
|
if not model_id or model_id in seen:
|
|
continue
|
|
seen.add(model_id)
|
|
source = normalize_source(args.get("source"))
|
|
engine = engine_for_model(model_id)
|
|
if engine != "vllm":
|
|
self.repo.upsert_candidate(
|
|
model_id=model_id,
|
|
source=source,
|
|
origin="model_download_success_page_gguf",
|
|
priority=PRIORITY_OTHERS_DOWNLOADED,
|
|
known_downloaded_by_others=True,
|
|
self_download_allowed=False,
|
|
notes=f"download task {record.get('id')} SUCCESS; GGUF/llamacpp candidate",
|
|
)
|
|
else:
|
|
self.repo.upsert_candidate(
|
|
model_id=model_id,
|
|
source=source,
|
|
origin="model_download_success_page",
|
|
priority=PRIORITY_OTHERS_DOWNLOADED,
|
|
known_downloaded_by_others=True,
|
|
self_download_allowed=False,
|
|
notes=f"download task {record.get('id')} SUCCESS",
|
|
)
|
|
imported += 1
|
|
|
|
if self.settings.download_success_max_pages and page >= self.settings.download_success_max_pages:
|
|
break
|
|
if total and page * self.settings.download_success_page_size >= total:
|
|
break
|
|
if len(records) < self.settings.download_success_page_size:
|
|
break
|
|
page += 1
|
|
|
|
LOG.info("download-success crawler imported/updated %s candidates", imported)
|
|
return imported
|