diff --git a/scripts/classify_problems.py b/scripts/classify_problems.py index ef1d3c1..bddad02 100644 --- a/scripts/classify_problems.py +++ b/scripts/classify_problems.py @@ -1,5 +1,6 @@ #!/usr/bin/env python3 import argparse +import contextlib import os import re import time @@ -36,6 +37,8 @@ class ProblemFetcher: @staticmethod def get_problem_source(url): """Determine the problem source based on URL""" + if "dmoj.ca" in url: + return "dmoj" return "kattis" if "kattis.com" in url else "codeforces" @staticmethod @@ -130,6 +133,24 @@ def extract_kattis_info(url): "url": f"https://open.kattis.com/problems/{problem_id}", } + @staticmethod + def extract_dmoj_info(url): + """Extract problem information from a DMOJ URL""" + # Bare codes are rejected here for the same reason the web ingestion + # rejects them: DMOJ codes are opaque, so any stray word would match. + clean_url = re.sub(r"^(https?:\/\/)?(www\.)?", "", url) + + match = re.search(r"^dmoj\.ca\/problem\/([a-z0-9_]+)$", clean_url) + + if not match: + return None + + problem_id = match.group(1) + return { + "problemId": problem_id, + "url": f"https://dmoj.ca/problem/{problem_id}", + } + @staticmethod def fetch_codeforces_problem(problem_info): """Fetch problem data from Codeforces API and website""" @@ -252,6 +273,36 @@ def fetch_kattis_problem(problem_info): print(f"Error fetching from Kattis: {e}") return None + @staticmethod + def fetch_dmoj_problem(problem_info): + """Fetch problem data from the DMOJ API + + DMOJ problem pages are behind bot protection, so metadata comes from + the public API instead. It carries no statement text, so classification + falls back to the name and types alone. + """ + try: + headers = { + "User-Agent": "Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 (KHTML, like Gecko) Chrome/91.0.4472.124 Safari/537.36" + } + api_url = f"https://dmoj.ca/api/v2/problem/{problem_info['problemId']}" + response = requests.get(api_url, headers=headers) + + if response.status_code != 200: + print(f"Failed to fetch DMOJ problem: HTTP {response.status_code}") + return None + + problem = response.json().get("data", {}).get("object", {}) + + name = (problem.get("name") or "").strip() or problem_info["problemId"] + types = [t for t in problem.get("types", []) if isinstance(t, str)] + + return {"name": name, "tags": types, "statement": ""} + + except Exception as e: + print(f"Error fetching from DMOJ: {e}") + return None + @classmethod def fetch_problem_details(cls, problem_id, url): """Fetch problem details based on URL""" @@ -271,9 +322,23 @@ def fetch_problem_details(cls, problem_id, url): return None return cls.fetch_kattis_problem(problem_info) + elif source == "dmoj": + problem_info = cls.extract_dmoj_info(url) + if not problem_info: + print(f"Invalid DMOJ URL: {url}") + return None + return cls.fetch_dmoj_problem(problem_info) + return None +def is_degraded_metadata(name, url): + """Whether a row still carries the bare problem code the submit flow stores + when provider metadata could not be fetched.""" + info = ProblemFetcher.extract_dmoj_info(url) + return bool(info) and name == info["problemId"] + + def classify_problem(name, tags, statement="", client=None): """Use Gemini to classify a problem based on its name, tags, and statement.""" # Truncate statement if it's too long (Gemini has context limits) @@ -284,17 +349,17 @@ def classify_problem(name, tags, statement="", client=None): prompt = f""" Given the following competitive programming problem, classify it into ONE of these types: {", ".join(PROBLEM_TYPES)} - + Problem name: {name} Problem tags: {", ".join(tags) if tags else "None"} - + Problem statement: {statement if statement else "Not available"} - + IMPORTANT: Choose the FIRST category in the list that applies to this problem. - For example, if a problem could be both "geometry" and "math", choose "geometry" + For example, if a problem could be both "geometry" and "math", choose "geometry" since it appears first in the list. - + Return only the type name, nothing else. """ @@ -342,18 +407,32 @@ def main(): default=2.0, help="Delay in seconds between API calls", ) + parser.add_argument( + "--backfill-only", + action="store_true", + help="Only repair DMOJ rows stored without provider metadata, leaving type untouched (needs no Gemini key)", + ) args = parser.parse_args() # Connect to the database using psycopg3 - with ( - genai.Client(api_key=GEMINI_API_KEY) as gemini_client, - psycopg.connect(SUPABASE_CONN) as conn, - ): + with contextlib.ExitStack() as stack: + gemini_client = ( + None + if args.backfill_only + else stack.enter_context(genai.Client(api_key=GEMINI_API_KEY)) + ) + conn = stack.enter_context(psycopg.connect(SUPABASE_CONN)) + # Create a cursor with conn.cursor() as cursor: try: # Query problems - if args.all: + if args.backfill_only: + # Only DMOJ rows can be degraded, so leave the rest alone. + cursor.execute( + "SELECT id, name, tags, url FROM problems WHERE url LIKE '%dmoj.ca%'" + ) + elif args.all: cursor.execute("SELECT id, name, tags, url FROM problems") else: cursor.execute( @@ -366,7 +445,8 @@ def main(): print("No problems found to classify.") return - print(f"Found {len(problems)} problems to classify.") + verb = "check" if args.backfill_only else "classify" + print(f"Found {len(problems)} problems to {verb}.") # Process problems in batches to avoid rate limiting for i in range(0, len(problems), args.batch_size): @@ -394,9 +474,23 @@ def main(): else: print(" Could not retrieve problem statement") + # A row whose name is still its bare problem code was + # stored without provider metadata, so refresh it here. + if details and is_degraded_metadata(name, url): + name = details["name"] + tags = details["tags"] + cursor.execute( + "UPDATE problems SET name = %s, tags = %s WHERE id = %s", + (name, tags, problem_id), + ) + print(f" → Backfilled metadata: {name}") + # Add a delay to avoid rate limiting time.sleep(args.delay) + if args.backfill_only: + continue + # Classify the problem problem_type = classify_problem( name, tags, statement, client=gemini_client @@ -413,7 +507,7 @@ def main(): conn.commit() print(f"Batch {i // args.batch_size + 1} completed and committed.") - print(f"Successfully classified {len(problems)} problems.") + print(f"Successfully processed {len(problems)} problems.") except Exception as e: # Rollback happens automatically with context manager on exception diff --git a/scripts/test_classify_problems.py b/scripts/test_classify_problems.py index 03231ba..b619e41 100644 --- a/scripts/test_classify_problems.py +++ b/scripts/test_classify_problems.py @@ -1,7 +1,13 @@ import unittest from types import SimpleNamespace +from unittest.mock import patch -from classify_problems import GEMINI_MODEL, classify_problem +from classify_problems import ( + GEMINI_MODEL, + ProblemFetcher, + classify_problem, + is_degraded_metadata, +) class FakeModels: @@ -45,5 +51,121 @@ def test_sdk_error_falls_back_to_misc(self): self.assertEqual(classify_problem("Example", [], client=client), "misc") +class FakeResponse: + def __init__(self, status_code=200, payload=None): + self.status_code = status_code + self._payload = payload if payload is not None else {} + + def json(self): + return self._payload + + +class DmojFetcherTest(unittest.TestCase): + def test_source_detection_covers_all_three_providers(self): + self.assertEqual( + ProblemFetcher.get_problem_source("https://dmoj.ca/problem/ciw26p2"), "dmoj" + ) + self.assertEqual( + ProblemFetcher.get_problem_source("https://open.kattis.com/problems/hello"), + "kattis", + ) + self.assertEqual( + ProblemFetcher.get_problem_source( + "https://codeforces.com/contest/1/problem/A" + ), + "codeforces", + ) + + def test_extracts_only_fully_qualified_dmoj_urls(self): + for url in [ + "https://dmoj.ca/problem/ciw26p2", + "https://www.dmoj.ca/problem/ciw26p2", + "dmoj.ca/problem/ciw26p2", + ]: + self.assertEqual( + ProblemFetcher.extract_dmoj_info(url), + { + "problemId": "ciw26p2", + "url": "https://dmoj.ca/problem/ciw26p2", + }, + ) + + for url in [ + "ciw26p2", + "https://dmoj.ca/problem/CIW26P2", + "https://evil.example/dmoj.ca/problem/ciw26p2", + "https://dmoj.ca/problem/ciw26p2/extra", + ]: + self.assertIsNone(ProblemFetcher.extract_dmoj_info(url), url) + + def test_reads_name_and_types_from_the_api_payload(self): + payload = { + "data": {"object": {"name": "CIW '26 P2", "types": ["Simulation", 7]}} + } + with patch( + "classify_problems.requests.get", return_value=FakeResponse(payload=payload) + ) as get: + details = ProblemFetcher.fetch_dmoj_problem({"problemId": "ciw26p2"}) + + self.assertEqual( + details, {"name": "CIW '26 P2", "tags": ["Simulation"], "statement": ""} + ) + self.assertEqual( + get.call_args[0][0], "https://dmoj.ca/api/v2/problem/ciw26p2" + ) + + def test_missing_name_falls_back_to_the_problem_code(self): + with patch( + "classify_problems.requests.get", + return_value=FakeResponse(payload={"data": {"object": {}}}), + ): + details = ProblemFetcher.fetch_dmoj_problem({"problemId": "ciw26p2"}) + + self.assertEqual(details["name"], "ciw26p2") + self.assertEqual(details["tags"], []) + + def test_upstream_and_transport_failures_return_none(self): + with patch( + "classify_problems.requests.get", return_value=FakeResponse(status_code=403) + ): + self.assertIsNone(ProblemFetcher.fetch_dmoj_problem({"problemId": "x"})) + + with patch( + "classify_problems.requests.get", side_effect=RuntimeError("offline") + ): + self.assertIsNone(ProblemFetcher.fetch_dmoj_problem({"problemId": "x"})) + + def test_routes_dmoj_urls_through_the_api_fetcher(self): + with patch.object( + ProblemFetcher, "fetch_dmoj_problem", return_value={"name": "ok"} + ) as fetch: + details = ProblemFetcher.fetch_problem_details( + 1, "https://dmoj.ca/problem/ciw26p2" + ) + + self.assertEqual(details, {"name": "ok"}) + self.assertEqual(fetch.call_args[0][0]["problemId"], "ciw26p2") + + self.assertIsNone( + ProblemFetcher.fetch_problem_details(1, "https://dmoj.ca/problem/BAD") + ) + + +class DegradedMetadataTest(unittest.TestCase): + def test_only_a_bare_problem_code_counts_as_degraded(self): + url = "https://dmoj.ca/problem/ciw26p2" + self.assertTrue(is_degraded_metadata("ciw26p2", url)) + self.assertFalse(is_degraded_metadata("CIW '26 P2 - Number Shuffle", url)) + + def test_non_dmoj_rows_are_never_treated_as_degraded(self): + # Guards the backfill from touching curated Kattis and Codeforces names. + self.assertFalse( + is_degraded_metadata("hello", "https://open.kattis.com/problems/hello") + ) + self.assertFalse( + is_degraded_metadata(None, "https://codeforces.com/contest/1/problem/A") + ) + + if __name__ == "__main__": unittest.main()