From 0cabc43abf212edfc60ee16d9cd5aa0abbbed8b1 Mon Sep 17 00:00:00 2001 From: CoolBoy Date: Sat, 1 Aug 2026 16:50:50 +0800 Subject: [PATCH] fix: coordinate concurrent account capacity filling --- .dockerignore | 2 + .gitignore | 2 + README.md | 26 ++- main.py | 8 +- modelhub_submmit_api/README.md | 12 +- modelhub_submmit_api/common.py | 103 +++++++- modelhub_submmit_api/daily_runner.py | 22 +- modelhub_submmit_api/history_stats.py | 10 +- modelhub_submmit_api/main.py | 80 +++++-- modelhub_submmit_api/modelhub_client.py | 272 +++++++++++++++++----- modelhub_submmit_api/outcome_tracker.py | 73 ++++-- modelhub_submmit_api/poll_runner.py | 27 ++- modelhub_submmit_api/submission_claims.py | 124 ++++++++++ tests/test_concurrency.py | 176 ++++++++++++++ 14 files changed, 794 insertions(+), 143 deletions(-) create mode 100644 modelhub_submmit_api/submission_claims.py create mode 100644 tests/test_concurrency.py diff --git a/.dockerignore b/.dockerignore index 39f942af..090c3843 100644 --- a/.dockerignore +++ b/.dockerignore @@ -27,3 +27,5 @@ modelhub_submmit_api/poll_runs/ modelhub_submmit_api/ledger/ modelhub_submmit_api/outcomes/ modelhub_submmit_api/history/ +.modelhub_state/ +modelhub_submmit_api/.modelhub_state/ diff --git a/.gitignore b/.gitignore index aa4525d5..c0661a88 100644 --- a/.gitignore +++ b/.gitignore @@ -24,3 +24,5 @@ modelhub_submmit_api/poll_runs/ modelhub_submmit_api/ledger/ modelhub_submmit_api/outcomes/ modelhub_submmit_api/history/ +.modelhub_state/ +modelhub_submmit_api/.modelhub_state/ diff --git a/README.md b/README.md index 6fef9acb..1fae39f7 100644 --- a/README.md +++ b/README.md @@ -26,15 +26,33 @@ the image. Optional tuning: -- `MODELHUB_AGENT_POLL_INTERVAL_SECONDS` default `300` -- `MODELHUB_AGENT_IDLE_INTERVAL_SECONDS` default `600` -- `MODELHUB_AGENT_POST_CYCLE_COOLDOWN_SECONDS` default `30` -- `MODELHUB_AGENT_MAX_SUBMITS_PER_RUN` default `5` +- `MODELHUB_AGENT_POLL_INTERVAL_SECONDS` default `15` +- `MODELHUB_AGENT_IDLE_INTERVAL_SECONDS` default `60` +- `MODELHUB_AGENT_POST_CYCLE_COOLDOWN_SECONDS` default `2` +- `MODELHUB_AGENT_MAX_SUBMITS_PER_RUN` default `0` (fill all currently available slots) +- `MODELHUB_AGENT_ACTIVE_TASK_CAP` default `100` per account +- `MODELHUB_AGENT_ACTIVE_COUNTS_TTL_SECONDS` default `15` +- `MODELHUB_AGENT_RESERVATION_TTL_SECONDS` default `120` +- `MODELHUB_AGENT_INSTANCE_ID` optional stable worker identity used to spread concurrent agents across accounts and candidates +- `MODELHUB_AGENT_CLAIMS_PATH` default `.modelhub_state/submission_claims.jsonl` - `MODELHUB_AGENT_DAILY_TARGET` - `MODELHUB_AGENT_MIN_DOWNLOADS` - `MODELHUB_AGENT_GPUS` - `MODELHUB_AGENT_EXTRA_ARGS` +## Concurrent Agents + +The token pool keeps a local reservation for every in-flight submission, so a +lagging platform count cannot send all concurrent requests to the same account. +If another process fills an account first, the submission is retried immediately +against another account with capacity. + +Workers that share a filesystem also coordinate model/GPU claims through +`.modelhub_state/submission_claims.jsonl`. Workers in isolated containers use +different candidate ordering (derived from `STRATEGY_ID`, instance ID, or +hostname), which reduces duplicate work while the platform remains the final +authority for account capacity and model/GPU uniqueness. + ## Deploy Create a tag and submit the repository URL plus tag in "我的适配智能体". diff --git a/main.py b/main.py index 0df81488..5610a6de 100644 --- a/main.py +++ b/main.py @@ -46,13 +46,13 @@ def _worker_command() -> list[str]: "-u", str(WORKER_SCRIPT), "--poll-interval-seconds", - os.getenv("MODELHUB_AGENT_POLL_INTERVAL_SECONDS", "300"), + os.getenv("MODELHUB_AGENT_POLL_INTERVAL_SECONDS", "15"), "--idle-interval-seconds", - os.getenv("MODELHUB_AGENT_IDLE_INTERVAL_SECONDS", "600"), + os.getenv("MODELHUB_AGENT_IDLE_INTERVAL_SECONDS", "60"), "--post-cycle-cooldown-seconds", - os.getenv("MODELHUB_AGENT_POST_CYCLE_COOLDOWN_SECONDS", "30"), + os.getenv("MODELHUB_AGENT_POST_CYCLE_COOLDOWN_SECONDS", "2"), "--max-submits-per-run", - os.getenv("MODELHUB_AGENT_MAX_SUBMITS_PER_RUN", "5"), + os.getenv("MODELHUB_AGENT_MAX_SUBMITS_PER_RUN", "0"), "--skip-history-archive", ] diff --git a/modelhub_submmit_api/README.md b/modelhub_submmit_api/README.md index 604598d9..cc6bc3a7 100644 --- a/modelhub_submmit_api/README.md +++ b/modelhub_submmit_api/README.md @@ -74,6 +74,9 @@ bash run_poll.sh --dry-run - The runner auto-discovers all safe GPU/template combinations from the public submit catalog. - Each model can be submitted at most once per GPU. - Multiple ModelHub tokens are pooled and used to route submissions to the account with available async capacity. +- Concurrent submissions reserve account slots locally, and an account-capacity race automatically falls through to another account. +- Concurrent local processes claim model/GPU pairs in `.modelhub_state/submission_claims.jsonl`; shared ledger, history, and outcome files use process locks and atomic replacement. +- Isolated agent containers diversify candidate order by instance identity to reduce cross-container duplicate submissions. - History stats are online-only. The local ledger is used for local accounting, but platform history is only used after the local ledger reaches the configured threshold. - The default history threshold is `500` records. @@ -96,15 +99,15 @@ Common flags: `run_poll.sh` adds: -- `--poll-interval-seconds`: sleep when all accounts are saturated -- `--idle-interval-seconds`: sleep when a cycle submits nothing +- `--poll-interval-seconds`: sleep when all accounts are saturated (default 15) +- `--idle-interval-seconds`: sleep when a cycle submits nothing (default 60) - `--max-scan-models`: hard cap on scanned HF models for this cycle (0 = auto) - `--scan-multiplier`: multiplier used for auto scan cap derivation from quota/queue capacity - `--max-submits-per-run`: max tasks to submit per poll cycle (0 = unlimited) - `--skip-outcome-sync`: skip outcome sync before scanning - `--skip-history-archive`: skip history archive download for this cycle - `--submit-concurrency`: concurrent task submission calls used by each cycle (0 = auto) -- `--post-cycle-cooldown-seconds`: pause after a successful cycle before next cycle (default 0) +- `--post-cycle-cooldown-seconds`: pause after a successful cycle before next cycle (default 2) - `--max-cycles`: optional hard stop for testing or batch windows ## Output @@ -140,3 +143,6 @@ python3 -m unittest discover -s tests -v - The platform still enforces per-account async capacity limits, so the poller can keep the queue close to full but cannot override the platform cap. - `bash run_poll.sh` now defaults to unlimited mode and keeps refilling until you stop the process manually. +- Queue polling defaults to 15 seconds, successful-cycle cooldown to 2 seconds, + and per-cycle submissions to all available slots. Override these values when + the platform requires a lower request rate. diff --git a/modelhub_submmit_api/common.py b/modelhub_submmit_api/common.py index 5e6539b9..b17221de 100644 --- a/modelhub_submmit_api/common.py +++ b/modelhub_submmit_api/common.py @@ -1,9 +1,15 @@ from __future__ import annotations +import fcntl +import hashlib import json +import os +import socket +import tempfile +from contextlib import contextmanager from datetime import datetime, timezone from pathlib import Path -from typing import Any +from typing import Any, Callable, Iterator def ensure_utc(value: datetime) -> datetime: @@ -54,16 +60,57 @@ def json_dumps(value: Any) -> str: return json.dumps(value, ensure_ascii=False, sort_keys=True) -def read_json(path: Path) -> Any: - return json.loads(path.read_text(encoding="utf-8")) +def runtime_instance_id() -> str: + """Return a non-secret identity that differs between concurrent workers.""" + base = ( + os.getenv("MODELHUB_AGENT_INSTANCE_ID") + or os.getenv("STRATEGY_ID") + or os.getenv("HOSTNAME") + or socket.gethostname() + ) + raw = f"{base}:{os.getpid()}" + return hashlib.sha256(raw.encode("utf-8")).hexdigest()[:20] -def write_json(path: Path, value: Any) -> None: +def _lock_path(path: Path) -> Path: + return path.with_name(f".{path.name}.lock") + + +@contextmanager +def file_lock(path: Path, *, shared: bool = False) -> Iterator[None]: + """Coordinate access to a path between processes on the same filesystem.""" path.parent.mkdir(parents=True, exist_ok=True) - path.write_text(json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n", encoding="utf-8") + lock_path = _lock_path(path) + with lock_path.open("a+b") as handle: + fcntl.flock(handle.fileno(), fcntl.LOCK_SH if shared else fcntl.LOCK_EX) + try: + yield + finally: + fcntl.flock(handle.fileno(), fcntl.LOCK_UN) -def read_jsonl(path: Path) -> list[dict[str, Any]]: +def _atomic_write_text(path: Path, content: str) -> None: + path.parent.mkdir(parents=True, exist_ok=True) + descriptor, temporary_name = tempfile.mkstemp( + dir=str(path.parent), + prefix=f".{path.name}.", + suffix=".tmp", + ) + try: + with os.fdopen(descriptor, "w", encoding="utf-8") as handle: + handle.write(content) + handle.flush() + os.fsync(handle.fileno()) + os.replace(temporary_name, path) + except Exception: + try: + os.unlink(temporary_name) + except FileNotFoundError: + pass + raise + + +def _read_jsonl_unlocked(path: Path) -> list[dict[str, Any]]: if not path.exists(): return [] rows: list[dict[str, Any]] = [] @@ -74,13 +121,45 @@ def read_jsonl(path: Path) -> list[dict[str, Any]]: return rows +def _jsonl_content(rows: list[dict[str, Any]]) -> str: + return "".join(json.dumps(row, ensure_ascii=False, sort_keys=True) + "\n" for row in rows) + + +def read_json(path: Path) -> Any: + with file_lock(path, shared=True): + return json.loads(path.read_text(encoding="utf-8")) + + +def write_json(path: Path, value: Any) -> None: + content = json.dumps(value, ensure_ascii=False, indent=2, sort_keys=True) + "\n" + with file_lock(path): + _atomic_write_text(path, content) + + +def read_jsonl(path: Path) -> list[dict[str, Any]]: + with file_lock(path, shared=True): + return _read_jsonl_unlocked(path) + + def write_jsonl(path: Path, rows: list[dict[str, Any]]) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - content = "".join(json.dumps(row, ensure_ascii=False, sort_keys=True) + "\n" for row in rows) - path.write_text(content, encoding="utf-8") + with file_lock(path): + _atomic_write_text(path, _jsonl_content(rows)) def append_jsonl(path: Path, row: dict[str, Any]) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - with path.open("a", encoding="utf-8") as handle: - handle.write(json.dumps(row, ensure_ascii=False, sort_keys=True) + "\n") + line = json.dumps(row, ensure_ascii=False, sort_keys=True) + "\n" + with file_lock(path): + with path.open("a", encoding="utf-8") as handle: + handle.write(line) + handle.flush() + + +def update_jsonl( + path: Path, + updater: Callable[[list[dict[str, Any]]], list[dict[str, Any]]], +) -> list[dict[str, Any]]: + """Atomically read, update, and replace a JSONL file under one lock.""" + with file_lock(path): + updated = updater(_read_jsonl_unlocked(path)) + _atomic_write_text(path, _jsonl_content(updated)) + return updated diff --git a/modelhub_submmit_api/daily_runner.py b/modelhub_submmit_api/daily_runner.py index 812d76a5..2ac970e1 100644 --- a/modelhub_submmit_api/daily_runner.py +++ b/modelhub_submmit_api/daily_runner.py @@ -8,10 +8,11 @@ from typing import Any, Callable from common import utc_now, write_json from hf_discovery import HuggingFaceDiscovery -from main import DEFAULT_LEDGER_PATH, DEFAULT_RUNS_DIR, run_submission +from main import DEFAULT_LEDGER_PATH, DEFAULT_RUNS_DIR, make_run_dir, run_submission from modelhub_client import ModelHubClient, ModelHubClientPool from outcome_tracker import OutcomeTracker from runner_common import DEFAULT_KEY_PATH, ensure_tokens +from submission_claims import DEFAULT_CLAIMS_PATH from template_selector import TemplateSelector @@ -81,6 +82,11 @@ def build_parser() -> argparse.ArgumentParser: parser.add_argument("--key-path", default=str(DEFAULT_KEY_PATH), help="Path to KEY.md containing MODELSCOPE_TOKEN/XC_TOKEN") parser.add_argument("--runs-dir", default=str(DEFAULT_RUNS_DIR), help=argparse.SUPPRESS) parser.add_argument("--ledger-path", default=str(DEFAULT_LEDGER_PATH), help=argparse.SUPPRESS) + parser.add_argument( + "--claims-path", + default=os.getenv("MODELHUB_AGENT_CLAIMS_PATH", str(DEFAULT_CLAIMS_PATH)), + help=argparse.SUPPRESS, + ) parser.add_argument("--history-archive-path", default="history/platform_tasks.jsonl", help=argparse.SUPPRESS) parser.add_argument("--history-archive-limit", type=int, default=5000, help=argparse.SUPPRESS) parser.add_argument("--daily-runs-dir", default=str(DEFAULT_DAILY_RUNS_DIR), help=argparse.SUPPRESS) @@ -117,6 +123,7 @@ def make_wave_namespace(base_args: argparse.Namespace, wave: WaveSpec) -> argpar runs_dir=base_args.runs_dir, ledger_path=base_args.ledger_path, outcomes_path=getattr(base_args, "outcomes_path", "outcomes/submissions.jsonl"), + claims_path=getattr(base_args, "claims_path", str(DEFAULT_CLAIMS_PATH)), history_archive_path=base_args.history_archive_path, history_archive_limit=base_args.history_archive_limit, hf_base_url=base_args.hf_base_url, @@ -132,7 +139,7 @@ def run_daily_batches( now=None, run_fn: Callable[..., dict[str, Any]] = run_submission, hf_discovery: HuggingFaceDiscovery | None = None, - modelhub_client: ModelHubClient | None = None, + modelhub_client: ModelHubClient | ModelHubClientPool | None = None, template_selector: TemplateSelector | None = None, outcome_tracker: OutcomeTracker | None = None, ) -> dict[str, Any]: @@ -140,15 +147,12 @@ def run_daily_batches( hf_discovery = hf_discovery or HuggingFaceDiscovery(base_url=base_args.hf_base_url) if modelhub_client is None: modelhub_tokens = list(getattr(base_args, "modelhub_tokens", None) or ([] if not base_args.modelhub_token else [base_args.modelhub_token])) - if len(modelhub_tokens) > 1: - clients = [ModelHubClient(token=token, base_url=base_args.modelhub_base_url) for token in modelhub_tokens] - modelhub_client = ModelHubClientPool(clients) - else: - modelhub_client = ModelHubClient(token=base_args.modelhub_token, base_url=base_args.modelhub_base_url) + token_values: list[str | None] = modelhub_tokens or [base_args.modelhub_token] + clients = [ModelHubClient(token=token, base_url=base_args.modelhub_base_url) for token in token_values] + modelhub_client = ModelHubClientPool(clients) template_selector = template_selector or TemplateSelector() - daily_run_dir = Path(base_args.daily_runs_dir) / now.strftime("%Y%m%dT%H%M%SZ") - daily_run_dir.mkdir(parents=True, exist_ok=True) + daily_run_dir = make_run_dir(Path(base_args.daily_runs_dir), now) log(f"[daily] daily_run_dir={daily_run_dir}") log( f"[daily] target={base_args.daily_target} rounds={base_args.rounds} " diff --git a/modelhub_submmit_api/history_stats.py b/modelhub_submmit_api/history_stats.py index 387dfa48..085f1f00 100644 --- a/modelhub_submmit_api/history_stats.py +++ b/modelhub_submmit_api/history_stats.py @@ -5,7 +5,7 @@ from datetime import datetime from pathlib import Path from typing import Any -from common import append_jsonl, ensure_utc, parse_datetime, read_jsonl, write_jsonl +from common import append_jsonl, ensure_utc, parse_datetime, read_jsonl, update_jsonl WAITING_STATUSES = {"waiting", "queued"} @@ -32,10 +32,10 @@ def update_history_archive( *, limit: int = HISTORY_ARCHIVE_LIMIT_DEFAULT, ) -> list[dict[str, Any]]: - merged = merge_history_records(load_history_archive(path), tasks, limit=limit) - path.parent.mkdir(parents=True, exist_ok=True) - write_jsonl(path, merged) - return merged + return update_jsonl( + path, + lambda existing: merge_history_records(existing, tasks, limit=limit), + ) def merge_history_records( diff --git a/modelhub_submmit_api/main.py b/modelhub_submmit_api/main.py index 0d6a8c09..9f2850f1 100644 --- a/modelhub_submmit_api/main.py +++ b/modelhub_submmit_api/main.py @@ -7,7 +7,7 @@ from datetime import timedelta from pathlib import Path from typing import Any -from common import parse_datetime, utc_now, write_json, write_jsonl +from common import parse_datetime, runtime_instance_id, utc_now, write_json, write_jsonl from hf_discovery import HuggingFaceDiscovery from history_stats import ( append_ledger_entry, @@ -19,6 +19,7 @@ from history_stats import ( from modelhub_client import ModelHubAPIError, ModelHubClient, ModelHubClientPool from models import CandidateModel, HFModelSummary, ModelInspection from outcome_tracker import DEFAULT_OUTCOMES_PATH, OutcomeTracker +from submission_claims import DEFAULT_CLAIMS_PATH, SubmissionClaimStore, diversify_candidates from task_registry import TASK_SPEC_BY_TYPE, all_task_types, choose_framework_for_task, choose_text_generation_framework, pipeline_tags_for_task_types, task_specs_for_model from template_selector import TemplateSelector @@ -71,6 +72,11 @@ def build_parser() -> argparse.ArgumentParser: parser.add_argument("--runs-dir", default=str(DEFAULT_RUNS_DIR), help=argparse.SUPPRESS) parser.add_argument("--ledger-path", default=str(DEFAULT_LEDGER_PATH), help=argparse.SUPPRESS) parser.add_argument("--outcomes-path", default=str(DEFAULT_OUTCOMES_PATH), help=argparse.SUPPRESS) + parser.add_argument( + "--claims-path", + default=os.getenv("MODELHUB_AGENT_CLAIMS_PATH", str(DEFAULT_CLAIMS_PATH)), + help=argparse.SUPPRESS, + ) parser.add_argument("--history-archive-path", default="history/platform_tasks.jsonl", help=argparse.SUPPRESS) parser.add_argument("--history-archive-limit", type=int, default=5000, help=argparse.SUPPRESS) parser.add_argument("--hf-base-url", default=os.getenv("MODELSCOPE_BASE_URL", "https://modelscope.cn"), help=argparse.SUPPRESS) @@ -216,9 +222,7 @@ def resolve_max_submit_count( remaining_daily_quota: int, ) -> int: explicit_limit = int(getattr(args, "max_submits_per_run", 0) or 0) - planned_submit_count = planned_count - if remaining_daily_quota > 0: - planned_submit_count = min(planned_submit_count, remaining_daily_quota) + planned_submit_count = min(planned_count, max(0, remaining_daily_quota)) if explicit_limit > 0: planned_submit_count = min(planned_submit_count, explicit_limit) return planned_submit_count @@ -328,16 +332,16 @@ def submit_candidate( def make_run_dir(runs_dir: Path, now) -> Path: + runs_dir.mkdir(parents=True, exist_ok=True) base_name = now.strftime("%Y%m%dT%H%M%SZ") - run_dir = runs_dir / base_name - if not run_dir.exists(): - return run_dir - - for suffix in range(1, 1000): - candidate = runs_dir / f"{base_name}.{suffix:03d}" - if not candidate.exists(): + for suffix in range(0, 1000): + candidate = runs_dir / (base_name if suffix == 0 else f"{base_name}.{suffix:03d}") + try: + candidate.mkdir(exist_ok=False) return candidate - return candidate + except FileExistsError: + continue + raise RuntimeError(f"Unable to allocate a unique run directory under {runs_dir}") def run_submission( @@ -359,17 +363,17 @@ def run_submission( hf_discovery = hf_discovery or HuggingFaceDiscovery(base_url=args.hf_base_url) if modelhub_client is None: modelhub_tokens = list(getattr(args, "modelhub_tokens", None) or ([] if not args.modelhub_token else [args.modelhub_token])) - if len(modelhub_tokens) > 1: - clients = [ModelHubClient(token=token, base_url=args.modelhub_base_url) for token in modelhub_tokens] - modelhub_client = ModelHubClientPool(clients) - else: - modelhub_client = ModelHubClient(token=args.modelhub_token, base_url=args.modelhub_base_url) + token_values: list[str | None] = modelhub_tokens or [args.modelhub_token] + clients = [ModelHubClient(token=token, base_url=args.modelhub_base_url) for token in token_values] + modelhub_client = ModelHubClientPool(clients) + + if hasattr(modelhub_client, "begin_cycle"): + modelhub_client.begin_cycle() runs_dir = Path(args.runs_dir) ledger_path = Path(args.ledger_path) history_archive_path = Path(args.history_archive_path) run_dir = make_run_dir(runs_dir, now) - run_dir.mkdir(parents=True, exist_ok=True) ledger_path.parent.mkdir(parents=True, exist_ok=True) history_archive_path.parent.mkdir(parents=True, exist_ok=True) @@ -387,8 +391,26 @@ def run_submission( day_start = now.replace(hour=0, minute=0, second=0, microsecond=0) # Count today's submissions from the local ledger (avoids expensive paginated API call) daily_snapshot = count_submissions_for_day(tasks=[], ledger_entries=ledger_entries, day_start=day_start, day_end=now) - # Fallback: if ledger has no entries yet, do a quick API check (limit to 1 page) - if daily_snapshot["totalCount"] <= 0: + # A configured daily target must include every pooled account and every + # concurrent worker, not only the first token's first page. + if args.daily_target > 0: + list_kwargs: dict[str, Any] = { + "page_size": 100, + "only_mine": True, + "begin_time": day_start, + "end_time": now, + } + if isinstance(modelhub_client, ModelHubClientPool): + list_kwargs["_fanout_all"] = True + today_tasks = modelhub_client.list_tasks(**list_kwargs) + daily_snapshot = count_submissions_for_day( + tasks=today_tasks, + ledger_entries=ledger_entries, + day_start=day_start, + day_end=now, + ) + # In unlimited mode, keep the inexpensive one-page fallback for an empty ledger. + elif daily_snapshot["totalCount"] <= 0: today_tasks_page = modelhub_client.list_tasks_page( current=1, page_size=20, @@ -533,7 +555,17 @@ def run_submission( planned_count=len(candidates), remaining_daily_quota=remaining_daily_quota, ) - planned_candidates = candidates[:planned_submit_count] + instance_id = runtime_instance_id() + diversified_candidates = diversify_candidates(candidates, instance_id=instance_id) + claim_store: SubmissionClaimStore | None = None + if args.dry_run: + planned_candidates = diversified_candidates[:planned_submit_count] + else: + claim_store = SubmissionClaimStore( + Path(getattr(args, "claims_path", DEFAULT_CLAIMS_PATH)), + owner_id=instance_id, + ) + planned_candidates = claim_store.claim(diversified_candidates, limit=planned_submit_count) submit_workers = 1 if not args.dry_run: submit_workers = resolve_submit_concurrency( @@ -609,6 +641,12 @@ def run_submission( submit_time=result["submitTime"], ) + if claim_store is not None: + submitted_candidates = [result["candidate"] for result in ordered_results.values() if result.get("outcome") == "submitted"] + failed_candidates = [result["candidate"] for result in ordered_results.values() if result.get("outcome") != "submitted"] + claim_store.mark_submitted(submitted_candidates) + claim_store.release(failed_candidates) + write_jsonl(run_dir / "submitted.jsonl", submitted) write_jsonl(run_dir / "skipped.jsonl", skipped) write_jsonl(run_dir / "failed.jsonl", failed) diff --git a/modelhub_submmit_api/modelhub_client.py b/modelhub_submmit_api/modelhub_client.py index e322395a..45fc4190 100644 --- a/modelhub_submmit_api/modelhub_client.py +++ b/modelhub_submmit_api/modelhub_client.py @@ -1,13 +1,14 @@ from __future__ import annotations +import hashlib import os -import time import threading +import time from concurrent.futures import ThreadPoolExecutor, as_completed from datetime import datetime, timedelta from typing import Any -from common import format_modelhub_datetime, parse_datetime +from common import format_modelhub_datetime, parse_datetime, runtime_instance_id from defaults import EMBEDDED_MODELHUB_XC_TOKEN from http_json import HttpJsonError, JsonHttpClient @@ -233,18 +234,63 @@ def is_active_task(task: dict[str, Any]) -> bool: return True +CAPACITY_ERROR_MARKERS = ( + "达到上限", + "达上限", + "任务数量已达", + "队列已满", + "queue is full", + "queue full", + "capacity", + "too many active", + "active task limit", +) + + +def is_capacity_error(error: ModelHubAPIError) -> bool: + if error.code in {409, 429}: + return True + message = str(error).strip().lower() + if isinstance(error.payload, dict): + message = f"{message} {error.payload.get('message') or ''}".lower() + return any(marker in message for marker in CAPACITY_ERROR_MARKERS) + + class ModelHubClientPool: - def __init__(self, clients: list[ModelHubClient], *, active_task_cap: int = 100) -> None: + def __init__( + self, + clients: list[ModelHubClient], + *, + active_task_cap: int | None = None, + active_counts_ttl: float | None = None, + reservation_ttl: float | None = None, + instance_id: str | None = None, + ) -> None: if not clients: raise ValueError("At least one ModelHub client is required") self.clients = clients - self.active_task_cap = active_task_cap - self._active_counts: list[int] = [0 for _ in clients] + configured_cap = active_task_cap if active_task_cap is not None else os.getenv("MODELHUB_AGENT_ACTIVE_TASK_CAP", "100") + self.active_task_cap = max(1, int(configured_cap)) + configured_ttl = active_counts_ttl if active_counts_ttl is not None else os.getenv("MODELHUB_AGENT_ACTIVE_COUNTS_TTL_SECONDS", "15") + configured_reservation_ttl = ( + reservation_ttl + if reservation_ttl is not None + else os.getenv("MODELHUB_AGENT_RESERVATION_TTL_SECONDS", "120") + ) + self._active_counts_ttl = max(1.0, float(configured_ttl)) + self._reservation_ttl = max(self._active_counts_ttl * 2, float(configured_reservation_ttl)) + self._remote_counts: list[int] = [0 for _ in clients] self._active_refresh_at: float = 0.0 - self._active_counts_ttl: float = 30.0 + self._counts_initialized = False self._state_lock = threading.Lock() - # Per-cycle cache for search_by_model_id results (model_id -> merged verify result map) - self._verify_cache: dict[str, dict[str, Any]] = {} + self._refresh_lock = threading.Lock() + self._reservations: list[dict[int, dict[str, Any]]] = [{} for _ in clients] + self._reservation_sequence = 0 + identity = instance_id or runtime_instance_id() + identity_hash = hashlib.sha256(identity.encode("utf-8")).hexdigest() + self._selection_cursor = int(identity_hash[:12], 16) % len(clients) + self._verify_cache: dict[str, tuple[float, dict[str, Any]]] = {} + self._verify_cache_ttl = max(1.0, float(os.getenv("MODELHUB_AGENT_VERIFY_CACHE_TTL_SECONDS", "30"))) # Single reader client to avoid fanout on read operations self._reader = clients[0] @@ -267,48 +313,86 @@ class ModelHubClientPool: except Exception: return [] + def _counts_are_fresh_locked(self, now: float) -> bool: + return self._counts_initialized and (now - self._active_refresh_at) < self._active_counts_ttl + def _refresh_active_counts(self, *, force: bool = False) -> None: - now = time.time() - if (not force) and self._active_counts and (now - self._active_refresh_at) < self._active_counts_ttl: - return + now = time.monotonic() + with self._state_lock: + if not force and self._counts_are_fresh_locked(now): + return - def _to_indexed_result(index: int, client: ModelHubClient) -> tuple[int, int]: - return index, self._safe_count_active_tasks(client) + # Do not hold the scheduler lock during remote I/O. One refresher is + # enough; all submitting threads can continue using their reservations. + with self._refresh_lock: + now = time.monotonic() + with self._state_lock: + if not force and self._counts_are_fresh_locked(now): + return - results: list[tuple[int, int]] = [] - with ThreadPoolExecutor(max_workers=min(len(self.clients), 12)) as executor: - futures = {executor.submit(_to_indexed_result, index, client): index for index, client in enumerate(self.clients)} - for future in as_completed(futures): - index = futures[future] - try: - results.append(future.result()) - except Exception: - results.append((index, self.active_task_cap)) + def _to_indexed_result(index: int, client: ModelHubClient) -> tuple[int, int]: + return index, self._safe_count_active_tasks(client) - self._active_counts = [0 for _ in self.clients] - for index, count in sorted(results, key=lambda item: item[0]): - self._active_counts[index] = count - # Use current time after refresh completes, not the stale 'now' from function start - self._active_refresh_at = time.time() + results: list[tuple[int, int]] = [] + with ThreadPoolExecutor(max_workers=min(len(self.clients), 12)) as executor: + futures = { + executor.submit(_to_indexed_result, index, client): index + for index, client in enumerate(self.clients) + } + for future in as_completed(futures): + index = futures[future] + try: + results.append(future.result()) + except Exception: + results.append((index, self.active_task_cap)) + + refreshed_at = time.monotonic() + with self._state_lock: + for index, count in results: + old_remote_count = self._remote_counts[index] + new_remote_count = min(self.active_task_cap, max(0, int(count))) + acknowledged = max(0, new_remote_count - old_remote_count) + completed_reservations = sorted( + ( + (reservation_id, reservation) + for reservation_id, reservation in self._reservations[index].items() + if not reservation["inflight"] + ), + key=lambda item: item[1]["updated_at"], + ) + for reservation_id, _reservation in completed_reservations[:acknowledged]: + self._reservations[index].pop(reservation_id, None) + + # A remote count can stay flat when one old task finishes as + # one new task appears. Expiry prevents that net-zero update + # from reserving a slot forever. + for reservation_id, reservation in list(self._reservations[index].items()): + if reservation["inflight"]: + continue + if refreshed_at - reservation["updated_at"] >= self._reservation_ttl: + self._reservations[index].pop(reservation_id, None) + self._remote_counts[index] = new_remote_count + self._counts_initialized = True + self._active_refresh_at = refreshed_at + + def _effective_count_locked(self, index: int) -> int: + return self._remote_counts[index] + len(self._reservations[index]) + + def _counts_snapshot_locked(self) -> list[int]: + return [min(self.active_task_cap, self._effective_count_locked(index)) for index in range(len(self.clients))] def active_task_counts(self) -> list[int]: + self._refresh_active_counts() with self._state_lock: - now = time.time() - if self._active_counts and (now - self._active_refresh_at) < self._active_counts_ttl: - pass # use cached counts - else: - self._refresh_active_counts() - return list(self._active_counts) + return self._counts_snapshot_locked() def available_submit_slots(self) -> int: + self._refresh_active_counts() with self._state_lock: - # Only refresh if cache is stale (respects TTL) or counts are empty - now = time.time() - if self._active_counts and (now - self._active_refresh_at) < self._active_counts_ttl: - pass # use cached counts - else: - self._refresh_active_counts() - return sum(max(0, self.active_task_cap - count) for count in self._active_counts) + return sum( + max(0, self.active_task_cap - self._effective_count_locked(index)) + for index in range(len(self.clients)) + ) def list_tasks(self, **kwargs): # noqa: ANN003, ANN001 # Default: use only the reader client to avoid fanout amplification. @@ -339,10 +423,19 @@ class ModelHubClientPool: cached = self._verify_cache.get(model_id) if cached is None: return None - return dict(cached) + cached_at, payload = cached + if time.monotonic() - cached_at >= self._verify_cache_ttl: + self._verify_cache.pop(model_id, None) + return None + return dict(payload) def _write_to_cache(self, model_id: str, payload: dict[str, Any]) -> None: - self._verify_cache[model_id] = payload + self._verify_cache[model_id] = (time.monotonic(), payload) + + def begin_cycle(self) -> None: + """Drop model verification cache entries from the previous scan cycle.""" + with self._state_lock: + self._verify_cache.clear() def search_by_model_id(self, model_id: str) -> dict[str, Any]: # Check cache first (per-cycle cache to avoid repeated API calls for the same model) @@ -368,30 +461,87 @@ class ModelHubClientPool: def processed_gpus_for_model(self, model_id: str) -> set[str]: return set(self.get_verify_result_map(model_id).keys()) - def add_task(self, payload: dict[str, Any]) -> dict[str, Any]: + def _reserve_account(self, excluded: set[int]) -> tuple[int, int] | None: with self._state_lock: - self._refresh_active_counts(force=True) - selected_index = None - best_remaining = -1 - for index, active_count in enumerate(self._active_counts): - remaining = self.active_task_cap - active_count - if remaining > best_remaining: - best_remaining = remaining - selected_index = index - if selected_index is None or best_remaining <= 0: + remaining_by_index = { + index: self.active_task_cap - self._effective_count_locked(index) + for index in range(len(self.clients)) + if index not in excluded + } + usable = {index: remaining for index, remaining in remaining_by_index.items() if remaining > 0} + if not usable: + return None + best_remaining = max(usable.values()) + tied = [index for index, remaining in usable.items() if remaining == best_remaining] + selected_index = min( + tied, + key=lambda index: (index - self._selection_cursor) % len(self.clients), + ) + self._selection_cursor = (selected_index + 1) % len(self.clients) + self._reservation_sequence += 1 + reservation_id = self._reservation_sequence + self._reservations[selected_index][reservation_id] = { + "inflight": True, + "updated_at": time.monotonic(), + } + return selected_index, reservation_id + + def _finish_reservation(self, index: int, reservation_id: int, *, succeeded: bool) -> None: + with self._state_lock: + reservation = self._reservations[index].get(reservation_id) + if reservation is None: + return + if not succeeded: + self._reservations[index].pop(reservation_id, None) + return + reservation["inflight"] = False + reservation["updated_at"] = time.monotonic() + + def _mark_account_saturated(self, index: int) -> None: + with self._state_lock: + self._remote_counts[index] = self.active_task_cap + self._counts_initialized = True + self._active_refresh_at = time.monotonic() + + def add_task(self, payload: dict[str, Any]) -> dict[str, Any]: + self._refresh_active_counts() + attempted_accounts: set[int] = set() + forced_refresh_done = False + last_capacity_error: ModelHubAPIError | None = None + + while True: + reservation = self._reserve_account(attempted_accounts) + if reservation is None: + if not forced_refresh_done: + self._refresh_active_counts(force=True) + forced_refresh_done = True + continue + if last_capacity_error is not None: + raise last_capacity_error raise ModelHubAPIError( f"当前等待中或运行中的异步模型验证任务数量已达上限({self.active_task_cap})" ) - self._active_counts[selected_index] += 1 - selected_client = self.clients[selected_index] - try: - response = selected_client.add_task(payload) - except Exception: - with self._state_lock: - self._active_counts[selected_index] -= 1 - raise - return response + selected_index, reservation_id = reservation + selected_client = self.clients[selected_index] + try: + response = selected_client.add_task(payload) + except ModelHubAPIError as exc: + self._finish_reservation(selected_index, reservation_id, succeeded=False) + if not is_capacity_error(exc): + raise + # Another process may have filled this account after our count + # refresh. Mark it full locally and immediately try another one. + self._mark_account_saturated(selected_index) + attempted_accounts.add(selected_index) + last_capacity_error = exc + continue + except Exception: + self._finish_reservation(selected_index, reservation_id, succeeded=False) + raise + + self._finish_reservation(selected_index, reservation_id, succeeded=True) + return response def list_tasks_page(self, **kwargs): # noqa: ANN003, ANN001 """Single-page task listing via the reader client (no fanout).""" diff --git a/modelhub_submmit_api/outcome_tracker.py b/modelhub_submmit_api/outcome_tracker.py index 584ed797..79e36893 100644 --- a/modelhub_submmit_api/outcome_tracker.py +++ b/modelhub_submmit_api/outcome_tracker.py @@ -5,7 +5,7 @@ from datetime import datetime, timedelta from pathlib import Path from typing import Any -from common import append_jsonl, parse_datetime, read_jsonl, utc_now, write_jsonl +from common import append_jsonl, parse_datetime, read_jsonl, update_jsonl, utc_now from history_stats import classify_failure, is_failure, is_success from modelhub_client import ModelHubClient, ModelHubClientPool @@ -25,9 +25,20 @@ class OutcomeTracker: self._by_model_gpu: dict[tuple[str, str], list[dict[str, Any]]] = defaultdict(list) self._failed_model_gpus: set[tuple[str, str]] = set() - loaded = read_jsonl(self.path) - for record in loaded: - self._records.append(record) + self._records = read_jsonl(self.path) + self._rebuild_indexes() + + last_sync_times = [ + parse_datetime(record.get("lastSyncTime")) + for record in self._records + if record.get("lastSyncTime") + ] + self._last_sync_time: datetime = max(last_sync_times) if last_sync_times else utc_now() - timedelta(days=7) + + def _rebuild_indexes(self) -> None: + self._by_task_id.clear() + self._by_model_gpu.clear() + for record in self._records: task_id = record.get("taskId") if task_id: self._by_task_id[str(task_id)] = record @@ -37,13 +48,6 @@ class OutcomeTracker: self._rebuild_failed_index() - last_sync_times = [ - parse_datetime(record.get("lastSyncTime")) - for record in self._records - if record.get("lastSyncTime") - ] - self._last_sync_time: datetime = max(last_sync_times) if last_sync_times else utc_now() - timedelta(days=7) - def record_submission( self, model_id: str, @@ -162,8 +166,23 @@ class OutcomeTracker: } def save(self) -> None: - self.path.parent.mkdir(parents=True, exist_ok=True) - write_jsonl(self.path, self._records) + local_records = list(self._records) + + def merge(existing: list[dict[str, Any]]) -> list[dict[str, Any]]: + merged = list(existing) + index_by_key = {_outcome_record_key(record): index for index, record in enumerate(merged)} + for record in local_records: + key = _outcome_record_key(record) + existing_index = index_by_key.get(key) + if existing_index is None: + index_by_key[key] = len(merged) + merged.append(record) + continue + merged[existing_index] = _prefer_newer_outcome(merged[existing_index], record) + return merged + + self._records = update_jsonl(self.path, merge) + self._rebuild_indexes() def _rebuild_failed_index(self) -> None: self._failed_model_gpus.clear() @@ -234,3 +253,31 @@ def _summarize(records: list[dict[str, Any]]) -> dict[str, Any]: "pendingRate": round(pending_count / total, 4) if total > 0 else 0.0, "failureBreakdown": dict(failure_breakdown), } + + +def _outcome_record_key(record: dict[str, Any]) -> str: + task_id = record.get("taskId") + if task_id is not None: + return f"task:{task_id}" + return "fallback:{model}|{gpu}|{time}".format( + model=record.get("modelId") or "", + gpu=record.get("targetGpu") or "", + time=record.get("submitTime") or "", + ) + + +def _outcome_version(record: dict[str, Any]) -> tuple[int, float, float]: + last_sync = parse_datetime(record.get("lastSyncTime")) + submit_time = parse_datetime(record.get("submitTime")) + outcome_rank = 1 if record.get("outcome") in {"success", "failed"} else 0 + return ( + outcome_rank, + last_sync.timestamp() if last_sync else 0.0, + submit_time.timestamp() if submit_time else 0.0, + ) + + +def _prefer_newer_outcome(existing: dict[str, Any], candidate: dict[str, Any]) -> dict[str, Any]: + if _outcome_version(candidate) >= _outcome_version(existing): + return candidate + return existing diff --git a/modelhub_submmit_api/poll_runner.py b/modelhub_submmit_api/poll_runner.py index 9f1a201e..a772645d 100644 --- a/modelhub_submmit_api/poll_runner.py +++ b/modelhub_submmit_api/poll_runner.py @@ -2,6 +2,7 @@ from __future__ import annotations import argparse import json +import os import sys import time from pathlib import Path @@ -10,10 +11,11 @@ from typing import Any, Callable from common import utc_now, write_json from daily_runner import DEFAULT_DAILY_RUNS_DIR, log, run_daily_batches from hf_discovery import HuggingFaceDiscovery -from main import DEFAULT_LEDGER_PATH, DEFAULT_RUNS_DIR +from main import DEFAULT_LEDGER_PATH, DEFAULT_RUNS_DIR, make_run_dir from modelhub_client import ModelHubClient, ModelHubClientPool from outcome_tracker import DEFAULT_OUTCOMES_PATH, OutcomeTracker from runner_common import DEFAULT_KEY_PATH, ensure_tokens +from submission_claims import DEFAULT_CLAIMS_PATH from template_selector import TemplateSelector @@ -58,6 +60,11 @@ def build_parser() -> argparse.ArgumentParser: parser.add_argument("--key-path", default=str(DEFAULT_KEY_PATH), help="Path to KEY.md containing MODELSCOPE_TOKEN/XC_TOKEN") parser.add_argument("--runs-dir", default=str(DEFAULT_RUNS_DIR), help=argparse.SUPPRESS) parser.add_argument("--ledger-path", default=str(DEFAULT_LEDGER_PATH), help=argparse.SUPPRESS) + parser.add_argument( + "--claims-path", + default=os.getenv("MODELHUB_AGENT_CLAIMS_PATH", str(DEFAULT_CLAIMS_PATH)), + help=argparse.SUPPRESS, + ) parser.add_argument("--history-archive-path", default="history/platform_tasks.jsonl", help=argparse.SUPPRESS) parser.add_argument("--history-archive-limit", type=int, default=5000, help=argparse.SUPPRESS) parser.add_argument("--daily-runs-dir", default=str(DEFAULT_DAILY_RUNS_DIR), help=argparse.SUPPRESS) @@ -68,9 +75,9 @@ def build_parser() -> argparse.ArgumentParser: parser.add_argument("--modelhub-token", default=None, help=argparse.SUPPRESS) parser.add_argument("--hf-token", default=None, help=argparse.SUPPRESS) parser.add_argument("--modelscope-token", default=None, help=argparse.SUPPRESS) - parser.add_argument("--poll-interval-seconds", type=int, default=60, help="Sleep between polling cycles when no slots are available") - parser.add_argument("--idle-interval-seconds", type=int, default=30, help="Sleep between cycles when a scan submits nothing") - parser.add_argument("--post-cycle-cooldown-seconds", type=int, default=0, help="Short sleep after a successful cycle") + parser.add_argument("--poll-interval-seconds", type=int, default=15, help="Sleep between polling cycles when no slots are available") + parser.add_argument("--idle-interval-seconds", type=int, default=60, help="Sleep between cycles when a scan submits nothing") + parser.add_argument("--post-cycle-cooldown-seconds", type=int, default=2, help="Short sleep after a successful cycle") parser.add_argument("--max-cycles", type=int, default=0, help="Optional hard stop after N cycles; 0 means run until quota is reached") parser.add_argument("--print-stats", action="store_true", help="Load outcomes, sync, print stats report, and exit") return parser @@ -85,12 +92,11 @@ def _make_cycle_args(base_args: argparse.Namespace) -> argparse.Namespace: return cycle_args -def _build_modelhub_client(base_args: argparse.Namespace) -> ModelHubClient | ModelHubClientPool: +def _build_modelhub_client(base_args: argparse.Namespace) -> ModelHubClientPool: modelhub_tokens = list(getattr(base_args, "modelhub_tokens", None) or ([] if not base_args.modelhub_token else [base_args.modelhub_token])) - if len(modelhub_tokens) > 1: - clients = [ModelHubClient(token=token, base_url=base_args.modelhub_base_url) for token in modelhub_tokens] - return ModelHubClientPool(clients) - return ModelHubClient(token=base_args.modelhub_token, base_url=base_args.modelhub_base_url) + token_values: list[str | None] = modelhub_tokens or [base_args.modelhub_token] + clients = [ModelHubClient(token=token, base_url=base_args.modelhub_base_url) for token in token_values] + return ModelHubClientPool(clients) def run_poll_loop( @@ -110,8 +116,7 @@ def run_poll_loop( poll_runs_dir = Path(base_args.poll_runs_dir) poll_runs_dir.mkdir(parents=True, exist_ok=True) - poll_run_dir = poll_runs_dir / now.strftime("%Y%m%dT%H%M%SZ") - poll_run_dir.mkdir(parents=True, exist_ok=True) + poll_run_dir = make_run_dir(poll_runs_dir, now) outcome_tracker = outcome_tracker or OutcomeTracker(Path(base_args.outcomes_path)) OUTCOME_SYNC_INTERVAL = 3 diff --git a/modelhub_submmit_api/submission_claims.py b/modelhub_submmit_api/submission_claims.py new file mode 100644 index 00000000..580f0968 --- /dev/null +++ b/modelhub_submmit_api/submission_claims.py @@ -0,0 +1,124 @@ +from __future__ import annotations + +import hashlib +from datetime import timedelta +from pathlib import Path +from typing import Any + +from common import parse_datetime, runtime_instance_id, update_jsonl, utc_now + + +DEFAULT_CLAIMS_PATH = Path(".modelhub_state/submission_claims.jsonl") + + +def candidate_key(candidate: dict[str, Any]) -> str: + model_id = candidate.get("repoId") or candidate.get("modelAddress") or "unknown" + target_gpu = candidate.get("targetGpu") or "unknown" + return f"{model_id}|{target_gpu}" + + +def diversify_candidates( + candidates: list[dict[str, Any]], + *, + instance_id: str, +) -> list[dict[str, Any]]: + """Give concurrent agents different deterministic candidate orders.""" + + def sort_key(candidate: dict[str, Any]) -> str: + value = f"{instance_id}|{candidate_key(candidate)}" + return hashlib.sha256(value.encode("utf-8")).hexdigest() + + return sorted(candidates, key=sort_key) + + +class SubmissionClaimStore: + """Small filesystem-backed lease store for local multi-process deduplication.""" + + def __init__( + self, + path: Path | str = DEFAULT_CLAIMS_PATH, + *, + owner_id: str | None = None, + claim_ttl_seconds: int = 600, + submitted_ttl_seconds: int = 24 * 60 * 60, + ) -> None: + self.path = Path(path) + self.owner_id = owner_id or runtime_instance_id() + self.claim_ttl_seconds = max(30, int(claim_ttl_seconds)) + self.submitted_ttl_seconds = max(self.claim_ttl_seconds, int(submitted_ttl_seconds)) + + def claim( + self, + candidates: list[dict[str, Any]], + *, + limit: int, + ) -> list[dict[str, Any]]: + if limit <= 0 or not candidates: + return [] + + selected: list[dict[str, Any]] = [] + now = utc_now() + expires_at = now + timedelta(seconds=self.claim_ttl_seconds) + + def update(records: list[dict[str, Any]]) -> list[dict[str, Any]]: + active = [record for record in records if not self._is_expired(record, now)] + claimed_keys = {str(record.get("key")) for record in active if record.get("key")} + for candidate in candidates: + key = candidate_key(candidate) + if key in claimed_keys: + continue + active.append( + { + "key": key, + "ownerId": self.owner_id, + "state": "claimed", + "claimedAt": now.isoformat(), + "expiresAt": expires_at.isoformat(), + } + ) + claimed_keys.add(key) + selected.append(candidate) + if len(selected) >= limit: + break + return active + + update_jsonl(self.path, update) + return selected + + def mark_submitted(self, candidates: list[dict[str, Any]]) -> None: + keys = {candidate_key(candidate) for candidate in candidates} + if not keys: + return + now = utc_now() + expires_at = now + timedelta(seconds=self.submitted_ttl_seconds) + + def update(records: list[dict[str, Any]]) -> list[dict[str, Any]]: + active = [record for record in records if not self._is_expired(record, now)] + for record in active: + if record.get("ownerId") == self.owner_id and record.get("key") in keys: + record["state"] = "submitted" + record["expiresAt"] = expires_at.isoformat() + return active + + update_jsonl(self.path, update) + + def release(self, candidates: list[dict[str, Any]]) -> None: + keys = {candidate_key(candidate) for candidate in candidates} + if not keys: + return + now = utc_now() + + def update(records: list[dict[str, Any]]) -> list[dict[str, Any]]: + return [ + record + for record in records + if not self._is_expired(record, now) + and not (record.get("ownerId") == self.owner_id and record.get("key") in keys) + ] + + update_jsonl(self.path, update) + + @staticmethod + def _is_expired(record: dict[str, Any], now) -> bool: # noqa: ANN001 + expires_at = parse_datetime(record.get("expiresAt")) + return expires_at is None or expires_at <= now diff --git a/tests/test_concurrency.py b/tests/test_concurrency.py new file mode 100644 index 00000000..3461cc12 --- /dev/null +++ b/tests/test_concurrency.py @@ -0,0 +1,176 @@ +from __future__ import annotations + +import sys +import tempfile +import threading +import unittest +from concurrent.futures import ThreadPoolExecutor +from datetime import datetime, timezone +from pathlib import Path + + +PACKAGE_DIR = Path(__file__).resolve().parents[1] / "modelhub_submmit_api" +if str(PACKAGE_DIR) in sys.path: + sys.path.remove(str(PACKAGE_DIR)) +sys.path.insert(0, str(PACKAGE_DIR)) + +from common import read_jsonl # noqa: E402 +from main import make_run_dir # noqa: E402 +from modelhub_client import ModelHubAPIError, ModelHubClientPool # noqa: E402 +from outcome_tracker import OutcomeTracker # noqa: E402 +from submission_claims import SubmissionClaimStore, candidate_key # noqa: E402 + + +class FakeClient: + def __init__(self, active_count: int = 0, *, reject_capacity: bool = False) -> None: + self.active_count = active_count + self.reject_capacity = reject_capacity + self.submitted: list[dict] = [] + self.count_calls = 0 + self.search_calls = 0 + self._lock = threading.Lock() + + def count_active_tasks(self, **_kwargs) -> int: # noqa: ANN003 + with self._lock: + self.count_calls += 1 + return self.active_count + + def add_task(self, payload: dict) -> dict: + if self.reject_capacity: + raise ModelHubAPIError("当前等待中或运行中的异步模型验证任务数量已达上限") + with self._lock: + self.submitted.append(payload) + return {"code": 0, "data": {"id": len(self.submitted)}} + + def search_by_model_id(self, _model_id: str) -> dict: + with self._lock: + self.search_calls += 1 + return {"code": 0, "data": {"verifyResult": {}}} + + def list_tasks(self, **_kwargs) -> list[dict]: # noqa: ANN003 + return [] + + def list_tasks_page(self, **_kwargs) -> dict: # noqa: ANN003 + return {"code": 0, "data": {"records": [], "pages": 0}} + + def find_recent_task_id(self, *_args, **_kwargs): # noqa: ANN002, ANN003 + return None + + +def make_candidate(index: int) -> dict: + return { + "repoId": f"owner/model-{index}", + "modelAddress": f"https://modelscope.cn/models/owner/model-{index}", + "targetGpu": "NVIDIA-A100", + } + + +class ClientPoolConcurrencyTests(unittest.TestCase): + def test_every_account_reaches_capacity_under_load(self) -> None: + clients = [FakeClient() for _ in range(12)] + pool = ModelHubClientPool( + clients, # type: ignore[arg-type] + active_task_cap=10, + active_counts_ttl=60, + instance_id="full-load-test", + ) + + with ThreadPoolExecutor(max_workers=24) as executor: + list(executor.map(lambda index: pool.add_task({"index": index}), range(120))) + + self.assertEqual([10] * 12, [len(client.submitted) for client in clients]) + self.assertEqual([10] * 12, pool.active_task_counts()) + self.assertEqual(0, pool.available_submit_slots()) + + def test_concurrent_submissions_reserve_and_balance_accounts(self) -> None: + clients = [FakeClient(), FakeClient()] + pool = ModelHubClientPool( + clients, # type: ignore[arg-type] + active_task_cap=4, + active_counts_ttl=60, + instance_id="balance-test", + ) + + with ThreadPoolExecutor(max_workers=8) as executor: + results = list(executor.map(lambda index: pool.add_task({"index": index}), range(8))) + + self.assertEqual(8, len(results)) + self.assertEqual([4, 4], [len(client.submitted) for client in clients]) + self.assertEqual([4, 4], pool.active_task_counts()) + self.assertEqual(0, pool.available_submit_slots()) + # A lagging remote count must not erase successful local reservations. + pool._refresh_active_counts(force=True) + self.assertEqual([4, 4], pool.active_task_counts()) + + def test_capacity_rejection_falls_through_to_another_account(self) -> None: + full_elsewhere = FakeClient(active_count=0, reject_capacity=True) + available = FakeClient(active_count=1) + pool = ModelHubClientPool( + [full_elsewhere, available], # type: ignore[arg-type] + active_task_cap=2, + instance_id="fallback-test", + ) + + response = pool.add_task({"model": "x"}) + + self.assertEqual(1, response["data"]["id"]) + self.assertEqual(0, len(full_elsewhere.submitted)) + self.assertEqual(1, len(available.submitted)) + + def test_verification_cache_is_reset_between_cycles(self) -> None: + client = FakeClient() + pool = ModelHubClientPool([client], active_task_cap=2) # type: ignore[arg-type] + pool.search_by_model_id("owner/model") + pool.search_by_model_id("owner/model") + self.assertEqual(1, client.search_calls) + pool.begin_cycle() + pool.search_by_model_id("owner/model") + self.assertEqual(2, client.search_calls) + + +class ProcessCoordinationTests(unittest.TestCase): + def test_concurrent_claim_stores_select_disjoint_candidates(self) -> None: + with tempfile.TemporaryDirectory() as temporary_dir: + path = Path(temporary_dir) / "claims.jsonl" + candidates = [make_candidate(index) for index in range(6)] + first = SubmissionClaimStore(path, owner_id="worker-a") + second = SubmissionClaimStore(path, owner_id="worker-b") + + with ThreadPoolExecutor(max_workers=2) as executor: + first_future = executor.submit(first.claim, candidates, limit=3) + second_future = executor.submit(second.claim, candidates, limit=3) + first_claims = first_future.result() + second_claims = second_future.result() + + first_keys = {candidate_key(candidate) for candidate in first_claims} + second_keys = {candidate_key(candidate) for candidate in second_claims} + self.assertEqual(3, len(first_claims)) + self.assertEqual(3, len(second_claims)) + self.assertTrue(first_keys.isdisjoint(second_keys)) + + def test_outcome_saves_merge_instead_of_overwriting_other_worker(self) -> None: + with tempfile.TemporaryDirectory() as temporary_dir: + path = Path(temporary_dir) / "outcomes.jsonl" + first = OutcomeTracker(path) + second = OutcomeTracker(path) + first.record_submission("model-a", "gpu", "vllm", "text-generation", "task-a", "2026-01-01T00:00:00+00:00") + second.record_submission("model-b", "gpu", "vllm", "text-generation", "task-b", "2026-01-01T00:00:01+00:00") + + first.save() + second.save() + + self.assertEqual({"task-a", "task-b"}, {row["taskId"] for row in read_jsonl(path)}) + + def test_run_directories_are_allocated_atomically(self) -> None: + with tempfile.TemporaryDirectory() as temporary_dir: + base = Path(temporary_dir) / "runs" + now = datetime(2026, 1, 1, tzinfo=timezone.utc) + with ThreadPoolExecutor(max_workers=8) as executor: + paths = list(executor.map(lambda _index: make_run_dir(base, now), range(8))) + + self.assertEqual(8, len(set(paths))) + self.assertTrue(all(path.is_dir() for path in paths)) + + +if __name__ == "__main__": + unittest.main()