Spaces:
Sleeping
Sleeping
| """The answering agent: intent routing -> retrieval -> grounded, cited generation. | |
| Pipeline per user message: | |
| 1. Classify intent (greeting / logistics / research-or-person) with cheap heuristics. | |
| 2. Retrieve passages (+ graph expansion) from the knowledge base. | |
| 3. Grounding guard: if nothing in the corpus is relevant, decline rather than | |
| hallucinate. | |
| 4. Otherwise build a numbered, source-attributed context and stream an answer | |
| that cites with [n]; a clickable Sources list is always appended by us so the | |
| links are correct even if the model forgets to cite. | |
| """ | |
| from __future__ import annotations | |
| import re | |
| from typing import Iterator, List, Tuple | |
| import config | |
| import llm | |
| from retriever import retrieve, RetrievalResult | |
| # --------------------------------------------------------------------------- # | |
| # Intent classification (cheap, no LLM call) | |
| # --------------------------------------------------------------------------- # | |
| _GREETING_RE = re.compile(r"^\s*(hi|hey|hello|yo|sup|good (morning|afternoon|evening)|thanks|thank you)\b", re.IGNORECASE) | |
| _LOGISTICS_RE = re.compile( | |
| r"\b(phd|ph\.d|student|apply|application|admission|join|position|opening|" | |
| r"intern|internship|recruit|advisor|advise|supervis|collaborat|partner|" | |
| r"contact|email|reach|hire|hiring|work with|opportunit)\w*", | |
| re.IGNORECASE, | |
| ) | |
| def classify_intent(message: str) -> str: | |
| text = (message or "").strip() | |
| if _GREETING_RE.search(text) and len(text) < 40: | |
| return "greeting" | |
| if _LOGISTICS_RE.search(text): | |
| return "logistics" | |
| return "research" | |
| def wants_to_connect(message: str) -> bool: | |
| """Heuristic: does this message look like a prospective student/collaborator?""" | |
| return classify_intent(message) == "logistics" | |
| def _content_to_text(content) -> str: | |
| """Coerce a Gradio message 'content' to plain text. | |
| Gradio 6 may hand history back with content as a string, a list of rich-text | |
| parts (e.g. [{'text': '...', 'type': 'text'}]), or a dict — but the LLM | |
| backends expect a plain string. Normalize all shapes here. | |
| """ | |
| if content is None: | |
| return "" | |
| if isinstance(content, str): | |
| return content | |
| if isinstance(content, dict): | |
| return content.get("text", "") | |
| if isinstance(content, list): | |
| parts = [] | |
| for p in content: | |
| if isinstance(p, str): | |
| parts.append(p) | |
| elif isinstance(p, dict) and p.get("text"): | |
| parts.append(p["text"]) | |
| return " ".join(parts).strip() | |
| return str(content) | |
| # --------------------------------------------------------------------------- # | |
| # Prompt construction | |
| # --------------------------------------------------------------------------- # | |
| def _persona_system() -> str: | |
| base_rules = ( | |
| "Rules:\n" | |
| "- Use ONLY the numbered context provided to state facts about the research.\n" | |
| "- Cite sources inline as [n] using the numbers in the context.\n" | |
| "- If the context does not contain the answer, say you don't have that " | |
| "information rather than guessing.\n" | |
| "- Be concise, professional, and engaging. Prefer 2-5 short paragraphs or " | |
| "bullet points.\n" | |
| "- Do not invent publication titles, links, dates, or numbers." | |
| ) | |
| if config.PERSONA == "first_person": | |
| return ( | |
| f"You are an AI assistant speaking as {config.PROFESSOR_NAME} " | |
| f"({config.PROFESSOR_TITLE}) on the lab website. Stay in character in " | |
| f"the first person, but you are an automated assistant, not the real " | |
| f"person. {base_rules}" | |
| ) | |
| return ( | |
| f"You are the AI research assistant for {config.LAB_NAME}. You help " | |
| f"prospective students, collaborators, and the public understand the work of " | |
| f"{config.PROFESSOR_NAME} ({config.PROFESSOR_TITLE}). Refer to him in the " | |
| f"third person as 'Prof. Zhao'. {base_rules}" | |
| ) | |
| def _build_context(result: RetrievalResult) -> Tuple[str, List[dict]]: | |
| """Turn hits into numbered context text + a deduped source list for rendering.""" | |
| sources: List[dict] = [] | |
| doc_to_num: dict = {} | |
| blocks: List[str] = [] | |
| for hit in result.hits: | |
| doc = hit.chunk["doc_id"] | |
| if doc not in doc_to_num: | |
| doc_to_num[doc] = len(sources) + 1 | |
| sources.append( | |
| { | |
| "n": doc_to_num[doc], | |
| "title": hit.chunk["title"], | |
| "url": hit.chunk["url"], | |
| "source": hit.chunk["source"], | |
| } | |
| ) | |
| n = doc_to_num[doc] | |
| tag = " (related work, via paper graph)" if hit.via_graph else "" | |
| blocks.append(f"[{n}]{tag} {hit.chunk['title']}\n{hit.chunk['text']}") | |
| return "\n\n".join(blocks), sources | |
| def _render_sources(sources: List[dict]) -> str: | |
| if not sources: | |
| return "" | |
| lines = ["\n\n---", "**Sources**"] | |
| for s in sources: | |
| # only publications get external links worth surfacing distinctly | |
| lines.append(f"{s['n']}. [{s['title']}]({s['url']})") | |
| return "\n".join(lines) | |
| _GROUNDING_GUARD = ( | |
| "I don't have that in my knowledge base, so I'd rather not guess. " | |
| "I can speak to Prof. Zhao's research on **graph neural networks**, " | |
| "**spatio-temporal / geospatial ML**, **retrieval-augmented and agentic LLMs**, " | |
| "**AI for science**, and **trustworthy & efficient AI**. Try asking about one of " | |
| "those — for example, *“What is Spatial-RAG?”* or *“Summarize the lab's work on " | |
| "graph RAG.”*\n\nIf you'd like to reach the lab directly, tell me you're a " | |
| "prospective student or collaborator and I can take your details." | |
| ) | |
| # --------------------------------------------------------------------------- # | |
| # Public: streaming answer | |
| # --------------------------------------------------------------------------- # | |
| def stream_answer(message: str, history: List[dict] | None = None) -> Iterator[str]: | |
| """Yield the cumulative answer text (suitable for Gradio streaming).""" | |
| history = history or [] | |
| intent = classify_intent(message) | |
| result = retrieve(message) | |
| # Grounding guard: refuse to fabricate when the corpus is irrelevant. | |
| if not result.is_grounded and intent != "greeting": | |
| yield _GROUNDING_GUARD | |
| return | |
| context, sources = _build_context(result) | |
| tier = "deep" if intent == "research" else llm.route(message) | |
| system = _persona_system() | |
| convo = [ | |
| {"role": m["role"], "content": _content_to_text(m.get("content"))} | |
| for m in history | |
| if m.get("role") in ("user", "assistant") | |
| ] | |
| convo = [m for m in convo if m["content"]][-6:] # drop empties, keep last 6 | |
| user_turn = ( | |
| f"Numbered context:\n{context}\n\n" | |
| f"Question: {message}\n\n" | |
| f"Answer using only the context above and cite with [n]." | |
| ) | |
| messages = [{"role": "system", "content": system}, *convo, {"role": "user", "content": user_turn}] | |
| body = "" | |
| for piece in llm.stream(messages, tier=tier): | |
| body += piece | |
| yield body | |
| # Always append correct, clickable sources. | |
| tail = _render_sources(sources) | |
| if tail: | |
| yield body + tail | |