108 lines
2.9 KiB
Python
108 lines
2.9 KiB
Python
import argparse
|
|||
|
|
import pathlib
|
||
|
|
import pytest
|
||
|
|
|
||
|
|
from collections import defaultdict
|
||
|
|
|
||
|
|
SCRIPT_DIR = pathlib.Path(__file__).parent.resolve()
|
||
|
|
|
||
|
|
TEST_GROUPS = {
|
||
|
|
'BASIC_API': {
|
||
|
|
'tests': [
|
||
|
|
'test_empty_data_dir',
|
||
|
|
'test_bad_request',
|
||
|
|
'test_nonexistent_image',
|
||
|
|
],
|
||
|
|
'points': 1
|
||
|
|
},
|
||
|
|
'BASIC_PROCESSING': {
|
||
|
|
'tests': [
|
||
|
|
'test_task_queue',
|
||
|
|
'test_single_image',
|
||
|
|
'test_multiple_images',
|
||
|
|
'test_captions_generated_on_workers'
|
||
|
|
],
|
||
|
|
'points': 3
|
||
|
|
},
|
||
|
|
'CALLBACKS': {
|
||
|
|
'tests': [
|
||
|
|
'test_multiple_images_no_listdir'
|
||
|
|
],
|
||
|
|
'points': 2
|
||
|
|
},
|
||
|
|
'FAULT_TOLERANCE_1': {
|
||
|
|
'tests': [
|
||
|
|
'test_heartbeats_timeout'
|
||
|
|
],
|
||
|
|
'points': 1
|
||
|
|
},
|
||
|
|
'FAULT_TOLERANCE_2': {
|
||
|
|
'tests': [
|
||
|
|
'test_publisher_confirms'
|
||
|
|
],
|
||
|
|
'points': 1
|
||
|
|
},
|
||
|
|
'FAULT_TOLERANCE_3': {
|
||
|
|
'tests': [
|
||
|
|
'test_faulty_worker',
|
||
|
|
'test_two_faulty_workers',
|
||
|
|
],
|
||
|
|
'points': 1
|
||
|
|
},
|
||
|
|
'FAULT_TOLERANCE_4': {
|
||
|
|
'tests': [
|
||
|
|
'test_faulty_worker_and_rabbit_restart',
|
||
|
|
'test_total_eclipse_of_the_heart',
|
||
|
|
'test_notifications_survive_rabbit_restart',
|
||
|
|
'test_duplicate_deliveries'
|
||
|
|
],
|
||
|
|
'points': 1
|
||
|
|
}
|
||
|
|
}
|
||
|
|
|
||
|
|
class PassedCounter:
|
||
|
|
def __init__(self):
|
||
|
|
self.test_to_group = {}
|
||
|
|
for group_name, group in TEST_GROUPS.items():
|
||
|
|
for test in group['tests']:
|
||
|
|
self.test_to_group[test] = group_name
|
||
|
|
self.passed_by_group = defaultdict(int)
|
||
|
|
|
||
|
|
def pytest_report_teststatus(self, report, config):
|
||
|
|
if report.when == 'call' and report.passed:
|
||
|
|
test = report.nodeid.split('::')[1].split('[')[0]
|
||
|
|
group_name = self.test_to_group[test]
|
||
|
|
self.passed_by_group[group_name] += 1
|
||
|
|
|
||
|
|
|
||
|
|
def main(argv=None):
|
||
|
|
parser = argparse.ArgumentParser()
|
||
|
|
parser.add_argument('--ci', action='store_true',
|
||
|
|
help='Fail unless all reference solution tests pass')
|
||
|
|
args = parser.parse_args(argv)
|
||
|
|
counter = PassedCounter()
|
||
|
|
test_status = pytest.main(
|
||
|
|
['-vs', '--tb=short', str(SCRIPT_DIR / 'test_server.py')], plugins=[counter])
|
||
|
|
|
||
|
|
score = 0
|
||
|
|
print()
|
||
|
|
for group_name, group in TEST_GROUPS.items():
|
||
|
|
total = len(group['tests'])
|
||
|
|
passed = counter.passed_by_group[group_name]
|
||
|
|
print(f'Test group {group_name}: passed {passed} of {total} tests')
|
||
|
|
if passed == total:
|
||
|
|
score += group['points']
|
||
|
|
|
||
|
|
print(f"\nSCORE: {score}")
|
||
|
|
if args.ci:
|
||
|
|
max_score = sum(group['points'] for group in TEST_GROUPS.values())
|
||
|
|
return int(test_status) or int(score != max_score)
|
||
|
|
# A failed student test still produces a valid partial score.
|
||
|
|
if test_status == pytest.ExitCode.TESTS_FAILED:
|
||
|
|
return 0
|
||
|
|
return int(test_status)
|
||
|
|
|
||
|
|
|
||
|
|
if __name__ == '__main__':
|
||
|
|
raise SystemExit(main())
|