57 lines
2.0 KiB
Python
57 lines
2.0 KiB
Python
from langchain_ollama import ChatOllama
|
|
from langchain.prompts import ChatPromptTemplate
|
|
from langchain.agents import create_agent
|
|
from langchain.tools import tool
|
|
import os
|
|
|
|
MODEL = os.getenv("OLLAMA_MODEL", "llama3")
|
|
BASE_URL = os.getenv("OLLAMA_BASE_URL", "http://localhost:11434/v1")
|
|
|
|
# Initialize LLM with local Ollama endpoint.
|
|
llm = ChatOllama(model=MODEL, base_url=BASE_URL)
|
|
|
|
SYSTEM_PROMPT = """You are an assistant that answers user queries using a knowledge base. Use the provided tools to search and add content."""
|
|
prompt = ChatPromptTemplate.from_messages([
|
|
("system", SYSTEM_PROMPT),
|
|
])
|
|
|
|
from .rag_tools import add_to_knowledge_base as rag_add, search_knowledge_base as rag_search
|
|
from .init_loader import load_documents
|
|
|
|
# Load existing documents from ./data directory at startup
|
|
load_documents("./data")
|
|
|
|
@tool
|
|
def add_to_knowledge_base(content: str, title: str = "Document") -> str:
|
|
"""Add content to the knowledge base."""
|
|
return rag_add(content=content, title=title)
|
|
|
|
@tool
|
|
def search_knowledge_base(query: str, max_results: int = 5) -> str:
|
|
"""Search the knowledge base for relevant chunks."""
|
|
results = rag_search(query=query, max_results=max_results)
|
|
formatted = "\n".join([f"{score:.4f}: {text[:200]}..." for text, score in results])
|
|
return formatted if formatted else "No results found."
|
|
|
|
tools = [add_to_knowledge_base, search_knowledge_base]
|
|
agent = create_agent(llm=llm, prompt=prompt, tools=tools)
|
|
|
|
executor = agent
|
|
|
|
if __name__ == "__main__":
|
|
print("RAG Agent Interactive Mode. Type /quit to exit.")
|
|
while True:
|
|
try:
|
|
user_input = input("User: ")
|
|
except (EOFError, KeyboardInterrupt):
|
|
print("\nGoodbye!")
|
|
break
|
|
if user_input.strip().lower() in {"/quit", "quit"}:
|
|
print("Goodbye!")
|
|
break
|
|
try:
|
|
response = executor.invoke({"input": user_input})
|
|
print("Assistant:", response)
|
|
except Exception as e:
|
|
print("Error:", str(e))
|