Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
118 changes: 106 additions & 12 deletions scripts/classify_problems.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,6 @@
#!/usr/bin/env python3
import argparse
import contextlib
import os
import re
import time
Expand Down Expand Up @@ -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
Expand Down Expand Up @@ -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"""
Expand Down Expand Up @@ -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"""
Expand All @@ -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)
Expand All @@ -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.
"""

Expand Down Expand Up @@ -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(
Expand All @@ -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):
Expand Down Expand Up @@ -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
Expand All @@ -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
Expand Down
124 changes: 123 additions & 1 deletion scripts/test_classify_problems.py
Original file line number Diff line number Diff line change
@@ -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:
Expand Down Expand Up @@ -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()
Loading