Files
OBLITERATUS/tests/test_watchtower_contracts.py
T

647 lines
21 KiB
Python

"""Deterministic transport, state, and scheduler contracts for Watchtower."""
from __future__ import annotations
import json
import math
import sys
import threading
from datetime import datetime, timezone
from pathlib import Path
from types import ModuleType, SimpleNamespace
import pytest
from obliteratus.service_contracts import (
SchedulerEvent,
SchedulerState,
normalize_watchtower_candidate,
scheduler_transition,
validate_day_window,
validate_scheduler_interval,
)
from obliteratus.watchtower import DiscoveredModel, Watchtower
NOW = datetime(2026, 8, 16, 12, 30, tzinfo=timezone.utc)
def _model(model_id: str, **overrides):
values = {
"id": model_id,
"downloads": 2_000,
"likes": 10,
"tags": ["instruct"],
"license": "apache-2.0",
"pipeline_tag": "text-generation",
}
values.update(overrides)
return SimpleNamespace(**values)
def _watchtower(tmp_path: Path, **kwargs) -> Watchtower:
return Watchtower(
tmp_path / "watchtower.json",
clock=lambda: NOW,
fetch_models=lambda: [],
**kwargs,
)
@pytest.mark.parametrize(
"state,event,expected",
[
(SchedulerState.STOPPED, SchedulerEvent.START, SchedulerState.RUNNING),
(SchedulerState.RUNNING, SchedulerEvent.REQUEST_STOP, SchedulerState.STOPPING),
(SchedulerState.STOPPING, SchedulerEvent.STOP_TIMEOUT, SchedulerState.STOPPING),
(SchedulerState.STOPPING, SchedulerEvent.STOP_CONFIRMED, SchedulerState.STOPPED),
],
)
def test_scheduler_state_machine_contract(state, event, expected) -> None:
assert scheduler_transition(state, event) is expected
@pytest.mark.parametrize(
"state,event",
[
(SchedulerState.STOPPED, SchedulerEvent.REQUEST_STOP),
(SchedulerState.STOPPED, SchedulerEvent.STOP_CONFIRMED),
(SchedulerState.RUNNING, SchedulerEvent.START),
(SchedulerState.RUNNING, SchedulerEvent.STOP_CONFIRMED),
(SchedulerState.STOPPING, SchedulerEvent.START),
(SchedulerState.STOPPING, SchedulerEvent.REQUEST_STOP),
],
)
def test_scheduler_state_machine_rejects_invalid_transitions(state, event) -> None:
with pytest.raises(ValueError, match="scheduler transition"):
scheduler_transition(state, event)
def test_service_validation_messages_and_zero_boundaries() -> None:
with pytest.raises(TypeError) as interval_type:
validate_scheduler_interval("60")
assert str(interval_type.value) == "scheduler interval must be a number"
with pytest.raises(ValueError) as interval_value:
validate_scheduler_interval(0)
assert str(interval_value.value) == (
"scheduler interval must be finite and greater than zero"
)
with pytest.raises(TypeError) as day_type:
validate_day_window("7")
assert str(day_type.value) == "days must be a non-negative integer"
with pytest.raises(ValueError) as day_value:
validate_day_window(-1)
assert str(day_value.value) == "days must be a non-negative integer"
assert validate_day_window(0) == 0
def test_candidate_normalization_defaults_boundaries_and_license_tags() -> None:
missing_counts = SimpleNamespace(
id="org/model",
tags=[],
license="mit",
pipeline_tag="text-generation",
)
assert normalize_watchtower_candidate(missing_counts) == (
"org/model", 0, 0, [], "mit", "text-generation"
)
assert normalize_watchtower_candidate(
_model("org/zero", downloads=0, likes=0)
) == ("org/zero", 0, 0, ["instruct"], "apache-2.0", "text-generation")
assert normalize_watchtower_candidate(_model("org/bool", likes=True)) is None
assert normalize_watchtower_candidate(_model("org/tuple", tags=("chat",))) is None
assert normalize_watchtower_candidate(_model("org/license", license=7)) is None
assert normalize_watchtower_candidate(_model("org/pipeline", pipeline_tag=7)) is None
tagged = SimpleNamespace(
modelId="org/tagged",
downloads=1,
likes=0,
tags=["license:custom:variant"],
pipeline_tag="",
)
assert normalize_watchtower_candidate(tagged) == (
"org/tagged", 1, 0, ["license:custom:variant"], "custom:variant", ""
)
def test_state_load_is_transactional_on_malformed_records(tmp_path: Path) -> None:
state = tmp_path / "watchtower.json"
state.write_text(
json.dumps(
{
"last_scan": "2026-08-15T00:00:00+00:00",
"scan_count": 9,
"models": {
"org/valid": {"model_id": "org/valid", "name": "valid", "org": "org"},
"org/bad": "not-an-object",
},
}
),
encoding="utf-8",
)
watchtower = Watchtower(state, clock=lambda: NOW, fetch_models=lambda: [])
assert watchtower.get_all_models() == []
assert watchtower.get_stats() == {
"total_tracked": 0,
"last_scan": None,
"scan_count": 0,
"by_status": {},
}
def test_model_state_rejects_non_object_and_missing_required_fields() -> None:
with pytest.raises(ValueError, match="must be an object"):
DiscoveredModel.from_dict("bad") # type: ignore[arg-type]
with pytest.raises(ValueError, match="missing required fields"):
DiscoveredModel.from_dict({})
@pytest.mark.parametrize(
"payload",
[
[],
{"last_scan": 7},
{"scan_count": True},
{"scan_count": -1},
{"models": []},
{
"models": {
"org/key": {"model_id": "org/other", "name": "other", "org": "org"}
}
},
],
)
def test_state_load_rejects_malformed_envelopes_transactionally(
payload: object,
tmp_path: Path,
) -> None:
state = tmp_path / "watchtower.json"
state.write_text(json.dumps(payload), encoding="utf-8")
watchtower = Watchtower(state, clock=lambda: NOW, fetch_models=lambda: [])
assert watchtower.get_all_models() == []
assert watchtower.get_stats()["last_scan"] is None
assert watchtower.get_stats()["scan_count"] == 0
def test_state_save_failure_cleans_up_without_masking_operation(
tmp_path: Path,
caplog: pytest.LogCaptureFixture,
) -> None:
blocked_parent = tmp_path / "not-a-directory"
blocked_parent.write_text("file", encoding="utf-8")
watchtower = Watchtower(
blocked_parent / "watchtower.json",
clock=lambda: NOW,
fetch_models=lambda: [],
)
watchtower._save_state()
assert "failed to save state" in caplog.text
def test_state_writes_serialize_snapshot_and_shared_temporary_path(
monkeypatch: pytest.MonkeyPatch,
tmp_path: Path,
) -> None:
state = tmp_path / "watchtower.json"
temporary = state.with_suffix(".tmp")
watchtower = Watchtower(state, clock=lambda: NOW, fetch_models=lambda: [])
first_write_started = threading.Event()
release_first_write = threading.Event()
overlapping_write = threading.Event()
active_writes = 0
active_lock = threading.Lock()
original_write_text = Path.write_text
def controlled_write(path: Path, data: str, **kwargs) -> int:
nonlocal active_writes
if path != temporary:
return original_write_text(path, data, **kwargs)
with active_lock:
active_writes += 1
if active_writes > 1:
overlapping_write.set()
first_write_started.set()
assert release_first_write.wait(timeout=1)
try:
return original_write_text(path, data, **kwargs)
finally:
with active_lock:
active_writes -= 1
monkeypatch.setattr(Path, "write_text", controlled_write)
first = threading.Thread(target=watchtower._save_state)
second = threading.Thread(target=watchtower._save_state)
first.start()
assert first_write_started.wait(timeout=1)
second.start()
overlapped = overlapping_write.wait(timeout=0.05)
release_first_write.set()
first.join(timeout=1)
second.join(timeout=1)
assert not overlapped
assert not first.is_alive()
assert not second.is_alive()
assert json.loads(state.read_text(encoding="utf-8"))["scan_count"] == 0
@pytest.mark.parametrize(
"bad_field,bad_value",
[
("status", "invented"),
("downloads_7d", "many"),
("obliteration_metrics", []),
("discovered_at", 7),
],
)
def test_state_load_rejects_malformed_model_fields_transactionally(
bad_field: str,
bad_value: object,
tmp_path: Path,
) -> None:
state = tmp_path / "watchtower.json"
record = {"model_id": "org/model", "name": "model", "org": "org"}
record[bad_field] = bad_value
state.write_text(
json.dumps({"last_scan": "prior", "scan_count": 2, "models": {"org/model": record}}),
encoding="utf-8",
)
watchtower = Watchtower(state, clock=lambda: NOW, fetch_models=lambda: [])
assert watchtower.get_all_models() == []
assert watchtower.get_stats()["last_scan"] is None
assert watchtower.get_stats()["scan_count"] == 0
def test_valid_state_round_trips_unknown_fields_and_defaults(tmp_path: Path) -> None:
state = tmp_path / "watchtower.json"
state.write_text(
json.dumps(
{
"last_scan": "2026-08-15T00:00:00+00:00",
"scan_count": 2,
"models": {
"org/model": {
"model_id": "org/model",
"name": "model",
"org": "org",
"unknown_future_field": True,
}
},
}
),
encoding="utf-8",
)
watchtower = Watchtower(state, clock=lambda: NOW, fetch_models=lambda: [])
model = watchtower.get_model("org/model")
assert model is not None
assert model.status == "new"
assert watchtower.get_stats()["scan_count"] == 2
def test_huggingface_adapter_is_deterministic_bounded_and_fault_isolated(
monkeypatch: pytest.MonkeyPatch,
) -> None:
calls: list[dict[str, object]] = []
class HfApi:
def list_models(self, **kwargs):
calls.append(kwargs)
if kwargs["author"] == "bad":
raise OSError("one organization is unavailable")
return [
_model(f"{kwargs['author']}/kept", downloads=1_000),
_model(f"{kwargs['author']}/low", downloads=999),
]
fake_module = ModuleType("huggingface_hub")
fake_module.HfApi = HfApi
monkeypatch.setitem(sys.modules, "huggingface_hub", fake_module)
records = Watchtower._fetch_models_from_hf(
orgs={"zeta", "bad", "Alpha"},
limit_per_org=3,
min_downloads=1_000,
)
assert [call["author"] for call in calls] == ["Alpha", "bad", "zeta"]
assert all(
call == {
"author": call["author"],
"pipeline_tag": "text-generation",
"sort": "downloads",
"direction": -1,
"limit": 3,
}
for call in calls
)
assert [record.id for record in records] == ["Alpha/kept", "zeta/kept"]
@pytest.mark.parametrize(
"license_id, expected",
[
(None, False),
(" MIT ", True),
("apache-2.0", True),
("apache-2.0-custom", True),
("mitigation-license", False),
("proprietary", False),
],
)
def test_open_license_matching_has_identifier_boundaries(license_id, expected) -> None:
assert Watchtower._has_open_license(license_id) is expected
@pytest.mark.parametrize(
"model_id,tags,expected",
[
("org/model-instruct", None, True),
("org/model-it", None, True),
("org/bitnet-base", None, False),
("org/plain", ["Conversational"], True),
("org/plain", ["pretraining"], False),
(7, [], False),
("org/plain", "chat", False),
("org/plain", [7], False),
],
)
def test_instruction_matching_has_token_boundaries(model_id, tags, expected) -> None:
assert Watchtower._is_instruction_tuned(model_id, tags) is expected
@pytest.mark.parametrize(
"model_id, expected",
[
("org/model-0.5B", "500M"),
("org/model-7B", "7B"),
("org/model-70b-instruct", "70B"),
("org/unsized", "unknown"),
],
)
def test_size_estimation_boundaries(model_id: str, expected: str) -> None:
assert Watchtower._estimate_size(model_id) == expected
def test_scan_filters_candidates_updates_existing_and_is_deterministic(tmp_path: Path) -> None:
candidates = [
object(),
_model("org/bad-downloads", downloads="many"),
_model("org/bad-likes", likes=-1),
_model("org/bad-tags", tags=[7]),
_model("org/bad-license", license=7),
SimpleNamespace(modelId="org/missing-id-fallback-7B", downloads=2_000, likes=1,
tags=["License: mit", "chat"], license="", pipeline_tag=""),
_model("org/wrong-task-7B", pipeline_tag="text-classification"),
_model("org/unpopular-7B", downloads=999),
_model("org/closed-7B", license="proprietary"),
_model("org/base-7B", tags=[]),
_model("Qwen/major-base-8B", tags=[]),
_model("org/fresh-instruct-7B"),
]
logs: list[str] = []
callback_ids: list[str] = []
watchtower = Watchtower(
tmp_path / "watchtower.json",
clock=lambda: NOW,
fetch_models=lambda: candidates,
)
watchtower.on_new_model(lambda model: callback_ids.append(model.model_id))
watchtower.on_new_model(lambda _model: (_ for _ in ()).throw(RuntimeError("ignored")))
discovered = watchtower.scan(on_log=logs.append)
assert [model.model_id for model in discovered] == [
"org/missing-id-fallback-7B",
"Qwen/major-base-8B",
"org/fresh-instruct-7B",
]
assert callback_ids == [model.model_id for model in discovered]
assert watchtower.get_stats()["scan_count"] == 1
assert watchtower.get_stats()["last_scan"] == NOW.isoformat()
assert logs[0] == "Watchtower: starting scan..."
assert logs[-1].endswith("3 new discoveries")
candidates[-1].downloads = 9_000
candidates[-1].likes = 99
assert watchtower.scan() == []
updated = watchtower.get_model("org/fresh-instruct-7B")
assert updated is not None
assert (updated.downloads_7d, updated.likes, updated.last_updated) == (
9_000,
99,
NOW.isoformat(),
)
def test_transport_failure_does_not_advance_scan_state(tmp_path: Path) -> None:
def fail():
raise OSError("offline")
watchtower = Watchtower(
tmp_path / "watchtower.json",
clock=lambda: NOW,
fetch_models=fail,
)
assert watchtower.scan() == []
assert watchtower.get_stats()["scan_count"] == 0
assert not watchtower.state_file.exists()
malformed = Watchtower(
tmp_path / "malformed.json",
clock=lambda: NOW,
fetch_models=lambda: None, # type: ignore[return-value]
)
assert malformed.scan() == []
assert malformed.get_stats()["scan_count"] == 0
def test_status_updates_validate_state_and_preserve_empty_metrics(tmp_path: Path) -> None:
watchtower = Watchtower(
tmp_path / "watchtower.json",
clock=lambda: NOW,
fetch_models=lambda: [_model("org/model-instruct-7B")],
)
watchtower.scan()
assert watchtower.set_status("missing", "queued") is False
with pytest.raises(ValueError, match="status"):
watchtower.set_status("org/model-instruct-7B", "invented")
assert watchtower.set_status("org/model-instruct-7B", "obliterated", {}) is True
model = watchtower.get_model("org/model-instruct-7B")
assert model is not None
assert model.status == "obliterated"
assert model.obliteration_metrics == {}
assert model.last_updated == NOW.isoformat()
assert [item.model_id for item in watchtower.get_obliterated()] == [model.model_id]
@pytest.mark.parametrize("interval", [0, -1, True, math.inf, math.nan, "1"])
def test_scheduler_rejects_non_positive_or_non_finite_intervals(
interval: object,
tmp_path: Path,
) -> None:
watchtower = _watchtower(tmp_path)
with pytest.raises((TypeError, ValueError), match="interval"):
watchtower.start_scheduler(interval=interval) # type: ignore[arg-type]
def test_scheduler_runs_immediately_and_stops_without_sleep_races(tmp_path: Path) -> None:
scanned = threading.Event()
def fetch():
scanned.set()
return []
watchtower = Watchtower(
tmp_path / "watchtower.json",
clock=lambda: NOW,
fetch_models=fetch,
)
watchtower.start_scheduler(interval=3_600)
assert scanned.wait(timeout=1)
assert watchtower.is_scanning
assert watchtower.stop_scheduler() is True
assert not watchtower.is_scanning
def test_scheduler_callback_can_request_its_own_stop_without_self_join(
tmp_path: Path,
) -> None:
callback_finished = threading.Event()
stop_results: list[bool] = []
watchtower = Watchtower(
tmp_path / "watchtower.json",
clock=lambda: NOW,
fetch_models=lambda: [_model("community/self-stop-instruct-7B")],
)
def request_stop(_model_info: DiscoveredModel) -> None:
stop_results.append(watchtower.stop_scheduler())
callback_finished.set()
watchtower.on_new_model(request_stop)
watchtower.start_scheduler(interval=3_600)
assert callback_finished.wait(timeout=1)
assert stop_results == [False]
assert watchtower._scheduler_state is SchedulerState.STOPPING
assert watchtower.stop_scheduler() is True
assert watchtower._scheduler_thread is None
assert watchtower._scheduler_state is SchedulerState.STOPPED
def test_scheduler_restart_and_start_failure_are_transactional(tmp_path: Path) -> None:
threads = []
class FakeThread:
def __init__(self, **kwargs):
assert kwargs["daemon"] is True
assert kwargs["name"] == "watchtower-scheduler"
self.alive = False
threads.append(self)
def start(self) -> None:
self.alive = True
def join(self, timeout: float) -> None:
assert timeout == 5.0
self.alive = False
def is_alive(self) -> bool:
return self.alive
watchtower = _watchtower(tmp_path, thread_factory=FakeThread)
watchtower.start_scheduler(interval=60)
assert watchtower.is_scanning
watchtower.start_scheduler(interval=120)
assert len(threads) == 2
assert not threads[0].is_alive()
assert threads[1].is_alive()
assert watchtower.stop_scheduler() is True
class FailingThread(FakeThread):
def start(self) -> None:
raise RuntimeError("cannot start")
failing = _watchtower(tmp_path, thread_factory=FailingThread)
with pytest.raises(RuntimeError, match="cannot start"):
failing.start_scheduler(interval=60)
assert failing._scheduler_thread is None
assert failing._scheduler_state is SchedulerState.STOPPED
assert not failing.is_scanning
def test_stuck_scheduler_is_retained_and_blocks_duplicate_start(tmp_path: Path) -> None:
class StuckThread:
def is_alive(self) -> bool:
return True
def join(self, timeout: float) -> None:
assert timeout == 0.01
watchtower = _watchtower(tmp_path, scheduler_join_timeout=0.01)
original = StuckThread()
watchtower._scheduler_thread = original
watchtower._scheduler_state = SchedulerState.RUNNING
assert watchtower.stop_scheduler() is False
assert watchtower._scheduler_thread is original
assert watchtower._scheduler_state is SchedulerState.STOPPING
with pytest.raises(RuntimeError, match="did not stop"):
watchtower.start_scheduler(interval=60)
assert watchtower._scheduler_thread is original
def test_queries_sort_limit_copy_and_format_rows(tmp_path: Path) -> None:
watchtower = _watchtower(tmp_path)
watchtower._models = {
"slow": DiscoveredModel(
model_id="org/slow",
name="slow",
org="org",
downloads_7d=10,
discovered_at="not-a-date-value",
),
"fast": DiscoveredModel(
model_id="org/fast",
name="fast",
org="org",
downloads_7d=20,
likes=3,
license="mit",
discovered_at=NOW.isoformat(),
),
}
assert [m.model_id for m in watchtower.get_trending(limit=1)] == ["org/fast"]
assert watchtower.get_model_choices() == ["org/fast", "org/slow"]
copied = watchtower.get_all_models()
copied.clear()
assert len(watchtower.get_all_models()) == 2
assert [m.model_id for m in watchtower.get_new_models()] == ["org/slow", "org/fast"]
rows = watchtower.format_table()
assert rows[0] == [
"org/fast", "org", "", "20", "3", "mit", "2026-08-16 12:30", "🆕 new"
]
assert rows[1][6] == "not-a-date-value"
def test_clock_must_be_timezone_aware(tmp_path: Path) -> None:
watchtower = Watchtower(
tmp_path / "watchtower.json",
clock=lambda: datetime(2026, 8, 16),
fetch_models=lambda: [],
)
with pytest.raises(ValueError, match="timezone-aware"):
watchtower.scan()