Spaces:
Runtime error
Runtime error
Download src/agents.py from MuhammadAzfar/ResearchMind-AI: direct link, hf CLI and curl.
- Browser
- Download file 9.26 kB
-
https://huggingface.co/spaces/MuhammadAzfar/ResearchMind-AI/resolve/main/src/agents.py
- Command line
-
hf download hf://spaces/MuhammadAzfar/ResearchMind-AI/src/agents.py
-
curl -L -o agents.py https://huggingface.co/spaces/MuhammadAzfar/ResearchMind-AI/resolve/main/src/agents.py
9.26 kB
| """The agents of ResearchMind. Each one is a small, testable function. | |
| Planner -> splits the topic into sub-questions | |
| Researcher -> searches the web for every sub-question (in parallel) | |
| Writer -> drafts a cited report from retrieved evidence | |
| Checker -> verifies every cited claim against the passages it cites | |
| Reviser -> rewrites flagged claims so they match the evidence | |
| Translator -> renders the verified report in another language | |
| """ | |
| import json | |
| import re | |
| from collections import Counter | |
| from concurrent.futures import ThreadPoolExecutor | |
| from src import config | |
| from src.llm import LLMError | |
| from src.prompts import ( | |
| SYSTEM_CHECKER, | |
| SYSTEM_PLANNER, | |
| SYSTEM_REVISER, | |
| SYSTEM_TRANSLATOR, | |
| SYSTEM_WRITER, | |
| build_check_prompt, | |
| build_planner_prompt, | |
| build_revise_prompt, | |
| build_translate_prompt, | |
| build_writer_prompt, | |
| ) | |
| from src.search import SearchError | |
| _CITATION = re.compile(r"\[(\d{1,2}(?:\s*,\s*\d{1,2})*)\]") | |
| VERDICTS = ("supported", "partial", "unsupported") | |
| # ------------------------------------------------------------------ helpers | |
| def parse_json(text: str): | |
| """Parses JSON even if the model wrapped it in prose or code fences.""" | |
| if not text: | |
| return None | |
| try: | |
| return json.loads(text) | |
| except Exception: | |
| pass | |
| match = re.search(r"\{.*\}", text, re.S) | |
| if match: | |
| try: | |
| return json.loads(match.group(0)) | |
| except Exception: | |
| return None | |
| return None | |
| def clean_citations(text: str, max_n: int) -> str: | |
| """Normalises [1, 2] to [1][2] and drops citation numbers the model invented (outside 1..max_n).""" | |
| def _replace(match): | |
| numbers = [int(n) for n in re.split(r"\s*,\s*", match.group(1))] | |
| return "".join(f"[{n}]" for n in numbers if 1 <= n <= max_n) | |
| return _CITATION.sub(_replace, text) | |
| # ------------------------------------------------------------------ Planner | |
| def plan(topic: str, n: int, complete) -> list: | |
| """Returns up to n search-friendly sub-questions; falls back to the topic itself.""" | |
| messages = [ | |
| {"role": "system", "content": SYSTEM_PLANNER}, | |
| {"role": "user", "content": build_planner_prompt(topic, n)}, | |
| ] | |
| data = parse_json(complete(messages, role="fast", temperature=0.2, json_mode=True)) | |
| questions = [] | |
| for q in (data or {}).get("sub_questions", []): | |
| if isinstance(q, str) and q.strip() and q.strip() not in questions: | |
| questions.append(q.strip()) | |
| return questions[:n] or [topic] | |
| # --------------------------------------------------------------- Researcher | |
| def research_web(questions: list, cfg: dict, searcher) -> tuple: | |
| """Searches every sub-question in parallel. Returns (unique_sources, error_messages).""" | |
| def one(question): | |
| try: | |
| return searcher( | |
| question, | |
| max_results=cfg["results_per_question"], | |
| search_depth=cfg["search_depth"], | |
| ), None | |
| except SearchError as e: | |
| return [], str(e) | |
| sources, errors, seen = [], [], set() | |
| with ThreadPoolExecutor(max_workers=min(len(questions), 5) or 1) as pool: | |
| for results, error in pool.map(one, questions): | |
| if error and error not in errors: | |
| errors.append(error) | |
| for item in results: | |
| key = item.get("url") or item["title"] | |
| if key in seen: | |
| continue | |
| seen.add(key) | |
| sources.append({**item, "kind": "web"}) | |
| return sources, errors | |
| # ------------------------------------------------------------------- Writer | |
| def format_context(selected: list, sources: list) -> str: | |
| """Renders retrieved passages grouped under their numbered source.""" | |
| blocks = [] | |
| for source in sources: | |
| texts = [c["text"] for c in selected if c["cite"] == source["n"]] | |
| body = "\n".join(f"- {t}" for t in texts) | |
| where = source["url"] if source.get("url") else "uploaded document" | |
| blocks.append(f"[{source['n']}] {source['title']} ({where})\n{body}") | |
| return "\n\n".join(blocks) | |
| def write_stream(topic, mode, questions, selected, sources, stream): | |
| """Yields text deltas of the first draft.""" | |
| messages = [ | |
| {"role": "system", "content": SYSTEM_WRITER}, | |
| { | |
| "role": "user", | |
| "content": build_writer_prompt(topic, mode, questions, format_context(selected, sources), len(sources)), | |
| }, | |
| ] | |
| yield from stream(messages, role="writer", temperature=0.3) | |
| # ------------------------------------------------------------------ Checker | |
| def extract_claims(report: str) -> tuple: | |
| """Splits a report into cited sentences (claims to verify) and counts uncited factual sentences.""" | |
| claims, uncited = [], 0 | |
| for raw in report.split("\n"): | |
| line = raw.strip() | |
| if not line or line.startswith(("#", ">", "---")): | |
| continue | |
| line = re.sub(r"^([-*\u2022]|\d+[.)])\s+", "", line) | |
| line = re.sub(r"\*\*|__", "", line) | |
| # "claim. [1]" -> "claim[1]." so the citation stays attached to its sentence | |
| line = re.sub(r"([.!?])\s*((?:\[\d{1,2}\])+)", r"\2\1", line) | |
| for sentence in re.split(r"(?<=[.!?])\s+", line): | |
| sentence = sentence.strip() | |
| cites = sorted({int(n) for n in re.findall(r"\[(\d{1,2})\]", sentence)}) | |
| if cites: | |
| claims.append({"id": len(claims) + 1, "text": sentence, "cites": cites}) | |
| elif len(sentence.split()) >= 8: | |
| uncited += 1 | |
| return claims, uncited | |
| def check_claims(claims: list, store, complete) -> dict: | |
| """Verifies claims in small batches against the passages they cite. Returns {claim_id: {verdict, reason}}.""" | |
| results = {} | |
| to_check = claims[: config.MAX_CLAIMS_CHECKED] | |
| for start in range(0, len(to_check), config.CHECK_BATCH_SIZE): | |
| batch = to_check[start : start + config.CHECK_BATCH_SIZE] | |
| items = [ | |
| {"id": c["id"], "text": c["text"], "evidence": store.support_for(c["text"], c["cites"], k=2)} | |
| for c in batch | |
| ] | |
| messages = [ | |
| {"role": "system", "content": SYSTEM_CHECKER}, | |
| {"role": "user", "content": build_check_prompt(items)}, | |
| ] | |
| try: | |
| data = parse_json(complete(messages, role="fast", temperature=0.0, json_mode=True)) | |
| except LLMError as e: | |
| print(f"[DEBUG WARNING] Checker batch skipped: {e}") | |
| continue | |
| for entry in (data or {}).get("results", []): | |
| try: | |
| claim_id = int(entry["id"]) | |
| except (KeyError, TypeError, ValueError): | |
| continue | |
| verdict = str(entry.get("verdict", "")).lower().strip() | |
| if verdict in VERDICTS and any(c["id"] == claim_id for c in batch): | |
| results[claim_id] = {"verdict": verdict, "reason": str(entry.get("reason", ""))[:200]} | |
| return results | |
| def summarize_check(claims: list, results: dict, uncited: int) -> dict: | |
| counts = Counter(results[c["id"]]["verdict"] for c in claims if c["id"] in results) | |
| checked = sum(counts.values()) | |
| total_sentences = len(claims) + uncited | |
| return { | |
| "claims": len(claims), | |
| "checked": checked, | |
| "supported": counts["supported"], | |
| "partial": counts["partial"], | |
| "unsupported": counts["unsupported"], | |
| "support_rate": (counts["supported"] / checked) if checked else None, | |
| "citation_coverage": (len(claims) / total_sentences) if total_sentences else None, | |
| } | |
| # ------------------------------------------------------------------ Reviser | |
| def revise(report: str, claims: list, results: dict, store, complete) -> str: | |
| """Rewrites the report so flagged claims match their evidence. Returns the original if the rewrite looks broken.""" | |
| issues = [] | |
| for claim in claims: | |
| verdict = results.get(claim["id"]) | |
| if verdict and verdict["verdict"] in ("unsupported", "partial"): | |
| issues.append({ | |
| "text": claim["text"], | |
| "verdict": verdict["verdict"], | |
| "reason": verdict["reason"] or "not backed by the cited source", | |
| "evidence": store.support_for(claim["text"], claim["cites"], k=2), | |
| }) | |
| if not issues: | |
| return report | |
| messages = [ | |
| {"role": "system", "content": SYSTEM_REVISER}, | |
| {"role": "user", "content": build_revise_prompt(report, issues)}, | |
| ] | |
| revised = complete(messages, role="writer", temperature=0.1).strip() | |
| # Guard against a stub or truncated answer replacing the whole report | |
| if len(revised) < 0.5 * len(report): | |
| return report | |
| return revised | |
| # --------------------------------------------------------------- Translator | |
| def translate(report: str, language: str, complete) -> str: | |
| messages = [ | |
| {"role": "system", "content": SYSTEM_TRANSLATOR}, | |
| {"role": "user", "content": build_translate_prompt(report, language)}, | |
| ] | |
| translated = complete(messages, role="writer", temperature=0.1).strip() | |
| if len(translated) < 0.3 * len(report): | |
| raise LLMError("The translation came back too short.") | |
| return translated | |