4040import argparse
4141import asyncio
4242import traceback
43- from langchain_core .prompts import PromptTemplate
44- from langchain_core .prompts import ChatPromptTemplate
45- from langchain_openai import OpenAI , ChatOpenAI
46- from openai import OpenAI as OpenAIClient
47- from langchain_core .messages import HumanMessage , SystemMessage , AIMessageChunk
48- from typing import TypedDict , List
43+
4944from fastapi import FastAPI , HTTPException , Request
5045from fastapi .responses import StreamingResponse
5146from pydantic import BaseModel
52- from langchain_openai import ChatOpenAI
5347from starlette .middleware .cors import CORSMiddleware
5448
55-
56- from langgraph .prebuilt import create_react_agent
57-
58- from langchain_core .prompts import MessagesPlaceholder
59- from mcp .client .sse import sse_client
60- from mcp .client .streamable_http import streamablehttp_client
61- from mcp import ClientSession
62- from langchain_mcp_adapters .tools import load_mcp_tools
63- from langchain_core .messages import AIMessage ,ToolMessage
6449from ..config import config
65- config .validate ()
66-
67-
68- model_name = config .INFO_LLM_MODEL_NAME
69- api_key = config .INFO_LLM_MODEL_API_KEY
70- api_base = config .INFO_LLM_MODEL_BASE_URL
71- ua = config .USER_AGENT
72-
73- model_name_writer = config .WRITER_LLM_MODEL_NAME
74- api_key_writer = config .WRITER_LLM_MODEL_API_KEY
75- api_base_writer = config .WRITER_LLM_MODEL_BASE_URL
76-
77- prompt_dir = config .PROMPT_DIRECTORY
78-
79- searchinrag_prompt_path = os .path .join (config .PROMPT_DIRECTORY , "searchinrag.md" )
80- searchinweb_prompt_path = os .path .join (config .PROMPT_DIRECTORY , "searchinweb.md" )
81- writenovel_prompt_path = os .path .join (config .PROMPT_DIRECTORY , "writenovel.md" )
50+ from .novel_service import NovelGenerationService
8251
52+ # Initialize Config
53+ config .validate ()
8354
84- # 初始化 LLM 模型
85- model = ChatOpenAI (
86- model_name = model_name ,
87- api_key = api_key ,
88- base_url = api_base ,
89- )
90- model_writer = ChatOpenAI (
91- model_name = model_name_writer ,
92- api_key = api_key_writer ,
93- base_url = api_base_writer ,
94- )
55+ # Initialize Service
56+ novel_service = NovelGenerationService ()
9557
9658app = FastAPI (title = "LangChain-Umamusume-Server" )
9759
98-
99-
10060# 前端报CORS时
10161app .add_middleware (
10262 CORSMiddleware ,
10666 allow_headers = ['*' ],
10767)
10868
109- def extract_tool_info (agent_response ):
110- tool_calls = []
111- tool_results = []
112-
113- for msg in agent_response ["messages" ]:
114- # 判断消息类型
115- if isinstance (msg , AIMessage ) and hasattr (msg , "tool_calls" ):
116- for tool_call in msg .tool_calls :
117- tool_calls .append ({
118- "name" : tool_call ["name" ],
119- "arguments" : tool_call ["args" ]
120- })
121-
122- elif isinstance (msg , ToolMessage ):
123- tool_results .append ({
124- "id" : msg .tool_call_id ,
125- "name" : msg .name ,
126- "content" : msg .content ,
127- "status" : msg .status
128- })
129-
130- # 提取最终回答
131- final_answer = None
132- for msg in reversed (agent_response ["messages" ]):
133- if isinstance (msg , AIMessage ):
134- final_answer = msg .content
135- break
136-
137- return {
138- "tool_calls" : tool_calls ,
139- "tool_results" : tool_results ,
140- "final_answer" : final_answer
141- }
142-
14369# 定义请求模型
14470class QuestionRequest (BaseModel ):
14571 question : str
@@ -157,80 +83,27 @@ async def ask_question(request: Request, user_request: QuestionRequest):
15783
15884 rag_url = request .app .state .rag_mcp_url
15985 web_url = request .app .state .web_mcp_url
86+
16087 # --- 可选:添加防御性检查 ---
16188 if not rag_url or not web_url :
16289 error_msg = "MCP server URLs are not configured. Please start the server with -r and -w arguments."
16390 print (f"[ERROR] { error_msg } " )
16491 raise HTTPException (status_code = 500 , detail = error_msg )
92+
16593 try :
166- # 第一阶段:使用 RAG Agent 获取基础信息
167- async with streamablehttp_client (rag_url ) as (read_stream , write_stream , get_session_id ):
168- # print('MCP server连接成功')
169- async with ClientSession (read_stream , write_stream ) as session :
170-
171- await session .initialize ()
172- rag_tools = await load_mcp_tools (session )
173- print ("RAG可用MCP工具:" , [tool .name for tool in rag_tools ])
174-
175- rag_agent = create_react_agent (model , rag_tools )
176- with open (searchinrag_prompt_path , "r" , encoding = "utf-8" ) as file :
177- template = file .read ()
178- first_input = template .format (user_question = user_question )
179- rag_result = await rag_agent .ainvoke ({"messages" : [HumanMessage (content = first_input )]})
180- base_info = rag_result ["messages" ][- 1 ].content
181- result1 = extract_tool_info (rag_result )
182- print (f"[第一阶段结果] 基础信息: { base_info } " )
183- print ("\n [第一阶段Tool Call] : " , result1 ["tool_calls" ])
184-
185- # 第二阶段:使用 Web Agent 结合基础信息回答问题
186- async with streamablehttp_client (web_url ) as (read_stream , write_stream , get_session_id ):
187- # print('MCP server连接成功')
188- async with ClientSession (read_stream , write_stream ) as session :
189- await session .initialize ()
190- web_tools = await load_mcp_tools (session )
191- print ("Web可用MCP工具:" , [tool .name for tool in web_tools ])
192-
193- web_agent = create_react_agent (model , web_tools )
194- with open (searchinweb_prompt_path , "r" , encoding = "utf-8" ) as file :
195- template = file .read ()
196- final_input = template .format (user_question = user_question ,base_info = base_info )
197- web_result = await web_agent .ainvoke (
198- {"messages" : [HumanMessage (content = final_input )]},
199- config = {"recursion_limit" : 75 }
200- )
201- final_answer = web_result ["messages" ][- 1 ].content
202- result2 = extract_tool_info (web_result )
203- print (f"[第二阶段结果] 最终回答: { final_answer } " )
204- print ("\n [第二阶段Tool Call]: " , result2 ["tool_calls" ])
205- print ("\n [第二阶段Tool Results]: " , result2 ["tool_results" ])
206- web_info = final_answer
207-
208- # return AnswerResponse(answer=final_answer)
209-
210- # 第三阶段:使用结合基础信息创作小说
211- with open (writenovel_prompt_path , "r" , encoding = "utf-8" ) as file :
212- template = file .read ()
213- final_input = template .format (user_question = user_question ,base_info = base_info ,web_info = web_info )
214- novel_agent = create_react_agent (model_writer ,tools = [])
215- result = await novel_agent .ainvoke ({"messages" : [HumanMessage (content = final_input )]})
216- final_answer = result ["messages" ][- 1 ].content
94+ final_answer = await novel_service .process_novel_generation (user_question , rag_url , web_url )
21795 print (" Final Answer:\n " , final_answer )
21896 return AnswerResponse (answer = final_answer )
21997
220-
221- except BaseException as e :
98+ except Exception as e :
22299 # --- 打印详细的错误信息 ---
223100 error_msg = f"[ERROR] Unhandled error during processing in /ask route: { type (e ).__name__ } : { e } "
224101 print (error_msg )
225102
226103 # --- 打印完整的堆栈跟踪 ---
227104 traceback_str = traceback .format_exc ()
228105 print (f"[ERROR] Full Traceback:\n { traceback_str } " )
229- # --- 打印结束 ---
230106
231- # --- 返回给客户端更具体的错误信息---
232- # 将 traceback_str 也包含在 detail 中可以帮助开发者快速定位问题
233- # 不要将敏感信息通过 detail 暴露给最终用户
234107 detailed_error = f"{ str (e )} \n (See server logs for full traceback)"
235108 raise HTTPException (status_code = 500 , detail = detailed_error )
236109
@@ -248,113 +121,10 @@ async def ask_question_stream(request: Request, user_request: QuestionRequest):
248121
249122 async def event_stream ():
250123 try :
251- # --- 第一阶段:RAG ---
252- yield json .dumps ({"event" : "status" , "data" : "正在RAG搜索中,查找相关角色信息..." }, ensure_ascii = False ) + "\n "
253- print ("[STREAM] 开始第一阶段:RAG 搜索" )
254-
255- try :
256- async with streamablehttp_client (rag_url ) as (read_stream , write_stream , get_session_id ):
257- async with ClientSession (read_stream , write_stream ) as session :
258- await session .initialize ()
259- rag_tools = await load_mcp_tools (session )
260- rag_agent = create_react_agent (model , rag_tools )
261-
262- with open (searchinrag_prompt_path , "r" , encoding = "utf-8" ) as file :
263- template = file .read ()
264- first_input = template .format (user_question = user_question )
265- rag_result = await rag_agent .ainvoke ({"messages" : [HumanMessage (content = first_input )]})
266- base_info = rag_result ["messages" ][- 1 ].content
267- print (f"[STREAM] RAG 阶段完成,基础信息长度: { len (base_info )} " )
268- yield json .dumps ({"event" : "status" , "data" : "RAG搜索完成!" }, ensure_ascii = False ) + "\n "
269- yield json .dumps ({"event" : "rag_result" , "data" : base_info }, ensure_ascii = False ) + "\n "
270- except asyncio .CancelledError :
271- print ("[STREAM] RAG 阶段被用户取消" )
272- yield json .dumps ({"event" : "error" , "data" : "生成已被取消" }, ensure_ascii = False ) + "\n "
273- return
124+ async for event in novel_service .process_novel_generation_stream (user_question , rag_url , web_url ):
125+ yield event
274126
275- # --- 第二阶段:Web Search ---
276- yield json .dumps ({"event" : "status" , "data" : "正在Web搜索中,获取更多相关信息..." }, ensure_ascii = False ) + "\n "
277- print ("[STREAM] 开始第二阶段:Web 搜索" )
278-
279- try :
280- async with streamablehttp_client (web_url ) as (read_stream , write_stream , get_session_id ):
281- async with ClientSession (read_stream , write_stream ) as session :
282- await session .initialize ()
283- web_tools = await load_mcp_tools (session )
284- web_agent = create_react_agent (model , web_tools )
285-
286- with open (searchinweb_prompt_path , "r" , encoding = "utf-8" ) as file :
287- template = file .read ()
288- final_input = template .format (user_question = user_question , base_info = base_info )
289- web_result = await web_agent .ainvoke (
290- {"messages" : [HumanMessage (content = final_input )]},
291- config = {"recursion_limit" : 75 }
292- )
293- web_info = web_result ["messages" ][- 1 ].content
294- print (f"[STREAM] Web 搜索阶段完成,信息长度: { len (web_info )} " )
295- yield json .dumps ({"event" : "status" , "data" : "Web搜索完成!" }, ensure_ascii = False ) + "\n "
296- yield json .dumps ({"event" : "web_result" , "data" : web_info }, ensure_ascii = False ) + "\n "
297- except asyncio .CancelledError :
298- print ("[STREAM] Web 搜索阶段被用户取消" )
299- yield json .dumps ({"event" : "error" , "data" : "生成已被取消" }, ensure_ascii = False ) + "\n "
300- return
301-
302- # --- 第三阶段:小说生成(流式)---
303- yield json .dumps ({"event" : "status" , "data" : "开始生成小说..." }, ensure_ascii = False ) + "\n "
304- print ("[STREAM] 开始第三阶段:小说生成" )
305- try :
306- with open (writenovel_prompt_path , "r" , encoding = "utf-8" ) as file :
307- template = file .read ()
308- final_input = template .format (user_question = user_question , base_info = base_info , web_info = web_info )
309-
310- # 使用流式调用模型生成小说
311- print ("[STREAM] 调用模型生成小说(流式)..." )
312-
313- # 直接使用 model_writer 的流式接口,而不是通过 agent
314- async for chunk in model_writer .astream ([HumanMessage (content = final_input )]):
315- # LangChain 的流式输出返回 AIMessageChunk 对象
316- if isinstance (chunk , AIMessageChunk ):
317- if chunk .content :
318- # 流式输出每个 token chunk
319- yield json .dumps ({"event" : "token" , "data" : chunk .content }, ensure_ascii = False ) + "\n "
320- elif hasattr (chunk , 'content' ) and chunk .content :
321- # 兼容其他格式
322- yield json .dumps ({"event" : "token" , "data" : chunk .content }, ensure_ascii = False ) + "\n "
323- elif isinstance (chunk , dict ) and 'content' in chunk :
324- yield json .dumps ({"event" : "token" , "data" : chunk ['content' ]}, ensure_ascii = False ) + "\n "
325-
326- print ("[STREAM] 模型流式调用完成" )
327-
328- # 完成信号
329- print ("[STREAM] 发送完成信号" )
330- yield json .dumps ({"event" : "done" , "data" : "" }, ensure_ascii = False ) + "\n "
331-
332- except asyncio .CancelledError :
333- print ("[STREAM] 小说生成阶段被用户取消" )
334- yield json .dumps ({"event" : "error" , "data" : "生成已被取消" }, ensure_ascii = False ) + "\n "
335- return
336- except Exception as e :
337- error_msg = f"Error in post_writer: { type (e ).__name__ } : { str (e )} "
338- print (f"[STREAM ERROR] { error_msg } " )
339- traceback_str = traceback .format_exc ()
340- print (f"[STREAM ERROR] Full Traceback:\n { traceback_str } " )
341- yield json .dumps ({"event" : "error" , "data" : error_msg }, ensure_ascii = False ) + "\n "
342-
343- except asyncio .CancelledError :
344- # 专门处理用户取消操作
345- print ("[STREAM] 流式生成被用户取消" )
346- yield json .dumps ({"event" : "cancelled" , "data" : "生成已被用户取消" }, ensure_ascii = False ) + "\n "
347- return
348- except ExceptionGroup as eg :
349- # 处理 ExceptionGroup
350- error_msg = f"ExceptionGroup: { len (eg .exceptions )} sub-exceptions"
351- print (f"[STREAM ERROR] { error_msg } " )
352- for i , exc in enumerate (eg .exceptions ):
353- print (f"[STREAM ERROR] Sub-exception { i } : { type (exc ).__name__ } : { str (exc )} " )
354- traceback_str = traceback .format_exc ()
355- print (f"[STREAM ERROR] Full Traceback:\n { traceback_str } " )
356- yield json .dumps ({"event" : "error" , "data" : error_msg }, ensure_ascii = False ) + "\n "
357- except BaseException as e :
127+ except Exception as e :
358128 error_msg = f"{ type (e ).__name__ } : { str (e )} "
359129 print (f"[STREAM ERROR] { error_msg } " )
360130 traceback_str = traceback .format_exc ()
@@ -363,4 +133,3 @@ async def event_stream():
363133
364134 # 返回流式响应
365135 return StreamingResponse (event_stream (), media_type = "text/event-stream" )
366-
0 commit comments