fix: coordinate concurrent account capacity filling
This commit is contained in:
124
modelhub_submmit_api/submission_claims.py
Normal file
124
modelhub_submmit_api/submission_claims.py
Normal file
@@ -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
|
||||
Reference in New Issue
Block a user