Spaces:
Runtime error
Runtime error
| import json | |
| import os | |
| from typing import List, Optional | |
| import pandas as pd | |
| import pytesseract | |
| import requests | |
| from dotenv import load_dotenv | |
| from langchain_community.document_loaders import WikipediaLoader | |
| from langchain_community.tools.tavily_search import TavilySearchResults | |
| from langchain_core.messages import HumanMessage, SystemMessage | |
| from langchain_core.tools import tool | |
| from langchain_groq import ChatGroq | |
| from langgraph.graph import START, MessagesState, StateGraph | |
| from langgraph.prebuilt import ToolNode, tools_condition | |
| from PIL import Image | |
| from code_interpreter import CodeInterpreter | |
| from extract_answer import extract_final_answer | |
| from files_util import format_user_message | |
| load_dotenv() | |
| ROOT_DIR = os.path.dirname(os.path.abspath(__file__)) | |
| PROMPT_PATH = os.path.join(ROOT_DIR, "system_prompt.txt") | |
| MAX_TOOL_ROUNDS = 8 | |
| RECURSION_LIMIT = 20 | |
| _interpreter = None | |
| def get_system_prompt() -> str: | |
| with open(PROMPT_PATH, "r", encoding="utf-8") as handle: | |
| return handle.read() | |
| def get_interpreter() -> CodeInterpreter: | |
| global _interpreter | |
| if _interpreter is None: | |
| _interpreter = CodeInterpreter() | |
| return _interpreter | |
| def wiki_search(query: str) -> str: | |
| """Search Wikipedia and return up to 2 page excerpts. | |
| Args: | |
| query: Search query, preferably a proper name or article title. | |
| """ | |
| docs = WikipediaLoader(query=query, load_max_docs=2).load() | |
| if not docs: | |
| return "No Wikipedia results." | |
| chunks = [] | |
| for doc in docs: | |
| source = doc.metadata.get("source", "") | |
| chunks.append(f'<Document source="{source}">\n{doc.page_content}\n</Document>') | |
| return "\n\n---\n\n".join(chunks) | |
| def web_search(query: str) -> str: | |
| """Search the public web via Tavily and return up to 3 results. | |
| Args: | |
| query: Search query. Be specific; include years, full names, and distinctive terms. | |
| """ | |
| results = TavilySearchResults(max_results=3).invoke(query) | |
| if not results: | |
| return "No web results." | |
| chunks = [] | |
| for doc in results: | |
| chunks.append( | |
| f'<Document source="{doc.get("url", "")}" title="{doc.get("title", "")}">\n' | |
| f'{doc.get("content", "")}\n</Document>' | |
| ) | |
| return "\n\n---\n\n".join(chunks) | |
| def execute_python(code: str) -> str: | |
| """Run Python for calculation, pandas/Excel/CSV analysis, dates, and string parsing. | |
| Args: | |
| code: Complete Python source. Print the values you need. Use absolute file paths from the question. | |
| """ | |
| result = get_interpreter().execute_code(code, language="python") | |
| lines = [] | |
| if result["status"] == "success": | |
| lines.append("Python execution succeeded.") | |
| if result.get("stdout"): | |
| lines.append("stdout:\n" + result["stdout"].strip()) | |
| if result.get("stderr"): | |
| lines.append("stderr:\n" + result["stderr"].strip()) | |
| if result.get("result") is not None: | |
| lines.append("result:\n" + str(result["result"]).strip()) | |
| for df_info in result.get("dataframes") or []: | |
| preview = pd.DataFrame(df_info["head"]) | |
| lines.append(f"DataFrame {df_info['name']} shape={df_info['shape']}\n{preview}") | |
| else: | |
| lines.append("Python execution failed.") | |
| if result.get("stderr"): | |
| lines.append(result["stderr"].strip()) | |
| return "\n\n".join(lines)[:12000] | |
| def download_file_from_url(url: str, filename: Optional[str] = None) -> str: | |
| """Download a file from an http(s) URL to a local path and return that path. | |
| Args: | |
| url: Direct download URL. | |
| filename: Optional filename to use in the temp directory. | |
| """ | |
| from urllib.parse import urlparse | |
| import tempfile | |
| import uuid | |
| try: | |
| if not filename: | |
| filename = os.path.basename(urlparse(url).path) or f"download_{uuid.uuid4().hex[:8]}" | |
| filepath = os.path.join(tempfile.gettempdir(), filename) | |
| response = requests.get(url, timeout=60, stream=True) | |
| response.raise_for_status() | |
| with open(filepath, "wb") as handle: | |
| for chunk in response.iter_content(chunk_size=8192): | |
| handle.write(chunk) | |
| return f"Downloaded to {filepath}" | |
| except Exception as exc: | |
| return f"Download failed: {exc}" | |
| def read_file(file_path: str, max_chars: int = 8000) -> str: | |
| """Read a local text, CSV, Excel, JSON, or PDF file and return a preview. | |
| Args: | |
| file_path: Absolute path on disk. | |
| max_chars: Truncate text-like previews to this many characters. | |
| """ | |
| if not os.path.exists(file_path): | |
| return f"File not found: {file_path}" | |
| ext = os.path.splitext(file_path)[1].lower() | |
| try: | |
| if ext in {".csv"}: | |
| df = pd.read_csv(file_path) | |
| return ( | |
| f"CSV rows={len(df)} cols={list(df.columns)}\n" | |
| f"{df.head(20).to_string()}\n\n{df.describe(include='all')}" | |
| )[:max_chars] | |
| if ext in {".xlsx", ".xls"}: | |
| df = pd.read_excel(file_path) | |
| return ( | |
| f"Excel rows={len(df)} cols={list(df.columns)}\n" | |
| f"{df.head(20).to_string()}\n\n{df.describe(include='all')}" | |
| )[:max_chars] | |
| if ext == ".pdf": | |
| try: | |
| import fitz | |
| except ImportError: | |
| return "PyMuPDF is not installed; use execute_python if available." | |
| doc = fitz.open(file_path) | |
| text = "\n".join(page.get_text() for page in doc) | |
| doc.close() | |
| return text[:max_chars] or "PDF had no extractable text." | |
| with open(file_path, "r", encoding="utf-8", errors="replace") as handle: | |
| return handle.read(max_chars) | |
| except Exception as exc: | |
| return f"Failed to read {file_path}: {exc}" | |
| def extract_text_from_image(image_path: str) -> str: | |
| """OCR text from a local image with Tesseract. | |
| Args: | |
| image_path: Absolute path to a jpg/png/webp/gif/bmp file. | |
| """ | |
| try: | |
| image = Image.open(image_path) | |
| text = pytesseract.image_to_string(image) | |
| return text.strip() or "OCR returned no text." | |
| except Exception as exc: | |
| return f"OCR failed: {exc}" | |
| TOOLS = [ | |
| web_search, | |
| wiki_search, | |
| execute_python, | |
| download_file_from_url, | |
| read_file, | |
| extract_text_from_image, | |
| ] | |
| def _tool_signature(message) -> Optional[str]: | |
| calls = getattr(message, "tool_calls", None) or [] | |
| if not calls: | |
| return None | |
| parts = [] | |
| for call in calls: | |
| name = call.get("name") if isinstance(call, dict) else getattr(call, "name", "") | |
| args = call.get("args") if isinstance(call, dict) else getattr(call, "args", {}) | |
| parts.append(f"{name}:{json.dumps(args, sort_keys=True, default=str)}") | |
| return "|".join(parts) | |
| def _should_force_final(messages) -> bool: | |
| tool_round_messages = [m for m in messages if getattr(m, "tool_calls", None)] | |
| if len(tool_round_messages) >= MAX_TOOL_ROUNDS: | |
| return True | |
| if len(tool_round_messages) >= 2: | |
| if _tool_signature(tool_round_messages[-1]) == _tool_signature(tool_round_messages[-2]): | |
| return True | |
| return False | |
| def require_env(name: str) -> str: | |
| value = os.getenv(name, "").strip() | |
| if value: | |
| return value | |
| raise RuntimeError( | |
| f"Missing {name}. " | |
| "Local: put it in `.env`. " | |
| "Hugging Face Space: Settings → Variables and secrets → New secret." | |
| ) | |
| def build_graph(provider: str = "groq", use_retriever: bool = False): | |
| """Compile the tool-calling graph. Retriever is off unless explicitly enabled.""" | |
| if provider != "groq": | |
| raise ValueError("Only provider='groq' is supported after the refactor.") | |
| api_key = require_env("GROQ_API_KEY") | |
| llm = ChatGroq(model="qwen/qwen3-32b", temperature=0, api_key=api_key) | |
| llm_with_tools = llm.bind_tools(TOOLS) | |
| retriever_store = _build_optional_retriever() if use_retriever else None | |
| def maybe_example(state: MessagesState): | |
| if retriever_store is None: | |
| return {} | |
| question = "" | |
| for message in state["messages"]: | |
| if isinstance(message, HumanMessage): | |
| question = message.content | |
| break | |
| hits = retriever_store.similarity_search(question, k=1) | |
| if not hits: | |
| return {} | |
| return { | |
| "messages": [ | |
| HumanMessage( | |
| content=( | |
| "Optional similar example for format only. " | |
| "Do not copy if it is a different question.\n\n" | |
| f"{hits[0].page_content}" | |
| ) | |
| ) | |
| ] | |
| } | |
| def assistant(state: MessagesState): | |
| messages = state["messages"] | |
| if _should_force_final(messages): | |
| stop = SystemMessage( | |
| content=( | |
| "Stop calling tools. Using only the information already in this " | |
| "conversation, reply with one line: FINAL ANSWER: <answer>" | |
| ) | |
| ) | |
| return {"messages": [llm.invoke(messages + [stop])]} | |
| return {"messages": [llm_with_tools.invoke(messages)]} | |
| builder = StateGraph(MessagesState) | |
| builder.add_node("maybe_example", maybe_example) | |
| builder.add_node("assistant", assistant) | |
| builder.add_node("tools", ToolNode(TOOLS)) | |
| builder.add_edge(START, "maybe_example") | |
| builder.add_edge("maybe_example", "assistant") | |
| builder.add_conditional_edges("assistant", tools_condition) | |
| builder.add_edge("tools", "assistant") | |
| return builder.compile() | |
| def _build_optional_retriever(): | |
| url = os.environ.get("SUPABASE_URL") | |
| key = os.environ.get("SUPABASE_SERVICE_ROLE_KEY") | |
| if not url or not key: | |
| print("Retriever requested but SUPABASE_* env vars are missing; skipping.") | |
| return None | |
| from langchain_community.vectorstores import SupabaseVectorStore | |
| from langchain_huggingface import HuggingFaceEmbeddings | |
| from supabase.client import create_client | |
| embeddings = HuggingFaceEmbeddings(model_name="sentence-transformers/all-mpnet-base-v2") | |
| client = create_client(url, key) | |
| return SupabaseVectorStore( | |
| client=client, | |
| embedding=embeddings, | |
| table_name="documents2", | |
| query_name="match_documents_2", | |
| ) | |
| def run_agent( | |
| question: str, | |
| file_paths: Optional[List[str]] = None, | |
| graph=None, | |
| use_retriever: bool = False, | |
| ) -> dict: | |
| """Run one question and return raw last message plus extracted FINAL ANSWER.""" | |
| if graph is None: | |
| graph = build_graph(use_retriever=use_retriever) | |
| payload = format_user_message(question, file_paths) | |
| result = graph.invoke( | |
| { | |
| "messages": [ | |
| SystemMessage(content=get_system_prompt()), | |
| HumanMessage(content=payload), | |
| ] | |
| }, | |
| {"recursion_limit": RECURSION_LIMIT}, | |
| ) | |
| last = result["messages"][-1] | |
| raw = last.content if hasattr(last, "content") else str(last) | |
| if isinstance(raw, list): | |
| raw = "".join( | |
| part.get("text", "") if isinstance(part, dict) else str(part) for part in raw | |
| ) | |
| return { | |
| "raw": raw, | |
| "final": extract_final_answer(raw), | |
| "messages": result["messages"], | |
| } | |
| if __name__ == "__main__": | |
| demo = "When was a picture of St. Thomas Aquinas first added to the Wikipedia page on the Principle of double effect?" | |
| output = run_agent(demo) | |
| print("RAW:\n", output["raw"]) | |
| print("FINAL:", output["final"]) | |