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

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 

19 

20import requests 

21 

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) 

30 

31logger = getLogger(__name__) 

32 

33_ASCII_SPACE = 0x20 

34_ASCII_DELETE = 0x7F 

35 

36 

37class NotLoggedInError(RuntimeError): 

38 pass 

39 

40 

41class _BaseProblem(Problem): 

42 def iter_system_cases(self) -> Iterator[TestCaseFile]: 

43 return iter_testcases(directory=self.test_directory) 

44 

45 def is_testdata_cached(self) -> bool: 

46 return any(self.iter_system_cases()) 

47 

48 def download_system_cases(self) -> Iterable[TestCaseData] | bool: 

49 test_directory = self.test_directory 

50 

51 if self.is_testdata_cached(): 

52 logger.info("download:already exists: %s", self.url) 

53 return True 

54 

55 self.problem_directory.mkdir(parents=True, exist_ok=True) 

56 

57 samples = list(self._download_cases()) 

58 

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 

66 

67 # write samples to files 

68 save_testcases(samples, directory=test_directory) 

69 return samples 

70 

71 @abstractmethod 

72 def _download_cases(self) -> Iterable[TestCaseData]: ... 

73 

74 

75class LibraryCheckerProblem(Problem): 

76 checker_exe_name: ClassVar[str] = ( 

77 "checker.exe" if sys.platform == "win32" else "checker" 

78 ) 

79 

80 def __init__(self, *, problem_id: str): 

81 self.problem_id = problem_id 

82 self._source_directory = None 

83 

84 def __hash__(self) -> int: 

85 return hash((self.problem_id, self.repo_path)) 

86 

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 

91 

92 @property 

93 def repo_path(self): 

94 return config.get_cache_dir() / "library-checker-problems" 

95 

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) 

104 

105 def is_testdata_cached(self) -> bool: 

106 try: 

107 return any(self.iter_system_cases()) 

108 except RuntimeError: 

109 return False 

110 

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 

115 

116 @property 

117 def checker(self) -> pathlib.Path | None: 

118 return self.source_directory / self.checker_exe_name 

119 

120 def generate_test_cases(self) -> None: 

121 self.update_cloned_repository() 

122 path = self.repo_path 

123 

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 

135 

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 

147 

148 @property 

149 def url(self) -> str: 

150 return f"https://judge.yosupo.jp/problem/{self.problem_id}" 

151 

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 

164 

165 _is_repository_updated: ClassVar[set[pathlib.Path]] = set() 

166 

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 

170 

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 

184 

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 ) 

203 

204 LibraryCheckerProblem._is_repository_updated.add(self.repo_path) 

205 

206 

207class _YukicoderProblemNo(int): 

208 def __new__(cls, value: int): 

209 return super().__new__(cls, value) 

210 

211 def __str__(self) -> str: 

212 return "no/" + super().__str__() 

213 

214 

215class _YukicoderProblemId(int): 

216 def __new__(cls, value: int): 

217 return super().__new__(cls, value) 

218 

219 

220class YukicoderProblem(_BaseProblem): 

221 problem: _YukicoderProblemNo | _YukicoderProblemId 

222 

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") 

230 

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 

240 

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 

250 

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") 

255 

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 ) 

261 

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 ) 

267 

268 if token != token.strip(): 

269 raise NotLoggedInError( 

270 "YUKICODER_TOKEN contains leading or trailing whitespace." 

271 ) 

272 

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})") 

285 

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 ) 

293 

294 raise NotLoggedInError( 

295 "YUKICODER_TOKEN contains control characters: " 

296 + ", ".join(bad_control_chars) 

297 + "." 

298 + extra 

299 ) 

300 

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}") 

306 

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 ) 

312 

313 return token 

314 

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}"} 

319 

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 

325 

326 headers = self.yukicoder_headers() 

327 if not self._is_logged_in(headers=headers): 

328 raise NotLoggedInError("Required: $YUKICODER_TOKEN environment variable") 

329 

330 self.problem_directory.parent.mkdir(parents=True, exist_ok=True) 

331 

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" 

340 

341 try: 

342 self._download_testcase_zip(zip_path, headers=headers) 

343 case_count = self._extract_testcase_zip(zip_path, staging_directory) 

344 

345 if case_count == 0: 

346 logger.error( 

347 "Sample not found", 

348 extra={"github": GitHubMessageParams()}, 

349 ) 

350 return False 

351 

352 self.problem_directory.mkdir(parents=True, exist_ok=True) 

353 

354 if test_directory.exists(): 

355 shutil.rmtree(test_directory) 

356 

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) 

362 

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") 

367 

368 self.problem_directory.parent.mkdir(parents=True, exist_ok=True) 

369 

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" 

377 

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] = {} 

383 

384 for info in fh.infolist(): 

385 filename = info.filename 

386 if filename.endswith("/"): 

387 continue 

388 

389 path = pathlib.PurePosixPath(filename) 

390 self._validate_zip_member_path(path) 

391 

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) 

396 

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) 

403 

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" 

411 

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 ) 

432 

433 logger.info("download:yukicoder testcase.zip: %s", url) 

434 

435 started_at = time.perf_counter() 

436 last_reported_at = started_at 

437 total = 0 

438 

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() 

447 

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 ) 

454 

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 ) 

461 

462 with destination.open("wb") as out: 

463 for chunk in resp.iter_content(chunk_size=chunk_size): 

464 if not chunk: 

465 continue 

466 

467 out.write(chunk) 

468 total += len(chunk) 

469 

470 now = time.perf_counter() 

471 elapsed = now - started_at 

472 

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 ) 

478 

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 

485 

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 ) 

501 

502 last_reported_at = now 

503 

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 ) 

513 

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] = {} 

521 

522 with zipfile.ZipFile(zip_path) as fh: 

523 for info in fh.infolist(): 

524 filename = info.filename 

525 if filename.endswith("/"): 

526 continue 

527 

528 path = pathlib.PurePosixPath(filename) 

529 self._validate_zip_member_path(path) 

530 

531 if filename.startswith("test_in/"): 

532 inputs[path.stem] = info 

533 elif filename.startswith("test_out/"): 

534 outputs[path.stem] = info 

535 

536 common_names = sorted(inputs.keys() & outputs.keys()) 

537 

538 if len(inputs) != len(common_names) or len(outputs) != len(common_names): 

539 logger.warning("dangling output case") 

540 

541 if not common_names: 

542 logger.warning("no cases found") 

543 return 0 

544 

545 destination.mkdir(parents=True, exist_ok=False) 

546 

547 for name in common_names: 

548 input_path = destination / _name_to_filename(name, "in") 

549 output_path = destination / _name_to_filename(name, "out") 

550 

551 with fh.open(inputs[name]) as src, input_path.open("wb") as out: 

552 shutil.copyfileobj(src, out) 

553 

554 with fh.open(outputs[name]) as src, output_path.open("wb") as out: 

555 shutil.copyfileobj(src, out) 

556 

557 return len(common_names) 

558 

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}") 

563 

564 @property 

565 def url(self) -> str: 

566 return f"https://yukicoder.me/problems/{self.problem}" 

567 

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 

585 

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) 

591 

592 

593class AOJProblem(_BaseProblem): 

594 def __init__(self, *, problem_id: str): 

595 self.problem_id = problem_id 

596 

597 def _download_cases(self) -> Iterable[TestCaseData]: 

598 return AOJProblem.download_cases(self.problem_id) 

599 

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) 

608 

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}" 

615 

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() 

620 

621 yield TestCaseData( 

622 header["name"], 

623 resp_in.content, 

624 resp_out.content, 

625 ) 

626 

627 @property 

628 def url(self) -> str: 

629 return f"http://judge.u-aizu.ac.jp/onlinejudge/description.jsp?id={self.problem_id}" 

630 

631 @classmethod 

632 def from_url(cls, url: str) -> Optional["AOJProblem"]: 

633 result = urllib.parse.urlparse(url) 

634 

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) 

647 

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) 

661 

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) 

672 

673 return None 

674 

675 

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 

682 

683 self._problem_id: str | None = None 

684 

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 

699 

700 def _download_cases(self) -> Iterable[TestCaseData]: 

701 return AOJProblem.download_cases(self.get_problem_id()) 

702 

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}" 

706 

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 

720 

721 

722@dataclass 

723class LocalProblem(TestCaseProvider): 

724 path: pathlib.Path 

725 

726 def download_system_cases(self) -> Iterable[TestCaseData] | bool: 

727 return bool(any(self.iter_system_cases())) 

728 

729 def iter_system_cases(self) -> Iterable[TestCaseFile]: 

730 return iter_testcases(directory=self.path, recursive=True) 

731 

732 

733def normalize_url_path(path: str) -> str: 

734 """A wrapper of posixpath.normpath. 

735 

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 

742 

743 

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) 

749 

750 

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 

756 

757 

758_InputOutput = TypeVar("_InputOutput") 

759 

760 

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") 

768 

769 if len(common_keys) == 0: 

770 logger.warning("no cases found") 

771 

772 for key in sorted(common_keys): 

773 yield (key, inputs[key], outputs[key]) 

774 

775 

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) 

782 

783 

784def _casename(path: pathlib.Path, *, directory: pathlib.Path) -> str: 

785 return path.relative_to(directory).with_suffix("").as_posix() 

786 

787 

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 "" 

794 

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 

801 

802 return merge_testcase_files(inputs, outputs) 

803 

804 

805def _name_to_filename(name: str, ext: str): 

806 return pathlib.Path(name).with_suffix(f".{ext}").name 

807 

808 

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) 

813 

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)