Files
task-6a1d75c5fd30e81cf3126ae7/main.py
T
2026-06-04 16:05:58 +00:00

131 lines
4.2 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
import os
import asyncio
import json
from pathlib import Path
from typing import List
import httpx
from langchain_openai import ChatOpenAI
from langchain_core.messages import HumanMessage
from langchain.tools import tool
from deepagents import create_deep_agent
from deepagents.backends import FilesystemBackend, LocalShellBackend, CompositeBackend
from langchain_chroma import Chroma
from langchain_ollama import OllamaEmbeddings
# ---------------------
# 1. Chroma DB helpers
# ---------------------
CHROMA_PATH = Path("./chroma_faq")
DATA_DIR = Path("./data")
def load_faq_to_chroma() -> None:
"""Load all .md files from data/ into a persistent Chroma store."""
if CHROMA_PATH.exists():
# already loaded
return
CHROMA_PATH.mkdir(parents=True, exist_ok=True)
embeddings = OllamaEmbeddings(model="nomic-embed-text")
db = Chroma.from_folder(str(DATA_DIR), embedding=embeddings, persist_directory=str(CHROMA_PATH))
db.persist()
@tool
def search_course_docs(query: str, k: int = 3) -> str:
"""Search local course FAQ documents in Chroma DB."""
load_faq_to_chroma()
embeddings = OllamaEmbeddings(model="nomic-embed-text")
db = Chroma(persist_directory=str(CHROMA_PATH), embedding=embeddings)
docs = db.similarity_search(query, k=k)
if not docs:
return "No relevant FAQ found."
return "\n\n---\n\n".join(doc.page_content for doc in docs)
# ---------------------
# 2. MCPstyle tool
# ---------------------
# For demo we use a local JSON file served by python -m http.server
# The file is located at ./meta/course_meta.json
META_URL = "http://localhost:8000/course_meta.json"
@tool
def fetch_course_meta(query: str) -> str:
"""Fetch course metadata (e.g., schedule) from a mock MCP server."""
try:
response = httpx.get(META_URL, timeout=5.0)
response.raise_for_status()
data = response.json()
except Exception as e:
return f"Error fetching metadata: {e}"
# Simple keyword search in the JSON
results: List[str] = []
for key, value in data.items():
if query.lower() in key.lower() or query.lower() in str(value).lower():
results.append(f"{key}: {value}")
return "\n".join(results) if results else "No metadata matches your query."
# ---------------------
# 3. Agent setup
# ---------------------
llm = ChatOpenAI(
model="openai/gpt-oss-20b:free",
base_url="https://openrouter.ai/api/v1",
api_key=os.getenv("OPENAI_API_KEY"),
temperature=0.0,
)
backend = CompositeBackend([
LocalShellBackend(workspace_dir="./workspace"),
FilesystemBackend(),
])
SYSTEM_PROMPT = (
"You are a helpful FAQ assistant for the course.\n"
"When a user asks a question, first decide whether the answer is best found in the local FAQ documents or in the course metadata.\n"
"If the answer is in the FAQ, use the tool `search_course_docs`.\n"
"If the answer requires schedule or other metadata, use the tool `fetch_course_meta`.\n"
"Do not call both tools unless absolutely necessary.\n"
"In your final response, prepend the source: `source: chroma` or `source: mcp_meta`."
)
agent = create_deep_agent(
model=llm,
tools=[search_course_docs, fetch_course_meta],
backend=backend,
system_prompt=SYSTEM_PROMPT,
)
# ---------------------
# 4. CLI
# ---------------------
PRESET_QUESTIONS = [
"What is the deadline for the final project?", # FAQ
"When does the next lecture on deep learning start?", # metadata
"Explain the concept of attention mechanism.", # FAQ
]
async def run_agent(question: str) -> str:
result = await agent.ainvoke(
{"messages": [HumanMessage(content=question)]},
{"configurable": {"thread_id": "session-1"}},
)
return result["messages"][-1].content
async def main():
print("--- FAQ Bot Demo ---\n")
for i, q in enumerate(PRESET_QUESTIONS, 1):
print(f"Q{i}: {q}")
ans = await run_agent(q)
print(f"A{i}: {ans}\n")
print("Enter your own question (or 'exit' to quit):")
while True:
user_q = input("> ")
if user_q.lower() in {"exit", "quit"}:
break
ans = await run_agent(user_q)
print(ans)
if __name__ == "__main__":
asyncio.run(main())