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

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 

18 

19import requests 

20 

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) 

29 

30logger = getLogger(__name__) 

31 

32 

33class NotLoggedInError(RuntimeError): 

34 pass 

35 

36 

37class _BaseProblem(Problem): 

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

39 return iter_testcases(directory=self.test_directory) 

40 

41 def is_testdata_cached(self) -> bool: 

42 test_directory = self.test_directory 

43 return test_directory.exists() and any(test_directory.iterdir()) 

44 

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

46 test_directory = self.test_directory 

47 

48 if self.is_testdata_cached(): 

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

50 return True 

51 

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

53 

54 samples = list(self._download_cases()) 

55 

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 

63 

64 # write samples to files 

65 save_testcases(samples, directory=test_directory) 

66 return samples 

67 

68 @abstractmethod 

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

70 

71 

72class LibraryCheckerProblem(Problem): 

73 checker_exe_name: ClassVar[str] = ( 

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

75 ) 

76 

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

78 self.problem_id = problem_id 

79 self._source_directory = None 

80 

81 def __hash__(self) -> int: 

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

83 

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 

88 

89 @property 

90 def repo_path(self): 

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

92 

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) 

101 

102 def is_testdata_cached(self) -> bool: 

103 try: 

104 return any(self.iter_system_cases()) 

105 except RuntimeError: 

106 return False 

107 

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 

112 

113 @property 

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

115 return self.source_directory / self.checker_exe_name 

116 

117 def generate_test_cases(self) -> None: 

118 self.update_cloned_repository() 

119 path = self.repo_path 

120 

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 

132 

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" 

137 

138 def sync_testdata(self) -> None: 

139 self.update_cloned_repository() 

140 

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 

146 

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 

158 

159 @property 

160 def url(self) -> str: 

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

162 

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 

175 

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

177 

178 def update_cloned_repository(self) -> None: 

179 if self.repo_path in self._is_repository_updated: 

180 return 

181 

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 

195 

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 ) 

214 

215 LibraryCheckerProblem._is_repository_updated.add(self.repo_path) 

216 

217 

218class _YukicoderProblemNo(int): 

219 def __new__(cls, value: int): 

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

221 

222 def __str__(self) -> str: 

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

224 

225 

226class _YukicoderProblemId(int): 

227 def __new__(cls, value: int): 

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

229 

230 

231class YukicoderProblem(_BaseProblem): 

232 problem: _YukicoderProblemNo | _YukicoderProblemId 

233 

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

241 

242 def _download_cases(self) -> list[TestCaseData]: 

243 """Download yukicoder problem. 

244 

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

251 

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) 

256 

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 ] 

273 

274 @property 

275 def url(self) -> str: 

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

277 

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 

295 

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) 

301 

302 

303class AOJProblem(_BaseProblem): 

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

305 self.problem_id = problem_id 

306 

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

308 return AOJProblem.download_cases(self.problem_id) 

309 

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) 

318 

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

325 

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

330 

331 yield TestCaseData( 

332 header["name"], 

333 resp_in.content, 

334 resp_out.content, 

335 ) 

336 

337 @property 

338 def url(self) -> str: 

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

340 

341 @classmethod 

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

343 result = urllib.parse.urlparse(url) 

344 

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) 

357 

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) 

371 

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) 

382 

383 return None 

384 

385 

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 

392 

393 self._problem_id: str | None = None 

394 

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 

409 

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

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

412 

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

416 

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 

430 

431 

432@dataclass 

433class LocalProblem(TestCaseProvider): 

434 path: pathlib.Path 

435 

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

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

438 

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

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

441 

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

454 

455 

456def _normpath(path: str) -> str: 

457 """A wrapper of posixpath.normpath. 

458 

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 

465 

466 

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) 

471 

472 

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 

478 

479 

480_InOut = TypeVar("_InOut") 

481 

482 

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

490 

491 if len(common_keys) == 0: 

492 logger.warning("no cases found") 

493 

494 for key in sorted(common_keys): 

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

496 

497 

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) 

504 

505 

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

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

508 

509 

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

516 

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 

523 

524 return merge_testcase_files(inputs, outputs) 

525 

526 

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

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

529 

530 

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) 

535 

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)