Coverage for src / competitive_verifier / oj / problem.py: 61%
463 statements
« prev ^ index » next coverage.py v7.13.1, created at 2026-10-03 23:40 +0900
« prev ^ index » next coverage.py v7.13.1, created at 2026-10-03 23:40 +0900
1import glob
2import json
3import os
4import pathlib
5import posixpath
6import re
7import shutil
8import subprocess
9import sys
10import tempfile
11import time
12import urllib.parse
13import zipfile
14from abc import abstractmethod
15from collections.abc import Iterable, Iterator
16from dataclasses import dataclass
17from logging import getLogger
18from typing import ClassVar, Optional, TypeVar
20import requests
22from competitive_verifier import config
23from competitive_verifier.log import GitHubMessageParams
24from competitive_verifier.models import (
25 Problem,
26 TestCaseData,
27 TestCaseFile,
28 TestCaseProvider,
29)
31logger = getLogger(__name__)
33_ASCII_SPACE = 0x20
34_ASCII_DELETE = 0x7F
37class NotLoggedInError(RuntimeError):
38 pass
41class _BaseProblem(Problem):
42 def iter_system_cases(self) -> Iterator[TestCaseFile]:
43 return iter_testcases(directory=self.test_directory)
45 def is_testdata_cached(self) -> bool:
46 return any(self.iter_system_cases())
48 def download_system_cases(self) -> Iterable[TestCaseData] | bool:
49 test_directory = self.test_directory
51 if self.is_testdata_cached():
52 logger.info("download:already exists: %s", self.url)
53 return True
55 self.problem_directory.mkdir(parents=True, exist_ok=True)
57 samples = list(self._download_cases())
59 # Check samples
60 if not samples: 60 ↛ 61line 60 didn't jump to line 61 because the condition on line 60 was never true
61 logger.error(
62 "Sample not found",
63 extra={"github": GitHubMessageParams()},
64 )
65 return False
67 # write samples to files
68 save_testcases(samples, directory=test_directory)
69 return samples
71 @abstractmethod
72 def _download_cases(self) -> Iterable[TestCaseData]: ...
75class LibraryCheckerProblem(Problem):
76 checker_exe_name: ClassVar[str] = (
77 "checker.exe" if sys.platform == "win32" else "checker"
78 )
80 def __init__(self, *, problem_id: str):
81 self.problem_id = problem_id
82 self._source_directory = None
84 def __hash__(self) -> int:
85 return hash((self.problem_id, self.repo_path))
87 def __eq__(self, value: object) -> bool:
88 if not isinstance(value, LibraryCheckerProblem): 88 ↛ 89line 88 didn't jump to line 89 because the condition on line 88 was never true
89 return False
90 return self.problem_id == value.problem_id and self.repo_path == value.repo_path
92 @property
93 def repo_path(self):
94 return config.get_cache_dir() / "library-checker-problems"
96 def iter_system_cases(self) -> Iterator[TestCaseFile]:
97 inputs: dict[str, pathlib.Path] = {}
98 outputs: dict[str, pathlib.Path] = {}
99 for path in self.source_directory.glob("in/*.in"):
100 inputs[path.stem] = path
101 for path in self.source_directory.glob("out/*.out"):
102 outputs[path.stem] = path
103 return merge_testcase_files(inputs, outputs)
105 def is_testdata_cached(self) -> bool:
106 try:
107 return any(self.iter_system_cases())
108 except RuntimeError:
109 return False
111 def download_system_cases(self) -> bool:
112 self.problem_directory.mkdir(parents=True, exist_ok=True)
113 self.generate_test_cases()
114 return True
116 @property
117 def checker(self) -> pathlib.Path | None:
118 return self.source_directory / self.checker_exe_name
120 def generate_test_cases(self) -> None:
121 self.update_cloned_repository()
122 path = self.repo_path
124 spec = str(self.source_directory / "info.toml")
125 command = [sys.executable, str(path / "generate.py"), spec]
126 logger.info("$ %s", " ".join(command))
127 try:
128 subprocess.check_call(command, stdout=sys.stderr, stderr=sys.stderr)
129 except subprocess.CalledProcessError:
130 logger.exception(
131 "the generate.py failed: check https://github.com/yosupo06/library-checker-problems/issues",
132 extra={"github": GitHubMessageParams()},
133 )
134 raise
136 @property
137 def source_directory(self):
138 if self._source_directory is None:
139 problem_id = self.problem_id
140 info_tomls = list(
141 self.repo_path.glob(f"**/{glob.escape(problem_id)}/info.toml")
142 )
143 if len(info_tomls) != 1:
144 raise RuntimeError(f"the problem {problem_id!r} not found or broken")
145 self._source_directory = info_tomls[0].parent
146 return self._source_directory
148 @property
149 def url(self) -> str:
150 return f"https://judge.yosupo.jp/problem/{self.problem_id}"
152 @classmethod
153 def from_url(cls, url: str) -> Optional["LibraryCheckerProblem"]:
154 # example: https://judge.yosupo.jp/problem/unionfind
155 result = urllib.parse.urlparse(url)
156 if result.scheme in ("", "http", "https") and result.netloc in (
157 "judge.yosupo.jp",
158 "old.yosupo.jp",
159 ):
160 m = re.match(r"/problem/(\w+)/?", result.path)
161 if m: 161 ↛ 163line 161 didn't jump to line 163 because the condition on line 161 was always true
162 return cls(problem_id=m.group(1))
163 return None
165 _is_repository_updated: ClassVar[set[pathlib.Path]] = set()
167 def update_cloned_repository(self) -> None:
168 if self.repo_path in self._is_repository_updated: 168 ↛ 169line 168 didn't jump to line 169 because the condition on line 168 was never true
169 return
171 try:
172 subprocess.check_call(
173 ["git", "--version"], # noqa: S607
174 stdout=sys.stderr,
175 stderr=sys.stderr,
176 )
177 except FileNotFoundError:
178 logger.exception(
179 "git command not found",
180 exc_info=False,
181 extra={"github": GitHubMessageParams()},
182 )
183 raise
185 path = self.repo_path
186 if not path.exists(): 186 ↛ 197line 186 didn't jump to line 197 because the condition on line 186 was always true
187 # init the problem repository
188 url = "https://github.com/yosupo06/library-checker-problems"
189 logger.info("$ git clone %s %s", url, path)
190 subprocess.check_call(
191 ["git", "clone", url, str(path)], # noqa: S607
192 stdout=sys.stderr,
193 stderr=sys.stderr,
194 )
195 else:
196 # sync the problem repository
197 logger.info("$ git -C %s pull", path)
198 subprocess.check_call(
199 ["git", "-C", str(path), "pull"], # noqa: S607
200 stdout=sys.stderr,
201 stderr=sys.stderr,
202 )
204 LibraryCheckerProblem._is_repository_updated.add(self.repo_path)
207class _YukicoderProblemNo(int):
208 def __new__(cls, value: int):
209 return super().__new__(cls, value)
211 def __str__(self) -> str:
212 return "no/" + super().__str__()
215class _YukicoderProblemId(int):
216 def __new__(cls, value: int):
217 return super().__new__(cls, value)
220class YukicoderProblem(_BaseProblem):
221 problem: _YukicoderProblemNo | _YukicoderProblemId
223 def __init__(self, *, problem_no: int | None = None, problem_id: int | None = None):
224 if problem_no is not None:
225 self.problem = _YukicoderProblemNo(problem_no)
226 elif problem_id is not None: 226 ↛ 229line 226 didn't jump to line 229 because the condition on line 226 was always true
227 self.problem = _YukicoderProblemId(problem_id)
228 else:
229 raise ValueError("Needs problem_no or problem_id")
231 @staticmethod
232 def _env_float(name: str, default: float) -> float:
233 value = os.environ.get(name)
234 if value is None or value == "":
235 return default
236 try:
237 return float(value)
238 except ValueError as e:
239 raise ValueError(f"{name} must be a float: {value!r}") from e
241 @staticmethod
242 def _env_int(name: str, default: int) -> int:
243 value = os.environ.get(name)
244 if value is None or value == "":
245 return default
246 try:
247 return int(value)
248 except ValueError as e:
249 raise ValueError(f"{name} must be an integer: {value!r}") from e
251 @staticmethod
252 def _validate_yukicoder_token(token: str) -> str:
253 if not token: 253 ↛ 254line 253 didn't jump to line 254 because the condition on line 253 was never true
254 raise NotLoggedInError("Required: $YUKICODER_TOKEN environment variable")
256 if token.startswith("YUKICODER_TOKEN="):
257 raise NotLoggedInError(
258 "YUKICODER_TOKEN must contain only the token value, "
259 "not a 'YUKICODER_TOKEN=...' assignment."
260 )
262 if token.lower().startswith("bearer "): 262 ↛ 263line 262 didn't jump to line 263 because the condition on line 262 was never true
263 raise NotLoggedInError(
264 "YUKICODER_TOKEN must contain only the token value, "
265 "not a 'Bearer ...' authorization header value."
266 )
268 if token != token.strip():
269 raise NotLoggedInError(
270 "YUKICODER_TOKEN contains leading or trailing whitespace."
271 )
273 bad_control_chars: list[str] = []
274 for index, ch in enumerate(token):
275 code = ord(ch)
276 if code < _ASCII_SPACE or code == _ASCII_DELETE:
277 name = {
278 0x09: "TAB",
279 0x0A: "LF",
280 0x0D: "CR",
281 0x1B: "ESC",
282 _ASCII_DELETE: "DEL",
283 }.get(code, "control")
284 bad_control_chars.append(f"offset {index}: 0x{code:02X} ({name})")
286 if bad_control_chars:
287 extra = ""
288 if "\x1b[" in token: 288 ↛ 294line 288 didn't jump to line 294 because the condition on line 288 was always true
289 extra = (
290 " It looks like an ANSI escape sequence was included, "
291 "possibly by pressing an arrow key while entering the token."
292 )
294 raise NotLoggedInError(
295 "YUKICODER_TOKEN contains control characters: "
296 + ", ".join(bad_control_chars)
297 + "."
298 + extra
299 )
301 non_visible_ascii: list[str] = []
302 for index, ch in enumerate(token):
303 code = ord(ch)
304 if code <= _ASCII_SPACE or code >= _ASCII_DELETE: 304 ↛ 305line 304 didn't jump to line 305 because the condition on line 304 was never true
305 non_visible_ascii.append(f"offset {index}: U+{code:04X}")
307 if non_visible_ascii: 307 ↛ 308line 307 didn't jump to line 308 because the condition on line 307 was never true
308 raise NotLoggedInError(
309 "YUKICODER_TOKEN contains characters that are not visible ASCII: "
310 + ", ".join(non_visible_ascii)
311 )
313 return token
315 @classmethod
316 def yukicoder_headers(cls) -> dict[str, str]:
317 token = cls._validate_yukicoder_token(os.environ.get("YUKICODER_TOKEN", ""))
318 return {"Authorization": f"Bearer {token}"}
320 def download_system_cases(self) -> Iterable[TestCaseData] | bool:
321 test_directory = self.test_directory
322 if test_directory.exists() and any(test_directory.iterdir()):
323 logger.info("download:already exists: %s", self.url)
324 return True
326 headers = self.yukicoder_headers()
327 if not self._is_logged_in(headers=headers):
328 raise NotLoggedInError("Required: $YUKICODER_TOKEN environment variable")
330 self.problem_directory.parent.mkdir(parents=True, exist_ok=True)
332 tmp_root = pathlib.Path(
333 tempfile.mkdtemp(
334 prefix=f"{self.hash_id}.",
335 dir=self.problem_directory.parent,
336 )
337 )
338 zip_path = tmp_root / "testcase.zip"
339 staging_directory = tmp_root / "test"
341 try:
342 self._download_testcase_zip(zip_path, headers=headers)
343 case_count = self._extract_testcase_zip(zip_path, staging_directory)
345 if case_count == 0:
346 logger.error(
347 "Sample not found",
348 extra={"github": GitHubMessageParams()},
349 )
350 return False
352 self.problem_directory.mkdir(parents=True, exist_ok=True)
354 if test_directory.exists():
355 shutil.rmtree(test_directory)
357 staging_directory.rename(test_directory)
358 logger.info("download:saved: %s cases: %s", case_count, self.url)
359 return True
360 finally:
361 shutil.rmtree(tmp_root, ignore_errors=True)
363 def _download_cases(self) -> list[TestCaseData]:
364 headers = self.yukicoder_headers()
365 if not self._is_logged_in(headers=headers):
366 raise NotLoggedInError("Required: $YUKICODER_TOKEN environment variable")
368 self.problem_directory.parent.mkdir(parents=True, exist_ok=True)
370 tmp_root = pathlib.Path(
371 tempfile.mkdtemp(
372 prefix=f"{self.hash_id}.",
373 dir=self.problem_directory.parent,
374 )
375 )
376 zip_path = tmp_root / "testcase.zip"
378 try:
379 self._download_testcase_zip(zip_path, headers=headers)
380 with zipfile.ZipFile(zip_path) as fh:
381 inputs: dict[str, bytes] = {}
382 outputs: dict[str, bytes] = {}
384 for info in fh.infolist():
385 filename = info.filename
386 if filename.endswith("/"):
387 continue
389 path = pathlib.PurePosixPath(filename)
390 self._validate_zip_member_path(path)
392 if filename.startswith("test_in/"):
393 inputs[path.stem] = fh.read(info)
394 elif filename.startswith("test_out/"):
395 outputs[path.stem] = fh.read(info)
397 return [
398 TestCaseData(name=name, input_data=i, output_data=o)
399 for name, i, o in enumerate_input_outputs(inputs, outputs)
400 ]
401 finally:
402 shutil.rmtree(tmp_root, ignore_errors=True)
404 def _download_testcase_zip(
405 self,
406 destination: pathlib.Path,
407 *,
408 headers: dict[str, str] | None,
409 ) -> None:
410 url = f"{self.url}/testcase.zip"
412 connect_timeout = self._env_float(
413 "COMPETITIVE_VERIFIER_YUKICODER_CONNECT_TIMEOUT",
414 10.0,
415 )
416 read_timeout = self._env_float(
417 "COMPETITIVE_VERIFIER_YUKICODER_READ_TIMEOUT",
418 30.0,
419 )
420 download_timeout = self._env_float(
421 "COMPETITIVE_VERIFIER_YUKICODER_DOWNLOAD_TIMEOUT",
422 0.0,
423 )
424 report_interval = self._env_float(
425 "COMPETITIVE_VERIFIER_YUKICODER_REPORT_INTERVAL",
426 5.0,
427 )
428 chunk_size = self._env_int(
429 "COMPETITIVE_VERIFIER_YUKICODER_CHUNK_SIZE",
430 1024 * 1024,
431 )
433 logger.info("download:yukicoder testcase.zip: %s", url)
435 started_at = time.perf_counter()
436 last_reported_at = started_at
437 total = 0
439 with requests.get(
440 url,
441 headers=headers,
442 allow_redirects=True,
443 stream=True,
444 timeout=(connect_timeout, read_timeout),
445 ) as resp:
446 resp.raise_for_status()
448 content_length_text = resp.headers.get("content-length")
449 content_length = (
450 int(content_length_text)
451 if content_length_text is not None and content_length_text.isdigit()
452 else None
453 )
455 logger.info(
456 "download:yukicoder response: status=%s, content-type=%s, content-length=%s",
457 resp.status_code,
458 resp.headers.get("content-type"),
459 content_length_text,
460 )
462 with destination.open("wb") as out:
463 for chunk in resp.iter_content(chunk_size=chunk_size):
464 if not chunk:
465 continue
467 out.write(chunk)
468 total += len(chunk)
470 now = time.perf_counter()
471 elapsed = now - started_at
473 if download_timeout > 0 and elapsed > download_timeout:
474 raise TimeoutError(
475 f"download timeout: {url}: "
476 f"{total} bytes in {download_timeout:.0f} sec"
477 )
479 if (
480 report_interval > 0
481 and now - last_reported_at >= report_interval
482 ):
483 mib = total / 1024 / 1024
484 speed = mib / elapsed if elapsed > 0 else 0.0
486 if content_length:
487 percent = total * 100.0 / content_length
488 logger.info(
489 "download:yukicoder progress: %.1f / %.1f MiB, %.1f%%, %.2f MiB/s",
490 mib,
491 content_length / 1024 / 1024,
492 percent,
493 speed,
494 )
495 else:
496 logger.info(
497 "download:yukicoder progress: %.1f MiB, %.2f MiB/s",
498 mib,
499 speed,
500 )
502 last_reported_at = now
504 elapsed = time.perf_counter() - started_at
505 mib = total / 1024 / 1024
506 speed = mib / elapsed if elapsed > 0 else 0.0
507 logger.info(
508 "download:yukicoder done: %.1f MiB, %.2f MiB/s, %.1f sec",
509 mib,
510 speed,
511 elapsed,
512 )
514 def _extract_testcase_zip(
515 self,
516 zip_path: pathlib.Path,
517 destination: pathlib.Path,
518 ) -> int:
519 inputs: dict[str, zipfile.ZipInfo] = {}
520 outputs: dict[str, zipfile.ZipInfo] = {}
522 with zipfile.ZipFile(zip_path) as fh:
523 for info in fh.infolist():
524 filename = info.filename
525 if filename.endswith("/"):
526 continue
528 path = pathlib.PurePosixPath(filename)
529 self._validate_zip_member_path(path)
531 if filename.startswith("test_in/"):
532 inputs[path.stem] = info
533 elif filename.startswith("test_out/"):
534 outputs[path.stem] = info
536 common_names = sorted(inputs.keys() & outputs.keys())
538 if len(inputs) != len(common_names) or len(outputs) != len(common_names):
539 logger.warning("dangling output case")
541 if not common_names:
542 logger.warning("no cases found")
543 return 0
545 destination.mkdir(parents=True, exist_ok=False)
547 for name in common_names:
548 input_path = destination / _name_to_filename(name, "in")
549 output_path = destination / _name_to_filename(name, "out")
551 with fh.open(inputs[name]) as src, input_path.open("wb") as out:
552 shutil.copyfileobj(src, out)
554 with fh.open(outputs[name]) as src, output_path.open("wb") as out:
555 shutil.copyfileobj(src, out)
557 return len(common_names)
559 @staticmethod
560 def _validate_zip_member_path(path: pathlib.PurePosixPath) -> None:
561 if path.is_absolute() or ".." in path.parts:
562 raise RuntimeError(f"unsafe path in testcase.zip: {path}")
564 @property
565 def url(self) -> str:
566 return f"https://yukicoder.me/problems/{self.problem}"
568 @classmethod
569 def from_url(cls, url: str) -> Optional["YukicoderProblem"]:
570 # example: https://yukicoder.me/problems/no/499
571 # example: http://yukicoder.me/problems/1476
572 result = urllib.parse.urlparse(url)
573 dirname, basename = posixpath.split(normalize_url_path(result.path))
574 if result.scheme in ("", "http", "https") and result.netloc == "yukicoder.me":
575 try:
576 n = int(basename)
577 except ValueError:
578 pass
579 else:
580 if dirname == "/problems/no":
581 return cls(problem_no=n)
582 if dirname == "/problems":
583 return cls(problem_id=n)
584 return None
586 def _is_logged_in(self, *, headers: dict[str, str] | None = None) -> bool:
587 url = "https://yukicoder.me"
588 resp = requests.get(url, headers=headers, allow_redirects=True, timeout=10)
589 resp.raise_for_status()
590 return "login-btn" not in str(resp.content)
593class AOJProblem(_BaseProblem):
594 def __init__(self, *, problem_id: str):
595 self.problem_id = problem_id
597 def _download_cases(self) -> Iterable[TestCaseData]:
598 return AOJProblem.download_cases(self.problem_id)
600 @staticmethod
601 def download_cases(problem_id: str) -> Iterable[TestCaseData]:
602 # get header
603 # reference: http://developers.u-aizu.ac.jp/api?key=judgedat%2Ftestcases%2F%7BproblemId%7D%2Fheader_GET
604 url = f"https://judgedat.u-aizu.ac.jp/testcases/{problem_id}/header"
605 resp = requests.get(url, allow_redirects=True, timeout=10)
606 resp.raise_for_status()
607 header_res = json.loads(resp.text)
609 # get testcases via the official API
610 for header in header_res["headers"]:
611 # 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)"
612 # 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)
613 serial = header["serial"]
614 url = f"https://judgedat.u-aizu.ac.jp/testcases/{problem_id}/{serial}"
616 resp_in = requests.get(url + "/in", allow_redirects=True, timeout=10)
617 resp_in.raise_for_status()
618 resp_out = requests.get(url + "/out", allow_redirects=True, timeout=10)
619 resp_out.raise_for_status()
621 yield TestCaseData(
622 header["name"],
623 resp_in.content,
624 resp_out.content,
625 )
627 @property
628 def url(self) -> str:
629 return f"http://judge.u-aizu.ac.jp/onlinejudge/description.jsp?id={self.problem_id}"
631 @classmethod
632 def from_url(cls, url: str) -> Optional["AOJProblem"]:
633 result = urllib.parse.urlparse(url)
635 # example: http://judge.u-aizu.ac.jp/onlinejudge/description.jsp?id=1169
636 # example: http://judge.u-aizu.ac.jp/onlinejudge/description.jsp?id=DSL_1_A&lang=jp
637 querystring = urllib.parse.parse_qs(result.query)
638 if (
639 result.scheme in ("", "http", "https")
640 and result.netloc == "judge.u-aizu.ac.jp"
641 and normalize_url_path(result.path) == "/onlinejudge/description.jsp"
642 and querystring.get("id")
643 and len(querystring["id"]) == 1
644 ):
645 (n,) = querystring["id"]
646 return cls(problem_id=n)
648 # example: https://onlinejudge.u-aizu.ac.jp/challenges/sources/JAG/Prelim/2881
649 # example: https://onlinejudge.u-aizu.ac.jp/courses/library/4/CGL/3/CGL_3_B
650 m = re.match(
651 r"^/(challenges|courses)/(sources|library/\d+|lesson/\d+)/(\w+)/(\w+)/(\w+)$",
652 normalize_url_path(result.path),
653 )
654 if (
655 result.scheme in ("", "http", "https")
656 and result.netloc == "onlinejudge.u-aizu.ac.jp"
657 and m
658 ):
659 n = m.group(5)
660 return cls(problem_id=n)
662 # example: https://onlinejudge.u-aizu.ac.jp/problems/0423
663 # example: https://onlinejudge.u-aizu.ac.jp/problems/CGL_3_B
664 m = re.match(r"^/problems/(\w+)$", normalize_url_path(result.path))
665 if (
666 result.scheme in ("", "http", "https")
667 and result.netloc == "onlinejudge.u-aizu.ac.jp"
668 and m
669 ):
670 n = m.group(1)
671 return cls(problem_id=n)
673 return None
676class AOJArenaProblem(_BaseProblem):
677 def __init__(self, *, arena_id: str, alphabet: str):
678 if len(alphabet) != 1 or not alphabet.isupper(): 678 ↛ 679line 678 didn't jump to line 679 because the condition on line 678 was never true
679 raise ValueError(arena_id, alphabet)
680 self.arena_id = arena_id
681 self.alphabet = alphabet
683 self._problem_id: str | None = None
685 def get_problem_id(self) -> str:
686 if self._problem_id is None:
687 url = f"https://judgeapi.u-aizu.ac.jp/arenas/{self.arena_id}/problems"
688 resp = requests.get(url, allow_redirects=True, timeout=10)
689 resp.raise_for_status()
690 problems = json.loads(resp.text)
691 for problem in problems:
692 if problem["id"] == self.alphabet:
693 p = problem["problemId"]
694 logger.debug("problem: %s", p)
695 self._problem_id = p
696 return p
697 raise ValueError("Problem is not found.")
698 return self._problem_id
700 def _download_cases(self) -> Iterable[TestCaseData]:
701 return AOJProblem.download_cases(self.get_problem_id())
703 @property
704 def url(self) -> str:
705 return f"https://onlinejudge.u-aizu.ac.jp/services/room.html#{self.arena_id}/problems/{self.alphabet}"
707 @classmethod
708 def from_url(cls, url: str) -> Optional["AOJArenaProblem"]:
709 # example: https://onlinejudge.u-aizu.ac.jp/services/room.html#RitsCamp19Day2/problems/A
710 result = urllib.parse.urlparse(url)
711 if (
712 result.scheme in ("", "http", "https")
713 and result.netloc == "onlinejudge.u-aizu.ac.jp"
714 and normalize_url_path(result.path) == "/services/room.html"
715 ):
716 fragment = result.fragment.split("/")
717 if len(fragment) == 3 and fragment[1] == "problems": # noqa: PLR2004 717 ↛ 719line 717 didn't jump to line 719 because the condition on line 717 was always true
718 return cls(arena_id=fragment[0], alphabet=fragment[2].upper())
719 return None
722@dataclass
723class LocalProblem(TestCaseProvider):
724 path: pathlib.Path
726 def download_system_cases(self) -> Iterable[TestCaseData] | bool:
727 return bool(any(self.iter_system_cases()))
729 def iter_system_cases(self) -> Iterable[TestCaseFile]:
730 return iter_testcases(directory=self.path, recursive=True)
733def normalize_url_path(path: str) -> str:
734 """A wrapper of posixpath.normpath.
736 posixpath.normpath doesn't collapse a leading duplicated slashes.
737 """
738 path = posixpath.normpath(path)
739 if path.startswith("//"):
740 path = "/" + path.lstrip("/")
741 return path
744def _subclasses_recursive(cls: type[object]) -> Iterable[type[Problem]]:
745 for ch in cls.__subclasses__():
746 if issubclass(ch, Problem): 746 ↛ 745line 746 didn't jump to line 745 because the condition on line 746 was always true
747 yield ch
748 yield from _subclasses_recursive(ch)
751def problem_from_url(url: str) -> Problem | None:
752 for ch in set(_subclasses_recursive(Problem)):
753 if (problem := ch.from_url(url)) is not None:
754 return problem
755 return None
758_InputOutput = TypeVar("_InputOutput")
761def enumerate_input_outputs(
762 inputs: dict[str, _InputOutput],
763 outputs: dict[str, _InputOutput],
764) -> Iterator[tuple[str, _InputOutput, _InputOutput]]:
765 common_keys = inputs.keys() & outputs.keys()
766 if len(inputs) != len(common_keys) or len(outputs) != len(common_keys):
767 logger.warning("dangling output case")
769 if len(common_keys) == 0:
770 logger.warning("no cases found")
772 for key in sorted(common_keys):
773 yield (key, inputs[key], outputs[key])
776def merge_testcase_files(
777 inputs: dict[str, pathlib.Path],
778 outputs: dict[str, pathlib.Path],
779) -> Iterator[TestCaseFile]:
780 for name, i, o in enumerate_input_outputs(inputs, outputs):
781 yield TestCaseFile(name=name, input_path=i, output_path=o)
784def _casename(path: pathlib.Path, *, directory: pathlib.Path) -> str:
785 return path.relative_to(directory).with_suffix("").as_posix()
788def iter_testcases(
789 *, directory: pathlib.Path, recursive: bool = False
790) -> Iterator[TestCaseFile]:
791 inputs: dict[str, pathlib.Path] = {}
792 outputs: dict[str, pathlib.Path] = {}
793 pre = "**/" if recursive else ""
795 for path in directory.glob(pre + "*.in"):
796 if path.is_file(): 796 ↛ 795line 796 didn't jump to line 795 because the condition on line 796 was always true
797 inputs[_casename(path, directory=directory)] = path
798 for path in directory.glob(pre + "*.out"):
799 if path.is_file(): 799 ↛ 798line 799 didn't jump to line 798 because the condition on line 799 was always true
800 outputs[_casename(path, directory=directory)] = path
802 return merge_testcase_files(inputs, outputs)
805def _name_to_filename(name: str, ext: str):
806 return pathlib.Path(name).with_suffix(f".{ext}").name
809def save_testcases(samples: Iterable[TestCaseData], *, directory: pathlib.Path):
810 for sample in samples:
811 for data, ext in [(sample.input_data, "in"), (sample.output_data, "out")]:
812 path = directory / _name_to_filename(sample.name, ext)
814 if path.exists(): 814 ↛ 815line 814 didn't jump to line 815 because the condition on line 814 was never true
815 logger.error("Failed to download since file already exists: %s", path)
816 path.parent.mkdir(parents=True, exist_ok=True)
817 path.write_bytes(data)
818 logger.debug("saved to: %s", path)