Skip to content

Commit a0ffe7e

Browse files
committed
refactor: Extract novel generation logic into a dedicated service
1 parent 0bad965 commit a0ffe7e

3 files changed

Lines changed: 201 additions & 245 deletions

File tree

readme.md

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -191,7 +191,8 @@ Custom Search JSON API 每天免费提供 100 次搜索查询。额外请求的
191191
| |
192192
| |-server/ # 服务端模块
193193
| | |-__init__.py # 初始化文件
194-
| | |-novel_generator.py # 小说生成逻辑实现
194+
| | |-novel_generator.py # 小说生成FASTAPI接口
195+
| | |-novel_service.py # 小说生成逻辑实现
195196
| | |-rag_query.py # RAG查询逻辑实现
196197
| | |-umamusume_create_novel.py # 服务端主程序入口
197198
| |

src/umamusume_novel/server/novel_generator.py

Lines changed: 13 additions & 244 deletions
Original file line numberDiff line numberDiff line change
@@ -40,63 +40,23 @@
4040
import argparse
4141
import asyncio
4242
import 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+
4944
from fastapi import FastAPI, HTTPException, Request
5045
from fastapi.responses import StreamingResponse
5146
from pydantic import BaseModel
52-
from langchain_openai import ChatOpenAI
5347
from 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
6449
from ..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

9658
app = FastAPI(title="LangChain-Umamusume-Server")
9759

98-
99-
10060
# 前端报CORS时
10161
app.add_middleware(
10262
CORSMiddleware,
@@ -106,40 +66,6 @@
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
# 定义请求模型
14470
class 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

Comments
 (0)