diff --git a/autosetter/cli.py b/autosetter/cli.py index e23bb6c..8d6fea0 100644 --- a/autosetter/cli.py +++ b/autosetter/cli.py @@ -6,9 +6,10 @@ Executes: 1. Intake: Read statement image/PDF. 2. Extraction: Extract and validate `problem.json` via Qwen vision model (1st model — UNCHANGED). -3. Generation: Produce artifacts via Ollama (validator.cpp, generator.cpp, solution.cpp, etc.) -4. Validation: Sandboxed compilation, sample verification, test case generation, and checker probing. -5. Packaging: Assemble release package bundle with manifest. +3. Similarity: Look up similar existing problems in the vector database (optional). +4. Generation: Produce artifacts via Ollama (validator.cpp, generator.cpp, solution.cpp, etc.) +5. Validation: Sandboxed compilation, sample verification, test case generation, and checker probing. +6. Packaging: Assemble release package bundle with manifest. """ from __future__ import annotations @@ -18,9 +19,9 @@ import logging import os import sys -from dataclasses import dataclass +from dataclasses import dataclass, field from pathlib import Path -from typing import Callable, List, Optional +from typing import Any, Callable, Dict, List, Optional from autosetter.config import ( DEFAULT_NUM_TESTS, @@ -29,6 +30,8 @@ DEFAULT_TEXT_MODEL, DEFAULT_VISION_MODEL, PROMPTS_DIR, + SIMILARITY_ENABLED, + SIMILARITY_TOP_K, ) from autosetter.extractor import ( JSONExtractionError, @@ -41,6 +44,11 @@ from autosetter.packager import Packager, PackagerError from autosetter.pipeline import PipelineError, TestPipeline, TestReport from autosetter.sandbox import SandboxError, SandboxLocalClient, ensure_testlib +from autosetter.similarity import ( + SimilaritySearchError, + find_similar_problems, + save_similar_problems, +) from autosetter.vision import ImageParsingError logger = logging.getLogger(__name__) @@ -62,6 +70,7 @@ class PipelineResult: package_dir: Path report: Optional[TestReport] = None validation_error: str = "" + similar_problems: List[Dict[str, Any]] = field(default_factory=list) @property def ready_for_release(self) -> bool: @@ -95,6 +104,8 @@ def generate_from_image( out_dir: str | Path = DEFAULT_OUT_DIR, prompts_dir: str | Path = PROMPTS_DIR, progress_callback: Optional[Callable[[str], None]] = None, + similarity_check: bool = SIMILARITY_ENABLED, + similar_k: int = SIMILARITY_TOP_K, ) -> PipelineResult: """ Execute the end-to-end AutoSetter problem packaging pipeline. @@ -119,6 +130,10 @@ def generate_from_image( Directory containing prompt templates. progress_callback : Optional[Callable[[str], None]] Progress reporting callback. + similarity_check : bool + If True, look up similar existing problems in the vector database (Qdrant). + similar_k : int + Number of similar problems to return. Returns ------- @@ -163,6 +178,18 @@ def generate_from_image( except JSONExtractionError as exc: raise AutoSetterError(f"Failed to save problem.json: {exc}") from exc + # 3b. Similar-problem lookup in the vector database (optional, never fatal) + similar_problems: List[Dict[str, Any]] = [] + if similarity_check: + logger_fn("Searching vector database for similar problems...") + try: + similar_problems = find_similar_problems(problem_data, k=similar_k) + save_similar_problems(similar_problems, output_root / "similar_problems.json") + for match in similar_problems: + logger_fn(f" {match['score']:.4f} {match['title']} {match['url'] or 'N/A'}") + except SimilaritySearchError as exc: + logger_fn(f"⚠️ Similar-problem search skipped: {exc}") + # 4. Generate Downstream Artifacts & Validation Loop test_report = None validation_error = "" @@ -287,6 +314,7 @@ def generate_from_image( package_dir=package_dir, report=test_report, validation_error=validation_error, + similar_problems=similar_problems, ) @@ -334,6 +362,18 @@ def build_arg_parser() -> argparse.ArgumentParser: default=False, help="Skip sandbox compilation and validation stage.", ) + parser.add_argument( + "--no-similarity", + action="store_true", + default=not SIMILARITY_ENABLED, + help="Skip the similar-problem search in the vector database.", + ) + parser.add_argument( + "--similar-k", + type=int, + default=SIMILARITY_TOP_K, + help=f"Number of similar problems to retrieve (default: {SIMILARITY_TOP_K}).", + ) parser.add_argument( "--out-dir", type=str, @@ -369,6 +409,8 @@ def main(argv: Optional[List[str]] = None) -> int: num_tests=args.num_tests, skip_validation=args.skip_validation, out_dir=args.out_dir, + similarity_check=not args.no_similarity, + similar_k=args.similar_k, ) except AutoSetterError as exc: print(f"Error: {exc}", file=sys.stderr) @@ -379,6 +421,10 @@ def main(argv: Optional[List[str]] = None) -> int: print(f"\nArtifacts written to: {result.generated_dir}") print(f"Package assembled at: {result.package_dir}") + if result.similar_problems: + print("Similar existing problems:") + for match in result.similar_problems: + print(f" {match['score']:.4f} {match['title']} {match['url'] or 'N/A'}") if result.ready_for_release: print(f"Ready for release — {result.summary}") diff --git a/autosetter/config.py b/autosetter/config.py index df2d373..b122883 100644 --- a/autosetter/config.py +++ b/autosetter/config.py @@ -65,6 +65,15 @@ SUPPORTED_PDF_EXTENSIONS = {".pdf"} SUPPORTED_EXTENSIONS = SUPPORTED_RASTER_EXTENSIONS | SUPPORTED_PDF_EXTENSIONS +# Similar-problem search (vector_database/ + Qdrant) +# Must match the model and collection used to build the vector database. +SIMILARITY_ENABLED = os.environ.get("AUTOSETTER_SIMILARITY", "1") != "0" +SIMILARITY_TOP_K = int(os.environ.get("AUTOSETTER_SIMILARITY_K", "5")) +QDRANT_URL = os.environ.get("QDRANT_URL", "http://localhost:6333") +QDRANT_COLLECTION = os.environ.get("QDRANT_COLLECTION", "competitive_programming_problems") +QDRANT_API_KEY = os.environ.get("QDRANT_API_KEY") or None +EMBEDDING_MODEL = os.environ.get("EMBEDDING_MODEL", "BAAI/bge-small-en-v1.5") + # Codeforces Polygon API POLYGON_API_URL = os.environ.get("POLYGON_API_URL", "https://polygon.codeforces.com/api/") POLYGON_API_KEY = os.environ.get("POLYGON_API_KEY", "") diff --git a/autosetter/similarity.py b/autosetter/similarity.py new file mode 100644 index 0000000..a21c766 --- /dev/null +++ b/autosetter/similarity.py @@ -0,0 +1,91 @@ +""" +autosetter.similarity +===================== +Similar-problem lookup against the vector database built by `vector_database/`. + +After extraction, the problem.json is embedded the same way the stored problems +were and searched in Qdrant, so the setter can see whether the problem (or a +close variant) already exists, with links to the matches. + +This stage is optional: the heavy dependencies (sentence-transformers, torch, +qdrant-client) live in `vector_database/requirements.txt`, and a missing +dependency or unreachable Qdrant only skips the check. +""" + +from __future__ import annotations + +import json +import sys +from pathlib import Path +from typing import Any, Dict, List + +from autosetter.config import ( + EMBEDDING_MODEL, + PROJECT_ROOT, + QDRANT_API_KEY, + QDRANT_COLLECTION, + QDRANT_URL, + SIMILARITY_TOP_K, +) + + +class SimilaritySearchError(Exception): + """Raised when the similar-problem lookup cannot run.""" + + +def _as_text(value: Any) -> str: + if isinstance(value, list): + return "\n".join(str(v) for v in value) + return str(value or "").strip() + + +def find_similar_problems( + problem_data: Dict[str, Any], + k: int = SIMILARITY_TOP_K, +) -> List[Dict[str, Any]]: + """ + Return the top-k stored problems most similar to `problem_data`. + + Each match is a dict with score, title, source, source_id, difficulty, tags and url. + """ + # vector_database/ sits at the repository root, next to autosetter/ + if str(PROJECT_ROOT) not in sys.path: + sys.path.append(str(PROJECT_ROOT)) + try: + from vector_database.processing.schema import Problem + from vector_database.search import ProblemSearcher + except ImportError as exc: + raise SimilaritySearchError( + f"vector database dependencies are not installed ({exc}); " + "run `pip install -r vector_database/requirements.txt`" + ) from exc + + problem = Problem( + id="autosetter_query", + source="autosetter", + source_id="", + title=_as_text(problem_data.get("title")), + statement=_as_text(problem_data.get("story")) or None, + input_format=_as_text(problem_data.get("input_format")) or None, + output_format=_as_text(problem_data.get("output_format")) or None, + constraints=_as_text(problem_data.get("constraints")) or None, + ) + + try: + searcher = ProblemSearcher( + url=QDRANT_URL, + collection_name=QDRANT_COLLECTION, + model_name=EMBEDDING_MODEL, + api_key=QDRANT_API_KEY, + ) + return searcher.search_problem(problem, k=k) + except Exception as exc: + raise SimilaritySearchError(f"search against Qdrant at {QDRANT_URL} failed: {exc}") from exc + + +def save_similar_problems(matches: List[Dict[str, Any]], output_path: str | Path) -> Path: + """Write the matches to a JSON file.""" + path = Path(output_path) + path.parent.mkdir(parents=True, exist_ok=True) + path.write_text(json.dumps(matches, indent=2, ensure_ascii=False), encoding="utf-8") + return path diff --git a/tests/test_cli.py b/tests/test_cli.py index ce2e3a9..17d117a 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -85,6 +85,7 @@ def test_generate_from_image_end_to_end(tmp_path: Path, monkeypatch, stub_client image_path=img_path, out_dir=out_dir, skip_validation=True, + similarity_check=False, ) assert result.generated_dir == out_dir / "generated" diff --git a/tests/test_similarity.py b/tests/test_similarity.py new file mode 100644 index 0000000..aa95d60 --- /dev/null +++ b/tests/test_similarity.py @@ -0,0 +1,91 @@ +""" +Unit tests for the similar-problem lookup (autosetter.similarity) and its +wiring into the CLI pipeline. Qdrant and the embedding model are stubbed. +""" + +from __future__ import annotations + +import json +from pathlib import Path + +from PIL import Image + +from autosetter.cli import build_arg_parser, generate_from_image +from autosetter.similarity import SimilaritySearchError, save_similar_problems +from tests.test_cli import SAMPLE_PROBLEM + +MATCHES = [ + { + "score": 0.93, + "title": "Theatre Square", + "source": "codeforces", + "source_id": "1A", + "difficulty": "1000", + "tags": ["math"], + "url": "https://codeforces.com/problemset/problem/1/A", + } +] + + +def _run_pipeline(tmp_path: Path, monkeypatch, stub_client, **kwargs): + img_path = tmp_path / "problem.png" + Image.new("RGB", (50, 50), color="white").save(img_path, format="PNG") + problem_reply = "```json\n" + json.dumps(SAMPLE_PROBLEM) + "\n```" + client = stub_client( + replies={ + "OCR specialist": problem_reply, + "competitive programming problem statement": problem_reply, + }, + default="```cpp\nint main() { return 0; }\n```", + ) + monkeypatch.setattr("autosetter.cli.OllamaClient", lambda **_: client) + return generate_from_image( + image_path=img_path, out_dir=tmp_path / "out", skip_validation=True, **kwargs + ) + + +def test_arg_parser_similarity_flags(): + args = build_arg_parser().parse_args(["p.png", "--no-similarity", "--similar-k", "3"]) + assert args.no_similarity is True + assert args.similar_k == 3 + + +def test_save_similar_problems(tmp_path: Path): + path = save_similar_problems(MATCHES, tmp_path / "similar_problems.json") + assert json.loads(path.read_text()) == MATCHES + + +def test_pipeline_saves_similar_problems(tmp_path: Path, monkeypatch, stub_client): + seen = {} + + def fake_find(problem_data, k): + seen["title"], seen["k"] = problem_data["title"], k + return MATCHES + + monkeypatch.setattr("autosetter.cli.find_similar_problems", fake_find) + result = _run_pipeline(tmp_path, monkeypatch, stub_client, similar_k=3) + + assert seen == {"title": SAMPLE_PROBLEM["title"], "k": 3} + assert result.similar_problems == MATCHES + saved = json.loads((tmp_path / "out" / "similar_problems.json").read_text()) + assert saved[0]["url"] == "https://codeforces.com/problemset/problem/1/A" + + +def test_pipeline_continues_when_search_fails(tmp_path: Path, monkeypatch, stub_client): + def failing_find(problem_data, k): + raise SimilaritySearchError("Qdrant unreachable") + + monkeypatch.setattr("autosetter.cli.find_similar_problems", failing_find) + result = _run_pipeline(tmp_path, monkeypatch, stub_client) + + assert result.similar_problems == [] + assert (tmp_path / "out" / "package" / "manifest.json").exists() + + +def test_pipeline_skips_search_when_disabled(tmp_path: Path, monkeypatch, stub_client): + def should_not_run(problem_data, k): + raise AssertionError("search should be skipped") + + monkeypatch.setattr("autosetter.cli.find_similar_problems", should_not_run) + result = _run_pipeline(tmp_path, monkeypatch, stub_client, similarity_check=False) + assert result.similar_problems == [] diff --git a/vector_database/.gitignore b/vector_database/.gitignore index 986f852..8f4774e 100644 --- a/vector_database/.gitignore +++ b/vector_database/.gitignore @@ -1,3 +1,4 @@ __pycache__ dataset/ .env +qdrant_storage/ diff --git a/vector_database/config.py b/vector_database/config.py index fa10e05..27a7694 100644 --- a/vector_database/config.py +++ b/vector_database/config.py @@ -32,7 +32,7 @@ QDRANT_URL = os.getenv("QDRANT_URL", "http://localhost:6333") QDRANT_COLLECTION = os.getenv("QDRANT_COLLECTION", "competitive_programming_problems") QDRANT_BATCH_SIZE = int(os.getenv("QDRANT_BATCH_SIZE", "64")) -QDRANT_API_KEY = os.getenv("QDRANT_API_KEY", None) +QDRANT_API_KEY = os.getenv("QDRANT_API_KEY") or None STATE_FILE = DATASET_DIR / "state.json" STATS_FILE = DATASET_DIR / "stats.json" diff --git a/vector_database/embeddings/embedder.py b/vector_database/embeddings/embedder.py index d4d40f8..ef76749 100644 --- a/vector_database/embeddings/embedder.py +++ b/vector_database/embeddings/embedder.py @@ -38,11 +38,12 @@ def build_embedding_text(problem: Problem) -> str: class ProblemEmbedder: - def __init__(self, model_name: str = "BAAI/bge-small-en-v1.5"): + def __init__(self, model_name: str = "BAAI/bge-small-en-v1.5", batch_size: int = 32): if SentenceTransformer is None: raise ImportError("sentence-transformers is not installed. Please install it.") self.model_name = model_name + self.batch_size = batch_size self.device = 'cuda' if torch.cuda.is_available() else 'cpu' print(f"Loading embedding model {model_name} on {self.device}...") self.model = SentenceTransformer(model_name, device=self.device) @@ -59,6 +60,8 @@ def embed_batch(self, problems: List[Problem]) -> List[List[float]]: texts = [build_embedding_text(p) for p in problems] # We normalize embeddings for cosine similarity - embeddings = self.model.encode(texts, normalize_embeddings=True, show_progress_bar=False) + embeddings = self.model.encode( + texts, batch_size=self.batch_size, normalize_embeddings=True, show_progress_bar=False + ) return embeddings.tolist() diff --git a/vector_database/pipeline.py b/vector_database/pipeline.py index 6852b91..7734d8d 100644 --- a/vector_database/pipeline.py +++ b/vector_database/pipeline.py @@ -156,7 +156,7 @@ def embed_and_upload(self): print("Starting Embedding and Qdrant Upload...") - embedder = ProblemEmbedder(model_name=config.EMBEDDING_MODEL) + embedder = ProblemEmbedder(model_name=config.EMBEDDING_MODEL, batch_size=config.EMBEDDING_BATCH_SIZE) qdrant = QdrantManager( url=config.QDRANT_URL, collection_name=config.QDRANT_COLLECTION, @@ -169,50 +169,64 @@ def embed_and_upload(self): batch_size = config.QDRANT_BATCH_SIZE current_batch_probs = [] - + # Lines of problems.jsonl consumed by the current batch (including blank/unparseable ones), + # so the resume offset always points at the first line not yet uploaded. + current_batch_lines = 0 + + # Number of lines of problems.jsonl already handled skip_count = self.state.state["embedded_and_uploaded"] - + failed_file = config.LOGS_DIR / "upload_failed.jsonl" - + with open(final_file, 'r', encoding='utf-8') as f, open(failed_file, 'a', encoding='utf-8') as err_f: lines = f.readlines() - + # Resume skipping lines_to_process = lines[skip_count:] if not lines_to_process: print("All problems already embedded and uploaded.") return - + for line in tqdm(lines_to_process, desc="Embedding & Uploading"): + current_batch_lines += 1 if not line.strip(): continue - + try: - prob = Problem(**json.loads(line)) - current_batch_probs.append(prob) - - if len(current_batch_probs) >= batch_size: - self._process_upload_batch(embedder, qdrant, current_batch_probs, err_f) - current_batch_probs = [] + current_batch_probs.append(Problem(**json.loads(line))) except Exception as e: self.stats.stats["failed_records"] += 1 err_f.write(json.dumps({"error": str(e), "line": line}) + "\n") - + + if len(current_batch_probs) >= batch_size: + if not self._process_upload_batch(embedder, qdrant, current_batch_probs, current_batch_lines, err_f): + return + current_batch_probs = [] + current_batch_lines = 0 + # Process remaining - if current_batch_probs: - self._process_upload_batch(embedder, qdrant, current_batch_probs, err_f) + if current_batch_lines: + self._process_upload_batch(embedder, qdrant, current_batch_probs, current_batch_lines, err_f) - def _process_upload_batch(self, embedder, qdrant, batch_probs, err_f): + def _process_upload_batch(self, embedder, qdrant, batch_probs, batch_lines, err_f) -> bool: + """ + Embeds and uploads a batch. On success advances the resume offset by `batch_lines`. + On failure returns False without advancing it, so the next run retries this batch. + """ try: - embeddings = embedder.embed_batch(batch_probs) - qdrant.upsert_batch(batch_probs, embeddings) - + if batch_probs: + embeddings = embedder.embed_batch(batch_probs) + qdrant.upsert_batch(batch_probs, embeddings) + self.stats.stats["embedding_count"] += len(batch_probs) - self.state.state["embedded_and_uploaded"] += len(batch_probs) - + self.state.state["embedded_and_uploaded"] += batch_lines + self.state.save() self.stats.save() + return True except Exception as e: - self.stats.stats["failed_records"] += len(batch_probs) for p in batch_probs: err_f.write(json.dumps({"problem_id": p.id, "error": str(e)}) + "\n") + print(f"Batch upload failed ({e}). Stopping; re-run to retry from this batch.") + self.stats.save() + return False diff --git a/vector_database/scripts/test_search.py b/vector_database/scripts/test_search.py index 59b86ac..0db7d53 100644 --- a/vector_database/scripts/test_search.py +++ b/vector_database/scripts/test_search.py @@ -1,99 +1,47 @@ import argparse -import os import sys from pathlib import Path -from qdrant_client import QdrantClient -from google import genai # Absolute path adjustments for local custom modules sys.path.append(str(Path(__file__).resolve().parent.parent.parent)) from vector_database import config -from vector_database.embeddings.embedder import ProblemEmbedder - - -def resolve_links_with_llm(results_summary: str) -> str: - """ - Sends the raw structural database matches to Gemini to reconstruct - the correct Codeforces URL variants. - """ - print("\nRouting search data to LLM for precise URL mapping...") - - # Initialize the standard Gemini client (automatically loads GEMINI_API_KEY from environment) - client = genai.Client() - - # FIX: Restored full explicit URL patterns so the LLM knows exactly how to build paths - prompt = f""" - You are an expert competitive programming assistant. Below are raw vector search results from a database containing Codeforces problems. - Analyze the 'ID', 'Title', and 'Source' properties to output the direct, clickable live link for every single match. - - Strict Structural Formatting Rules: - 1. Standard problem IDs (e.g., 1234/A) must use the standard public problemset view path: - https://codeforces.com - 2. Gym IDs (e.g., 987654/B or any ID where the contest part is 6 digits or more) must route exactly to: - https://codeforces.com - - Print the cleanly formatted results exactly as provided in the raw data layout, but replace the text placeholder blocks with the actual verified live URLs. - - Raw Results Data: - {results_summary} - """ - - try: - response = client.models.generate_content( - model='gemini-2.5-flash', - contents=prompt, - ) - # Ensure we always return a string fallback if response content is unexpectedly empty - return response.text if response.text else results_summary - except Exception as e: - return f"\n[LLM Error]: Could not resolve links dynamically ({e})\n{results_summary}" +from vector_database.search import ProblemSearcher def main(): - parser = argparse.ArgumentParser(description="Test similarity search in Qdrant with LLM link formatting.") - parser.add_argument("query", type=str, help="Problem description to search for") + parser = argparse.ArgumentParser(description="Test similarity search in Qdrant.") + parser.add_argument("query", type=str, nargs="?", help="Problem description to search for") + parser.add_argument("--file", type=str, help="Read the query (e.g. a full problem statement) from a file") parser.add_argument("--k", type=int, default=5, help="Number of results to retrieve") args = parser.parse_args() - - print(f"Loading embedder ({config.EMBEDDING_MODEL})...") - embedder = ProblemEmbedder(model_name=config.EMBEDDING_MODEL) - - print("Generating embedding for query...") - query_vector = embedder.model.encode([args.query], normalize_embeddings=True).tolist()[0] - - print(f"Connecting to Qdrant at {config.QDRANT_URL}...") - client = QdrantClient(url=config.QDRANT_URL, api_key=config.QDRANT_API_KEY) - - print(f"Searching for top {args.k} matches...") - results = client.query_points( + + if args.file: + query = Path(args.file).read_text(encoding="utf-8") + elif args.query: + query = args.query + else: + parser.error("provide a query string or --file") + + print(f"Loading embedder ({config.EMBEDDING_MODEL}) and connecting to Qdrant at {config.QDRANT_URL}...") + searcher = ProblemSearcher( + url=config.QDRANT_URL, collection_name=config.QDRANT_COLLECTION, - query=query_vector, - limit=args.k - ).points - - # Construct an unformatted text summary block to pass to the LLM - results_summary = "" - for i, res in enumerate(results, 1): - payload = res.payload - prob_id = payload.get('problem_id', '').replace('cf_', '') - tags = payload.get('tags') - tags_str = ', '.join(tags) if isinstance(tags, list) else tags - - results_summary += f""" -[{i}] Score: {res.score:.4f} -ID: {prob_id} -Title: {payload.get('title')} -Source: {payload.get('source')} ({payload.get('source_id')}) -Difficulty: {payload.get('difficulty')} -Tags: {tags_str} -URL: [Insert Live Codeforces URL here] -""" + model_name=config.EMBEDDING_MODEL, + api_key=config.QDRANT_API_KEY, + ) - # Pass the text to the LLM to get the final output containing functional links - final_output = resolve_links_with_llm(results_summary) - - print("\n=== Final Formatted Matches ===") - print(final_output) + print(f"Searching for top {args.k} matches...") + results = searcher.search_text(query, k=args.k) + + print("\n=== Matches ===") + for i, res in enumerate(results, 1): + print(f""" +[{i}] Score: {res['score']:.4f} +Title: {res['title']} +Source: {res['source']} ({res['source_id']}) +Difficulty: {res['difficulty']} +Tags: {', '.join(res['tags'])} +URL: {res['url'] or 'N/A'}""") if __name__ == "__main__": diff --git a/vector_database/search.py b/vector_database/search.py new file mode 100644 index 0000000..1c1ef69 --- /dev/null +++ b/vector_database/search.py @@ -0,0 +1,68 @@ +import re +from typing import Any, Dict, List, Optional + +from qdrant_client import QdrantClient + +from .embeddings.embedder import ProblemEmbedder +from .processing.schema import Problem + + +def build_codeforces_url(source_id: str) -> Optional[str]: + """ + Deterministically builds a Codeforces URL from an ID like '1234A', '1234/A' or '1352C1'. + Contest IDs of 6+ digits are gym contests. + """ + match = re.fullmatch(r"(\d+)/?([A-Za-z]\d*)", source_id or "") + if not match: + return None + contest_id, index = match.groups() + if len(contest_id) >= 6: + return f"https://codeforces.com/gym/{contest_id}/problem/{index}" + return f"https://codeforces.com/problemset/problem/{contest_id}/{index}" + + +def resolve_url(payload: Dict[str, Any]) -> Optional[str]: + """Prefer the URL stored in Qdrant; fall back to building it from the ID.""" + if payload.get("url"): + return payload["url"] + if payload.get("source") == "codeforces": + return build_codeforces_url(payload.get("source_id", "")) + return None + + +class ProblemSearcher: + """Semantic search over the problems stored in Qdrant.""" + + def __init__(self, url: str, collection_name: str, model_name: str, api_key: Optional[str] = None): + self.embedder = ProblemEmbedder(model_name=model_name) + self.client = QdrantClient(url=url, api_key=api_key) + self.collection_name = collection_name + + def search_text(self, query: str, k: int = 5) -> List[Dict[str, Any]]: + """Search with free text (e.g. a pasted problem statement).""" + vector = self.embedder.model.encode([query], normalize_embeddings=True).tolist()[0] + return self._search(vector, k) + + def search_problem(self, problem: Problem, k: int = 5) -> List[Dict[str, Any]]: + """Search with a structured problem, embedded in the same format as the stored problems.""" + return self._search(self.embedder.embed_batch([problem])[0], k) + + def _search(self, vector: List[float], k: int) -> List[Dict[str, Any]]: + points = self.client.query_points( + collection_name=self.collection_name, + query=vector, + limit=k + ).points + + return [ + { + "score": round(p.score, 4), + "title": p.payload.get("title"), + "source": p.payload.get("source"), + "source_id": p.payload.get("source_id"), + "difficulty": p.payload.get("difficulty"), + "tags": p.payload.get("tags") or [], + "url": resolve_url(p.payload), + } + for p in points + ] diff --git a/vector_database/sources/codeforces.py b/vector_database/sources/codeforces.py index fc8de10..c2dac73 100644 --- a/vector_database/sources/codeforces.py +++ b/vector_database/sources/codeforces.py @@ -1,11 +1,15 @@ import json import requests -import time -from typing import Iterator +from typing import Iterator, Optional from pathlib import Path from .base import BaseSource from ..processing.schema import Problem + +def _as_difficulty(value) -> Optional[str]: + """Stringify a rating/difficulty, keeping missing values as None instead of 'None'.""" + return str(value) if value not in (None, "") else None + class CodeforcesSource(BaseSource): """ Codeforces source collector. @@ -37,7 +41,7 @@ def _collect_from_dataset(self, limit: int = None) -> Iterator[Problem]: source_id=str(data.get('id', data.get('source_id', ''))), title=data.get('title', ''), url=data.get('url'), - difficulty=str(data.get('rating', data.get('difficulty'))), + difficulty=_as_difficulty(data.get('rating', data.get('difficulty'))), rating=data.get('rating'), tags=data.get('tags', []), statement=data.get('statement', data.get('description')), @@ -50,7 +54,7 @@ def _collect_from_dataset(self, limit: int = None) -> Iterator[Problem]: count += 1 def _collect_from_api(self, limit: int = None) -> Iterator[Problem]: - url = "https://codeforces.com/api/problemset.problems" + url = "https://codeforces.com/api/problemset.problems?lang=en" response = requests.get(url) if response.status_code != 200: raise Exception(f"Failed to fetch from Codeforces API: {response.status_code}") @@ -87,5 +91,3 @@ def _collect_from_api(self, limit: int = None) -> Iterator[Problem]: examples=[] ) count += 1 - # Be polite to API - time.sleep(0.1) diff --git a/vector_database/sources/leetcode.py b/vector_database/sources/leetcode.py index c075ed8..aae6ee3 100644 --- a/vector_database/sources/leetcode.py +++ b/vector_database/sources/leetcode.py @@ -35,7 +35,7 @@ def collect(self, limit: int = None) -> Iterator[Problem]: source_id=source_id, title=data.get('title', ''), url=data.get('url'), - difficulty=str(data.get('difficulty', '')), + difficulty=str(data['difficulty']) if data.get('difficulty') else None, rating=None, # LC generally uses Easy/Medium/Hard instead of numerical rating tags=data.get('tags', []), statement=data.get('statement'),