t2av_eval / backend /test_assignment.py
Youngsun's picture
Implement multi-round balanced assignment and MTurk completion codes
23f9faf
Raw
History Blame
21.2 kB
"""Unit tests for the multi-round balanced assignment + completion-code
feature.
Runs offline (no network / no huggingface_hub calls): assignment.py's core
functions are pure, and the round-lifecycle integration tests patch main.py's
HF-facing I/O functions with an in-memory fake store.
Run with: python -m unittest backend.test_assignment (from repo root)
or: python -m unittest test_assignment (from backend/)
"""
import json
import os
import random
import sys
import unittest
from datetime import datetime, timedelta, timezone
from unittest import mock
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
import assignment # noqa: E402
def make_record(user, video_id, created_at, extra=None):
record = {
"id": "x",
"video_id": video_id,
"user": user,
"created_at": created_at,
"annotations": {"responses": {}},
}
if extra:
record.update(extra)
return record
def iso(dt):
return dt.isoformat().replace("+00:00", "Z")
class DedupeLatestTests(unittest.TestCase):
def test_existing_annotations_preserved_and_counted(self):
records = [
make_record("Adi", "v1", "2026-01-01T00:00:00Z"),
make_record("Adi", "v2", "2026-01-01T00:00:00Z"),
make_record("youngsun", "v1", "2026-01-01T00:00:00Z"),
]
original_len = len(records)
deduped, skipped = assignment.dedupe_latest(records)
self.assertEqual(len(records), original_len, "dedupe_latest must not mutate its input")
self.assertEqual(skipped, 0)
self.assertEqual(len(deduped), 3)
by_video = assignment.completed_by_video(deduped)
self.assertEqual(by_video["v1"], {"adi", "youngsun"})
self.assertEqual(by_video["v2"], {"adi"})
def test_same_annotator_same_video_counted_once_latest_kept(self):
records = [
make_record("Adi", "v1", "2026-01-01T00:00:00Z", {"annotations": {"marker": "old"}}),
make_record("Adi", "v1", "2026-01-02T00:00:00Z", {"annotations": {"marker": "new"}}),
make_record("adi", "v1", "2026-01-03T00:00:00Z", {"annotations": {"marker": "newest"}}),
]
deduped, skipped = assignment.dedupe_latest(records)
self.assertEqual(skipped, 0)
self.assertEqual(len(deduped), 1, "same annotator + same video, across case variants, is one record")
self.assertEqual(deduped[("adi", "v1")]["annotations"]["marker"], "newest")
def test_malformed_or_partial_records_skipped_not_counted(self):
records = [
make_record("Adi", "v1", "2026-01-01T00:00:00Z"),
{"id": "y", "video_id": "v2", "user": "Adi"}, # missing created_at (autosave/reserve-like)
{"id": "z", "user": "Adi", "created_at": "2026-01-01T00:00:00Z"}, # missing video_id
{"id": "w", "video_id": "v3", "user": "", "created_at": "2026-01-01T00:00:00Z"}, # empty user
{"id": "q", "video_id": "v4", "user": "Adi", "created_at": "not-a-date"}, # unparsable timestamp
]
deduped, skipped = assignment.dedupe_latest(records)
self.assertEqual(len(deduped), 1)
self.assertEqual(skipped, 4)
class PickBalancedTests(unittest.TestCase):
def test_prefers_lowest_coverage(self):
coverage = {"v1": 4, "v2": 0, "v3": 2}
picked = assignment.pick_balanced(list(coverage), coverage.get, 2, random.Random(0), target=5)
self.assertEqual(set(picked), {"v2", "v3"})
def test_excludes_at_or_over_target_when_enough_under_target_remain(self):
coverage = {"v1": 5, "v2": 5, "v3": 1}
picked = assignment.pick_balanced(list(coverage), coverage.get, 1, random.Random(0), target=5)
self.assertEqual(picked, ["v3"])
def test_falls_back_to_at_target_videos_when_not_enough_under_target(self):
coverage = {"v1": 5, "v2": 5, "v3": 1}
picked = assignment.pick_balanced(list(coverage), coverage.get, 3, random.Random(0), target=5)
self.assertEqual(set(picked), {"v1", "v2", "v3"})
def test_never_returns_more_than_requested_or_duplicates(self):
coverage = {f"v{i}": i % 3 for i in range(30)}
picked = assignment.pick_balanced(list(coverage), coverage.get, 7, random.Random(1), target=5)
self.assertEqual(len(picked), 7)
self.assertEqual(len(set(picked)), 7)
class CompletionCodeTests(unittest.TestCase):
def test_format(self):
code = assignment.generate_completion_code()
self.assertTrue(code.startswith("T2AV-"))
suffix = code[len("T2AV-"):]
self.assertEqual(len(suffix), 12)
self.assertTrue(all(c in assignment.COMPLETION_CODE_ALPHABET for c in suffix))
def test_codes_are_not_trivially_repeated(self):
codes = {assignment.generate_completion_code() for _ in range(200)}
self.assertEqual(len(codes), 200, "200 draws from a 36^12 space should never collide")
class BuildRoundTests(unittest.TestCase):
def test_always_excludes_already_completed_never_seeds_with_them(self):
"""Core policy change: a round is exactly N NEW videos, never padded
with history, regardless of how much history exists."""
catalog = [f"v{i}" for i in range(1, 201)] # 200 videos, like production
already_completed = set(catalog[:51]) # mirrors real production data (Adi: 51 completed)
round_record = assignment.build_round(
"adi", "Adi", catalog, already_completed, {}, {}, random.Random(0), round_number=1,
)
self.assertEqual(round_record["actual_size"], 20)
self.assertEqual(len(round_record["video_ids"]), 20)
self.assertEqual(
already_completed.intersection(round_record["video_ids"]), set(),
"none of the 51 previously completed videos may appear in the new round",
)
self.assertEqual(len(set(round_record["video_ids"])), 20, "no duplicates within the round")
self.assertEqual(round_record["round_number"], 1)
self.assertEqual(round_record["requested_size"], 20)
def test_one_historical_completion_gets_twenty_new_not_nineteen(self):
catalog = [f"v{i}" for i in range(1, 201)]
already_completed = {"v1"}
round_record = assignment.build_round(
"youngsun", "youngsun_0715", catalog, already_completed, {}, {}, random.Random(0), round_number=1,
)
self.assertEqual(round_record["actual_size"], 20)
self.assertNotIn("v1", round_record["video_ids"])
def test_brand_new_annotator_gets_exactly_twenty(self):
catalog = [f"v{i}" for i in range(1, 201)]
round_record = assignment.build_round(
"new", "New", catalog, set(), {}, {}, random.Random(0), round_number=1,
)
self.assertEqual(round_record["actual_size"], 20)
def test_reduced_size_when_not_enough_eligible_remain(self):
catalog = [f"v{i}" for i in range(1, 26)] # only 25 videos total
already_completed = set(catalog[:15]) # 15 done, only 10 eligible left
round_record = assignment.build_round(
"x", "X", catalog, already_completed, {}, {}, random.Random(0), round_number=1,
)
self.assertEqual(round_record["actual_size"], 10)
self.assertEqual(round_record["requested_size"], 20)
self.assertEqual(len(round_record["video_ids"]), 10)
self.assertEqual(already_completed.intersection(round_record["video_ids"]), set())
def test_second_round_excludes_first_rounds_completions_too(self):
"""After Adi finishes 51 (history) + 20 (round 1) = 71, round 2 must
exclude all 71, selecting from the remaining 129."""
catalog = [f"v{i}" for i in range(1, 201)]
historical = set(catalog[:51])
round1 = assignment.build_round(
"adi", "Adi", catalog, historical, {}, {}, random.Random(0), round_number=1,
)
after_round1 = historical | set(round1["video_ids"])
self.assertEqual(len(after_round1), 71)
round2 = assignment.build_round(
"adi", "Adi", catalog, after_round1, {}, {}, random.Random(1), round_number=2,
)
self.assertEqual(round2["actual_size"], 20)
self.assertEqual(after_round1.intersection(round2["video_ids"]), set())
remaining_pool_size = len(catalog) - len(after_round1)
self.assertEqual(remaining_pool_size, 129)
self.assertEqual(round2["round_number"], 2)
def test_existing_completions_affect_new_round_priority(self):
catalog = ["v1", "v2", "v3"]
completed_by_video_map = {"v1": {"p1", "p2", "p3", "p4"}} # coverage 4, still < target 5
round_record = assignment.build_round(
"new", "New", catalog, set(), completed_by_video_map, {}, random.Random(0),
round_number=1, videos_per_annotator=2,
)
self.assertNotIn("v1", round_record["video_ids"], "v2/v3 have lower coverage and must be preferred")
self.assertEqual(set(round_record["video_ids"]), {"v2", "v3"})
class ReservedByVideoTests(unittest.TestCase):
def test_counts_open_uncompleted_unexpired_reservations(self):
now = datetime(2026, 1, 10, tzinfo=timezone.utc)
assignments = [
{"annotator_id": "a", "video_ids": ["v1", "v2"], "created_at": "2026-01-09T00:00:00Z"},
{"annotator_id": "b", "video_ids": ["v1"], "created_at": "2026-01-09T00:00:00Z"},
]
completed_by_annotator = {"a": {"v2"}}
reserved = assignment.reserved_by_video(assignments, completed_by_annotator, now, ttl_seconds=7 * 24 * 3600)
self.assertEqual(reserved.get("v1"), 2)
self.assertNotIn("v2", reserved)
def test_excludes_expired_reservations_without_needing_to_delete_anything(self):
now = datetime(2026, 1, 10, tzinfo=timezone.utc)
assignments = [{"annotator_id": "a", "video_ids": ["v1"], "created_at": "2025-01-01T00:00:00Z"}]
reserved = assignment.reserved_by_video(assignments, {}, now, ttl_seconds=24 * 3600)
self.assertEqual(reserved.get("v1", 0), 0)
def test_excludes_the_requesting_annotator_from_their_own_reservations(self):
now = datetime(2026, 1, 10, tzinfo=timezone.utc)
assignments = [{"annotator_id": "a", "video_ids": ["v1"], "created_at": "2026-01-09T00:00:00Z"}]
reserved = assignment.reserved_by_video(assignments, {}, now, ttl_seconds=24 * 3600, exclude_annotator="a")
self.assertEqual(reserved.get("v1", 0), 0)
class SequentialBalanceTests(unittest.TestCase):
def test_sequential_rounds_stay_reasonably_balanced(self):
catalog = [f"v{i}" for i in range(1, 81)] # 80 videos
reserved_by_video_map = {}
rng = random.Random(42)
for i in range(10): # 10 annotators * 20 slots = 200 slots over 80 videos
round_record = assignment.build_round(
f"person{i}", f"person{i}", catalog, set(), {}, reserved_by_video_map, rng, round_number=1,
)
for video_id in round_record["video_ids"]:
reserved_by_video_map[video_id] = reserved_by_video_map.get(video_id, 0) + 1
counts = list(reserved_by_video_map.values())
self.assertEqual(len(reserved_by_video_map), 80, "every video should have been touched")
self.assertLessEqual(max(counts) - min(counts), 2, "balancing should keep coverage tight")
class FakeRepoStore:
"""Stands in for the annotations dataset repo's assignments/ and
completions/ prefixes."""
def __init__(self):
self.files = {}
def read(self, path):
return self.files.get(path)
def write_kwargs(self, **kwargs):
path = kwargs["path_in_repo"]
content = json.loads(kwargs["path_or_fileobj"].getvalue().decode("utf-8"))
self.files[path] = content
def list_active_round_records(self):
records = []
for path, pointer in self.files.items():
if not path.endswith("/current.json"):
continue
annotator_id = path.split("/")[1]
round_path = f"assignments/{annotator_id}/rounds/{pointer['assignment_id']}.json"
completion_path = f"completions/{annotator_id}/{pointer['assignment_id']}.json"
round_record = self.files.get(round_path)
if round_record is None or completion_path in self.files:
continue
records.append(round_record)
return records
class RoundLifecycleIntegrationTests(unittest.TestCase):
"""Exercises main.get_or_create_current_round() with every HF network
call replaced by an in-memory fake, so this stays fully offline."""
def setUp(self):
import main as backend_main
self.backend_main = backend_main
self.store = FakeRepoStore()
self.catalog = [f"v{i}" for i in range(1, 201)] # 200 videos, like production
self.completed_records = [] # raw annotation records fed to both list_annotation_records variants
patches = [
mock.patch.object(backend_main, "_read_json_from_repo", side_effect=self.store.read),
# Both the cached and fresh variants must be mocked - the completion
# check in get_or_create_current_round() deliberately bypasses the
# cache (fresh=True) so a real annotator sees their code promptly;
# leaving list_annotation_records() unmocked would silently fall
# through to a real network call here.
mock.patch.object(
backend_main, "list_annotation_records_cached", side_effect=lambda: self.completed_records
),
mock.patch.object(
backend_main, "list_annotation_records", side_effect=lambda: self.completed_records
),
mock.patch.object(
backend_main, "list_active_round_records", side_effect=self.store.list_active_round_records
),
mock.patch.object(backend_main, "ensure_dataset_exists", return_value=None),
mock.patch.object(backend_main, "ensure_dataset_configured", return_value=None),
mock.patch.object(backend_main, "_current_catalog_video_ids", return_value=self.catalog),
mock.patch.object(backend_main, "upload_with_retry", side_effect=self.store.write_kwargs),
]
for patcher in patches:
patcher.start()
self.addCleanup(patcher.stop)
def _complete_round(self, annotator_raw, round_record):
"""Simulates the annotator finishing every video in a round by
appending matching annotation records - exactly what save_annotation
would have produced."""
now = datetime.now(timezone.utc)
for video_id in round_record["video_ids"]:
self.completed_records.append(make_record(annotator_raw, video_id, iso(now)))
def test_brand_new_annotator_gets_round_one_of_twenty(self):
result = self.backend_main.get_or_create_current_round("Scratch Tester")
self.assertEqual(result["status"], "in_progress")
self.assertEqual(result["round_number"], 1)
self.assertEqual(len(result["video_ids"]), 20)
self.assertEqual(len(set(result["video_ids"])), 20)
def test_refresh_while_incomplete_returns_identical_round(self):
first = self.backend_main.get_or_create_current_round("Scratch Tester")
second = self.backend_main.get_or_create_current_round("Scratch Tester")
self.assertEqual(first["video_ids"], second["video_ids"])
self.assertEqual(first["assignment_id"], second["assignment_id"])
def test_no_completion_code_at_nineteen_of_twenty(self):
first = self.backend_main.get_or_create_current_round("Scratch Tester")
for video_id in first["video_ids"][:19]:
self.completed_records.append(make_record("Scratch Tester", video_id, iso(datetime.now(timezone.utc))))
result = self.backend_main.get_or_create_current_round("Scratch Tester")
self.assertEqual(result["status"], "in_progress")
self.assertNotIn("completion_code", result)
def test_completion_code_persisted_once_at_twenty_of_twenty(self):
first = self.backend_main.get_or_create_current_round("Scratch Tester")
self._complete_round("Scratch Tester", first)
result_a = self.backend_main.get_or_create_current_round("Scratch Tester")
result_b = self.backend_main.get_or_create_current_round("Scratch Tester")
self.assertEqual(result_a["status"], "completed")
self.assertTrue(result_a["completion_code"].startswith("T2AV-"))
self.assertEqual(result_a["completion_code"], result_b["completion_code"], "code must be stable on refresh")
completion_files = [p for p in self.store.files if p.startswith("completions/")]
self.assertEqual(len(completion_files), 1, "completion record written exactly once")
def test_within_grace_period_shows_same_completed_round_not_a_new_one(self):
first = self.backend_main.get_or_create_current_round("Scratch Tester")
self._complete_round("Scratch Tester", first)
self.backend_main.get_or_create_current_round("Scratch Tester") # triggers completion-record creation
# Directly age the completion record to just inside the grace window.
completion_path = f"completions/scratch-tester/{first['assignment_id']}.json"
recent = datetime.now(timezone.utc) - timedelta(
seconds=assignment.COMPLETION_GRACE_PERIOD_SECONDS - 30
)
self.store.files[completion_path]["completed_at"] = iso(recent)
result = self.backend_main.get_or_create_current_round("Scratch Tester")
self.assertEqual(result["status"], "completed")
self.assertEqual(result["round_number"], 1)
def test_past_grace_period_auto_creates_next_round_with_no_extra_call(self):
first = self.backend_main.get_or_create_current_round("Scratch Tester")
self._complete_round("Scratch Tester", first)
self.backend_main.get_or_create_current_round("Scratch Tester") # triggers completion-record creation
# Directly age the completion record past the grace window - simulates "came back later".
completion_path = f"completions/scratch-tester/{first['assignment_id']}.json"
stale = datetime.now(timezone.utc) - timedelta(
seconds=assignment.COMPLETION_GRACE_PERIOD_SECONDS + 60
)
self.store.files[completion_path]["completed_at"] = iso(stale)
result = self.backend_main.get_or_create_current_round("Scratch Tester")
self.assertEqual(result["status"], "in_progress")
self.assertEqual(result["round_number"], 2)
self.assertEqual(set(result["video_ids"]).intersection(first["video_ids"]), set())
def test_never_assigns_the_same_video_twice_across_rounds(self):
first = self.backend_main.get_or_create_current_round("Scratch Tester")
self._complete_round("Scratch Tester", first)
self.backend_main.get_or_create_current_round("Scratch Tester")
completion_path = f"completions/scratch-tester/{first['assignment_id']}.json"
stale = datetime.now(timezone.utc) - timedelta(seconds=assignment.COMPLETION_GRACE_PERIOD_SECONDS + 60)
self.store.files[completion_path]["completed_at"] = iso(stale)
second = self.backend_main.get_or_create_current_round("Scratch Tester")
all_seen = first["video_ids"] + second["video_ids"]
self.assertEqual(len(all_seen), len(set(all_seen)))
def test_different_annotators_each_get_their_own_round(self):
a = self.backend_main.get_or_create_current_round("Person A")
b = self.backend_main.get_or_create_current_round("Person B")
self.assertEqual(len(a["video_ids"]), 20)
self.assertEqual(len(b["video_ids"]), 20)
round_files = [p for p in self.store.files if "/rounds/" in p]
self.assertEqual(len(round_files), 2)
def test_simultaneous_requests_do_not_create_duplicate_rounds(self):
"""The lock fully serializes get_or_create_current_round, so we
simulate "simultaneous" by calling it back-to-back and asserting only
one round/pointer combination was ever written - the real concurrency
guarantee is exercised live against a scratch HF repo separately."""
for _ in range(5):
self.backend_main.get_or_create_current_round("Scratch Tester")
round_files = [p for p in self.store.files if "/rounds/" in p]
self.assertEqual(len(round_files), 1)
def test_partial_records_never_mistakenly_counted_as_completed(self):
first = self.backend_main.get_or_create_current_round("Scratch Tester")
# An autosave/reservation-like record missing created_at for every video.
for video_id in first["video_ids"]:
self.completed_records.append({"id": "x", "video_id": video_id, "user": "Scratch Tester"})
result = self.backend_main.get_or_create_current_round("Scratch Tester")
self.assertEqual(result["status"], "in_progress")
if __name__ == "__main__":
unittest.main()