forked from fangtianchen/algonotes_rag
59 lines
1.5 KiB
Python
59 lines
1.5 KiB
Python
# scripts/chat.py
|
|
# CLI: interactive RAG chat with context preservation.
|
|
|
|
import argparse
|
|
|
|
from src.rag.agent import create_rag_agent
|
|
|
|
DEFAULT_THREAD_ID = "default"
|
|
|
|
|
|
_parser = argparse.ArgumentParser(prog="algonotes chat")
|
|
|
|
|
|
def chat_loop(thread_id):
|
|
|
|
agent = create_rag_agent()
|
|
config = {"configurable": {"thread_id": thread_id}}
|
|
|
|
print("RAG Agent 交互式问答")
|
|
print("输入 exit 退出\n")
|
|
|
|
while True:
|
|
query = input("问题: ")
|
|
if query.lower() in ("exit", "quit"):
|
|
break
|
|
|
|
print("思考中...\n")
|
|
|
|
full_response = ""
|
|
for event in agent.stream(
|
|
{"messages": [("user", query)]},
|
|
config,
|
|
):
|
|
if "model" in event:
|
|
content = event["model"]["messages"][-1].content
|
|
if content:
|
|
print(content, flush=True)
|
|
full_response += content
|
|
elif "tools" in event:
|
|
tool_msg = event["tools"]["messages"][-1]
|
|
print(
|
|
f"\n[Tool] {tool_msg.name} 返回 {len(tool_msg.content)} 字符\n",
|
|
flush=True,
|
|
)
|
|
|
|
if full_response:
|
|
print()
|
|
print("-" * 50)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
_parser.add_argument(
|
|
"--thread-id",
|
|
default=DEFAULT_THREAD_ID,
|
|
help="会话标识符,用于区分不同对话",
|
|
)
|
|
args = _parser.parse_args()
|
|
chat_loop(args.thread_id)
|