89 lines
2.6 KiB
Python
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())
|