feat: reserve queue capacity for recent models

This commit is contained in:
CoolBoy
2026-08-11 01:29:54 +08:00
parent 38ab25fc3c
commit f12d96b138
13 changed files with 927 additions and 46 deletions

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
import unittest
import sys
from datetime import datetime, timezone
from pathlib import Path
from typing import Any
@@ -12,16 +13,28 @@ if str(PACKAGE_DIR) in sys.path:
sys.path.insert(0, str(PACKAGE_DIR))
from modelhub_client import ModelHubClient, ModelHubClientPool # noqa: E402
from queue_cleanup import OwnedTask, cleanup_certain_oom_tasks, find_certain_oom_tasks # noqa: E402
from queue_cleanup import ( # noqa: E402
OwnedTask,
cleanup_certain_oom_tasks,
find_certain_oom_tasks,
find_old_overflow_tasks,
)
GIB = 1024**3
class FakeQueueClient:
def __init__(self, records: list[dict[str, Any]], *, disappear_on_recheck: bool = False) -> None:
def __init__(
self,
records: list[dict[str, Any]],
*,
disappear_on_recheck: bool = False,
drop_first_on_recheck: bool = False,
) -> None:
self.records = list(records)
self.disappear_on_recheck = disappear_on_recheck
self.drop_first_on_recheck = drop_first_on_recheck
self.waiting_reads = 0
self.stopped: list[list[int]] = []
@@ -32,6 +45,8 @@ class FakeQueueClient:
records: list[dict[str, Any]] = []
else:
records = [record for record in self.records if record["status"] == "waiting"]
if self.drop_first_on_recheck and self.waiting_reads >= 2:
records = sorted(records, key=lambda item: int(item["taskId"]))[1:]
else:
records = [record for record in self.records if record["status"] == status]
return {"code": 0, "data": {"records": records, "pages": 1}}
@@ -47,8 +62,13 @@ class FakeQueueClient:
class FakeDiscovery:
def __init__(self, sizes: dict[str, int | None]) -> None:
def __init__(
self,
sizes: dict[str, int | None],
last_modified: dict[str, datetime | None] | None = None,
) -> None:
self.sizes = sizes
self.last_modified = last_modified or {}
def list_repo_tree(self, repo_id: str) -> list[dict[str, Any]]:
size = self.sizes[repo_id]
@@ -56,6 +76,9 @@ class FakeDiscovery:
return [{"Path": "model.safetensors"}]
return [{"Path": "model.safetensors", "Size": size}]
def get_model_last_modified(self, repo_id: str) -> datetime | None:
return self.last_modified.get(repo_id)
class RecordingHttpClient:
def __init__(self) -> None:
@@ -74,6 +97,151 @@ class RecordingHttpClient:
class QueueCleanupTests(unittest.TestCase):
def test_old_models_are_selected_only_after_each_accounts_first_eighty_tasks(self) -> None:
now = datetime(2026, 8, 11, tzinfo=timezone.utc)
tasks = [
OwnedTask(0, index, "owner/old", "Iluvatar_bi-100", "waiting")
for index in range(1, 82)
]
tasks.append(OwnedTask(0, 82, "owner/recent", "Iluvatar_bi-100", "waiting"))
selected, skipped = find_old_overflow_tasks(
tasks,
model_last_modified={
"owner/old": datetime(2026, 7, 1, tzinfo=timezone.utc),
"owner/recent": datetime(2026, 8, 10, tzinfo=timezone.utc),
},
queue_threshold=80,
recent_model_days=7,
reference_time=now,
)
self.assertEqual([81], [item["taskId"] for item in selected])
self.assertEqual(81, selected[0]["queuePosition"])
self.assertEqual(1, skipped["recentOverflowTasks"])
def test_old_overflow_task_is_not_stopped_if_it_moves_into_first_eighty(self) -> None:
records = [
{
"taskId": index,
"modelId": "owner/old",
"gpuType": "Iluvatar_bi-100",
"status": "waiting",
}
for index in range(1, 82)
]
client = FakeQueueClient(records, drop_first_on_recheck=True)
pool = ModelHubClientPool(
[client], # type: ignore[list-item]
active_task_cap=100,
old_model_queue_threshold=80,
)
summary = cleanup_certain_oom_tasks(
pool,
FakeDiscovery(
{"owner/old": 1 * GIB},
{"owner/old": datetime(2026, 7, 1, tzinfo=timezone.utc)},
), # type: ignore[arg-type]
reference_time=datetime(2026, 8, 11, tzinfo=timezone.utc),
log=lambda _message: None,
)
self.assertEqual(1, summary["oldOverflowCount"])
self.assertEqual(0, summary["cancelledCount"])
self.assertEqual(1, summary["policyNoLongerAppliesCount"])
self.assertEqual([], client.stopped)
def test_cleanup_stops_old_task_beyond_eightieth_position(self) -> None:
records = [
{
"taskId": index,
"modelId": "owner/old",
"gpuType": "Iluvatar_bi-100",
"status": "waiting",
}
for index in range(1, 82)
]
client = FakeQueueClient(records)
pool = ModelHubClientPool(
[client], # type: ignore[list-item]
active_task_cap=100,
old_model_queue_threshold=80,
)
summary = cleanup_certain_oom_tasks(
pool,
FakeDiscovery(
{"owner/old": 1 * GIB},
{"owner/old": datetime(2026, 7, 1, tzinfo=timezone.utc)},
), # type: ignore[arg-type]
reference_time=datetime(2026, 8, 11, tzinfo=timezone.utc),
log=lambda _message: None,
)
self.assertEqual(1, summary["oldOverflowCount"])
self.assertEqual(1, summary["cancelledCount"])
self.assertEqual([[81]], client.stopped)
def test_dynamic_cleanup_keeps_positions_through_ninety_five(self) -> None:
records = [
{
"taskId": index,
"modelId": "owner/old",
"gpuType": "Iluvatar_bi-100",
"status": "waiting",
}
for index in range(1, 97)
]
client = FakeQueueClient(records)
pool = ModelHubClientPool(
[client], # type: ignore[list-item]
active_task_cap=100,
old_model_queue_threshold=80,
)
summary = cleanup_certain_oom_tasks(
pool,
FakeDiscovery(
{"owner/old": 1 * GIB},
{"owner/old": datetime(2026, 7, 1, tzinfo=timezone.utc)},
), # type: ignore[arg-type]
age_queue_threshold=95,
reference_time=datetime(2026, 8, 11, tzinfo=timezone.utc),
log=lambda _message: None,
)
self.assertEqual(95, summary["oldModelQueueThreshold"])
self.assertEqual([96], [item["taskId"] for item in summary["oldOverflowTasks"]])
self.assertEqual([[96]], client.stopped)
def test_oom_is_removed_before_recalculating_old_overflow_positions(self) -> None:
records = [
{
"taskId": index,
"modelId": "owner/large" if index == 1 else "owner/old",
"gpuType": "Iluvatar_bi-100",
"status": "waiting",
}
for index in range(1, 83)
]
client = FakeQueueClient(records)
pool = ModelHubClientPool(
[client], # type: ignore[list-item]
active_task_cap=100,
old_model_queue_threshold=80,
)
summary = cleanup_certain_oom_tasks(
pool,
FakeDiscovery(
{"owner/large": 40 * GIB, "owner/old": 1 * GIB},
{"owner/old": datetime(2026, 7, 1, tzinfo=timezone.utc)},
), # type: ignore[arg-type]
reference_time=datetime(2026, 8, 11, tzinfo=timezone.utc),
log=lambda _message: None,
)
self.assertEqual(1, summary["certainOomCount"])
self.assertEqual([82], [item["taskId"] for item in summary["oldOverflowTasks"]])
self.assertEqual([[1], [82]], client.stopped)
def test_stop_tasks_uses_documented_put_endpoint_and_integer_ids(self) -> None:
http = RecordingHttpClient()
client = ModelHubClient(http_client=http) # type: ignore[arg-type]