diff --git a/ARCHITECTURE.md b/ARCHITECTURE.md index 776ffed..d344a58 100644 --- a/ARCHITECTURE.md +++ b/ARCHITECTURE.md @@ -47,6 +47,7 @@ The repository adopts responsibility-based modular organization: | [`prompts.py`](file:///Users/shreshthdhimole/AutoSetter/autosetter/prompts.py) | Template loader and safe `{JSON}` substitution without string format hazards. | | [`extractor.py`](file:///Users/shreshthdhimole/AutoSetter/autosetter/extractor.py) | Extraction of `problem.json` via VLM, markdown code fence stripping, schema validation. | | [`generator.py`](file:///Users/shreshthdhimole/AutoSetter/autosetter/generator.py) | Downstream generation orchestrator for statement markdown and testlib C++ code. | +| `testgen/` | Z3 test generation from `test_spec.json`: spec validation, Z3 integer solving per seed strategy, array/string/tree/graph builders, sample checking. | | [`sandbox.py`](file:///Users/shreshthdhimole/AutoSetter/autosetter/sandbox.py) | C++ compilation and execution backends (host `g++` and Docker+NsJail HTTP). | | [`pipeline.py`](file:///Users/shreshthdhimole/AutoSetter/autosetter/pipeline.py) | Test generation, validation, jury answer solving, attribution, and checker probing. | | [`packager.py`](file:///Users/shreshthdhimole/AutoSetter/autosetter/packager.py) | Assembly of release bundle, test pairing, and manifest creation. | @@ -70,7 +71,7 @@ The repository adopts responsibility-based modular organization: - Serializes `problem.json` and renders five specialized prompt templates: 1. `statement.txt` ➔ `generated/statement.md` 2. `validator.txt` ➔ `generated/validator.cpp` (uses `testlib.h`) - 3. `generator.txt` ➔ `generated/generator.cpp` (uses `testlib.h`) + 3. `test_spec.txt` ➔ `generated/test_spec.json` (input spec for the Z3 test generator in `autosetter/testgen/`; checked against the official samples and retried with the errors if it rejects them) 4. `solution.txt` ➔ `generated/solution.cpp` (optimal C++17 solution) 5. `checker.txt` ➔ `generated/checker.cpp` (uses `testlib.h`) - Runs text inference against a local coding model (`qwen2.5-coder:7b`). diff --git a/README.md b/README.md index ad87a54..43342b0 100644 --- a/README.md +++ b/README.md @@ -7,7 +7,8 @@ statement image / PDF │ Qwen-VL ▼ problem.json ──► statement.md, solution.cpp, validator.cpp, - │ generator.cpp, checker.cpp (one Ollama call each) + │ test_spec.json, checker.cpp (one Ollama call each) + │ test_spec.json ──► Z3 ──► test inputs ▼ validate (compile, generate, validate, solve, check, probe) ▼ @@ -198,6 +199,28 @@ When all C++ files are generated by a single model from the same JSON, they can --- +## Z3 Test Generation + +Instead of writing a generator program, the text model writes `test_spec.json`: a declarative description of the input (integers with bounds, arrays, strings, permutations, matrices, rows of queries, trees, graphs), per-test constraints such as `k <= n`, file-wide constraints such as `sum(n) <= 200000`, and the line layout. The prompt is `autosetter/prompts/test_spec.txt`; the engine is `autosetter/testgen/`. + +- **Z3 solves the integers.** Every integer of every test case in the file is a Z3 variable with its bounds and constraints. They are fixed one at a time: Z3 reports the feasible range, and the seed's strategy picks within it, so relations and sums always stay satisfiable. +- **Builders fill in the bulk.** Arrays, strings, trees and graphs are built directly at the sizes Z3 chose, so max-size tests (e.g. n = 2·10⁵) take about a second. +- **Each seed has a strategy**: 1 = all minimum, 2 = all maximum, then random, small, near-maximum, log-scaled, and so on (`testgen.engine.PLAN`), with matching shapes (sorted arrays, path/star trees, ...). +- **The spec is checked against the official samples.** Samples are parsed with the spec's layout; a spec that rejects a sample is sent back to the model with the exact problem, like a C++ compile error. +- **Every generated input is parsed back and re-checked** against the spec before it reaches the validator. + +Try a spec on its own: +```bash +python -m autosetter.testgen out/generated/test_spec.json 1 10 -o /tmp/tests +python -m autosetter.testgen out/generated/test_spec.json --check-samples out/problem.json +``` + +Settings: `AUTOSETTER_Z3_TIMEOUT_MS` (per solver call, default 10000) and `AUTOSETTER_Z3_MAX_CASES` (most test cases per multi-test file, default 30). + +Limits: the spec describes value ranges and structure, not properties that relate elements to each other ("exactly one pair sums to target", "the answer exists"). The validator still rejects tests that break such guarantees, and the pipeline reports the generator as at fault. Z3-generated tests ship as files in `package/tests/`; Polygon cannot run Z3, so no generator script is emitted. A `generator.py` or `generator.cpp` in `out/generated/` is still used if there is no `test_spec.json`. + +--- + ## Testing ```bash diff --git a/autosetter/cli.py b/autosetter/cli.py index 07dca4e..e23bb6c 100644 --- a/autosetter/cli.py +++ b/autosetter/cli.py @@ -228,7 +228,15 @@ def generate_from_image( # Analyze failure for next iteration targets = [] feedback_context = {} - + + # Files that failed to build (including an unusable test spec) + for name, error in test_report.compilation.errors.items(): + if name in ("validator", "generator", "solution", "checker"): + targets.append(name) + feedback_context[name] = ( + f"Your file could not be used by the pipeline:\n{error[:2000]}" + ) + if "validator rejects official samples" in test_report.diagnosis: targets.append("validator") feedback_context["validator"] = "The validator you generated rejected the official problem samples provided in the problem description." diff --git a/autosetter/config.py b/autosetter/config.py index 54675c7..95dcf78 100644 --- a/autosetter/config.py +++ b/autosetter/config.py @@ -53,6 +53,12 @@ DEFAULT_EXECUTION_TIMEOUT = int(os.environ.get("AUTOSETTER_TIMEOUT", "5")) DEFAULT_COMPILE_TIMEOUT = int(os.environ.get("AUTOSETTER_COMPILE_TIMEOUT", "60")) +# Z3 test generation (autosetter.testgen) +# Per solver call; one test makes a few calls per integer variable. +Z3_TIMEOUT_MS = int(os.environ.get("AUTOSETTER_Z3_TIMEOUT_MS", "10000")) +# Most test cases packed into one multi-test input file (each adds solver variables). +Z3_MAX_CASES = int(os.environ.get("AUTOSETTER_Z3_MAX_CASES", "30")) + # Vision / Image Processing PDF_RENDER_DPI = int(os.environ.get("AUTOSETTER_PDF_DPI", "200")) SUPPORTED_RASTER_EXTENSIONS = {".png", ".jpg", ".jpeg"} diff --git a/autosetter/generator.py b/autosetter/generator.py index 5ee3c4e..28a58a5 100644 --- a/autosetter/generator.py +++ b/autosetter/generator.py @@ -37,6 +37,15 @@ # ───────────────────────────────────────────────────────────────────────────── from autosetter.prompts import PromptError, load_and_render_prompt +from autosetter.extractor import JSONExtractionError, parse_model_json +from autosetter.testgen import ( + SpecError, + TestGenError, + check_samples, + generate_test, + load_spec, +) + # ============================================================================= # Exception Classes @@ -71,6 +80,7 @@ class ArtifactSpec: is_cpp: bool = False # True if the artifact is C++ source code is_testlib: bool = False # True if the artifact requires testlib.h strip_code_fence: bool = True # True to strip ```...``` code fences from LLM output + is_test_spec: bool = False # True for the Z3 test spec (JSON, checked against samples) ARTIFACTS: List[ArtifactSpec] = [ @@ -92,15 +102,17 @@ class ArtifactSpec: is_testlib=True, strip_code_fence=True, ), - # ── Ollama-routed artifact: Test case generator (UNCHANGED) ── - # This artifact continues to use the existing Ollama backend. + # ── Test generator: a declarative input spec that autosetter.testgen + # turns into tests with Z3. Named "generator" so validation feedback and + # self-healing target it like any other generator. ── ArtifactSpec( name="generator", - prompt_template="generator.txt", - output_filename="generator.py", + prompt_template="test_spec.txt", + output_filename="test_spec.json", is_cpp=False, is_testlib=False, - strip_code_fence=True, + strip_code_fence=False, + is_test_spec=True, ), ArtifactSpec( name="solution", @@ -269,6 +281,42 @@ def check_cpp_syntax(code: str, include_dir: Path) -> Tuple[bool, str]: return (False, str(exc)) +def prepare_test_spec(raw_reply: str, json_payload: str) -> Tuple[str, str]: + """ + Parse and verify a model-written test spec. + + Returns (content_to_write, error). The spec must be valid, accept every + official sample input, and generate the smallest and largest tests + (seeds 1 and 2) without error. `error` is empty on success. + """ + try: + data = parse_model_json(raw_reply) + except JSONExtractionError as exc: + return raw_reply.strip() + "\n", f"The reply is not a JSON object: {str(exc)[:500]}" + + content = json.dumps(data, indent=2, ensure_ascii=False) + "\n" + try: + test_spec = load_spec(data) + except SpecError as exc: + return content, f"The spec is invalid:\n{exc}" + + samples = (json.loads(json_payload) or {}).get("samples") or [] + problems = check_samples(test_spec, samples) + if problems: + return content, ( + "The spec rejects the problem's official sample input(s), so it does not " + "describe the input correctly:\n" + "\n".join(f"- {p}" for p in problems) + ) + + for seed in (1, 2): + try: + generate_test(test_spec, seed) + except TestGenError as exc: + return content, f"Generating a test from the spec failed (seed {seed}):\n{exc}" + + return content, "" + + # ============================================================================= # Single Artifact Generation — Ollama Backend (UNCHANGED LOGIC) # ============================================================================= @@ -316,6 +364,7 @@ def generate_single_artifact( ) last_error = "" + base_prompt = current_prompt # ── Retry loop: generate → check syntax → repair if needed ── for attempt in range(max_retries + 1): @@ -331,6 +380,19 @@ def generate_single_artifact( f"Ollama text inference failed while generating '{spec.name}': {exc}" ) from exc + # ── Test spec: validate against the schema and the official samples ── + if spec.is_test_spec: + content, last_error = prepare_test_spec(raw_reply, json_payload) + if not last_error: + break + current_prompt = ( + f"{base_prompt}\n\n=========================================\n" + f"Your previous test spec was rejected:\n{last_error[:2000]}\n\n" + f"--- PREVIOUS SPEC ---\n{content.strip()[:4000]}\n\n" + "Fix every problem listed above. Output ONLY the corrected JSON spec." + ) + continue + # ── Post-process: strip code fences and sanitize C++ headers ── content = strip_code_fence(raw_reply) if spec.strip_code_fence else raw_reply if spec.is_cpp: diff --git a/autosetter/packager.py b/autosetter/packager.py index ba2d37d..3c26917 100644 --- a/autosetter/packager.py +++ b/autosetter/packager.py @@ -11,7 +11,7 @@ │ └── solution.cpp # Reference solution ├── files/ │ ├── validator.cpp # Input validator (testlib.h) -│ ├── generator.cpp # Test generator (testlib.h) +│ ├── test_spec.json # Z3 test spec (or generator.py / generator.cpp) │ ├── checker.cpp # Output checker (testlib.h) │ └── testlib.h # Bundled testlib header ├── tests/ @@ -108,15 +108,20 @@ def build( _log("Packaging testlib files...") files_dir = self.package_dir / "files" files_dir.mkdir(exist_ok=True) - gen_found = False - for gen_name in ("generator.py", "generator.cpp"): + gen_found = "" + for gen_name in ("test_spec.json", "generator.py", "generator.cpp"): gen_src = self.generated_dir / gen_name if gen_src.exists(): shutil.copy2(gen_src, files_dir / gen_name) - gen_found = True + gen_found = gen_name break if not gen_found: - _log(" ⚠️ generator.py / generator.cpp not found, skipping") + _log(" ⚠️ test_spec.json / generator.py / generator.cpp not found, skipping") + + # Only a testlib C++ generator can be re-run by Polygon from a script. + # Z3 (test_spec.json) and Python generators need libraries Polygon does + # not have, so their tests ship as files. + uses_script = gen_found == "generator.cpp" for name in ("validator.cpp", "checker.cpp"): src = self.generated_dir / name @@ -143,7 +148,7 @@ def build( generated_indices = set() report_src = self.tests_dir / "validation_report.json" - if report_src.exists(): + if uses_script and report_src.exists(): try: report_data = json.loads(report_src.read_text(encoding="utf-8")) for tc in report_data.get("test_cases", []): @@ -182,10 +187,12 @@ def build( _log(" ⚠️ Package contains no tests") # 7. Generate script file for Polygon (if tests were generated) - _log("Generating test script...") script_content = "" report_src = self.tests_dir / "validation_report.json" - if report_src.exists(): + if not uses_script: + _log("Tests are shipped as files (no Polygon generator script).") + elif report_src.exists(): + _log("Generating test script...") try: report_data = json.loads(report_src.read_text(encoding="utf-8")) for tc in report_data.get("test_cases", []): @@ -196,7 +203,7 @@ def build( if script_content: (self.package_dir / "script").write_text(script_content, encoding="utf-8") - else: + elif uses_script: _log(" ⚠️ Could not generate script, missing validation report or test_cases") # 8. Generate manifest.json diff --git a/autosetter/pipeline.py b/autosetter/pipeline.py index 1740b7c..4f792a5 100644 --- a/autosetter/pipeline.py +++ b/autosetter/pipeline.py @@ -29,6 +29,19 @@ from autosetter.config import DEFAULT_EXECUTION_TIMEOUT, DEFAULT_NUM_TESTS from autosetter.sandbox import ExecutionResult, SandboxError, SandboxLocalClient +from autosetter.testgen import ( + SpecError, + TestGenError, + TestSpec, + check_samples, + generate_test, + load_spec_file, + strategy_for, +) + +# The Z3 test spec written by the generator stage; preferred over a +# generator.py / generator.cpp program when present. +TEST_SPEC_FILENAME = "test_spec.json" class PipelineError(Exception): @@ -57,6 +70,7 @@ class TestCase: index: int seed: str + strategy: str = "" input_data: str = "" expected_output: str = "" generator_ok: bool = False @@ -181,6 +195,7 @@ def to_dict(self) -> Dict[str, Any]: { "index": tc.index, "seed": tc.seed, + "strategy": tc.strategy, "generator_ok": tc.generator_ok, "validator_ok": tc.validator_ok, "solution_ok": tc.solution_ok, @@ -243,6 +258,7 @@ def __init__( self.time_limit = time_limit self._log = progress_callback or (lambda msg: None) self.samples = samples or [] + self._test_spec: Optional[TestSpec] = None def run(self) -> TestReport: """Execute the full validation pipeline.""" @@ -256,11 +272,13 @@ def run(self) -> TestReport: if not report.compilation.solution: self._log("❌ Solution failed to compile — cannot validate.") + report.diagnosis = self._diagnose(report) report.duration_ms = int((time.monotonic() - start) * 1000) return report if not report.compilation.generator: self._log("❌ Generator failed to compile — cannot produce tests.") + report.diagnosis = self._diagnose(report) report.duration_ms = int((time.monotonic() - start) * 1000) return report @@ -289,7 +307,11 @@ def run(self) -> TestReport: self._log(f"Running {self.num_tests} test cases...") for i in range(1, self.num_tests + 1): seed = str(i) - tc = TestCase(index=i, seed=seed) + tc = TestCase( + index=i, + seed=seed, + strategy=strategy_for(i) if self._test_spec else "", + ) try: # Generate test input @@ -387,6 +409,8 @@ def _compile_all(self) -> CompilationReport: """Compile each generated C++ artifact.""" comp = CompilationReport() for name, (filename, needs_testlib) in self.ARTIFACTS.items(): + if name == "generator" and self._load_test_spec(comp): + continue source = self.generated_dir / filename if not source.exists(): if name == "generator": @@ -432,6 +456,42 @@ def _compile_all(self) -> CompilationReport: return comp + def _load_test_spec(self, comp: CompilationReport) -> bool: + """ + Load test_spec.json as the generator, if there is one. + + The spec plays the role the generator's compilation plays for C++: it + must load, and it must accept the official samples -- a spec that + rejects a known-good input would produce tests for a different problem. + Returns False when no spec exists, so a generator program is used. + """ + spec_path = self.generated_dir / TEST_SPEC_FILENAME + if not spec_path.exists(): + return False + + try: + spec = load_spec_file(spec_path) + problems = check_samples(spec, self.samples) + if problems: + raise SpecError( + "the test spec rejects official samples:\n" + + "\n".join(f"- {p}" for p in problems) + ) + except SpecError as exc: + comp.errors["generator"] = f"Invalid {TEST_SPEC_FILENAME}: {exc}" + error_path = self.generated_dir / "generator_compile_error.txt" + try: + error_path.write_text(comp.errors["generator"], encoding="utf-8") + except OSError as e_write: + self._log(f" ⚠️ Failed to write error note for generator: {e_write}") + self._log(f" ❌ generator: {TEST_SPEC_FILENAME} is unusable – see {error_path.name}") + return True + + self._test_spec = spec + comp.generator = True + self._log(f" ✅ generator: {TEST_SPEC_FILENAME} loaded (Z3 test generation)") + return True + def _check_samples(self) -> List[SampleCheck]: """Validate official problem samples using the validator.""" checks: List[SampleCheck] = [] @@ -540,6 +600,8 @@ def _diagnose(self, report: TestReport) -> str: if not report.compilation.solution: return "The solution does not compile; nothing downstream can be trusted." if not report.compilation.generator: + if (self.generated_dir / TEST_SPEC_FILENAME).exists(): + return f"The generator's {TEST_SPEC_FILENAME} is unusable, so there are no tests." return "The generator does not compile, so there are no tests." if not report.compilation.validator: return "The validator does not compile, so inputs were never validated." @@ -554,6 +616,12 @@ def _diagnose(self, report: TestReport) -> str: ) if not report.compilation.checker: return "The checker does not compile, so outputs were never judged." + if report.total_tests and not report.passed_tests: + first_error = next((tc.error for tc in report.test_cases if tc.error), "") + return ( + f"None of the {report.total_tests} generated tests is usable, so the " + f"checker could not be probed. First failure: {first_error}" + ).strip() if not report.checker_trusted: accepted = [ p.name @@ -573,6 +641,12 @@ def _diagnose(self, report: TestReport) -> str: def _generate_test(self, seed: str) -> str: """Run the generator with a seed to produce test input.""" + if self._test_spec is not None: + try: + return generate_test(self._test_spec, int(seed)).text + except TestGenError as exc: + raise SandboxError(f"Generator failed (seed={seed}): {exc}") from exc + generator_script = self.generated_dir / "generator.py" if generator_script.exists(): result = self.sandbox.run_binary( diff --git a/autosetter/prompts/test_spec.txt b/autosetter/prompts/test_spec.txt new file mode 100644 index 0000000..5c4801a --- /dev/null +++ b/autosetter/prompts/test_spec.txt @@ -0,0 +1,120 @@ +You are an expert competitive-programming problem setter. Your job is to describe the INPUT FORMAT and INPUT CONSTRAINTS of a problem as a JSON "test spec". A Z3-based test generator will read your spec and build valid test inputs from it, and it will also parse the problem's official sample inputs with your spec, so the spec must describe the input exactly. + +You do NOT write code. You do NOT solve the problem. You output one JSON object. + +PROBLEM SPECIFICATION: +{JSON} + +==================================================================== +OUTPUT FORMAT +==================================================================== + +Output ONLY a JSON object with these keys (no comments, no markdown, no prose): + +{ + "multi_test": null OR {"count": "", "min": , "max": }, + "variables": [ , ... ], + "constraints": [ "", ... ], + "global_constraints": [ "", ... ], + "layout": [ [ "", ... ], ... ] +} + +- "multi_test": use it when the first line of the input is the number of test cases (t). "count" is its name, "min"/"max" its bounds. The count is printed automatically on the first line; do NOT list it in "variables" or "layout". Everything else describes ONE test case. +- "variables": every value in one test case, in the order they are read. +- "constraints": extra per-test-case conditions between integer variables, e.g. "k <= n", "a + b <= n", "n % 2 == 0". +- "global_constraints": conditions across ALL test cases of a file. Wrap per-test variables in sum(...), e.g. "sum(n) <= 200000" for "the sum of n over all test cases does not exceed 2*10^5". Use [] if there are none. +- "layout": the lines of ONE test case, top to bottom. Each line is a list of variable names printed on that line, separated by spaces. + +==================================================================== +VARIABLE TYPES +==================================================================== + +Every variable has "name" (a plain identifier) and "type". Bounds, lengths and sizes are integers or expressions over INTEGER variables declared EARLIER (and the multi_test count). + +1. Integer + {"name": "n", "type": "int", "min": 1, "max": 200000} + +2. Array of integers (printed space-separated on one line) + {"name": "a", "type": "array", "length": "n", "min": 1, "max": 1000000000} + Optional: "distinct": true, "order": "ascending" | "descending" | "non_decreasing" | "non_increasing" + +3. Permutation of 1..length (printed on one line) + {"name": "p", "type": "permutation", "length": "n"} + Optional: "indexing": 0 for a permutation of 0..length-1 + +4. String (one token, no spaces) + {"name": "s", "type": "string", "length": "n", "alphabet": "a-z"} + "alphabet" lists allowed characters; ranges allowed: "a-z", "01", "a-zA-Z0-9", "()". + +5. Matrix / grid (one row per line; must be alone on its layout line) + Integers: {"name": "g", "type": "matrix", "rows": "n", "cols": "m", "min": 0, "max": 9} + Characters (each row printed as one string): {"name": "g", "type": "matrix", "rows": "n", "cols": "m", "alphabet": ".#"} + +6. Rows: a list of lines with a fixed number of integers each, such as queries or segments (one row per line; must be alone on its layout line) + {"name": "queries", "type": "rows", "count": "q", + "fields": [ {"name": "l", "min": 1, "max": "n"}, {"name": "r", "min": "l", "max": "n"} ]} + A field's bounds may use integer variables and EARLIER fields of the same row. + +7. Tree on "nodes" vertices (must be alone on its layout line) + {"name": "tree", "type": "tree", "nodes": "n", "format": "edges"} + "format": "edges" = n-1 lines "u v"; "parents" = one line p_2 ... p_n (parent of each vertex 2..n, each smaller than the vertex). + Optional: "indexing": 0, "weight": {"min": 1, "max": 1000000000} (edges become "u v w"). + +8. Graph (m lines "u v"; must be alone on its layout line) + {"name": "g", "type": "graph", "nodes": "n", "edges": "m", + "directed": false, "connected": false, "self_loops": false, "multi_edges": false} + Optional: "indexing": 0, "weight": {"min": 1, "max": 1000000000}. + Put the edge-count limits in "min"/"max" of m, e.g. a connected simple graph needs "min": "n - 1", "max": "min(200000, n*(n-1)//2)". + +==================================================================== +EXPRESSIONS +==================================================================== + +Integers and + - * // % **, comparisons (< <= > >= == !=, chains like "1 <= k <= n"), and / or / not, and min(...), max(...), abs(...). Write 10^9 as 1000000000 or "10**9". Do not use variables that are not integers declared earlier. + +==================================================================== +RULES +==================================================================== + +1. Copy every bound EXACTLY from the constraints. Do not shrink them. +2. Every variable must appear exactly once in "layout". +3. A size (length, rows, cols, count, nodes, edges) may only use integers printed EARLIER in the layout. +4. The official sample inputs MUST be valid under your spec: parse each sample in your head token by token and confirm every value is within its bounds. +5. If the input has several test cases but no count line, or another format this schema cannot express, describe the closest faithful format; never invent values that are not in the input. +6. Do not describe the OUTPUT. + +==================================================================== +EXAMPLES +==================================================================== + +Problem: first line t (1 <= t <= 10^4). Each test: a line with n and k (1 <= k <= n <= 2*10^5), then n integers a_i (1 <= a_i <= 10^9). Sum of n over all tests <= 2*10^5. +{ + "multi_test": {"count": "t", "min": 1, "max": 10000}, + "variables": [ + {"name": "n", "type": "int", "min": 1, "max": 200000}, + {"name": "k", "type": "int", "min": 1, "max": "n"}, + {"name": "a", "type": "array", "length": "n", "min": 1, "max": 1000000000} + ], + "constraints": [], + "global_constraints": ["sum(n) <= 200000"], + "layout": [["n", "k"], ["a"]] +} + +Problem: n vertices and m edges of a connected undirected graph without loops or multiple edges (2 <= n <= 10^5, n-1 <= m <= 2*10^5), edges "u v w" with 1 <= w <= 10^9; then q queries (1 <= q <= 10^5), each "l r" with 1 <= l <= r <= n. +{ + "multi_test": null, + "variables": [ + {"name": "n", "type": "int", "min": 2, "max": 100000}, + {"name": "m", "type": "int", "min": "n - 1", "max": "min(200000, n*(n-1)//2)"}, + {"name": "edges", "type": "graph", "nodes": "n", "edges": "m", "connected": true, + "weight": {"min": 1, "max": 1000000000}}, + {"name": "q", "type": "int", "min": 1, "max": 100000}, + {"name": "queries", "type": "rows", "count": "q", + "fields": [{"name": "l", "min": 1, "max": "n"}, {"name": "r", "min": "l", "max": "n"}]} + ], + "constraints": [], + "global_constraints": [], + "layout": [["n", "m"], ["edges"], ["q"], ["queries"]] +} + +Now output the JSON test spec for the problem above. diff --git a/autosetter/sandbox.py b/autosetter/sandbox.py index fc7e4c4..99a0d24 100644 --- a/autosetter/sandbox.py +++ b/autosetter/sandbox.py @@ -255,6 +255,8 @@ def compile_file( f"Compilation failed for {source.name}:\n{result.stderr}" ) + if not out_binary.exists() and out_binary.with_suffix(".exe").exists(): + return out_binary.with_suffix(".exe") return out_binary def run_binary( @@ -266,6 +268,9 @@ def run_binary( ) -> ExecutionResult: """Execute a compiled binary with optional stdin, args, and timeout.""" binary = Path(binary_path) + if not binary.exists() and binary.with_suffix(".exe").exists(): + # On Windows, g++ appends .exe to the -o path we pass it. + binary = binary.with_suffix(".exe") if not binary.exists(): raise SandboxError(f"Binary not found: {binary}") diff --git a/autosetter/testgen/__init__.py b/autosetter/testgen/__init__.py new file mode 100644 index 0000000..15b43ee --- /dev/null +++ b/autosetter/testgen/__init__.py @@ -0,0 +1,34 @@ +""" +autosetter.testgen +================== +Constraint-driven test generation with the Z3 SMT solver. + +The text model describes the problem's input as a JSON test spec +(`test_spec.json`); this package turns the spec into input files: + +- `spec` loads and validates the spec +- `engine` solves the integer variables with Z3, one strategy per seed +- `builders` builds arrays, strings, permutations, matrices, trees, graphs +- `layout` renders input text, parses it back, and checks it (including + the problem's official samples) +""" + +from autosetter.testgen.builders import TestGenError +from autosetter.testgen.engine import PLAN, GeneratedTest, generate_test, strategy_for +from autosetter.testgen.layout import LayoutError, check_input, check_samples +from autosetter.testgen.spec import SpecError, TestSpec, load_spec, load_spec_file + +__all__ = [ + "PLAN", + "GeneratedTest", + "LayoutError", + "SpecError", + "TestGenError", + "TestSpec", + "check_input", + "check_samples", + "generate_test", + "load_spec", + "load_spec_file", + "strategy_for", +] diff --git a/autosetter/testgen/__main__.py b/autosetter/testgen/__main__.py new file mode 100644 index 0000000..347bd6c --- /dev/null +++ b/autosetter/testgen/__main__.py @@ -0,0 +1,72 @@ +""" +Generate or check tests from a spec without running the whole pipeline. + + python -m autosetter.testgen test_spec.json 7 # print the test for seed 7 + python -m autosetter.testgen test_spec.json 1 10 -o tests # write tests/001.in .. 010.in + python -m autosetter.testgen test_spec.json --check-samples out/problem.json +""" + +from __future__ import annotations + +import argparse +import json +import sys +from pathlib import Path +from typing import List, Optional + +from autosetter.testgen import ( + SpecError, + TestGenError, + check_samples, + generate_test, + load_spec_file, +) + + +def main(argv: Optional[List[str]] = None) -> int: + parser = argparse.ArgumentParser(prog="python -m autosetter.testgen") + parser.add_argument("spec", help="path to test_spec.json") + parser.add_argument("seed", nargs="?", type=int, help="seed, or first seed of a range") + parser.add_argument("last_seed", nargs="?", type=int, help="last seed of a range") + parser.add_argument("-o", "--out-dir", help="write NNN.in files here instead of printing") + parser.add_argument("--check-samples", metavar="PROBLEM_JSON", + help="check the samples in problem.json against the spec") + args = parser.parse_args(argv) + + try: + spec = load_spec_file(args.spec) + except SpecError as exc: + print(f"Invalid spec:\n{exc}", file=sys.stderr) + return 1 + + if args.check_samples: + problem = json.loads(Path(args.check_samples).read_text(encoding="utf-8")) + problems = check_samples(spec, problem.get("samples") or []) + for p in problems: + print(p, file=sys.stderr) + print("samples OK" if not problems else f"{len(problems)} problem(s)") + if problems or args.seed is None: + return 1 if problems else 0 + + if args.seed is None: + parser.error("a seed is required unless --check-samples is given") + + last = args.last_seed if args.last_seed is not None else args.seed + for seed in range(args.seed, last + 1): + try: + test = generate_test(spec, seed) + except TestGenError as exc: + print(f"seed {seed}: {exc}", file=sys.stderr) + return 1 + if args.out_dir: + out = Path(args.out_dir) + out.mkdir(parents=True, exist_ok=True) + (out / f"{seed:03d}.in").write_text(test.text, encoding="utf-8") + print(f"seed {seed}: {test.strategy}, {test.num_cases} case(s)", file=sys.stderr) + else: + sys.stdout.write(test.text) + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/autosetter/testgen/builders.py b/autosetter/testgen/builders.py new file mode 100644 index 0000000..7e714cb --- /dev/null +++ b/autosetter/testgen/builders.py @@ -0,0 +1,381 @@ +""" +autosetter.testgen.builders +=========================== +Constructive builders for everything that is not a scalar integer. + +Z3 decides the sizes (n, m, ...); these builders fill in arrays, strings, +permutations, matrices, rows, trees and graphs of those sizes directly. +Encoding 2*10^5 array elements as solver variables would not finish, while +building them here is instant -- and every result is re-checked against the +spec afterwards anyway. + +`mode` is the per-seed strategy (see `engine.PLAN`) and picks the shape: +all-minimum, all-maximum, sorted, star tree, path tree, and so on. +""" + +from __future__ import annotations + +import math +import random +from typing import Any, Dict, List, Tuple + +from autosetter.testgen.expr import eval_int +from autosetter.testgen.spec import Variable + +MODES = ("min", "max", "random", "small", "near_max", "log") + +# Refuse inputs so large that building or writing them would stall the pipeline. +MAX_ITEMS = 5_000_000 + + +class TestGenError(Exception): + """Raised when a test cannot be generated from an otherwise valid spec.""" + + +def pick_int(lo: int, hi: int, mode: str, rng: random.Random) -> int: + """Choose an integer in [lo, hi] according to `mode`.""" + if lo > hi: + raise TestGenError(f"empty range [{lo}, {hi}]") + if mode == "min": + return lo + if mode == "max": + return hi + if mode == "small": + return rng.randint(lo, min(hi, lo + 9)) + if mode == "near_max": + return rng.randint(max(lo, hi - max(9, (hi - lo) // 20)), hi) + if mode == "log": + span = hi - lo + offset = int(round(math.exp(rng.uniform(0, math.log(span + 1))))) - 1 + return lo + min(max(offset, 0), span) + return rng.randint(lo, hi) + + +def mode_for(strategy: str, rng: random.Random) -> str: + """Per-item mode: fixed strategies stay fixed, 'random' mixes all of them.""" + if strategy == "random": + return rng.choice(("random", "random", "log", "small", "near_max", "min", "max")) + return strategy + + +def _size(var: Variable, key: str, env: Dict[str, Any]) -> int: + value = eval_int(var.exprs[key], env, f"{var.name}.{key}") + if value < 0: + raise TestGenError(f"{var.name}.{key} evaluated to {value}, which is negative") + if value > MAX_ITEMS: + raise TestGenError(f"{var.name}.{key} = {value} exceeds the generator limit {MAX_ITEMS}") + return value + + +def _bounds(var: Variable, env: Dict[str, Any]) -> Tuple[int, int]: + lo = eval_int(var.exprs["min"], env, f"{var.name}.min") + hi = eval_int(var.exprs["max"], env, f"{var.name}.max") + if lo > hi: + raise TestGenError(f"{var.name}: min {lo} is greater than max {hi}") + return lo, hi + + +# --------------------------------------------------------------------------- +# Arrays, permutations, strings, matrices, rows +# --------------------------------------------------------------------------- + +def _int_list(n: int, lo: int, hi: int, distinct: bool, strategy: str, rng: random.Random) -> List[int]: + if distinct and hi - lo + 1 < n: + raise TestGenError(f"cannot pick {n} distinct values from [{lo}, {hi}]") + + if strategy == "min": + pattern = "all_min" + elif strategy == "max": + pattern = "all_max" + elif strategy == "near_max": + pattern = "random_high" + elif strategy == "small": + pattern = "random_small" + elif strategy == "log": + pattern = "random" + else: + pattern = rng.choice(( + "random", "random", "equal", "sorted", "reversed", "extremes", + "few_distinct", "random_high", "random_small", + )) + + if distinct: + if pattern == "all_min": + return list(range(lo, lo + n)) + if pattern == "all_max": + return list(range(hi, hi - n, -1)) + if pattern in ("random_small", "random_high"): + width = min(hi - lo + 1, max(2 * n, n + 10)) + start = lo if pattern == "random_small" else hi - width + 1 + values = rng.sample(range(start, start + width), n) + else: + values = rng.sample(range(lo, hi + 1), n) + if pattern == "sorted": + values.sort() + elif pattern == "reversed": + values.sort(reverse=True) + return values + + if pattern == "all_min": + return [lo] * n + if pattern == "all_max": + return [hi] * n + if pattern == "equal": + return [rng.randint(lo, hi)] * n + if pattern == "extremes": + return [rng.choice((lo, hi)) for _ in range(n)] + if pattern == "few_distinct": + pool = [rng.randint(lo, hi) for _ in range(rng.randint(2, 3))] + return [rng.choice(pool) for _ in range(n)] + if pattern == "random_small": + top = min(hi, lo + 9) + return [rng.randint(lo, top) for _ in range(n)] + if pattern == "random_high": + bottom = max(lo, hi - max(9, (hi - lo) // 20)) + return [rng.randint(bottom, hi) for _ in range(n)] + values = [rng.randint(lo, hi) for _ in range(n)] + if pattern == "sorted": + values.sort() + elif pattern == "reversed": + values.sort(reverse=True) + return values + + +def build_array(var: Variable, env: Dict[str, Any], strategy: str, rng: random.Random) -> List[int]: + n = _size(var, "length", env) + lo, hi = _bounds(var, env) + order = var.opt("order") + distinct = bool(var.opt("distinct")) or order in ("ascending", "descending") + values = _int_list(n, lo, hi, distinct, strategy, rng) + if order in ("ascending", "non_decreasing"): + values.sort() + elif order in ("descending", "non_increasing"): + values.sort(reverse=True) + return values + + +def build_permutation(var: Variable, env: Dict[str, Any], strategy: str, rng: random.Random) -> List[int]: + n = _size(var, "length", env) + values = list(range(var.indexing, var.indexing + n)) + shape = {"min": "identity", "max": "reversed"}.get(strategy) or rng.choice( + ("random", "random", "random", "identity", "reversed", "rotated") + ) + if shape == "reversed": + values.reverse() + elif shape == "rotated" and n: + k = rng.randrange(n) + values = values[k:] + values[:k] + elif shape == "random": + rng.shuffle(values) + return values + + +def _string(length: int, alphabet: str, strategy: str, rng: random.Random) -> str: + if strategy == "min": + return alphabet[0] * length + if strategy == "max": + return alphabet[-1] * length + if strategy == "small": + shape = "few" + elif strategy in ("near_max", "log"): + shape = "random" + else: + shape = rng.choice(("random", "random", "same", "alternating", "few", "palindrome", "blocks")) + + if shape == "same": + return rng.choice(alphabet) * length + if shape == "alternating": + a, b = rng.choice(alphabet), rng.choice(alphabet) + return "".join(a if i % 2 == 0 else b for i in range(length)) + if shape == "few": + pool = rng.sample(alphabet, min(2, len(alphabet))) + return "".join(rng.choice(pool) for _ in range(length)) + if shape == "palindrome": + half = "".join(rng.choice(alphabet) for _ in range((length + 1) // 2)) + return (half + half[::-1][length % 2:])[:length] + if shape == "blocks": + out: List[str] = [] + while len(out) < length: + out.extend(rng.choice(alphabet) * rng.randint(1, max(1, length // 4))) + return "".join(out[:length]) + return "".join(rng.choice(alphabet) for _ in range(length)) + + +def build_string(var: Variable, env: Dict[str, Any], strategy: str, rng: random.Random) -> str: + return _string(_size(var, "length", env), var.alphabet, strategy, rng) + + +def build_matrix(var: Variable, env: Dict[str, Any], strategy: str, rng: random.Random) -> List[Any]: + rows, cols = _size(var, "rows", env), _size(var, "cols", env) + if rows * cols > MAX_ITEMS: + raise TestGenError(f"{var.name}: {rows}x{cols} exceeds the generator limit {MAX_ITEMS}") + if var.alphabet: + return [_string(cols, var.alphabet, strategy, rng) for _ in range(rows)] + lo, hi = _bounds(var, env) + return [_int_list(cols, lo, hi, False, strategy, rng) for _ in range(rows)] + + +def build_rows(var: Variable, env: Dict[str, Any], strategy: str, rng: random.Random) -> List[Tuple[int, ...]]: + count = _size(var, "count", env) + rows: List[Tuple[int, ...]] = [] + for _ in range(count): + row_env = dict(env) + values = [] + mode = mode_for(strategy, rng) + for f in var.fields: + lo = eval_int(f.min, row_env, f"{var.name}.{f.name}.min") + hi = eval_int(f.max, row_env, f"{var.name}.{f.name}.max") + if lo > hi: + raise TestGenError( + f"{var.name}: field '{f.name}' has empty range [{lo}, {hi}] " + f"(earlier fields in this row: {values})" + ) + value = pick_int(lo, hi, mode, rng) + row_env[f.name] = value + values.append(value) + rows.append(tuple(values)) + return rows + + +# --------------------------------------------------------------------------- +# Trees and graphs +# --------------------------------------------------------------------------- + +TREE_SHAPES = ("random", "path", "star", "binary", "caterpillar", "deep_random") + + +def _tree_shape(strategy: str, rng: random.Random) -> str: + return {"min": "path", "max": "star", "near_max": "deep_random", "small": "binary"}.get( + strategy + ) or rng.choice(TREE_SHAPES) + + +def _parents(n: int, shape: str, rng: random.Random) -> List[int]: + """parent[i] for nodes 2..n (1-based), always with parent[i] < i.""" + parents: List[int] = [] + spine = max(1, n // 2) + for i in range(2, n + 1): + if shape == "path": + p = i - 1 + elif shape == "star": + p = 1 + elif shape == "binary": + p = i // 2 + elif shape == "caterpillar": + p = i - 1 if i <= spine else rng.randint(1, spine) + elif shape == "deep_random": + p = rng.randint(max(1, i - 3), i - 1) + else: + p = rng.randint(1, i - 1) + parents.append(p) + return parents + + +def _weights(var: Variable, env: Dict[str, Any], count: int, strategy: str, rng: random.Random) -> List[int]: + if var.weight is None: + return [] + lo = eval_int(var.weight["min"], env, f"{var.name}.weight.min") + hi = eval_int(var.weight["max"], env, f"{var.name}.weight.max") + return [pick_int(lo, hi, mode_for(strategy, rng), rng) for _ in range(count)] + + +def _attach_weights(edges: List[Tuple[int, int]], weights: List[int]) -> List[Tuple[int, ...]]: + if not weights: + return list(edges) + return [edge + (w,) for edge, w in zip(edges, weights)] + + +def build_tree(var: Variable, env: Dict[str, Any], strategy: str, rng: random.Random) -> List[Any]: + n = _size(var, "nodes", env) + shape = _tree_shape(strategy, rng) + parents = _parents(n, shape, rng) + base = var.indexing - 1 # shift from 1-based construction labels + + if var.opt("format", "edges") == "parents": + return [p + base for p in parents] + + labels = list(range(1, n + 1)) + rng.shuffle(labels) + edges: List[Tuple[int, int]] = [] + for child, parent in enumerate(parents, start=2): + u, v = labels[child - 1] + base, labels[parent - 1] + base + edges.append((u, v) if rng.random() < 0.5 else (v, u)) + rng.shuffle(edges) + return _attach_weights(edges, _weights(var, env, len(edges), strategy, rng)) + + +def graph_capacity(n: int, directed: bool, self_loops: bool) -> int: + pairs = n * (n - 1) if directed else n * (n - 1) // 2 + return pairs + (n if self_loops else 0) + + +def build_graph(var: Variable, env: Dict[str, Any], strategy: str, rng: random.Random) -> List[Any]: + n, m = _size(var, "nodes", env), _size(var, "edges", env) + directed = bool(var.opt("directed", False)) + connected = bool(var.opt("connected", False)) + loops = bool(var.opt("self_loops", False)) + multi = bool(var.opt("multi_edges", False)) + base = var.indexing - 1 + + if m and n == 0: + raise TestGenError(f"{var.name}: {m} edges requested on a graph with 0 nodes") + capacity = graph_capacity(n, directed, loops) + if not multi and m > capacity: + raise TestGenError(f"{var.name}: {m} edges exceed the {capacity} possible without multi-edges") + if not multi and not loops and n == 1 and m: + raise TestGenError(f"{var.name}: a single node cannot have edges without self-loops") + if connected and n > 0 and m < n - 1: + raise TestGenError(f"{var.name}: a connected graph on {n} nodes needs at least {n - 1} edges, got {m}") + + def key(u: int, v: int) -> Tuple[int, int]: + return (u, v) if directed else (min(u, v), max(u, v)) + + edges: List[Tuple[int, int]] = [] + seen = set() + + def add(u: int, v: int) -> None: + seen.add(key(u, v)) + if not directed and rng.random() < 0.5: + u, v = v, u + edges.append((u, v)) + + if connected and n > 1: + labels = list(range(1, n + 1)) + rng.shuffle(labels) + for child, parent in enumerate(_parents(n, _tree_shape(strategy, rng), rng), start=2): + add(labels[child - 1], labels[parent - 1]) + + remaining = m - len(edges) + if remaining > 0 and not multi and remaining > (capacity - len(seen)) // 2 and capacity <= 3_000_000: + # Dense: enumerate the free pairs rather than rejection-sample them. + free = [ + (u, v) + for u in range(1, n + 1) + for v in (range(1, n + 1) if directed else range(u, n + 1)) + if (u != v or loops) and key(u, v) not in seen + ] + for u, v in rng.sample(free, remaining): + add(u, v) + else: + while len(edges) < m: + u, v = rng.randint(1, n), rng.randint(1, n) + if u == v and not loops: + continue + if not multi and key(u, v) in seen: + continue + add(u, v) + + rng.shuffle(edges) + shifted = [(u + base, v + base) for u, v in edges] + return _attach_weights(shifted, _weights(var, env, len(shifted), strategy, rng)) + + +BUILDERS = { + "array": build_array, + "permutation": build_permutation, + "string": build_string, + "matrix": build_matrix, + "rows": build_rows, + "tree": build_tree, + "graph": build_graph, +} diff --git a/autosetter/testgen/engine.py b/autosetter/testgen/engine.py new file mode 100644 index 0000000..ca4e2be --- /dev/null +++ b/autosetter/testgen/engine.py @@ -0,0 +1,221 @@ +""" +autosetter.testgen.engine +========================= +Z3-driven test generation from a `TestSpec`. + +For each seed: +1. Pick a strategy from `PLAN` (minimum values, maximum values, random, ...). +2. Build a Z3 model of every integer variable in every test case of the file: + its bounds, the per-test constraints and the cross-test constraints such + as ``sum(n) <= 200000``. +3. Fix the integers one at a time. For each, Z3 reports the smallest and + largest value still feasible given the choices so far; the strategy picks + within that range. Choosing in this order means relations like + ``k <= n`` and file-wide sums always stay satisfiable, and different seeds + give genuinely different values (Z3's own models tend to sit on bounds). +4. Build arrays, strings, trees, ... from the fixed integers (`builders`). +5. Render the input, parse it back and re-check it against the spec. +""" + +from __future__ import annotations + +import random +from dataclasses import dataclass +from typing import Any, Dict, List, Optional + +import z3 + +from autosetter.config import Z3_MAX_CASES, Z3_TIMEOUT_MS +from autosetter.testgen.builders import BUILDERS, TestGenError, mode_for, pick_int +from autosetter.testgen.expr import ExprError, eval_int, evaluate +from autosetter.testgen.layout import LayoutError, check, parse_input, render +from autosetter.testgen.spec import TestSpec + +# Strategy per seed. Seed 1 is the smallest test, seed 2 the largest. +PLAN = ( + "min", "max", "random", "small", "near_max", + "random", "log", "near_max", "random", "random", +) + +# "min" and "max" are deterministic, so repeating them would duplicate a test; +# later passes through PLAN use a randomised neighbour instead. +_REPEAT = {"min": "small", "max": "near_max"} + + +@dataclass +class GeneratedTest: + text: str + strategy: str + num_cases: int + + +def strategy_for(seed: int) -> str: + index = max(abs(seed), 1) - 1 + strategy = PLAN[index % len(PLAN)] + return _REPEAT.get(strategy, strategy) if index >= len(PLAN) else strategy + + +def _as_constraint(value: Any) -> Any: + return z3.BoolVal(value) if isinstance(value, bool) else value + + +def _extreme(opt: z3.Optimize, var: z3.ArithRef, minimize: bool) -> int: + opt.push() + if minimize: + opt.minimize(var) + else: + opt.maximize(var) + result = opt.check() + if result != z3.sat: + opt.pop() + raise TestGenError( + f"Z3 could not {'minimise' if minimize else 'maximise'} {var} ({result}); " + "the constraints may be too complex or the solver timed out" + ) + value = opt.model().eval(var, model_completion=True).as_long() + opt.pop() + return value + + +def _choose(opt: z3.Optimize, var: z3.ArithRef, mode: str, rng: random.Random) -> int: + if mode in ("min", "max"): + return _extreme(opt, var, minimize=(mode == "min")) + + lo = _extreme(opt, var, minimize=True) + hi = _extreme(opt, var, minimize=False) + target = pick_int(lo, hi, mode, rng) + if target in (lo, hi): + return target + + opt.push() + opt.add(var == target) + feasible = opt.check() == z3.sat + opt.pop() + if feasible: + return target + + # The feasible set has holes (e.g. "n is even"): take the closest value + # below the target, which exists because `lo` is feasible. + opt.push() + opt.add(var <= target) + value = _extreme(opt, var, minimize=False) + opt.pop() + return value + + +def _solve_for_count( + spec: TestSpec, + num_cases: int, + strategy: str, + rng: random.Random, + timeout_ms: int, +) -> Optional[List[Dict[str, int]]]: + """Fix every integer of `num_cases` test cases; None if infeasible.""" + opt = z3.Optimize() + opt.set("timeout", timeout_ms) + + top: Dict[str, Any] = {spec.multi.count: num_cases} if spec.multi else {} + case_envs: List[Dict[str, Any]] = [] + for index in range(num_cases): + env = dict(top) + for var in spec.variables: + if var.type != "int": + continue + symbol = z3.Int(f"{var.name}#{index}") + opt.add(_as_constraint(symbol >= evaluate(var.exprs["min"], env))) + opt.add(_as_constraint(symbol <= evaluate(var.exprs["max"], env))) + env[var.name] = symbol + for node in spec.constraints: + opt.add(_as_constraint(evaluate(node, env))) + case_envs.append(env) + + def sum_fn(arg: Any) -> Any: + return sum((evaluate(arg, env) for env in case_envs), 0) + + for node in spec.global_constraints: + opt.add(_as_constraint(evaluate(node, top, sum_fn))) + + result = opt.check() + if result == z3.unsat: + return None + if result != z3.sat: + raise TestGenError(f"Z3 returned '{result}' (timeout {timeout_ms} ms) for {num_cases} test case(s)") + + solved: List[Dict[str, int]] = [] + for env in case_envs: + values: Dict[str, int] = {} + for name in spec.int_names: + symbol = env[name] + value = _choose(opt, symbol, mode_for(strategy, rng), rng) + opt.add(symbol == value) + values[name] = value + solved.append(values) + return solved + + +def solve_integers( + spec: TestSpec, + strategy: str, + rng: random.Random, + timeout_ms: int = Z3_TIMEOUT_MS, + max_cases: int = Z3_MAX_CASES, +) -> tuple[Dict[str, int], List[Dict[str, int]]]: + """Choose the number of test cases and every integer in each of them.""" + if not spec.multi: + solved = _solve_for_count(spec, 1, strategy, rng, timeout_ms) + if solved is None: + raise TestGenError("the integer constraints are unsatisfiable") + return {}, solved + + lo = eval_int(spec.multi.min, {}, "multi_test.min") + hi = eval_int(spec.multi.max, {}, "multi_test.max") + if lo > hi: + raise TestGenError(f"multi_test.min {lo} is greater than multi_test.max {hi}") + # Every test case adds solver variables, so cap the count per file. + cap = min(hi, max(lo, max_cases)) + count = pick_int(lo, cap, "small" if strategy == "log" else mode_for(strategy, rng), rng) + + # Too many test cases can make a file-wide sum infeasible; back off. + while True: + solved = _solve_for_count(spec, count, strategy, rng, timeout_ms) + if solved is not None: + return {spec.multi.count: count}, solved + if count == lo: + raise TestGenError("the integer constraints are unsatisfiable") + count = max(lo, count // 2) + + +def generate_test( + spec: TestSpec, + seed: int, + timeout_ms: int = Z3_TIMEOUT_MS, + max_cases: int = Z3_MAX_CASES, +) -> GeneratedTest: + """Generate one input file for `seed`. Deterministic for a given spec and seed.""" + rng = random.Random(seed) + strategy = strategy_for(seed) + + try: + top, int_cases = solve_integers(spec, strategy, rng, timeout_ms, max_cases) + cases: List[Dict[str, Any]] = [] + for ints in int_cases: + env: Dict[str, Any] = dict(top) + env.update(ints) + for var in spec.variables: + if var.type != "int": + env[var.name] = BUILDERS[var.type](var, env, strategy, rng) + cases.append(env) + except ExprError as exc: + raise TestGenError(str(exc)) from exc + + text = render(spec, top, cases) + + try: + parsed_top, parsed_cases = parse_input(spec, text) + except LayoutError as exc: + raise TestGenError(f"generated input does not match the layout: {exc}") from exc + problems = check(spec, parsed_top, parsed_cases) + if problems: + raise TestGenError("generated input breaks the spec:\n" + "\n".join(problems[:10])) + + return GeneratedTest(text=text, strategy=strategy, num_cases=len(cases)) diff --git a/autosetter/testgen/expr.py b/autosetter/testgen/expr.py new file mode 100644 index 0000000..de5ae88 --- /dev/null +++ b/autosetter/testgen/expr.py @@ -0,0 +1,255 @@ +""" +autosetter.testgen.expr +======================= +Safe arithmetic/boolean expressions for test specs. + +Spec bounds and constraints are short Python-like expressions such as +``"n - 1"``, ``"k <= n"`` or ``"sum(n) <= 200000"``. They are parsed with +`ast` and evaluated by a small interpreter that accepts only arithmetic, +comparisons, boolean logic and a few functions -- never `eval`, since specs +are written by a model. + +The same interpreter runs over two kinds of values: +- plain Python ints, when checking a concrete test or computing a length; +- Z3 integer expressions, when building the constraint model to solve. +""" + +from __future__ import annotations + +import ast +from typing import Any, Callable, Dict, Mapping, Optional, Set + +import z3 + + +class ExprError(Exception): + """Raised when an expression is malformed, unsupported, or cannot be evaluated.""" + + +FUNCTIONS = {"min", "max", "abs", "sum"} + +_ALLOWED_NODES = ( + ast.Expression, ast.Constant, ast.Name, ast.Load, ast.BinOp, ast.UnaryOp, + ast.BoolOp, ast.Compare, ast.Call, + ast.Add, ast.Sub, ast.Mult, ast.FloorDiv, ast.Div, ast.Mod, ast.Pow, + ast.USub, ast.UAdd, ast.Not, ast.And, ast.Or, + ast.Eq, ast.NotEq, ast.Lt, ast.LtE, ast.Gt, ast.GtE, +) + +# Keeps `10**9`-style constants cheap while refusing absurd powers. +_MAX_EXPONENT = 64 + + +def parse_expr(source: Any) -> ast.expr: + """Parse a bound or constraint (a number or an expression string) into an AST.""" + if isinstance(source, bool): + raise ExprError(f"expected a number or expression, got boolean {source!r}") + if isinstance(source, (int, float)): + source = repr(_as_int(source)) + if not isinstance(source, str) or not source.strip(): + raise ExprError(f"expected a number or expression, got {source!r}") + + text = source.strip() + try: + tree = ast.parse(text, mode="eval") + except SyntaxError as exc: + raise ExprError(f"cannot parse expression {text!r}: {exc.msg}") from exc + + for node in ast.walk(tree): + if not isinstance(node, _ALLOWED_NODES): + raise ExprError( + f"unsupported syntax {type(node).__name__} in expression {text!r}" + ) + if isinstance(node, ast.Call): + if not isinstance(node.func, ast.Name) or node.func.id not in FUNCTIONS: + raise ExprError( + f"unsupported function call in {text!r}; allowed: {sorted(FUNCTIONS)}" + ) + if node.keywords: + raise ExprError(f"keyword arguments are not allowed in {text!r}") + if isinstance(node, ast.Constant): + _as_int(node.value) # validates the literal + return tree.body + + +def names_in(node: ast.AST, *, include_sum_args: bool = True) -> Set[str]: + """Variable names referenced by an expression (function names excluded).""" + found: Set[str] = set() + + def visit(n: ast.AST) -> None: + if isinstance(n, ast.Call): + if not include_sum_args and n.func.id == "sum": # type: ignore[attr-defined] + return + for arg in n.args: + visit(arg) + return + if isinstance(n, ast.Name): + found.add(n.id) + return + for child in ast.iter_child_nodes(n): + visit(child) + + visit(node) + return found + + +def sum_args(node: ast.AST) -> list[ast.expr]: + """The argument expressions of every ``sum(...)`` call inside `node`.""" + return [ + n.args[0] + for n in ast.walk(node) + if isinstance(n, ast.Call) and n.func.id == "sum" and n.args # type: ignore[attr-defined] + ] + + +def _as_int(value: Any) -> int: + if isinstance(value, bool): + return int(value) + if isinstance(value, int): + return value + if isinstance(value, float) and value.is_integer(): + return int(value) + raise ExprError(f"only integer literals are supported, got {value!r}") + + +def _is_z3(value: Any) -> bool: + return isinstance(value, z3.ExprRef) + + +def _and(values: list) -> Any: + if any(_is_z3(v) for v in values): + return z3.And(*values) + return all(values) + + +def _or(values: list) -> Any: + if any(_is_z3(v) for v in values): + return z3.Or(*values) + return any(values) + + +def _min(values: list) -> Any: + result = values[0] + for v in values[1:]: + result = z3.If(v < result, v, result) if _is_z3(v) or _is_z3(result) else min(result, v) + return result + + +def _max(values: list) -> Any: + result = values[0] + for v in values[1:]: + result = z3.If(v > result, v, result) if _is_z3(v) or _is_z3(result) else max(result, v) + return result + + +_COMPARE: Dict[type, Callable[[Any, Any], Any]] = { + ast.Eq: lambda a, b: a == b, + ast.NotEq: lambda a, b: a != b, + ast.Lt: lambda a, b: a < b, + ast.LtE: lambda a, b: a <= b, + ast.Gt: lambda a, b: a > b, + ast.GtE: lambda a, b: a >= b, +} + + +def evaluate( + node: ast.AST, + env: Mapping[str, Any], + sum_fn: Optional[Callable[[ast.expr], Any]] = None, +) -> Any: + """ + Evaluate an expression AST. + + `env` maps names to ints or Z3 expressions. `sum_fn` implements + ``sum(expr)`` (a sum over all test cases in the file); without it, + ``sum`` is an error. + """ + if isinstance(node, ast.Constant): + return _as_int(node.value) + + if isinstance(node, ast.Name): + if node.id not in env: + raise ExprError(f"unknown variable '{node.id}'") + return env[node.id] + + if isinstance(node, ast.UnaryOp): + operand = evaluate(node.operand, env, sum_fn) + if isinstance(node.op, ast.USub): + return -operand + if isinstance(node.op, ast.UAdd): + return operand + return z3.Not(operand) if _is_z3(operand) else not operand + + if isinstance(node, ast.BinOp): + left = evaluate(node.left, env, sum_fn) + right = evaluate(node.right, env, sum_fn) + op = node.op + if isinstance(op, ast.Add): + return left + right + if isinstance(op, ast.Sub): + return left - right + if isinstance(op, ast.Mult): + return left * right + if isinstance(op, (ast.FloorDiv, ast.Div)): + if _is_z3(left) or _is_z3(right): + return left / right # integer division for Z3 Int sorts + if right == 0: + raise ExprError("division by zero") + return left // right + if isinstance(op, ast.Mod): + if not (_is_z3(left) or _is_z3(right)) and right == 0: + raise ExprError("modulo by zero") + return left % right + if isinstance(op, ast.Pow): + if _is_z3(left) or _is_z3(right): + raise ExprError("'**' is only supported between constants (e.g. 10**9)") + if right < 0 or right > _MAX_EXPONENT: + raise ExprError(f"exponent {right} is out of the supported range") + return left ** right + raise ExprError(f"unsupported operator {type(op).__name__}") + + if isinstance(node, ast.BoolOp): + values = [evaluate(v, env, sum_fn) for v in node.values] + return _and(values) if isinstance(node.op, ast.And) else _or(values) + + if isinstance(node, ast.Compare): + left = evaluate(node.left, env, sum_fn) + parts = [] + for op, comparator in zip(node.ops, node.comparators): + right = evaluate(comparator, env, sum_fn) + parts.append(_COMPARE[type(op)](left, right)) + left = right + return parts[0] if len(parts) == 1 else _and(parts) + + if isinstance(node, ast.Call): + name = node.func.id # type: ignore[attr-defined] + if name == "sum": + if sum_fn is None: + raise ExprError("sum(...) is only allowed in global_constraints") + if len(node.args) != 1: + raise ExprError("sum(...) takes exactly one argument") + return sum_fn(node.args[0]) + args = [evaluate(a, env, sum_fn) for a in node.args] + if not args: + raise ExprError(f"{name}() needs at least one argument") + if name == "abs": + if len(args) != 1: + raise ExprError("abs() takes exactly one argument") + x = args[0] + return z3.If(x < 0, -x, x) if _is_z3(x) else abs(x) + return _min(args) if name == "min" else _max(args) + + raise ExprError(f"unsupported syntax {type(node).__name__}") + + +def eval_int(node: ast.AST, env: Mapping[str, Any], what: str = "expression") -> int: + """Evaluate to a concrete Python int, with a readable error otherwise.""" + value = evaluate(node, env) + if _is_z3(value) or isinstance(value, bool) or not isinstance(value, int): + raise ExprError(f"{what} did not evaluate to an integer") + return value + + +def source_of(node: ast.AST) -> str: + """Round-trip an AST back to readable source for messages.""" + return ast.unparse(node) diff --git a/autosetter/testgen/layout.py b/autosetter/testgen/layout.py new file mode 100644 index 0000000..7b4af9d --- /dev/null +++ b/autosetter/testgen/layout.py @@ -0,0 +1,351 @@ +""" +autosetter.testgen.layout +========================= +Render values to input text, read input text back into values, and check +values against the spec. + +Reading is what makes the spec trustworthy: the problem's official samples +are parsed with the spec's layout and checked against its constraints. A spec +that rejects an official sample is wrong, and is sent back to the model +before any test is generated. Generated tests go through the same +read-and-check round trip, so a builder bug can never ship an invalid test. +""" + +from __future__ import annotations + +import re +from typing import Any, Dict, List, Tuple + +from autosetter.testgen.builders import MAX_ITEMS +from autosetter.testgen.expr import ExprError, eval_int, evaluate, source_of +from autosetter.testgen.spec import BLOCK_TYPES, TestSpec, Variable + +Values = Dict[str, Any] + +_INT_RE = re.compile(r"^[+-]?\d+$") + + +class LayoutError(Exception): + """Raised when input text does not match the spec's layout.""" + + +# --------------------------------------------------------------------------- +# Rendering +# --------------------------------------------------------------------------- + +def _inline(value: Any) -> str: + if isinstance(value, list): + return " ".join(map(str, value)) + return str(value) + + +def _block(var: Variable, value: Any) -> List[str]: + if var.type == "tree" and var.opt("format", "edges") == "parents": + return [" ".join(map(str, value))] + if var.type == "matrix" and var.alphabet: + return list(value) + return [" ".join(map(str, row)) for row in value] + + +def render(spec: TestSpec, top: Values, cases: List[Values]) -> str: + """Format one input file.""" + lines: List[str] = [] + if spec.multi: + lines.append(str(top[spec.multi.count])) + for values in cases: + for line in spec.layout: + first = spec.by_name[line[0]] + if len(line) == 1 and first.type in BLOCK_TYPES: + lines.extend(_block(first, values[first.name])) + else: + lines.append(" ".join(_inline(values[name]) for name in line).strip()) + return "\n".join(lines) + "\n" + + +# --------------------------------------------------------------------------- +# Parsing +# --------------------------------------------------------------------------- + +class _Tokens: + def __init__(self, text: str) -> None: + self.tokens = text.split() + self.pos = 0 + + def token(self, what: str) -> str: + if self.pos >= len(self.tokens): + raise LayoutError(f"input ended while reading {what}") + self.pos += 1 + return self.tokens[self.pos - 1] + + def int(self, what: str) -> int: + tok = self.token(what) + if not _INT_RE.match(tok): + raise LayoutError(f"expected an integer for {what}, got {tok!r}") + return int(tok) + + def ints(self, count: int, what: str) -> List[int]: + return [self.int(what) for _ in range(count)] + + +def _size(var: Variable, key: str, env: Values) -> int: + try: + value = eval_int(var.exprs[key], env, f"{var.name}.{key}") + except ExprError as exc: + raise LayoutError(str(exc)) from exc + if value < 0 or value > MAX_ITEMS: + raise LayoutError(f"{var.name}.{key} evaluated to {value}, which is not a usable size") + return value + + +def _read(var: Variable, env: Values, tok: _Tokens) -> Any: + name, vtype = var.name, var.type + if vtype == "int": + return tok.int(name) + if vtype in ("array", "permutation"): + return tok.ints(_size(var, "length", env), f"elements of {name}") + if vtype == "string": + return tok.token(name) if _size(var, "length", env) > 0 else "" + if vtype == "matrix": + rows, cols = _size(var, "rows", env), _size(var, "cols", env) + if var.alphabet: + return [tok.token(f"a row of {name}") for _ in range(rows)] + return [tok.ints(cols, f"a row of {name}") for _ in range(rows)] + if vtype == "rows": + count, width = _size(var, "count", env), len(var.fields) + return [tuple(tok.ints(width, f"a row of {name}")) for _ in range(count)] + width = 3 if var.weight is not None else 2 + if vtype == "tree": + n = _size(var, "nodes", env) + if var.opt("format", "edges") == "parents": + return tok.ints(max(n - 1, 0), f"parents in {name}") + return [tuple(tok.ints(width, f"an edge of {name}")) for _ in range(max(n - 1, 0))] + m = _size(var, "edges", env) + return [tuple(tok.ints(width, f"an edge of {name}")) for _ in range(m)] + + +def parse_input(spec: TestSpec, text: str) -> Tuple[Values, List[Values]]: + """Read an input file into (top-level values, per-test-case values).""" + tok = _Tokens(text) + top: Values = {} + num_cases = 1 + if spec.multi: + num_cases = tok.int(spec.multi.count) + if num_cases < 0 or num_cases > MAX_ITEMS: + raise LayoutError(f"{spec.multi.count} = {num_cases} is not a usable test count") + top[spec.multi.count] = num_cases + + cases: List[Values] = [] + for _ in range(num_cases): + env = dict(top) + for line in spec.layout: + for name in line: + env[name] = _read(spec.by_name[name], env, tok) + cases.append(env) + + extra = len(tok.tokens) - tok.pos + if extra: + raise LayoutError( + f"{extra} unread token(s) after the last test case, starting with {tok.tokens[tok.pos]!r}" + ) + return top, cases + + +# --------------------------------------------------------------------------- +# Checking +# --------------------------------------------------------------------------- + +def _in_range(value: int, lo: int, hi: int) -> bool: + return lo <= value <= hi + + +class _DSU: + def __init__(self) -> None: + self.parent: Dict[int, int] = {} + + def find(self, x: int) -> int: + self.parent.setdefault(x, x) + while self.parent[x] != x: + self.parent[x] = self.parent[self.parent[x]] + x = self.parent[x] + return x + + def union(self, a: int, b: int) -> bool: + ra, rb = self.find(a), self.find(b) + if ra == rb: + return False + self.parent[ra] = rb + return True + + +def _check_weights(var: Variable, env: Values, edges: List[Tuple[int, ...]]) -> List[str]: + if var.weight is None: + return [] + lo = eval_int(var.weight["min"], env) + hi = eval_int(var.weight["max"], env) + for edge in edges: + if not _in_range(edge[2], lo, hi): + return [f"{var.name}: edge weight {edge[2]} is outside [{lo}, {hi}]"] + return [] + + +def _check_var(var: Variable, env: Values) -> List[str]: + name, value = var.name, env[var.name] + + if var.type == "int": + lo, hi = eval_int(var.exprs["min"], env), eval_int(var.exprs["max"], env) + return [] if _in_range(value, lo, hi) else [f"{name} = {value} is outside [{lo}, {hi}]"] + + if var.type == "array": + n = eval_int(var.exprs["length"], env) + lo, hi = eval_int(var.exprs["min"], env), eval_int(var.exprs["max"], env) + if len(value) != n: + return [f"{name} has {len(value)} elements, expected {n}"] + for i, x in enumerate(value): + if not _in_range(x, lo, hi): + return [f"{name}[{i}] = {x} is outside [{lo}, {hi}]"] + order = var.opt("order") + if var.opt("distinct") and len(set(value)) != len(value): + return [f"{name} must have distinct elements"] + pairs = list(zip(value, value[1:])) + ok = { + "ascending": all(a < b for a, b in pairs), + "descending": all(a > b for a, b in pairs), + "non_decreasing": all(a <= b for a, b in pairs), + "non_increasing": all(a >= b for a, b in pairs), + }.get(order, True) + return [] if ok else [f"{name} is not {order.replace('_', '-')}"] + + if var.type == "permutation": + n = eval_int(var.exprs["length"], env) + expected = list(range(var.indexing, var.indexing + n)) + return [] if sorted(value) == expected else [ + f"{name} is not a permutation of {var.indexing}..{var.indexing + n - 1}" + ] + + if var.type == "string": + n = eval_int(var.exprs["length"], env) + if len(value) != n: + return [f"{name} has length {len(value)}, expected {n}"] + bad = set(value) - set(var.alphabet) + return [f"{name} contains characters {sorted(bad)} outside its alphabet"] if bad else [] + + if var.type == "matrix": + rows, cols = eval_int(var.exprs["rows"], env), eval_int(var.exprs["cols"], env) + if len(value) != rows or any(len(r) != cols for r in value): + return [f"{name} is not {rows}x{cols}"] + if var.alphabet: + bad = set("".join(value)) - set(var.alphabet) + return [f"{name} contains characters {sorted(bad)} outside its alphabet"] if bad else [] + lo, hi = eval_int(var.exprs["min"], env), eval_int(var.exprs["max"], env) + for r, row in enumerate(value): + for c, x in enumerate(row): + if not _in_range(x, lo, hi): + return [f"{name}[{r}][{c}] = {x} is outside [{lo}, {hi}]"] + return [] + + if var.type == "rows": + for r, row in enumerate(value): + row_env = dict(env) + for f, x in zip(var.fields, row): + lo, hi = eval_int(f.min, row_env), eval_int(f.max, row_env) + if not _in_range(x, lo, hi): + return [f"{name} row {r + 1}: {f.name} = {x} is outside [{lo}, {hi}]"] + row_env[f.name] = x + return [] + + lo_label = var.indexing + if var.type == "tree": + n = eval_int(var.exprs["nodes"], env) + hi_label = lo_label + n - 1 + if var.opt("format", "edges") == "parents": + dsu = _DSU() + for offset, p in enumerate(value): + child = lo_label + 1 + offset + if not _in_range(p, lo_label, hi_label): + return [f"{name}: parent {p} of node {child} is not a node"] + if var.opt("parent_before_child", True) and p >= child: + return [f"{name}: parent {p} of node {child} is not smaller than the node"] + if not dsu.union(child, p): + return [f"{name}: parent pointers form a cycle"] + return [] + dsu = _DSU() + for edge in value: + u, v = edge[0], edge[1] + if not (_in_range(u, lo_label, hi_label) and _in_range(v, lo_label, hi_label)): + return [f"{name}: edge ({u}, {v}) uses a node outside {lo_label}..{hi_label}"] + if not dsu.union(u, v): + return [f"{name}: edge ({u}, {v}) closes a cycle, so it is not a tree"] + return _check_weights(var, env, value) + + # graph + n = eval_int(var.exprs["nodes"], env) + hi_label = lo_label + n - 1 + directed = bool(var.opt("directed", False)) + seen = set() + dsu = _DSU() + for edge in value: + u, v = edge[0], edge[1] + if not (_in_range(u, lo_label, hi_label) and _in_range(v, lo_label, hi_label)): + return [f"{name}: edge ({u}, {v}) uses a node outside {lo_label}..{hi_label}"] + if u == v and not var.opt("self_loops", False): + return [f"{name}: self-loop ({u}, {v}) is not allowed"] + key = (u, v) if directed else (min(u, v), max(u, v)) + if key in seen and not var.opt("multi_edges", False): + return [f"{name}: repeated edge ({u}, {v}) is not allowed"] + seen.add(key) + dsu.union(u, v) + if var.opt("connected", False) and n > 1: + roots = {dsu.find(x) for x in range(lo_label, hi_label + 1)} + if len(roots) > 1: + return [f"{name} is not connected"] + return _check_weights(var, env, value) + + +def check(spec: TestSpec, top: Values, cases: List[Values]) -> List[str]: + """Every way the values break the spec (empty when they satisfy it).""" + problems: List[str] = [] + try: + if spec.multi: + count = top[spec.multi.count] + lo, hi = eval_int(spec.multi.min, {}), eval_int(spec.multi.max, {}) + if not _in_range(count, lo, hi): + problems.append(f"{spec.multi.count} = {count} is outside [{lo}, {hi}]") + + for index, env in enumerate(cases, start=1): + prefix = f"test case {index}: " if spec.multi else "" + for var in spec.variables: + problems.extend(prefix + p for p in _check_var(var, env)) + for node in spec.constraints: + if not evaluate(node, env): + problems.append(f"{prefix}constraint '{source_of(node)}' is violated") + + for node in spec.global_constraints: + def sum_fn(arg: Any) -> int: + return sum(evaluate(arg, env) for env in cases) + + if not evaluate(node, top, sum_fn): + problems.append(f"global constraint '{source_of(node)}' is violated") + except ExprError as exc: + problems.append(str(exc)) + return problems + + +def check_input(spec: TestSpec, text: str) -> List[str]: + """Parse and check an input file; returns the problems found.""" + try: + top, cases = parse_input(spec, text) + except LayoutError as exc: + return [f"does not match the layout: {exc}"] + return check(spec, top, cases) + + +def check_samples(spec: TestSpec, samples: List[Dict[str, Any]]) -> List[str]: + """Check the problem's official sample inputs against the spec.""" + problems: List[str] = [] + for index, sample in enumerate(samples or [], start=1): + text = str((sample or {}).get("input") or "") + if not text.strip(): + continue + for problem in check_input(spec, text)[:5]: + problems.append(f"sample {index}: {problem}") + return problems diff --git a/autosetter/testgen/spec.py b/autosetter/testgen/spec.py new file mode 100644 index 0000000..7d6121e --- /dev/null +++ b/autosetter/testgen/spec.py @@ -0,0 +1,408 @@ +""" +autosetter.testgen.spec +======================= +The test spec: a declarative description of a problem's input. + +The text model writes the spec (``test_spec.json``) from ``problem.json``; +this module loads it and rejects anything malformed with messages specific +enough to feed back to the model for a retry. + +Shape:: + + { + "multi_test": {"count": "t", "min": 1, "max": 10000}, # or null + "variables": [ # one test case + {"name": "n", "type": "int", "min": 1, "max": 200000}, + {"name": "k", "type": "int", "min": 1, "max": "n"}, + {"name": "a", "type": "array", "length": "n", "min": 1, "max": 1000000000} + ], + "constraints": ["k <= n"], # per test case + "global_constraints": ["sum(n) <= 200000"], # across the file + "layout": [["n", "k"], ["a"]] # lines of one test case + } + +Integer variables are solved by Z3; every other type is built directly from +the solved integers (see `builders`). +""" + +from __future__ import annotations + +import ast +import json +import keyword +from dataclasses import dataclass, field +from pathlib import Path +from typing import Any, Dict, List, Optional, Set + +from autosetter.testgen.expr import FUNCTIONS, ExprError, names_in, parse_expr, sum_args + +TYPES = {"int", "array", "permutation", "string", "matrix", "rows", "tree", "graph"} + +# Types rendered over several lines; they must sit alone on a layout line. +BLOCK_TYPES = {"matrix", "rows", "tree", "graph"} + +ORDERS = {"ascending", "descending", "non_decreasing", "non_increasing"} + +# Size fields per type: what must be known before the value can be read. +SIZE_FIELDS = { + "array": ("length",), + "permutation": ("length",), + "string": ("length",), + "matrix": ("rows", "cols"), + "rows": ("count",), + "tree": ("nodes",), + "graph": ("nodes", "edges"), +} + +DEFAULT_ALPHABET = "a-z" + + +class SpecError(Exception): + """Raised when a test spec is malformed. The message lists every problem found.""" + + +@dataclass +class RowField: + name: str + min: ast.expr + max: ast.expr + + +@dataclass +class Variable: + name: str + type: str + raw: Dict[str, Any] + exprs: Dict[str, ast.expr] = field(default_factory=dict) + fields: List[RowField] = field(default_factory=list) + alphabet: str = "" + weight: Optional[Dict[str, ast.expr]] = None + + def opt(self, key: str, default: Any = None) -> Any: + return self.raw.get(key, default) + + @property + def indexing(self) -> int: + return int(self.raw.get("indexing", 1)) + + +@dataclass +class MultiTest: + count: str + min: ast.expr + max: ast.expr + + +@dataclass +class TestSpec: + raw: Dict[str, Any] + multi: Optional[MultiTest] + variables: List[Variable] + constraints: List[ast.expr] + global_constraints: List[ast.expr] + layout: List[List[str]] + + def __post_init__(self) -> None: + self.by_name: Dict[str, Variable] = {v.name: v for v in self.variables} + + @property + def int_names(self) -> List[str]: + return [v.name for v in self.variables if v.type == "int"] + + def to_json(self) -> str: + return json.dumps(self.raw, indent=2, ensure_ascii=False) + "\n" + + +def expand_alphabet(spec: str) -> str: + """Expand ``"a-z0-9"``-style character sets; a literal '-' may lead or trail.""" + chars: List[str] = [] + i = 0 + while i < len(spec): + if i + 2 < len(spec) and spec[i + 1] == "-": + lo, hi = spec[i], spec[i + 2] + if ord(lo) > ord(hi): + raise SpecError(f"invalid alphabet range '{lo}-{hi}'") + chars.extend(chr(c) for c in range(ord(lo), ord(hi) + 1)) + i += 3 + else: + chars.append(spec[i]) + i += 1 + unique = "".join(dict.fromkeys(chars)) + if not unique or any(c.isspace() for c in unique): + raise SpecError(f"alphabet {spec!r} must be non-empty and contain no whitespace") + return unique + + +class _Collector: + """Accumulates validation errors so the model sees all of them at once.""" + + def __init__(self) -> None: + self.errors: List[str] = [] + + def add(self, message: str) -> None: + self.errors.append(message) + + def expr(self, where: str, source: Any, allowed: Set[str]) -> Optional[ast.expr]: + try: + node = parse_expr(source) + except ExprError as exc: + self.add(f"{where}: {exc}") + return None + if sum_args(node): + self.add(f"{where}: sum(...) is only allowed in global_constraints") + return None + unknown = names_in(node) - allowed + if unknown: + self.add( + f"{where}: refers to {sorted(unknown)}, which are not integer " + f"variables defined before it (available: {sorted(allowed) or 'none'})" + ) + return None + return node + + +def _check_name(c: _Collector, name: Any, where: str, taken: Set[str]) -> bool: + if not isinstance(name, str) or not name.isidentifier() or keyword.iskeyword(name): + c.add(f"{where}: name {name!r} must be a plain identifier") + return False + if name in FUNCTIONS: + c.add(f"{where}: name '{name}' is reserved") + return False + if name in taken: + c.add(f"{where}: name '{name}' is defined more than once") + return False + taken.add(name) + return True + + +def _bool_opt(c: _Collector, var: Dict[str, Any], key: str, where: str) -> None: + if key in var and not isinstance(var[key], bool): + c.add(f"{where}: '{key}' must be true or false") + + +def load_spec(source: str | Dict[str, Any]) -> TestSpec: + """Parse and validate a spec given as a dict or JSON text. Raises SpecError.""" + if isinstance(source, str): + try: + source = json.loads(source) + except json.JSONDecodeError as exc: + raise SpecError(f"spec is not valid JSON: {exc}") from exc + if not isinstance(source, dict): + raise SpecError("spec must be a JSON object") + + c = _Collector() + taken: Set[str] = set() + known_unknown = set(source) - { + "multi_test", "variables", "constraints", "global_constraints", "layout", "notes" + } + if known_unknown: + c.add(f"unknown top-level keys {sorted(known_unknown)}") + + # ── multi_test ── + multi: Optional[MultiTest] = None + raw_multi = source.get("multi_test") + if raw_multi is not None: + if not isinstance(raw_multi, dict): + c.add("multi_test must be an object or null") + elif _check_name(c, raw_multi.get("count"), "multi_test.count", taken): + lo = c.expr("multi_test.min", raw_multi.get("min"), set()) + hi = c.expr("multi_test.max", raw_multi.get("max"), set()) + if lo is not None and hi is not None: + multi = MultiTest(raw_multi["count"], lo, hi) + + scope: Set[str] = {multi.count} if multi else set() + base_scope = set(scope) + + # ── variables ── + variables: List[Variable] = [] + raw_vars = source.get("variables") + if not isinstance(raw_vars, list) or not raw_vars: + c.add("variables must be a non-empty list") + raw_vars = [] + + for index, raw in enumerate(raw_vars): + where = f"variables[{index}]" + if not isinstance(raw, dict): + c.add(f"{where} must be an object") + continue + name, vtype = raw.get("name"), raw.get("type") + where = f"variable '{name}'" + if not _check_name(c, name, f"variables[{index}]", taken): + continue + if vtype not in TYPES: + c.add(f"{where}: type must be one of {sorted(TYPES)}, got {vtype!r}") + continue + + var = Variable(name=name, type=vtype, raw=raw) + + def need(key: str, allowed: Set[str] = scope) -> None: + if key not in raw: + c.add(f"{where}: missing '{key}'") + return + node = c.expr(f"{where}.{key}", raw[key], set(allowed)) + if node is not None: + var.exprs[key] = node + + if vtype == "int": + need("min") + need("max") + elif vtype == "array": + need("length") + need("min") + need("max") + _bool_opt(c, raw, "distinct", where) + if "order" in raw and raw["order"] not in ORDERS: + c.add(f"{where}: order must be one of {sorted(ORDERS)}") + elif vtype == "permutation": + need("length") + elif vtype == "string": + need("length") + elif vtype == "matrix": + need("rows") + need("cols") + if "alphabet" not in raw: + need("min") + need("max") + elif vtype == "rows": + need("count") + raw_fields = raw.get("fields") + if not isinstance(raw_fields, list) or not raw_fields: + c.add(f"{where}: 'fields' must be a non-empty list of {{name, min, max}}") + else: + row_scope = set(scope) + for f_index, rf in enumerate(raw_fields): + f_where = f"{where}.fields[{f_index}]" + if not isinstance(rf, dict): + c.add(f"{f_where} must be an object") + continue + if not _check_name(c, rf.get("name"), f_where, taken): + continue + lo = c.expr(f"{f_where}.min", rf.get("min"), row_scope) + hi = c.expr(f"{f_where}.max", rf.get("max"), row_scope) + if lo is not None and hi is not None: + var.fields.append(RowField(rf["name"], lo, hi)) + row_scope.add(rf["name"]) + elif vtype == "tree": + need("nodes") + if raw.get("format", "edges") not in ("edges", "parents"): + c.add(f"{where}: format must be 'edges' or 'parents'") + elif vtype == "graph": + need("nodes") + need("edges") + for key in ("directed", "connected", "self_loops", "multi_edges"): + _bool_opt(c, raw, key, where) + + if vtype in ("string", "matrix") and ("alphabet" in raw or vtype == "string"): + alphabet = raw.get("alphabet", DEFAULT_ALPHABET) + if not isinstance(alphabet, str): + c.add(f"{where}: alphabet must be a string such as \"a-z\" or \"01\"") + else: + try: + var.alphabet = expand_alphabet(alphabet) + except SpecError as exc: + c.add(f"{where}: {exc}") + + if vtype in ("permutation", "tree", "graph") and raw.get("indexing", 1) not in (0, 1): + c.add(f"{where}: indexing must be 0 or 1") + + if vtype in ("tree", "graph") and raw.get("weight") is not None: + weight = raw["weight"] + if not isinstance(weight, dict): + c.add(f"{where}: weight must be an object {{min, max}} or null") + else: + lo = c.expr(f"{where}.weight.min", weight.get("min"), scope) + hi = c.expr(f"{where}.weight.max", weight.get("max"), scope) + if lo is not None and hi is not None: + var.weight = {"min": lo, "max": hi} + + variables.append(var) + if vtype == "int": + scope.add(name) + + int_scope = set(scope) + + # ── constraints ── + constraints: List[ast.expr] = [] + for index, text in enumerate(source.get("constraints") or []): + node = c.expr(f"constraints[{index}]", text, int_scope) + if node is not None: + constraints.append(node) + + global_constraints: List[ast.expr] = [] + per_test_ints = int_scope - base_scope + for index, text in enumerate(source.get("global_constraints") or []): + where = f"global_constraints[{index}]" + try: + node = parse_expr(text) + except ExprError as exc: + c.add(f"{where}: {exc}") + continue + outside = names_in(node, include_sum_args=False) - base_scope + if outside: + c.add( + f"{where}: {sorted(outside)} must be wrapped in sum(...), " + "e.g. \"sum(n) <= 200000\"" + ) + continue + bad = set().union(*(names_in(a) for a in sum_args(node))) - per_test_ints + if bad: + c.add(f"{where}: sum(...) may only use per-test integer variables, not {sorted(bad)}") + continue + global_constraints.append(node) + + # ── layout ── + layout: List[List[str]] = [] + raw_layout = source.get("layout") + if not isinstance(raw_layout, list) or not raw_layout: + c.add("layout must be a non-empty list of lines, each a list of variable names") + raw_layout = [] + by_name = {v.name: v for v in variables} + seen: List[str] = [] + for index, line in enumerate(raw_layout): + if isinstance(line, str): + line = [line] + if not isinstance(line, list) or not line: + c.add(f"layout[{index}] must be a non-empty list of variable names") + continue + for name in line: + var = by_name.get(name) if isinstance(name, str) else None + if var is None: + hint = " (the multi_test count is printed automatically)" if multi and name == multi.count else "" + c.add(f"layout[{index}]: '{name}' is not a declared variable{hint}") + continue + if var.type in BLOCK_TYPES and len(line) > 1: + c.add(f"layout[{index}]: '{name}' is a multi-line {var.type} and must be alone on its line") + needed = set().union(*(names_in(var.exprs[k]) for k in SIZE_FIELDS.get(var.type, ()) if k in var.exprs)) - base_scope + missing = needed - set(seen) + if missing: + c.add( + f"layout[{index}]: the size of '{name}' depends on {sorted(missing)}, " + "which must be printed earlier in the layout" + ) + if name in seen: + c.add(f"layout: '{name}' appears more than once") + seen.append(name) + layout.append(list(line)) + for var in variables: + if var.name not in seen: + c.add(f"layout: variable '{var.name}' is declared but never printed") + + if c.errors: + raise SpecError("\n".join(f"- {e}" for e in c.errors)) + + return TestSpec( + raw=source, + multi=multi, + variables=variables, + constraints=constraints, + global_constraints=global_constraints, + layout=layout, + ) + + +def load_spec_file(path: str | Path) -> TestSpec: + try: + text = Path(path).read_text(encoding="utf-8") + except OSError as exc: + raise SpecError(f"cannot read test spec {path}: {exc}") from exc + return load_spec(text) diff --git a/requirements.txt b/requirements.txt index f433b0d..d6e699b 100644 --- a/requirements.txt +++ b/requirements.txt @@ -4,6 +4,7 @@ # text chat calls against the local Ollama server) # Pillow - image loading/validation/normalization for png/jpg/jpeg # PyMuPDF - offline PDF -> image rasterization (imported as `fitz`) +# z3-solver - SMT solver behind test generation (autosetter/testgen) ollama>=0.4.0 Pillow>=10.0.0 diff --git a/tests/fixtures.py b/tests/fixtures.py index 6dd1e7b..16b0f2a 100644 --- a/tests/fixtures.py +++ b/tests/fixtures.py @@ -99,3 +99,18 @@ """ SAMPLES = [{"input": "5\n", "output": "10\n", "explanation": ""}] + +# Z3 test specs for the same problem (see autosetter.testgen). +TEST_SPEC = { + "multi_test": None, + "variables": [{"name": "n", "type": "int", "min": 1, "max": 100}], + "constraints": [], + "global_constraints": [], + "layout": [["n"]], +} + +# Claims n >= 10, which the official sample (n = 5) contradicts. +TEST_SPEC_REJECTS_SAMPLE = { + **TEST_SPEC, + "variables": [{"name": "n", "type": "int", "min": 10, "max": 100}], +} diff --git a/tests/test_generator.py b/tests/test_generator.py index 4131b25..f85feed 100644 --- a/tests/test_generator.py +++ b/tests/test_generator.py @@ -7,6 +7,8 @@ from pathlib import Path import pytest +import json + from autosetter.generator import ( ARTIFACTS, generate_all_artifacts, @@ -14,6 +16,32 @@ ) from tests.conftest import StubOllamaClient +TWO_SUM_SPEC = { + "multi_test": None, + "variables": [ + {"name": "n", "type": "int", "min": 2, "max": 1000}, + {"name": "target", "type": "int", "min": -2000000000, "max": 2000000000}, + {"name": "nums", "type": "array", "length": "n", "min": -1000000000, "max": 1000000000}, + ], + "constraints": [], + "global_constraints": [], + "layout": [["n", "target"], ["nums"]], +} + + +class SequencedStub(StubOllamaClient): + """Replies to test-spec prompts from a queue; everything else gets C++.""" + + def __init__(self, spec_replies): + super().__init__(default="```cpp\nint main() { return 0; }\n```") + self.spec_replies = list(spec_replies) + + def chat_text(self, prompt, model=None, temperature=0.2): + self.prompts.append(prompt) + if "JSON \"test spec\"" in prompt and self.spec_replies: + return self.spec_replies.pop(0) + return self.default + SAMPLE_DATA = { "title": "Two Sum", "story": "Find pair.", @@ -56,3 +84,26 @@ def test_generate_all_artifacts(tmp_path: Path): artifact_path = results[spec.name] assert artifact_path.exists() assert any(f"Generating {spec.name}" in m for m in messages) + + +def test_test_spec_is_retried_until_it_accepts_the_samples(tmp_path: Path): + # The sample has n = 4, so a spec claiming n >= 5 must be sent back. + wrong = json.loads(json.dumps(TWO_SUM_SPEC)) + wrong["variables"][0]["min"] = 5 + client = SequencedStub([ + "Here is the spec:\n```json\n" + json.dumps(wrong) + "\n```", + json.dumps(TWO_SUM_SPEC), + ]) + + results = generate_all_artifacts( + problem_data=SAMPLE_DATA, + generated_dir=tmp_path, + client=client, + targets=["generator"], + ) + + assert results["generator"].name == "test_spec.json" + assert json.loads(results["generator"].read_text()) == TWO_SUM_SPEC + assert len(client.prompts) == 2 + assert "rejects the problem's official sample" in client.prompts[1] + assert "n = 4 is outside [5, 1000]" in client.prompts[1] diff --git a/tests/test_packager.py b/tests/test_packager.py index 98913ff..9bfea57 100644 --- a/tests/test_packager.py +++ b/tests/test_packager.py @@ -64,6 +64,24 @@ def test_complete_pairs_are_packaged(tmp_path: Path): assert manifest["ready_for_release"] is True +def test_z3_generated_tests_ship_as_files(tmp_path: Path): + generated, tests = make_dirs(tmp_path) + (generated / "generator.cpp").unlink() + (generated / "test_spec.json").write_text("{}") + for i in (1, 2): + (tests / f"{i:03d}.in").write_text("5\n") + (tests / f"{i:03d}.ans").write_text("10\n") + write_report(tests, test_cases=[{"index": 1, "seed": "1"}, {"index": 2, "seed": "2"}]) + + manifest = build(tmp_path) + + # Polygon cannot run Z3, so the tests must not be left to a generator script. + assert manifest["packaged_tests"] == 2 + assert manifest["ready_for_release"] is True + assert "files/test_spec.json" in [f.replace("\\", "/") for f in manifest["files"]] + assert not (tmp_path / "package" / "script").exists() + + def test_input_without_an_answer_is_excluded(tmp_path: Path): _, tests = make_dirs(tmp_path) (tests / "001.in").write_text("5\n") diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py index bd1094a..6630ff5 100644 --- a/tests/test_pipeline.py +++ b/tests/test_pipeline.py @@ -4,6 +4,7 @@ from __future__ import annotations +import json from pathlib import Path import pytest @@ -19,6 +20,8 @@ GENERATOR_OUT_OF_RANGE, SAMPLES, SOLUTION, + TEST_SPEC, + TEST_SPEC_REJECTS_SAMPLE, VALIDATOR, ) @@ -34,9 +37,12 @@ def build_pipeline( checker: str = CHECKER, samples=None, num_tests: int = 3, + test_spec=None, ) -> TestPipeline: generated = tmp_path / "generated" generated.mkdir(parents=True, exist_ok=True) + if test_spec is not None: + (generated / "test_spec.json").write_text(json.dumps(test_spec)) (generated / "validator.cpp").write_text(validator) if "import " in generator or "sys." in generator: (generated / "generator.py").write_text(generator) @@ -102,6 +108,8 @@ def test_out_of_range_generator_blames_generator(tmp_path: Path): assert report.passed_tests == 0 assert not report.all_passed assert all("generator is the file at fault" in tc.error for tc in report.test_cases) + assert "None of the 3 generated tests is usable" in report.diagnosis + assert "generator is the file at fault" in report.diagnosis def test_broken_validator_blames_validator(tmp_path: Path): @@ -162,3 +170,28 @@ def test_pipeline_with_python_generator(tmp_path: Path): assert report.checker_trusted assert report.compilation.generator assert report.passed_tests == report.total_tests == 3 + + +def test_pipeline_with_z3_test_spec(tmp_path: Path): + # The spec is preferred over the generator.cpp written alongside it; the + # out-of-range C++ generator would fail every test if it were used. + report = build_pipeline( + tmp_path, test_spec=TEST_SPEC, generator=GENERATOR_OUT_OF_RANGE, num_tests=5 + ).run() + + assert report.all_passed, report.diagnosis + assert report.passed_tests == report.total_tests == 5 + assert [tc.strategy for tc in report.test_cases] == ["min", "max", "random", "small", "near_max"] + tests = tmp_path / "tests" + assert (tests / "001.in").read_text() == "1\n" # "min" strategy + assert (tests / "002.in").read_text() == "100\n" # "max" strategy + assert (tests / "002.ans").read_text().strip() == "200" + + +def test_test_spec_rejecting_official_sample_is_unusable(tmp_path: Path): + report = build_pipeline(tmp_path, test_spec=TEST_SPEC_REJECTS_SAMPLE).run() + + assert not report.compilation.generator + assert "rejects official samples" in report.compilation.errors["generator"] + assert "test_spec.json is unusable" in report.diagnosis + assert not report.all_passed diff --git a/tests/test_testgen.py b/tests/test_testgen.py new file mode 100644 index 0000000..dadbc24 --- /dev/null +++ b/tests/test_testgen.py @@ -0,0 +1,173 @@ +""" +Unit tests for autosetter.testgen (Z3 spec-driven test generation). +""" + +from __future__ import annotations + +import pytest + +from autosetter.testgen import ( + SpecError, + TestGenError, + check_input, + check_samples, + generate_test, + load_spec, +) +from autosetter.testgen.expr import ExprError, parse_expr +from autosetter.testgen.layout import parse_input + +MULTI_SPEC = { + "multi_test": {"count": "t", "min": 1, "max": 10000}, + "variables": [ + {"name": "n", "type": "int", "min": 1, "max": 200000}, + {"name": "k", "type": "int", "min": 1, "max": "n"}, + {"name": "a", "type": "array", "length": "n", "min": 1, "max": 1000000000}, + ], + "constraints": ["k <= n"], + "global_constraints": ["sum(n) <= 200000"], + "layout": [["n", "k"], ["a"]], +} + +STRUCTURES_SPEC = { + "variables": [ + {"name": "n", "type": "int", "min": 2, "max": 300}, + {"name": "m", "type": "int", "min": "n - 1", "max": "min(n*(n-1)//2, 2000)"}, + {"name": "g", "type": "graph", "nodes": "n", "edges": "m", "connected": True, + "weight": {"min": 1, "max": 100}}, + {"name": "tr", "type": "tree", "nodes": "n"}, + {"name": "par", "type": "tree", "nodes": "n", "format": "parents"}, + {"name": "p", "type": "permutation", "length": "n"}, + {"name": "s", "type": "string", "length": "n", "alphabet": "a-c"}, + {"name": "b", "type": "array", "length": "n", "min": -5, "max": 1000, + "distinct": True, "order": "ascending"}, + {"name": "grid", "type": "matrix", "rows": 3, "cols": "n", "alphabet": ".#"}, + {"name": "q", "type": "int", "min": 1, "max": 50}, + {"name": "queries", "type": "rows", "count": "q", + "fields": [{"name": "l", "min": 1, "max": "n"}, {"name": "r", "min": "l", "max": "n"}]}, + ], + "layout": [["n", "m"], ["g"], ["tr"], ["par"], ["p"], ["s"], ["b"], ["grid"], ["q"], ["queries"]], +} + + +def test_expressions_reject_anything_but_arithmetic(): + for unsafe in ("__import__('os')", "n.real", "[n]", "open('x')", "lambda: 1"): + with pytest.raises(ExprError): + parse_expr(unsafe) + parse_expr("min(n*(n-1)//2, 10**5) + abs(-k)") + + +def test_spec_errors_are_all_reported_together(): + with pytest.raises(SpecError) as excinfo: + load_spec({ + "variables": [ + {"name": "a", "type": "array", "length": "n", "min": 1, "max": 5}, + {"name": "n", "type": "int", "min": 1}, + {"name": "x", "type": "blob"}, + ], + "layout": [["a", "n"], ["ghost"]], + }) + message = str(excinfo.value) + assert "variable 'a'.length" in message # n is declared after a + assert "missing 'max'" in message + assert "type must be one of" in message + assert "'ghost' is not a declared variable" in message + + +def test_size_must_be_printed_before_the_value_it_sizes(): + with pytest.raises(SpecError, match="must be printed earlier"): + load_spec({ + "variables": [ + {"name": "n", "type": "int", "min": 1, "max": 5}, + {"name": "a", "type": "array", "length": "n", "min": 1, "max": 5}, + ], + "layout": [["a"], ["n"]], + }) + + +def test_global_constraints_need_sum(): + with pytest.raises(SpecError, match=r"wrapped in sum"): + load_spec({**MULTI_SPEC, "global_constraints": ["n <= 200000"]}) + + +@pytest.mark.parametrize("seed", range(1, 11)) +def test_multi_test_output_satisfies_every_constraint(seed): + spec = load_spec(MULTI_SPEC) + test = generate_test(spec, seed) + + assert check_input(spec, test.text) == [] + top, cases = parse_input(spec, test.text) + assert top["t"] == len(cases) == test.num_cases + assert sum(c["n"] for c in cases) <= 200000 + assert all(1 <= c["k"] <= c["n"] for c in cases) + + +def test_generation_is_deterministic_per_seed_and_varies_across_seeds(): + spec = load_spec(MULTI_SPEC) + assert generate_test(spec, 3).text == generate_test(spec, 3).text + texts = {generate_test(spec, seed).text for seed in range(1, 13)} + assert len(texts) == 12 + + +def test_min_and_max_strategies_hit_the_bounds(): + spec = load_spec({ + "variables": [ + {"name": "n", "type": "int", "min": 3, "max": 1000}, + {"name": "k", "type": "int", "min": 1, "max": "n"}, + ], + "layout": [["n", "k"]], + }) + assert generate_test(spec, 1).text == "3 1\n" # strategy "min" + assert generate_test(spec, 2).text == "1000 1000\n" # strategy "max" + + +def test_values_respect_constraints_with_holes(): + spec = load_spec({ + "variables": [{"name": "n", "type": "int", "min": 1, "max": 1000}], + "constraints": ["n % 7 == 3"], + "layout": [["n"]], + }) + for seed in range(1, 15): + assert int(generate_test(spec, seed).text) % 7 == 3 + + +@pytest.mark.parametrize("seed", range(1, 11)) +def test_structures_are_valid(seed): + spec = load_spec(STRUCTURES_SPEC) + assert check_input(spec, generate_test(spec, seed).text) == [] + + +def test_unsatisfiable_constraints_raise(): + spec = load_spec({ + "variables": [{"name": "n", "type": "int", "min": 1, "max": 5}], + "constraints": ["n > 10"], + "layout": [["n"]], + }) + with pytest.raises(TestGenError, match="unsatisfiable"): + generate_test(spec, 1) + + +def test_impossible_distinct_array_raises(): + spec = load_spec({ + "variables": [ + {"name": "n", "type": "int", "min": 10, "max": 10}, + {"name": "a", "type": "array", "length": "n", "min": 1, "max": 3, "distinct": True}, + ], + "layout": [["n"], ["a"]], + }) + with pytest.raises(TestGenError, match="distinct"): + generate_test(spec, 1) + + +def test_samples_are_checked_against_the_spec(): + spec = load_spec(MULTI_SPEC) + good = {"input": "2\n3 2\n1 2 3\n1 1\n7\n"} + out_of_range = {"input": "1\n3 5\n1 2 3\n"} # k > n + short = {"input": "1\n3 1\n1 2\n"} # a has 2 of 3 elements + extra = {"input": "1\n1 1\n5 6\n"} # a trailing token + + assert check_samples(spec, [good]) == [] + problems = check_samples(spec, [good, out_of_range, short, extra]) + assert any(p.startswith("sample 2:") and "k = 5" in p for p in problems) + assert any(p.startswith("sample 3:") and "input ended" in p for p in problems) + assert any(p.startswith("sample 4:") and "unread token" in p for p in problems)