ruby56/Citation-Benchmark
0
1import sqlite32import os3from typing import Optional, List, Dict, Any, Union, Tuple4from pydantic import BaseModel, Field5import tasks6 7class Observation(BaseModel):8 task_instruction: str9 search_results: Optional[List[Dict[str, str]]] = None10 last_abstract: Optional[str] = None11 citations_data: Optional[List[Dict[str, str]]] = None12 message: Optional[str] = None13 step_count: int = 014 15class Action(BaseModel):16 action_type: str = Field(..., description="One of: 'search', 'read_abstract', 'get_citations', 'submit_id'")17 query: Optional[str] = Field(None, description="Search term (used only for 'search')")18 paper_id: Optional[str] = Field(None, description="Paper ID (corpus_id or arxiv_id) to read, cite, or submit")19 20class Reward(BaseModel):21 value: float22 is_terminal: bool23 details: str24 25class CitationEnv:26 def __init__(self, task_id: str = "T001"):27 self.task = next((t for t in tasks.TASKS if t.id == task_id), tasks.TASKS[0])28 self.grader = tasks.Grader(self.task)29 self.step_count = 030 self.max_steps = 1531 self.state_history = []32 self._current_obs = None33 34 db_path = 'citation_db.sqlite'35 if not os.path.exists(db_path):36 print(f"Database {db_path} not found locally. Downloading from HF Datasets...")37 try:38 from huggingface_hub import hf_hub_download39 # Pulls directly from the Dataset instead of the Space40 db_path = hf_hub_download(repo_id="ruby56/Citation-Database", filename="citation_db.sqlite", repo_type="dataset" , local_dir=".")41 print("Database successfully mounted!")42 except Exception as e:43 print(f"WARNING: Could not download DB: {e}")44 45 self.db = sqlite3.connect(db_path)46 47 def reset(self) -> Observation:48 self.step_count = 049 self.state_history = []50 self._current_obs = Observation(51 task_instruction=self.task.claim,52 message=f"Environment initialized. Goal: {self.task.claim}"53 )54 self.state_history.append(self._current_obs)55 return self._current_obs56 57 def state(self) -> Observation:58 if not self._current_obs:59 return self.reset()60 return self._current_obs61 62 def step(self, action: Action) -> Tuple[Observation, Reward, bool, dict]:63 self.step_count += 164 done = False65 info = {"task": self.task.id}66 67 reward_value = -0.0568 reward_details = "Step taken."69 70 if self.step_count >= self.max_steps and action.action_type != "submit_id":71 return self._current_obs, Reward(value=-1.0, is_terminal=True, details="Max steps reached."), True, info72 73 message = ""74 search_results = self._current_obs.search_results if self._current_obs else None75 last_abstract = self._current_obs.last_abstract if self._current_obs else None76 citations_data = self._current_obs.citations_data if self._current_obs else None77 78 cursor = self.db.cursor()79 80 if action.action_type == "search":81 q = (action.query or "")82 query = "SELECT corpus_id, arxiv_id, title, year FROM s2_papers WHERE title LIKE ? LIMIT 10"83 cursor.execute(query, (f"%{q}%",))84 results = []85 for row in cursor.fetchall():86 results.append({"corpus_id": str(row[0]), "arxiv_id": str(row[1]), "title": row[2], "year": str(row[3])})87 search_results = results88 message = f"Found {len(results)} papers."89 90 elif action.action_type == "read_abstract":91 pid = action.paper_id92 cursor.execute("SELECT abstract FROM arxiv_metadata WHERE arxiv_id = ?", (pid,))93 res = cursor.fetchone()94 if res:95 last_abstract = res[0]96 message = f"Abstract for {pid} loaded."97 else:98 message = f"Abstract not found for {pid}. Make sure to pass a valid arxiv_id."99 reward_value -= 0.1100 101 elif action.action_type == "get_citations":102 pid = action.paper_id103 query = """104 SELECT c.contexts, c.intent, ar.title, ar.abstract105 FROM citations c106 JOIN s2_papers s2_citing ON c.citing_corpus_id = s2_citing.corpus_id107 JOIN arxiv_metadata ar ON s2_citing.arxiv_id = ar.arxiv_id108 WHERE c.cited_corpus_id = ?109 LIMIT 5110 """111 cursor.execute(query, (pid,))112 res = cursor.fetchall()113 if res:114 c_data = []115 for row in res:116 c_data.append({"contexts": row[0], "intent": row[1], "citing_title": row[2], "citing_abstract": row[3]})117 citations_data = c_data118 message = f"Found {len(res)} valid citations with abstracts for corpus_id {pid}."119 else:120 message = f"No citation graph data found for corpus_id {pid}."121 122 elif action.action_type == "submit_id":123 pid = action.paper_id or ""124 score = self.grader.score(pid, self.db)125 reward_value = float(score) # final reward is the grader score126 reward_details = f"Final evaluation. Score: {score}"127 message = f"ID submitted: {pid}"128 done = True129 130 else:131 message = "Invalid action."132 reward_value -= 0.2133 134 new_obs = Observation(135 task_instruction=self.task.claim,136 search_results=search_results,137 last_abstract=last_abstract,138 citations_data=citations_data,139 message=message,140 step_count=self.step_count141 )142 self._current_obs = new_obs143 self.state_history.append(new_obs)144 info["history_len"] = len(self.state_history)145 146 return new_obs, Reward(value=reward_value, is_terminal=done, details=reward_details), done, info147 