Coverage for src / competitive_verifier / oj / problem.py: 81%
332 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-10-04 07:40 +0900
« prev ^ index » next coverage.py v7.13.1, created at 2026-10-04 07:40 +0900
1import glob
2import hashlib
3import json
4import os
5import pathlib
6import posixpath
7import re
8import subprocess
9import sys
10import urllib.parse
11import zipfile
12from abc import abstractmethod
13from collections.abc import Iterable, Iterator
14from dataclasses import dataclass
15from io import BytesIO
16from logging import getLogger
17from typing import ClassVar, Optional, TypeVar
19import requests
21from competitive_verifier import config
22from competitive_verifier.log import GitHubMessageParams
23from competitive_verifier.models import (
24 Problem,
25 TestCaseData,
26 TestCaseFile,
27 TestCaseProvider,
28)
30logger = getLogger(__name__)
33class NotLoggedInError(RuntimeError):
34 pass
37class _BaseProblem(Problem):
38 def iter_system_cases(self) -> Iterator[TestCaseFile]:
39 return iter_testcases(directory=self.test_directory)
41 def is_testdata_cached(self) -> bool:
42 test_directory = self.test_directory
43 return test_directory.exists() and any(test_directory.iterdir())
45 def download_system_cases(self) -> Iterable[TestCaseData] | bool:
46 test_directory = self.test_directory
48 if self.is_testdata_cached():
49 logger.info("download:already exists: %s", self.url)
50 return True
52 self.problem_directory.mkdir(parents=True, exist_ok=True)
54 samples = list(self._download_cases())
56 # Check samples
57 if not samples: 57 ↛ 58line 57 didn't jump to line 58 because the condition on line 57 was never true
58 logger.error(
59 "Sample not found",
60 extra={"github": GitHubMessageParams()},
61 )
62 return False
64 # write samples to files
65 save_testcases(samples, directory=test_directory)
66 return samples
68 @abstractmethod
69 def _download_cases(self) -> Iterable[TestCaseData]: ...
72class LibraryCheckerProblem(Problem):
73 checker_exe_name: ClassVar[str] = (
74 "checker.exe" if sys.platform == "win32" else "checker"
75 )
77 def __init__(self, *, problem_id: str):
78 self.problem_id = problem_id
79 self._source_directory = None
81 def __hash__(self) -> int:
82 return hash((self.problem_id, self.repo_path))
84 def __eq__(self, value: object) -> bool:
85 if not isinstance(value, LibraryCheckerProblem): 85 ↛ 86line 85 didn't jump to line 86 because the condition on line 85 was never true
86 return False
87 return self.problem_id == value.problem_id and self.repo_path == value.repo_path
89 @property
90 def repo_path(self):
91 return config.get_cache_dir() / "library-checker-problems"
93 def iter_system_cases(self) -> Iterator[TestCaseFile]:
94 inputs: dict[str, pathlib.Path] = {}
95 outputs: dict[str, pathlib.Path] = {}
96 for path in self.source_directory.glob("in/*.in"):
97 inputs[path.stem] = path
98 for path in self.source_directory.glob("out/*.out"):
99 outputs[path.stem] = path
100 return merge_testcase_files(inputs, outputs)
102 def is_testdata_cached(self) -> bool:
103 try:
104 return any(self.iter_system_cases())
105 except RuntimeError:
106 return False
108 def download_system_cases(self) -> bool:
109 self.problem_directory.mkdir(parents=True, exist_ok=True)
110 self.generate_test_cases()
111 return True
113 @property
114 def checker(self) -> pathlib.Path | None:
115 return self.source_directory / self.checker_exe_name
117 def generate_test_cases(self) -> None:
118 self.update_cloned_repository()
119 path = self.repo_path
121 spec = str(self.source_directory / "info.toml")
122 command = [sys.executable, str(path / "generate.py"), spec]
123 logger.info("$ %s", " ".join(command))
124 try:
125 subprocess.check_call(command, stdout=sys.stderr, stderr=sys.stderr)
126 except subprocess.CalledProcessError:
127 logger.exception(
128 "the generate.py failed: check https://github.com/yosupo06/library-checker-problems/issues",
129 extra={"github": GitHubMessageParams()},
130 )
131 raise
133 @property
134 def hash_json(self) -> pathlib.Path:
135 """The committed per-case digests of the generated test data."""
136 return self.source_directory / "hash.json"
138 def sync_testdata(self) -> None:
139 self.update_cloned_repository()
141 def testdata_hash(self) -> str | None:
142 try:
143 return hashlib.sha256(self.hash_json.read_bytes()).hexdigest()
144 except (OSError, RuntimeError):
145 return None
147 @property
148 def source_directory(self):
149 if self._source_directory is None:
150 problem_id = self.problem_id
151 info_tomls = list(
152 self.repo_path.glob(f"**/{glob.escape(problem_id)}/info.toml")
153 )
154 if len(info_tomls) != 1:
155 raise RuntimeError(f"the problem {problem_id!r} not found or broken")
156 self._source_directory = info_tomls[0].parent
157 return self._source_directory
159 @property
160 def url(self) -> str:
161 return f"https://judge.yosupo.jp/problem/{self.problem_id}"
163 @classmethod
164 def from_url(cls, url: str) -> Optional["LibraryCheckerProblem"]:
165 # example: https://judge.yosupo.jp/problem/unionfind
166 result = urllib.parse.urlparse(url)
167 if result.scheme in ("", "http", "https") and result.netloc in (
168 "judge.yosupo.jp",
169 "old.yosupo.jp",
170 ):
171 m = re.match(r"/problem/(\w+)/?", result.path)
172 if m: 172 ↛ 174line 172 didn't jump to line 174 because the condition on line 172 was always true
173 return cls(problem_id=m.group(1))
174 return None
176 _is_repository_updated: ClassVar[set[pathlib.Path]] = set()
178 def update_cloned_repository(self) -> None:
179 if self.repo_path in self._is_repository_updated:
180 return
182 try:
183 subprocess.check_call(
184 ["git", "--version"], # noqa: S607
185 stdout=sys.stderr,
186 stderr=sys.stderr,
187 )
188 except FileNotFoundError:
189 logger.exception(
190 "git command not found",
191 exc_info=False,
192 extra={"github": GitHubMessageParams()},
193 )
194 raise
196 path = self.repo_path
197 if not path.exists(): 197 ↛ 208line 197 didn't jump to line 208 because the condition on line 197 was always true
198 # init the problem repository
199 url = "https://github.com/yosupo06/library-checker-problems"
200 logger.info("$ git clone %s %s", url, path)
201 subprocess.check_call(
202 ["git", "clone", url, str(path)], # noqa: S607
203 stdout=sys.stderr,
204 stderr=sys.stderr,
205 )
206 else:
207 # sync the problem repository
208 logger.info("$ git -C %s pull", path)
209 subprocess.check_call(
210 ["git", "-C", str(path), "pull"], # noqa: S607
211 stdout=sys.stderr,
212 stderr=sys.stderr,
213 )
215 LibraryCheckerProblem._is_repository_updated.add(self.repo_path)
218class _YukicoderProblemNo(int):
219 def __new__(cls, value: int):
220 return super().__new__(cls, value)
222 def __str__(self) -> str:
223 return "no/" + super().__str__()
226class _YukicoderProblemId(int):
227 def __new__(cls, value: int):
228 return super().__new__(cls, value)
231class YukicoderProblem(_BaseProblem):
232 problem: _YukicoderProblemNo | _YukicoderProblemId
234 def __init__(self, *, problem_no: int | None = None, problem_id: int | None = None):
235 if problem_no is not None:
236 self.problem = _YukicoderProblemNo(problem_no)
237 elif problem_id is not None: 237 ↛ 240line 237 didn't jump to line 240 because the condition on line 237 was always true
238 self.problem = _YukicoderProblemId(problem_id)
239 else:
240 raise ValueError("Needs problem_no or problem_id")
242 def _download_cases(self) -> list[TestCaseData]:
243 """Download yukicoder problem.
245 Raises:
246 NotLoggedInError: If the `cargo metadata` command fails
247 """
248 headers: dict[str, str] | None = None
249 if yukicoder_token := os.environ.get("YUKICODER_TOKEN"):
250 headers = {"Authorization": f"Bearer {yukicoder_token}"}
252 if not self._is_logged_in(headers=headers):
253 raise NotLoggedInError("Required: $YUKICODER_TOKEN environment variable")
254 url = f"{self.url}/testcase.zip"
255 resp = requests.get(url, headers=headers, allow_redirects=True, timeout=10)
257 with zipfile.ZipFile(BytesIO(resp.content)) as fh:
258 inputs: dict[str, bytes] = {}
259 outputs: dict[str, bytes] = {}
260 for filename in fh.namelist():
261 if filename.endswith("/"):
262 continue
263 file = fh.read(filename)
264 path = pathlib.Path(filename)
265 if filename.startswith("test_in/"):
266 inputs[path.stem] = file
267 elif filename.startswith("test_out/"):
268 outputs[path.stem] = file
269 return [
270 TestCaseData(name=name, input_data=i, output_data=o)
271 for name, i, o in enumerate_inouts(inputs, outputs)
272 ]
274 @property
275 def url(self) -> str:
276 return f"https://yukicoder.me/problems/{self.problem}"
278 @classmethod
279 def from_url(cls, url: str) -> Optional["YukicoderProblem"]:
280 # example: https://yukicoder.me/problems/no/499
281 # example: http://yukicoder.me/problems/1476
282 result = urllib.parse.urlparse(url)
283 dirname, basename = posixpath.split(_normpath(result.path))
284 if result.scheme in ("", "http", "https") and result.netloc == "yukicoder.me":
285 try:
286 n = int(basename)
287 except ValueError:
288 pass
289 else:
290 if dirname == "/problems/no":
291 return cls(problem_no=n)
292 if dirname == "/problems":
293 return cls(problem_id=n)
294 return None
296 def _is_logged_in(self, *, headers: dict[str, str] | None = None) -> bool:
297 url = "https://yukicoder.me"
298 resp = requests.get(url, headers=headers, allow_redirects=True, timeout=10)
299 resp.raise_for_status()
300 return "login-btn" not in str(resp.content)
303class AOJProblem(_BaseProblem):
304 def __init__(self, *, problem_id: str):
305 self.problem_id = problem_id
307 def _download_cases(self) -> Iterable[TestCaseData]:
308 return AOJProblem.download_cases(self.problem_id)
310 @staticmethod
311 def download_cases(problem_id: str) -> Iterable[TestCaseData]:
312 # get header
313 # reference: http://developers.u-aizu.ac.jp/api?key=judgedat%2Ftestcases%2F%7BproblemId%7D%2Fheader_GET
314 url = f"https://judgedat.u-aizu.ac.jp/testcases/{problem_id}/header"
315 resp = requests.get(url, allow_redirects=True, timeout=10)
316 resp.raise_for_status()
317 header_res = json.loads(resp.text)
319 # get testcases via the official API
320 for header in header_res["headers"]:
321 # NOTE: the endpoints are not same to http://developers.u-aizu.ac.jp/api?key=judgedat%2Ftestcases%2F%7BproblemId%7D%2F%7Bserial%7D_GET since the json API often says "..... (terminated because of the limitation)"
322 # NOTE: even when using https://judgedat.u-aizu.ac.jp/testcases/PROBLEM_ID/SERIAL, there is the 1G limit (see https://twitter.com/beet_aizu/status/1194947611100188672)
323 serial = header["serial"]
324 url = f"https://judgedat.u-aizu.ac.jp/testcases/{problem_id}/{serial}"
326 resp_in = requests.get(url + "/in", allow_redirects=True, timeout=10)
327 resp_in.raise_for_status()
328 resp_out = requests.get(url + "/out", allow_redirects=True, timeout=10)
329 resp_out.raise_for_status()
331 yield TestCaseData(
332 header["name"],
333 resp_in.content,
334 resp_out.content,
335 )
337 @property
338 def url(self) -> str:
339 return f"http://judge.u-aizu.ac.jp/onlinejudge/description.jsp?id={self.problem_id}"
341 @classmethod
342 def from_url(cls, url: str) -> Optional["AOJProblem"]:
343 result = urllib.parse.urlparse(url)
345 # example: http://judge.u-aizu.ac.jp/onlinejudge/description.jsp?id=1169
346 # example: http://judge.u-aizu.ac.jp/onlinejudge/description.jsp?id=DSL_1_A&lang=jp
347 querystring = urllib.parse.parse_qs(result.query)
348 if (
349 result.scheme in ("", "http", "https")
350 and result.netloc == "judge.u-aizu.ac.jp"
351 and _normpath(result.path) == "/onlinejudge/description.jsp"
352 and querystring.get("id")
353 and len(querystring["id"]) == 1
354 ):
355 (n,) = querystring["id"]
356 return cls(problem_id=n)
358 # example: https://onlinejudge.u-aizu.ac.jp/challenges/sources/JAG/Prelim/2881
359 # example: https://onlinejudge.u-aizu.ac.jp/courses/library/4/CGL/3/CGL_3_B
360 m = re.match(
361 r"^/(challenges|courses)/(sources|library/\d+|lesson/\d+)/(\w+)/(\w+)/(\w+)$",
362 _normpath(result.path),
363 )
364 if (
365 result.scheme in ("", "http", "https")
366 and result.netloc == "onlinejudge.u-aizu.ac.jp"
367 and m
368 ):
369 n = m.group(5)
370 return cls(problem_id=n)
372 # example: https://onlinejudge.u-aizu.ac.jp/problems/0423
373 # example: https://onlinejudge.u-aizu.ac.jp/problems/CGL_3_B
374 m = re.match(r"^/problems/(\w+)$", _normpath(result.path))
375 if (
376 result.scheme in ("", "http", "https")
377 and result.netloc == "onlinejudge.u-aizu.ac.jp"
378 and m
379 ):
380 n = m.group(1)
381 return cls(problem_id=n)
383 return None
386class AOJArenaProblem(_BaseProblem):
387 def __init__(self, *, arena_id: str, alphabet: str):
388 if len(alphabet) != 1 or not alphabet.isupper(): 388 ↛ 389line 388 didn't jump to line 389 because the condition on line 388 was never true
389 raise ValueError(arena_id, alphabet)
390 self.arena_id = arena_id
391 self.alphabet = alphabet
393 self._problem_id: str | None = None
395 def get_problem_id(self) -> str:
396 if self._problem_id is None:
397 url = f"https://judgeapi.u-aizu.ac.jp/arenas/{self.arena_id}/problems"
398 resp = requests.get(url, allow_redirects=True, timeout=10)
399 resp.raise_for_status()
400 problems = json.loads(resp.text)
401 for problem in problems:
402 if problem["id"] == self.alphabet:
403 p = problem["problemId"]
404 logger.debug("problem: %s", p)
405 self._problem_id = p
406 return p
407 raise ValueError("Problem is not found.")
408 return self._problem_id
410 def _download_cases(self) -> Iterable[TestCaseData]:
411 return AOJProblem.download_cases(self.get_problem_id())
413 @property
414 def url(self) -> str:
415 return f"https://onlinejudge.u-aizu.ac.jp/services/room.html#{self.arena_id}/problems/{self.alphabet}"
417 @classmethod
418 def from_url(cls, url: str) -> Optional["AOJArenaProblem"]:
419 # example: https://onlinejudge.u-aizu.ac.jp/services/room.html#RitsCamp19Day2/problems/A
420 result = urllib.parse.urlparse(url)
421 if (
422 result.scheme in ("", "http", "https")
423 and result.netloc == "onlinejudge.u-aizu.ac.jp"
424 and _normpath(result.path) == "/services/room.html"
425 ):
426 fragment = result.fragment.split("/")
427 if len(fragment) == 3 and fragment[1] == "problems": # noqa: PLR2004 427 ↛ 429line 427 didn't jump to line 429 because the condition on line 427 was always true
428 return cls(arena_id=fragment[0], alphabet=fragment[2].upper())
429 return None
432@dataclass
433class LocalProblem(TestCaseProvider):
434 path: pathlib.Path
436 def download_system_cases(self) -> Iterable[TestCaseData] | bool:
437 return bool(any(self.iter_system_cases()))
439 def iter_system_cases(self) -> Iterable[TestCaseFile]:
440 return iter_testcases(directory=self.path, recursive=True)
442 def testdata_hash(self) -> str | None:
443 if not self.path.is_dir():
444 return None
445 digest = hashlib.sha256()
446 for case in sorted(self.iter_system_cases(), key=lambda c: c.name):
447 digest.update(case.name.encode())
448 digest.update(b"\0")
449 digest.update(case.input_path.read_bytes())
450 digest.update(b"\0")
451 digest.update(case.output_path.read_bytes())
452 digest.update(b"\0")
453 return digest.hexdigest()
456def _normpath(path: str) -> str:
457 """A wrapper of posixpath.normpath.
459 posixpath.normpath doesn't collapse a leading duplicated slashes.
460 """
461 path = posixpath.normpath(path)
462 if path.startswith("//"):
463 path = "/" + path.lstrip("/")
464 return path
467def _subclasses_recursive(cls: type[Problem]) -> Iterable[type[Problem]]:
468 yield from (children := cls.__subclasses__())
469 for ch in children:
470 yield from _subclasses_recursive(ch)
473def problem_from_url(url: str) -> Problem | None:
474 for ch in set(_subclasses_recursive(Problem)):
475 if (problem := ch.from_url(url)) is not None:
476 return problem
477 return None
480_InOut = TypeVar("_InOut")
483def enumerate_inouts(
484 inputs: dict[str, _InOut],
485 outputs: dict[str, _InOut],
486) -> Iterator[tuple[str, _InOut, _InOut]]:
487 common_keys = inputs.keys() & outputs.keys()
488 if len(inputs) != len(common_keys) or len(outputs) != len(common_keys):
489 logger.warning("dangling output case")
491 if len(common_keys) == 0:
492 logger.warning("no cases found")
494 for key in sorted(common_keys):
495 yield (key, inputs[key], outputs[key])
498def merge_testcase_files(
499 inputs: dict[str, pathlib.Path],
500 outputs: dict[str, pathlib.Path],
501) -> Iterator[TestCaseFile]:
502 for name, i, o in enumerate_inouts(inputs, outputs):
503 yield TestCaseFile(name=name, input_path=i, output_path=o)
506def _casename(path: pathlib.Path, *, directory: pathlib.Path) -> str:
507 return path.relative_to(directory).with_suffix("").as_posix()
510def iter_testcases(
511 *, directory: pathlib.Path, recursive: bool = False
512) -> Iterator[TestCaseFile]:
513 inputs: dict[str, pathlib.Path] = {}
514 outputs: dict[str, pathlib.Path] = {}
515 pre = "**/" if recursive else ""
517 for path in directory.glob(pre + "*.in"):
518 if path.is_file(): 518 ↛ 517line 518 didn't jump to line 517 because the condition on line 518 was always true
519 inputs[_casename(path, directory=directory)] = path
520 for path in directory.glob(pre + "*.out"):
521 if path.is_file(): 521 ↛ 520line 521 didn't jump to line 520 because the condition on line 521 was always true
522 outputs[_casename(path, directory=directory)] = path
524 return merge_testcase_files(inputs, outputs)
527def _name_to_filename(name: str, ext: str):
528 return pathlib.Path(name).with_suffix(f".{ext}").name
531def save_testcases(samples: Iterable[TestCaseData], *, directory: pathlib.Path):
532 for sample in samples:
533 for data, ext in [(sample.input_data, "in"), (sample.output_data, "out")]:
534 path = directory / _name_to_filename(sample.name, ext)
536 if path.exists(): 536 ↛ 537line 536 didn't jump to line 537 because the condition on line 536 was never true
537 logger.error("Failed to download since file already exists: %s", path)
538 path.parent.mkdir(parents=True, exist_ok=True)
539 path.write_bytes(data)
540 logger.debug("saved to: %s", path)