137 lines
5.1 KiB
Python
137 lines
5.1 KiB
Python
import cv2
|
|
import numpy as np
|
|
import pytest
|
|
import template_finder
|
|
from template_finder import TemplateMatch
|
|
from utils.misc import is_in_roi
|
|
import screen
|
|
import utils.download_test_assets # downloads assets if they don't already exist, doesn't need to be called
|
|
|
|
screen.set_window_position(0, 0)
|
|
|
|
def test_search():
|
|
"""
|
|
Test default search behavior (first match)
|
|
- searches first for cross, which doesn't perfectly match but should reach above threshold
|
|
- if cross matches above threshold as expected, then it won't bother to search for slash, which has a perfect match on the image
|
|
- test passes if the template match score is not perfect
|
|
"""
|
|
image = cv2.imread("test/assets/stash_slots.png")
|
|
slash = cv2.imread("test/assets/stash_slot_slash.png")
|
|
cross = cv2.imread("test/assets/stash_slot_cross.png")
|
|
threshold=0.6
|
|
match = template_finder.search([cross, slash], image, threshold)
|
|
assert threshold <= match.score < 1
|
|
|
|
def test_search_best_match():
|
|
"""
|
|
Test search "best_match" behavior
|
|
- searches first for cross, which doesn't perfectly match
|
|
- searches next for slash, which perfectly matches on image
|
|
- test passes if the center of the template match lies within the expected region of the slash
|
|
"""
|
|
image = cv2.imread("test/assets/stash_slots.png")
|
|
slash = cv2.imread("test/assets/stash_slot_slash.png")
|
|
cross = cv2.imread("test/assets/stash_slot_cross.png")
|
|
slash_expected_roi = [38, 0, 38, 38]
|
|
match = template_finder.search([cross, slash], image, threshold=0.6, best_match=True)
|
|
assert is_in_roi(slash_expected_roi, match.center)
|
|
|
|
def test_search_all():
|
|
"""
|
|
Test all matches for a single template in argument
|
|
- searches for empty slots with high threshold
|
|
- test passes if 3 matches result
|
|
"""
|
|
image = cv2.imread("test/assets/stash_slots.png")
|
|
empty = cv2.imread("test/assets/stash_slot_empty.png")
|
|
matches = template_finder.search_all(empty, image, threshold=0.98)
|
|
assert len(matches) == 3
|
|
|
|
def test_search_all_multiple_templates():
|
|
"""
|
|
Test all matches with multiple templates in argument
|
|
- searches for empty slots and slash with high threshold
|
|
- test passes if 4 matches result
|
|
"""
|
|
image = cv2.imread("test/assets/stash_slots.png")
|
|
empty = cv2.imread("test/assets/stash_slot_empty.png")
|
|
slash = cv2.imread("test/assets/stash_slot_slash.png")
|
|
matches = template_finder.search_all([empty, slash], image, threshold=0.98)
|
|
assert len(matches) == 4
|
|
|
|
def test_search_and_wait_stable_requires_consecutive_match(monkeypatch):
|
|
img = np.ones((720, 1280, 3), dtype=np.uint8) * 255
|
|
calls = []
|
|
matches = [
|
|
TemplateMatch(name="A", score=0.9, valid=True),
|
|
TemplateMatch(name="A", score=0.91, valid=True),
|
|
]
|
|
|
|
monkeypatch.setattr(template_finder, "grab", lambda force_new=False: img)
|
|
|
|
def fake_search(*args, **kwargs):
|
|
calls.append(kwargs)
|
|
return matches.pop(0)
|
|
|
|
monkeypatch.setattr(template_finder, "search", fake_search)
|
|
|
|
match = template_finder.search_and_wait_stable("A", timeout=1, confirmations=2, interval=0)
|
|
assert match.valid
|
|
assert match.name == "A"
|
|
assert len(calls) == 2
|
|
|
|
def test_search_and_wait_stable_rejects_single_frame_match(monkeypatch):
|
|
img = np.ones((720, 1280, 3), dtype=np.uint8) * 255
|
|
matches = [
|
|
TemplateMatch(name="A", score=0.9, valid=True),
|
|
TemplateMatch(name=None, score=0.1, valid=False),
|
|
]
|
|
|
|
monkeypatch.setattr(template_finder, "grab", lambda force_new=False: img)
|
|
|
|
def fake_search(*args, **kwargs):
|
|
return matches.pop(0) if matches else TemplateMatch(name=None, score=0.1, valid=False)
|
|
|
|
monkeypatch.setattr(template_finder, "search", fake_search)
|
|
|
|
match = template_finder.search_and_wait_stable("A", timeout=0.01, confirmations=2, interval=0)
|
|
assert not match.valid
|
|
assert match.name == "A"
|
|
|
|
def test_search_and_wait_stable_saves_missing_debug(monkeypatch):
|
|
img = np.ones((720, 1280, 3), dtype=np.uint8) * 255
|
|
saved = []
|
|
|
|
monkeypatch.setattr(template_finder, "grab", lambda force_new=False: img)
|
|
monkeypatch.setattr(template_finder, "search", lambda *args, **kwargs: TemplateMatch(name="A", score=0.2, valid=False))
|
|
monkeypatch.setattr(template_finder, "safe_imwrite", lambda path, image: saved.append((path, image.shape)) or True)
|
|
|
|
match = template_finder.search_and_wait_stable(
|
|
"A",
|
|
roi=[10, 20, 30, 40],
|
|
timeout=0.01,
|
|
confirmations=2,
|
|
interval=0,
|
|
suppress_debug=True,
|
|
save_debug=True,
|
|
)
|
|
|
|
assert not match.valid
|
|
assert len(saved) == 2
|
|
assert saved[0][0].endswith(".png")
|
|
assert saved[1][0].endswith("_roi.png")
|
|
assert saved[1][1] == (40, 30, 3)
|
|
|
|
if __name__ == "__main__":
|
|
image = cv2.imread("test/assets/stash_slots.png")
|
|
empty = cv2.imread("test/assets/stash_slot_empty.png")
|
|
slash = cv2.imread("test/assets/stash_slot_slash.png")
|
|
cross = cv2.imread("test/assets/stash_slot_cross.png")
|
|
slash_expected_roi = [38, 0, 38, 38]
|
|
|
|
|
|
matches = template_finder.search_all([empty, slash], image, threshold=0.98)
|
|
print(len(matches))
|
|
print(matches)
|