170 lines
5.9 KiB
Python
170 lines
5.9 KiB
Python
import argparse # 解析命令行参数
|
||
import asyncio # 异步事件循环
|
||
import json # 处理 JSON 文本
|
||
import os # 读取环境变量
|
||
import threading # 后台线程
|
||
import time # 简单的等待/演示主线程未阻塞
|
||
import uuid # 生成随机 chat/session id
|
||
|
||
import aiohttp # 异步 HTTP 客户端
|
||
import requests # 同步 HTTP 客户端
|
||
|
||
# 下面的默认参数可以通过环境变量覆盖,避免把机密写死在代码里。
|
||
API_URL = os.environ.get(
|
||
"RAGFLOW_URL",
|
||
"http://10.0.0.202:9099/system/ragflow/converse_with_chat_assistant",
|
||
)
|
||
AUTH_TOKEN = os.environ.get(
|
||
"RAGFLOW_TOKEN",
|
||
"Bearer ",
|
||
)
|
||
DEFAULT_CHAT_ID = os.environ.get(
|
||
"RAGFLOW_CHAT_ID", "db4bb966895b11f08cda0242ac130006"
|
||
)
|
||
DEFAULT_SESSION_ID = os.environ.get(
|
||
"RAGFLOW_SESSION_ID", "38d765e48a3811f0be310242ac130006"
|
||
)
|
||
DEFAULT_QUESTION = os.environ.get(
|
||
"RAGFLOW_QUESTION", "你好,请用简洁的语言介绍你自己"
|
||
)
|
||
LOGIN_URL = os.environ.get(
|
||
"RAGFLOW_LOGIN_URL",
|
||
"http://10.0.0.202:9099/login",
|
||
)
|
||
|
||
|
||
async def login_get_token() -> str:
|
||
payload = {
|
||
"username": "admin",
|
||
"password": "admin123"
|
||
}
|
||
headers = {
|
||
"Authorization": "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.eyJ1c2VyX2lkIjoiMSIsInVzZXJfbmFtZSI6ImFkbWluIiwiZGVwdF9uYW1lIjoiXHU3ODE0XHU1M2QxXHU5MGU4XHU5NWU4Iiwic2Vzc2lvbl9pZCI6IjYwNjdlMzkxLTRiNGMtNGUwYy1hZjM5LTVmNzczZjIzMmE3OSIsImxvZ2luX2luZm8iOnsiaXBhZGRyIjpudWxsLCJsb2dpbkxvY2F0aW9uIjoiXHU2NzJhXHU3N2U1IiwiYnJvd3NlciI6Ik90aGVyIiwib3MiOiJPdGhlciIsImxvZ2luVGltZSI6IjIwMjUtMTItMDIgMTA6NTk6NDAifSwiZXhwIjoxNzY3MjM2MzgwfQ.NM9j0emNz8trxU1DX87PVZZaWWHlFFwF7o0Pxop6B5c"
|
||
}
|
||
async with aiohttp.ClientSession() as session:
|
||
# 使用data参数发送form-data而不是json
|
||
async with session.post(
|
||
LOGIN_URL, data=payload, headers=headers, timeout=60
|
||
) as resp:
|
||
response_data = await resp.json() # 将响应转换为JSON
|
||
token = response_data.get("token") # 获取token字段
|
||
return token
|
||
|
||
|
||
async def stream_chat_async(question: str, chat_id: str, session_id: str) -> str:
|
||
"""
|
||
异步版(aiohttp),适合已有事件循环的应用,不阻塞主线程。
|
||
"""
|
||
token = await login_get_token()
|
||
|
||
payload = {
|
||
"chatId": chat_id,
|
||
"question": question,
|
||
"stream": True,
|
||
"sessionId": session_id,
|
||
}
|
||
headers = {
|
||
"Authorization": f"Bearer {token}",
|
||
"Accept": "text/event-stream",
|
||
}
|
||
|
||
answer_parts = []
|
||
event_type = "message"
|
||
|
||
async with aiohttp.ClientSession() as session:
|
||
async with session.post(
|
||
API_URL, json=payload, headers=headers, timeout=60
|
||
) as resp:
|
||
print(
|
||
f"# status={resp.status} content-type={resp.headers.get('content-type')}")
|
||
resp.raise_for_status()
|
||
|
||
while True:
|
||
raw_line = await resp.content.readline()
|
||
if not raw_line:
|
||
break
|
||
line = raw_line.decode(errors="ignore").strip() # 逐行读取 SSE 数据
|
||
if not line or line.startswith(":"):
|
||
continue
|
||
if line.startswith("event:"):
|
||
event_type = line.split(":", 1)[1].strip() or "message"
|
||
continue
|
||
if not line.startswith("data:"):
|
||
print(f"[skip] {line}")
|
||
continue
|
||
data_str = line[len("data:"):].strip()
|
||
if not data_str:
|
||
continue
|
||
try:
|
||
payload_obj = json.loads(data_str)
|
||
except json.JSONDecodeError:
|
||
print(f"[{event_type}] {data_str}")
|
||
event_type = "message"
|
||
continue
|
||
|
||
if event_type == "end" or payload_obj.get("status") == "completed":
|
||
print("\n[stream] completed")
|
||
break
|
||
if payload_obj.get("data") is True:
|
||
event_type = "message"
|
||
continue
|
||
if "answer" in payload_obj:
|
||
piece = payload_obj.get("answer", "")
|
||
answer_parts.append(piece)
|
||
# 流式输出:每个片段都单独打印,并且添加延时让效果更明显
|
||
print(piece, end="", flush=True)
|
||
time.sleep(0.1) # 每个片段间隔100ms,让流式效果更明显
|
||
else:
|
||
print(f"[{event_type}] {payload_obj}")
|
||
event_type = "message"
|
||
|
||
print("\n\n[stream completed] - Answer received in real-time via streaming!")
|
||
full_answer = "".join(answer_parts)
|
||
return full_answer
|
||
|
||
|
||
if __name__ == "__main__":
|
||
parser = argparse.ArgumentParser(
|
||
description="Call ragflow chat endpoint with SSE streaming."
|
||
)
|
||
parser.add_argument(
|
||
"question",
|
||
nargs="?",
|
||
help="要提的问题,未提供则用环境变量 RAGFLOW_QUESTION",
|
||
)
|
||
parser.add_argument(
|
||
"-q",
|
||
"--question",
|
||
dest="question_opt",
|
||
help="与位置参数等价,优先级更高。",
|
||
)
|
||
parser.add_argument(
|
||
"--chat-id",
|
||
default=DEFAULT_CHAT_ID,
|
||
help="chatId,默认取 RAGFLOW_CHAT_ID。",
|
||
)
|
||
parser.add_argument(
|
||
"--session-id",
|
||
default=DEFAULT_SESSION_ID or uuid.uuid4().hex,
|
||
help="sessionId,默认取 RAGFLOW_SESSION_ID,否则随机生成。",
|
||
)
|
||
parser.add_argument(
|
||
"--url",
|
||
default=API_URL,
|
||
help="接口地址,默认取 RAGFLOW_URL。",
|
||
)
|
||
args = parser.parse_args()
|
||
|
||
# Allow either positional or -q/--question to override env default.
|
||
question = args.question_opt or args.question or DEFAULT_QUESTION
|
||
|
||
# Override globals if provided via CLI.
|
||
API_URL = args.url # type: ignore
|
||
|
||
asyncio.run(
|
||
stream_chat_async(
|
||
question=question,
|
||
chat_id=args.chat_id,
|
||
session_id=args.session_id,
|
||
)
|
||
) |