kangda-robot-backend/ruoyi-fastapi-backend/test/test_sse.py

170 lines
5.9 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 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,
)
)