Skip to content

Commit 64624cb

Browse files
committed
Refactor grouped error detection
- Introduce an `ErrorMultiplicity` enum and return it from `get_expected_errors()`. - Extract `determine_group_error()` from `diff_expected_errors()` and simplify logic. - Rerun black on `main.py`.
1 parent 8e71a2e commit 64624cb

1 file changed

Lines changed: 73 additions & 33 deletions

File tree

‎conformance/src/main.py‎

Lines changed: 73 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -6,7 +6,8 @@
66
import re
77
import sys
88
import tomllib
9-
from collections.abc import Sequence
9+
from collections.abc import Container, Sequence
10+
from enum import Enum
1011
from pathlib import Path
1112
from time import time
1213

@@ -63,17 +64,24 @@ def run_tests(
6364
update_type_checker_info(type_checker, root_dir)
6465

6566

67+
class ErrorMultiplicity(Enum):
68+
"""How many lines with the same tag can have an error."""
69+
70+
SINGLE = 1 # exactly one line must have an error
71+
MULTI = 2 # at least one line must have an error
72+
REQUIRE_SUCCESS = 3 # at least one line must not have an error
73+
74+
6675
def get_expected_errors(test_case: Path) -> tuple[
6776
dict[int, tuple[int, int]],
68-
dict[str, tuple[list[int], bool, bool]],
77+
dict[str, tuple[list[int], ErrorMultiplicity]],
6978
]:
7079
"""Return the line numbers where type checkers are expected to produce an error.
7180
7281
The return value is a tuple of two dictionaries:
7382
- The format of the first is {line number: (number of required errors, number of optional errors)}.
74-
- The format of the second is {error tag: ([lines where the error may appear], allow multiple, require success)}.
75-
If require success is True, at least one line can't raise an error; otherwise, if allow multiple is True, the
76-
error may appear on multiple lines; otherwise, it must appear exactly once.
83+
- The format of the second is {error tag: ([lines where the error may appear], error multiplicity)}.
84+
See the ErrorMultiplicy enum for the second argument.
7785
7886
For example, the following test case:
7987
@@ -92,7 +100,7 @@ def f(): pass # E[final]
92100
with open(test_case, "r", encoding="utf-8") as f:
93101
lines = f.readlines()
94102
output: dict[int, tuple[int, int]] = {}
95-
groups: dict[str, tuple[list[int], bool, bool]] = {}
103+
groups: dict[str, tuple[list[int], ErrorMultiplicity]] = {}
96104
for i, line in enumerate(lines, start=1):
97105
line_without_comment, *_ = line.split("#")
98106
# Ignore lines with no non-comment content. This allows commenting out test cases.
@@ -110,27 +118,26 @@ def f(): pass # E[final]
110118
for match in re.finditer(r"# E\[([^\]]+)\]", line):
111119
tag = match.group(1)
112120
if tag.endswith("+"):
113-
allow_multiple = True
114-
require_success = False
121+
multiplicity: ErrorMultiplicity = ErrorMultiplicity.MULTI
115122
tag = tag[:-1]
116123
elif tag.endswith("!"):
117-
allow_multiple = True
118-
require_success = True
124+
multiplicity = ErrorMultiplicity.REQUIRE_SUCCESS
119125
tag = tag[:-1]
120126
else:
121-
allow_multiple = False
122-
require_success = False
127+
multiplicity = ErrorMultiplicity.SINGLE
123128
if tag not in groups:
124-
groups[tag] = ([i], allow_multiple, require_success)
129+
groups[tag] = ([i], multiplicity)
125130
else:
126-
if groups[tag][1] != allow_multiple:
127-
raise ValueError(f"Error group {tag} has inconsistent allow_multiple value in {test_case}")
128-
if groups[tag][2] != require_success:
129-
raise ValueError(f"Error group {tag} has inconsistent require_success value in {test_case}")
131+
if groups[tag][1] != multiplicity:
132+
raise ValueError(
133+
f"Error group {tag} has inconsistent multiplicity value in {test_case}"
134+
)
130135
groups[tag][0].append(i)
131136
for group, linenos in groups.items():
132137
if len(linenos) == 1:
133-
raise ValueError(f"Error group {group} only appears on a single line in {test_case}")
138+
raise ValueError(
139+
f"Error group {group} only appears on a single line in {test_case}"
140+
)
134141
return output, groups
135142

136143

@@ -148,34 +155,63 @@ def diff_expected_errors(
148155
lineno: [
149156
error
150157
for error in errors_list
151-
if not any(ignored in error for ignored in ignored_errors)]
158+
if not any(ignored in error for ignored in ignored_errors)
159+
]
152160
for lineno, errors_list in errors.items()
153161
}
154-
errors = {lineno: errors_list for lineno, errors_list in errors.items() if errors_list}
162+
errors = {
163+
lineno: errors_list for lineno, errors_list in errors.items() if errors_list
164+
}
155165

156166
differences: list[str] = []
157167
for expected_lineno, (expected_count, _) in expected_errors.items():
158168
if expected_lineno not in errors and expected_count > 0:
159-
differences.append(f"Line {expected_lineno}: Expected {expected_count} errors")
169+
differences.append(
170+
f"Line {expected_lineno}: Expected {expected_count} errors"
171+
)
160172
# We don't report an issue if the count differs, because type checkers may produce
161173
# multiple error messages for a single line.
162174
linenos_used_by_groups: set[int] = set()
163-
for group, (linenos, allow_multiple, require_success) in error_groups.items():
164-
num_errors = sum(1 for lineno in linenos if lineno in errors)
165-
if require_success and num_errors == len(linenos):
166-
differences.append(f"Lines {', '.join(map(str, linenos))}: Expected at least one success (tag {group!r})")
167-
elif num_errors == 0 and not require_success:
168-
differences.append(f"Lines {', '.join(map(str, linenos))}: Expected error (tag {group!r})")
169-
elif num_errors == 1 or allow_multiple or require_success:
175+
for group, (linenos, multiplicity) in error_groups.items():
176+
error = determine_group_error(group, linenos, multiplicity, errors)
177+
if error is None:
170178
linenos_used_by_groups.update(linenos)
171179
else:
172-
differences.append(f"Lines {', '.join(map(str, linenos))}: Expected exactly one error (tag {group!r})")
180+
differences.append(error)
173181
for actual_lineno, actual_errors in errors.items():
174-
if actual_lineno not in expected_errors and actual_lineno not in linenos_used_by_groups:
175-
differences.append(f"Line {actual_lineno}: Unexpected errors {actual_errors}")
182+
if (
183+
actual_lineno not in expected_errors
184+
and actual_lineno not in linenos_used_by_groups
185+
):
186+
differences.append(
187+
f"Line {actual_lineno}: Unexpected errors {actual_errors}"
188+
)
176189
return "".join(f"{diff}\n" for diff in differences)
177190

178191

192+
def determine_group_error(
193+
group: str,
194+
group_linenos: Sequence[int],
195+
multiplicity: ErrorMultiplicity,
196+
error_linenos: Container[int],
197+
) -> str | None:
198+
"""Return the error message for the given group or None."""
199+
num_errors = sum(1 for lineno in group_linenos if lineno in error_linenos)
200+
match multiplicity:
201+
case ErrorMultiplicity.SINGLE:
202+
if num_errors == 0:
203+
return f"Lines {', '.join(map(str, group_linenos))}: Expected error (tag {group!r})"
204+
elif num_errors > 1:
205+
return f"Lines {', '.join(map(str, group_linenos))}: Expected exactly one error (tag {group!r})"
206+
case ErrorMultiplicity.MULTI:
207+
if num_errors == 0:
208+
return f"Lines {', '.join(map(str, group_linenos))}: Expected error (tag {group!r})"
209+
case ErrorMultiplicity.REQUIRE_SUCCESS:
210+
if num_errors == len(group_linenos):
211+
return f"Lines {', '.join(map(str, group_linenos))}: Expected at least one success (tag {group!r})"
212+
return None
213+
214+
179215
def update_output_for_test(
180216
type_checker: TypeChecker,
181217
results_dir: Path,
@@ -201,7 +237,9 @@ def update_output_for_test(
201237
existing_results = {}
202238

203239
ignored_errors = existing_results.get("ignore_errors", [])
204-
errors_diff = "\n" + diff_expected_errors(type_checker, test_case, output, ignored_errors)
240+
errors_diff = "\n" + diff_expected_errors(
241+
type_checker, test_case, output, ignored_errors
242+
)
205243
old_errors_diff = "\n" + existing_results.get("errors_diff", "")
206244

207245
if errors_diff != old_errors_diff:
@@ -294,7 +332,9 @@ def main():
294332
if not type_checker.install():
295333
print(f"Skipping tests for {type_checker.name}")
296334
else:
297-
run_tests(root_dir, type_checker, test_cases, verbose=options.verbose)
335+
run_tests(
336+
root_dir, type_checker, test_cases, verbose=options.verbose
337+
)
298338

299339
# Generate a summary report.
300340
generate_summary(root_dir)

0 commit comments

Comments
 (0)