"""MongoDB store for syncing and querying experiment results.""" import logging from datetime import datetime, timezone from typing import Any from pymongo import MongoClient from pymongo.collection import Collection logger = logging.getLogger(__name__) class MongoDBStore: """Read/write store for solar-eval experiment data in MongoDB. Collections: - experiments: One doc per (project, run_id) — metadata + aggregate scores. - sample_results: One doc per (experiment_id, sample_idx) — per-sample inference + eval. - diff_details: One doc per correction diff — category-level detail for qualitative view. """ def __init__(self, connection_uri: str, database: str = "solar_eval"): self.client = MongoClient(connection_uri) self.db = self.client[database] self.experiments: Collection = self.db["experiments"] self.sample_results: Collection = self.db["sample_results"] self.diff_details: Collection = self.db["diff_details"] def ensure_indexes(self) -> None: """Create indexes for efficient querying.""" # experiments self.experiments.create_index([("project", 1), ("dataset_group", 1)]) self.experiments.create_index("prompt_id") self.experiments.create_index("model") self.experiments.create_index([("project", 1), ("run_id", 1)], unique=True) # sample_results self.sample_results.create_index( [("experiment_id", 1), ("sample_idx", 1)], unique=True ) # diff_details self.diff_details.create_index([("experiment_id", 1), ("sample_idx", 1)]) self.diff_details.create_index("category") def upsert_experiment(self, experiment: dict[str, Any]) -> str: """Insert or update an experiment document. Returns the _id.""" key = {"project": experiment["project"], "run_id": experiment["run_id"]} result = self.experiments.update_one( key, {"$set": {**experiment, "synced_at": datetime.now(timezone.utc)}}, upsert=True, ) if result.upserted_id: return str(result.upserted_id) doc = self.experiments.find_one(key) return str(doc["_id"]) if doc else "" def bulk_insert_samples(self, samples: list[dict[str, Any]]) -> int: """Bulk upsert sample results. Returns count of upserted/modified docs.""" if not samples: return 0 from pymongo import UpdateOne ops = [ UpdateOne( {"experiment_id": s["experiment_id"], "sample_idx": s["sample_idx"]}, {"$set": s}, upsert=True, ) for s in samples ] result = self.db["sample_results"].bulk_write(ops) return result.upserted_count + result.modified_count def bulk_insert_diffs(self, diffs: list[dict[str, Any]]) -> int: """Bulk insert diff details. Returns count inserted.""" if not diffs: return 0 # Clear existing diffs for this experiment first if diffs: exp_id = diffs[0].get("experiment_id") if exp_id: self.diff_details.delete_many({"experiment_id": exp_id}) result = self.diff_details.insert_many(diffs) return len(result.inserted_ids) def get_experiments(self, project: str, **filters: Any) -> list[dict[str, Any]]: """Query experiments with optional filters (dataset_group, model, etc.).""" query: dict[str, Any] = {"project": project} query.update(filters) return list(self.experiments.find(query).sort("created_at", -1)) def get_samples(self, experiment_id: str, **filters: Any) -> list[dict[str, Any]]: """Get per-sample results for an experiment.""" query: dict[str, Any] = {"experiment_id": experiment_id} query.update(filters) return list(self.sample_results.find(query).sort("sample_idx", 1)) def get_diffs(self, experiment_id: str, sample_idx: int | None = None) -> list[dict[str, Any]]: """Get diff details for an experiment, optionally filtered by sample.""" query: dict[str, Any] = {"experiment_id": experiment_id} if sample_idx is not None: query["sample_idx"] = sample_idx return list(self.diff_details.find(query).sort("sample_idx", 1)) def close(self) -> None: """Close the MongoDB connection.""" self.client.close()