Files

89 lines
2.6 KiB
Python

import asyncio
from pathlib import Path
from langchain_ollama import ChatOllama, OllamaEmbeddings
from langchain.embeddings import OpenAIEmbeddings
from langchain.vectorstores import Chroma
from langchain.agents import Tool, AgentExecutor, initialize_agent, AgentType
from langchain.tools import BaseTool
import httpx
import os
# Load FAQ data
DATA_DIR = Path("data")
# Embedding model
embeddings = OllamaEmbeddings(model="nomic-embed-text")
# Create Chroma store
def load_faq_to_chroma() -> Chroma:
from langchain.document_loaders import TextLoader
from langchain.text_splitter import RecursiveCharacterTextSplitter
docs = []
for md_file in DATA_DIR.glob("*.md"):
loader = TextLoader(str(md_file))
docs.extend(loader.load_and_split(RecursiveCharacterTextSplitter(chunk_size=500, chunk_overlap=50)))
db = Chroma.from_documents(docs, embeddings, persist_directory="./chroma_faq")
db.persist()
return db
chroma_db = load_faq_to_chroma()
# Tool: search in FAQ
class SearchFAQTool(BaseTool):
name = "search_course_docs"
description = "Search local FAQ docs. Use query string. Returns top k results."
def _run(self, query: str, k: int = 3):
results = chroma_db.similarity_search_with_score(query, k)
return "\n".join([f"{i+1}. {r[0].page_content[:200]}... (score: {r[1]:.4f})" for i, r in enumerate(results)])
search_tool = SearchFAQTool()
# Tool: fetch course metadata (MCP style)
class FetchMetaTool(BaseTool):
name = "fetch_course_meta"
description = "Fetch course metadata via HTTP. Use query string. Returns JSON string."
def _run(self, query: str):
# For demo, use local JSON file or mock endpoint
url = f"http://localhost:8000/meta?query={query}"
try:
resp = httpx.get(url, timeout=5)
resp.raise_for_status()
return resp.text
except Exception as e:
return f"Error fetching meta: {e}"
meta_tool = FetchMetaTool()
# LLM
llm = ChatOllama(model="llama3")
# Agent
tools = [search_tool, meta_tool]
agent = initialize_agent(
tools,
llm,
agent=AgentType.ZERO_SHOT_REACT_DESCRIPTION,
verbose=True,
handle_parsing_errors=True,
)
async def main():
# Simple CLI with predefined questions
questions = [
"What is the deadline for assignment 3?",
"How to use ChromaDB with LangChain?",
"What is the schedule for next week?",
]
for q in questions:
print("\nQuestion:", q)
result = await agent.arun(input=q)
print("Answer:\n", result)
if __name__ == "__main__":
asyncio.run(main())