Spaces:
Running on Zero
Running on Zero
| """ | |
| retrieval.py | |
| ------------------- | |
| Retrieval nodes and pipeline orchestration. | |
| """ | |
| import time | |
| import numpy as np | |
| from tavily import TavilyClient | |
| import vectorstore as vs_module | |
| from config import CONFIDENCE_THRESHOLD, TAVILY_API_KEY | |
| from logging_config import get_logger | |
| logger = get_logger(__name__) | |
| # Node 1: FAISS retrieval | |
| def retrieve_node(state: dict) -> dict: | |
| """ | |
| Search the FAISS index and decide whether web search is needed. | |
| Confidence scoring: sim = 1 / (1 + L2_distance). If best similarity | |
| < CONFIDENCE_THRESHOLD, sets needs_web=True to route to web_node. | |
| """ | |
| query = state["query"] | |
| # Access the current FAISS index via the module — never cache this | |
| # as a local variable because chunk_and_index() updates it in place. | |
| vs = vs_module.vectorstore | |
| if vs is None: | |
| logger.info("FAISS store is empty -> routing to web search") | |
| return { | |
| "retrieved_docs": "", | |
| "needs_web": True, | |
| "confidence": 0.0, | |
| "source_type": "none", | |
| } | |
| try: | |
| docs_with_scores = vs.similarity_search_with_score(query, k=6) | |
| except Exception as e: | |
| logger.warning(f"FAISS search failed: {e} -> web fallback") | |
| return { | |
| "retrieved_docs": "", | |
| "needs_web": True, | |
| "confidence": 0.0, | |
| "source_type": "none", | |
| } | |
| if not docs_with_scores: | |
| logger.info("No matches -> routing to web search") | |
| return { | |
| "retrieved_docs": "", | |
| "needs_web": True, | |
| "confidence": 0.0, | |
| "source_type": "none", | |
| } | |
| distances = np.array([score for _, score in docs_with_scores]) | |
| similarities = 1.0 / (1.0 + distances) | |
| best_score = float(np.max(similarities)) | |
| logger.info(f"best_sim={best_score:.3f} (threshold={CONFIDENCE_THRESHOLD})") | |
| docs_text = "\n\n".join( | |
| f"[Source: {doc.metadata.get('source', 'unknown')}]\n" | |
| f"{doc.page_content[:800]}" | |
| for doc, _ in docs_with_scores | |
| ) | |
| needs_web = best_score < CONFIDENCE_THRESHOLD | |
| if needs_web: | |
| logger.info("Low confidence -> fetching from Tavily") | |
| else: | |
| logger.info("Sufficient confidence -> skipping web search") | |
| return { | |
| "retrieved_docs": docs_text, | |
| "needs_web": needs_web, | |
| "confidence": best_score, | |
| "source_type": "faiss", | |
| } | |
| # Node 2: Tavily web search | |
| def web_node(state: dict) -> dict: | |
| """ | |
| Fetch fresh results from Tavily and index them into FAISS. | |
| Full raw content is indexed for future queries; a capped display | |
| snippet (1500 chars/source) goes into the answer prompt. | |
| """ | |
| query = state["query"] | |
| logger.info(f"Searching Tavily: {query!r}") | |
| if not TAVILY_API_KEY: | |
| logger.warning("TAVILY_API_KEY not set — skipping web search") | |
| return {"web_results": "", "sources_count": 0, "source_type": "none"} | |
| try: | |
| client = TavilyClient(api_key=TAVILY_API_KEY) | |
| response = client.search( | |
| query=query, | |
| max_results=5, | |
| search_depth="basic", | |
| include_raw_content=True, | |
| include_answer=False, | |
| ) | |
| except Exception as e: | |
| logger.error(f"Tavily failed: {e}") | |
| return {"web_results": "", "sources_count": 0, "source_type": "none"} | |
| results = response.get("results", []) | |
| if not results: | |
| logger.info("Tavily returned no results") | |
| return {"web_results": "", "sources_count": 0, "source_type": "none"} | |
| parts = [] | |
| indexed_count = 0 | |
| for item in results: | |
| title = item.get("title", "Untitled") | |
| url = item.get("url", "") | |
| raw = (item.get("raw_content") or item.get("content", "")).strip() | |
| display = raw[:1500] | |
| if not display: | |
| continue | |
| if url and len(raw) > 100: | |
| try: | |
| vs_module.chunk_and_index(raw, source_url=url) | |
| indexed_count += 1 | |
| except Exception as e: | |
| logger.warning(f"Indexing failed for {url}: {e}") | |
| parts.append(f"TITLE: {title}\nURL: {url}\n\n{display}") | |
| logger.info(f"{len(parts)} source(s) fetched | {indexed_count} indexed into FAISS") | |
| separator = "\n\n" + ("─" * 60) + "\n\n" | |
| web_results = separator.join(parts) if parts else "" | |
| return { | |
| "web_results": web_results, | |
| "source_type": "web", | |
| "sources_count": len(parts), | |
| } | |
| # Pipeline orchestrator | |
| def run_pipeline_phase1(query: str, history: list = None) -> dict: | |
| """ | |
| Call retrieve_node and (optionally) web_node in sequence. | |
| Direct node calls (not app.invoke()) let app.py intercept between | |
| retrieval and generation — necessary to show a 'Generating...' | |
| status update to the user before the (now much faster, but still | |
| non-zero) GPU generation call runs. | |
| """ | |
| from ui_helpers import build_retrieval_query | |
| if history is None: | |
| history = [] | |
| retrieval_query = build_retrieval_query(query, history) | |
| if retrieval_query != query: | |
| logger.debug(f"Follow-up enriched: {retrieval_query!r}") | |
| state = { | |
| "query": retrieval_query, | |
| "retrieved_docs": "", | |
| "web_results": "", | |
| "answer": "", | |
| "needs_web": False, | |
| "confidence": 0.0, | |
| "source_type": "none", | |
| "sources_count": 0, | |
| "prompt_tokens": 0, | |
| "t_retrieve": 0.0, | |
| "t_web": 0.0, | |
| "history": history, | |
| } | |
| t0 = time.perf_counter() | |
| state.update(retrieve_node(state)) | |
| state["t_retrieve"] = time.perf_counter() - t0 | |
| if state.get("needs_web", False): | |
| t1 = time.perf_counter() | |
| state.update(web_node(state)) | |
| state["t_web"] = time.perf_counter() - t1 | |
| # Restore original query — the model answers what the user typed, | |
| # not the internally enriched search string | |
| state["query"] = query | |
| return state | |