mirror of
https://github.com/jihe520/MathModelAgent.git
synced 2026-10-02 02:44:56 +08:00
modify
modify
This commit is contained in:
@@ -24,10 +24,13 @@
|
||||
## ✨ 功能特性
|
||||
|
||||
- 🔍 自动分析问题,数学建模,编写代码,纠正错误,撰写论文
|
||||
- 💻 本地代码解释器
|
||||
- 💻 Code Interperter
|
||||
- loacl Interperter: 基于 jupyter , 代码保存为 notebook 方便再编辑
|
||||
- 云端 code interperter: [E2B](https://e2b.dev/) 和 [daytona](https://app.daytona.io/)
|
||||
- 📝 生成一份编排好格式的论文
|
||||
- 🤝 muti-agents: ~~建模手~~,代码手(反思模块,本地代码解释器),论文手
|
||||
- 🤝 muti-agents: ~~建模手~~,代码手,论文手
|
||||
- 🔄 muti-llms: 每个agent设置不同的模型
|
||||
- 支持所有模型: [litellm](https://docs.litellm.ai/docs/providers)
|
||||
- 💰 成本低 agentless(单次任务成本约 1 rmb)
|
||||
|
||||
## 🚀 后期计划
|
||||
@@ -55,6 +58,10 @@
|
||||
> 项目处于实验探索迭代demo阶段,有许多需要改进优化改进地方,我(项目作者)很忙,有时间会优化更新
|
||||
> 欢迎贡献
|
||||
|
||||
|
||||
案例参考 ./demo 文件夹下
|
||||
如果你有好的案例可以提交 PR 在该目录下
|
||||
|
||||
## 📖 使用教程
|
||||
|
||||
> 确保电脑中安装好 Python, Nodejs, **Redis** 环境
|
||||
@@ -66,16 +73,9 @@
|
||||
1. 配置模型
|
||||
|
||||
复制`/backend/.env.dev.example`到`/backend/.env.dev`(删除`.example` 后缀)
|
||||
填写配置模型和 APIKEY
|
||||
**配置环境变量**
|
||||
推荐模型能力较强的、参数量大的模型。
|
||||
|
||||
```bash
|
||||
# support all model, check out https://docs.litellm.ai/docs/
|
||||
API_KEY=
|
||||
# gpt-4.1,deepseek/deepseek-chat,gemini/gemini-2.5-flash-preview-04-17
|
||||
MODEL=
|
||||
# 确保安装 Redis
|
||||
```
|
||||
|
||||
复制`/fronted/.env.example`到`/fronted/.env`(删除`.example` 后缀)
|
||||
|
||||
@@ -135,6 +135,7 @@ clone 项目后,下载 **Todo Tree** 插件,可以查看代码中所有具
|
||||
## 📄 版权License
|
||||
|
||||
个人免费使用,请勿商业用途,商业用途联系我(作者)
|
||||
禁止闭源分发
|
||||
|
||||
## 🙏 Reference
|
||||
|
||||
@@ -147,9 +148,17 @@ Thanks to the following projects:
|
||||
|
||||
## 其他
|
||||
|
||||
### Sponsor
|
||||
|
||||
<div align="center">
|
||||
<img src="./docs/sponser.png" alt="Buy Me a Coffee" width="280"/>
|
||||
</div>
|
||||
|
||||
感谢赞助
|
||||
[danmo-tyc](https://github.com/danmo-tyc)
|
||||
|
||||
### GROUP
|
||||
|
||||
有问题可以进群问
|
||||
[QQ 群:699970403](http://qm.qq.com/cgi-bin/qm/qr?_wv=1027&k=rFKquDTSxKcWpEhRgpJD-dPhTtqLwJ9r&authKey=xYKvCFG5My4uYZTbIIoV5MIPQedW7hYzf0%2Fbs4EUZ100UegQWcQ8xEEgTczHsyU6&noverify=0&group_code=699970403)
|
||||
|
||||
|
||||
@@ -3,7 +3,7 @@ ENV=dev
|
||||
API_KEY=
|
||||
# gpt-4.1,deepseek/deepseek-chat
|
||||
MODEL=
|
||||
# BASE_URL= 不需要填
|
||||
# BASE_URL= 默认不需要填
|
||||
|
||||
# 模型最大问答次数
|
||||
MAX_CHAT_TURNS=60
|
||||
|
||||
@@ -25,7 +25,6 @@ firstPage = """
|
||||
{问题}
|
||||
{模型的建立与求解}
|
||||
"""
|
||||
|
||||
RepeatQues = """
|
||||
# 一、问题重述
|
||||
## 1.1 问题背景
|
||||
@@ -97,12 +96,10 @@ symbol = """
|
||||
|
||||
{模型的建立与求解}
|
||||
"""
|
||||
|
||||
eda = """
|
||||
## 4.2 描述性统计
|
||||
大约200字
|
||||
"""
|
||||
|
||||
ques1 = """模板和要求如下
|
||||
# 五、模型的建立与求解
|
||||
## 5.1 问题一模型的建立与求解
|
||||
@@ -122,7 +119,6 @@ ques2 = """参考模板
|
||||
模型的求解过程
|
||||
大约600字
|
||||
"""
|
||||
|
||||
ques3 = """参考模板
|
||||
## 5.3 问题三模型的建立与求解
|
||||
### 5.3.1 模型的建立
|
||||
@@ -132,7 +128,6 @@ ques3 = """参考模板
|
||||
模型的求解过程
|
||||
大约600字
|
||||
"""
|
||||
|
||||
ques4 = """参考模板
|
||||
## 5.4 问题四模型的建立与求解
|
||||
### 5.4.1 模型的建立
|
||||
@@ -142,7 +137,6 @@ ques4 = """参考模板
|
||||
模型的求解过程
|
||||
大约600字
|
||||
"""
|
||||
|
||||
ques5 = """参考模板
|
||||
## 5.5 问题五模型的建立与求解
|
||||
### 5.5.1 模型的建立
|
||||
@@ -152,7 +146,6 @@ ques5 = """参考模板
|
||||
模型的求解过程
|
||||
大约600字
|
||||
"""
|
||||
|
||||
ques6 = """参考模板
|
||||
## 5.6 问题六模型的建立与求解
|
||||
### 5.6.1 模型的建立
|
||||
@@ -162,12 +155,10 @@ ques6 = """参考模板
|
||||
模型的求解过程
|
||||
大约600字
|
||||
"""
|
||||
|
||||
sensitivity_analysis = """参考模板
|
||||
# 六、模型的分析与检验
|
||||
## 6.1 灵敏度分析
|
||||
"""
|
||||
|
||||
judge = """参考模板和要求
|
||||
# 七、模型的评价、改进与推广
|
||||
## 7.1 模型的优点
|
||||
@@ -175,15 +166,4 @@ judge = """参考模板和要求
|
||||
## 7.3 模型的改进与推广
|
||||
优点数量要多于缺点,缺点大约2/3个
|
||||
大约200字
|
||||
"""
|
||||
|
||||
reference = """
|
||||
# 参考文献
|
||||
[1] 作者. (年份). 题目. 期刊, 卷(期), 页码.
|
||||
[2] 作者. (年份). 题目. 期刊, 卷(期), 页码.
|
||||
[3] 作者. (年份). 题目. 期刊, 卷(期), 页码.
|
||||
|
||||
例子
|
||||
[1] 刘培杰,李军英.体验式营销模式对农业旅游经济发展的影响[J].山西农经,2024,(16):67-69.DOI:10.16675/j.cnki.cn14-1065/f.2024.16.019.
|
||||
[2] ..
|
||||
"""
|
||||
@@ -17,9 +17,27 @@ def parse_cors(value: str) -> list[str]:
|
||||
|
||||
class Settings(BaseSettings):
|
||||
ENV: str
|
||||
API_KEY: str
|
||||
MODEL: str
|
||||
BASE_URL: Optional[str] = None
|
||||
|
||||
COORDINATOR_API_KEY: str
|
||||
COORDINATOR_MODEL: str
|
||||
COORDINATOR_BASE_URL: Optional[str] = None
|
||||
|
||||
MODELER_API_KEY: str
|
||||
MODELER_MODEL: str
|
||||
MODELER_BASE_URL: Optional[str] = None
|
||||
|
||||
CODER_API_KEY: str
|
||||
CODER_MODEL: str
|
||||
CODER_BASE_URL: Optional[str] = None
|
||||
|
||||
WRITER_API_KEY: str
|
||||
WRITER_MODEL: str
|
||||
WRITER_BASE_URL: Optional[str] = None
|
||||
|
||||
DEFAULT_API_KEY: str
|
||||
DEFAULT_MODEL: str
|
||||
DEFAULT_BASE_URL: Optional[str] = None
|
||||
|
||||
MAX_CHAT_TURNS: int
|
||||
MAX_RETRIES: int
|
||||
E2B_API_KEY: Optional[str] = None
|
||||
|
||||
@@ -0,0 +1,11 @@
|
||||
from .coder_agent import CoderAgent
|
||||
from .writer_agent import WriterAgent
|
||||
from .coordinator_agent import CoordinatorAgent
|
||||
from .modeler_agent import ModelerAgent
|
||||
|
||||
__all__ = [
|
||||
"CoderAgent",
|
||||
"WriterAgent",
|
||||
"CoordinatorAgent",
|
||||
"ModelerAgent",
|
||||
]
|
||||
@@ -0,0 +1,63 @@
|
||||
from app.core.llm.llm import LLM
|
||||
from app.utils.log_util import logger
|
||||
|
||||
|
||||
class Agent:
|
||||
def __init__(
|
||||
self,
|
||||
task_id: str,
|
||||
model: LLM,
|
||||
max_chat_turns: int = 30, # 单个agent最大对话轮次
|
||||
max_memory: int = 25, # 最大记忆轮次
|
||||
) -> None:
|
||||
self.task_id = task_id
|
||||
self.model = model
|
||||
self.chat_history: list[dict] = [] # 存储对话历史
|
||||
self.max_chat_turns = max_chat_turns # 最大对话轮次
|
||||
self.current_chat_turns = 0 # 当前对话轮次计数器
|
||||
self.max_memory = max_memory # 最大记忆轮次
|
||||
|
||||
async def run(self, prompt: str, system_prompt: str, sub_title: str) -> str:
|
||||
"""
|
||||
执行agent的对话并返回结果和总结
|
||||
|
||||
Args:
|
||||
prompt: 输入的提示
|
||||
|
||||
Returns:
|
||||
str: 模型的响应
|
||||
"""
|
||||
try:
|
||||
logger.info(f"{self.__class__.__name__}:开始:执行对话")
|
||||
self.current_chat_turns = 0 # 重置对话轮次计数器
|
||||
|
||||
# 更新对话历史
|
||||
self.append_chat_history({"role": "system", "content": system_prompt})
|
||||
self.append_chat_history({"role": "user", "content": prompt})
|
||||
|
||||
# 获取历史消息用于本次对话
|
||||
response = await self.model.chat(
|
||||
history=self.chat_history,
|
||||
agent_name=self.__class__.__name__,
|
||||
sub_title=sub_title,
|
||||
)
|
||||
response_content = response.choices[0].message.content
|
||||
self.chat_history.append({"role": "assistant", "content": response_content})
|
||||
logger.info(f"{self.__class__.__name__}:完成:执行对话")
|
||||
return response_content
|
||||
except Exception as e:
|
||||
error_msg = f"执行过程中遇到错误: {str(e)}"
|
||||
logger.error(f"Agent执行失败: {str(e)}")
|
||||
return error_msg
|
||||
|
||||
def append_chat_history(self, msg: dict) -> None:
|
||||
self.clear_memory()
|
||||
self.chat_history.append(msg)
|
||||
|
||||
def clear_memory(self):
|
||||
if len(self.chat_history) <= self.max_memory:
|
||||
return
|
||||
logger.info(f"{self.__class__.__name__}:清除记忆")
|
||||
|
||||
# 使用切片保留第一条和最后两条消息
|
||||
self.chat_history = self.chat_history[:2] + self.chat_history[-5:]
|
||||
@@ -1,100 +1,16 @@
|
||||
import json
|
||||
from app.core.llm import LLM
|
||||
from app.core.prompts import (
|
||||
get_completion_check_prompt,
|
||||
get_reflection_prompt,
|
||||
get_writer_prompt,
|
||||
CODER_PROMPT,
|
||||
MODELER_PROMPT,
|
||||
)
|
||||
from app.core.functions import coder_tools, writer_tools
|
||||
from app.models.model import CoderToWriter
|
||||
from app.models.user_output import UserOutput
|
||||
from app.utils.enums import CompTemplate, FormatOutPut
|
||||
from app.utils.log_util import logger
|
||||
from app.core.agents.agent import Agent
|
||||
from app.config.setting import settings
|
||||
from app.utils.common_utils import get_current_files
|
||||
from app.utils.log_util import logger
|
||||
from app.utils.redis_manager import redis_manager
|
||||
from app.schemas.response import SystemMessage
|
||||
from app.tools.base_interpreter import BaseCodeInterpreter
|
||||
from app.tools.openalex_scholar import OpenAlexScholar
|
||||
from icecream import ic
|
||||
|
||||
|
||||
class Agent:
|
||||
def __init__(
|
||||
self,
|
||||
task_id: str,
|
||||
model: LLM,
|
||||
max_chat_turns: int = 30, # 单个agent最大对话轮次
|
||||
user_output: UserOutput = None,
|
||||
max_memory: int = 25, # 最大记忆轮次
|
||||
) -> None:
|
||||
self.task_id = task_id
|
||||
self.model = model
|
||||
self.chat_history: list[dict] = [] # 存储对话历史
|
||||
self.max_chat_turns = max_chat_turns # 最大对话轮次
|
||||
self.current_chat_turns = 0 # 当前对话轮次计数器
|
||||
self.user_output = user_output
|
||||
self.max_memory = max_memory # 最大记忆轮次
|
||||
|
||||
async def run(self, prompt: str, system_prompt: str, sub_title: str) -> str:
|
||||
"""
|
||||
执行agent的对话并返回结果和总结
|
||||
|
||||
Args:
|
||||
prompt: 输入的提示
|
||||
|
||||
Returns:
|
||||
str: 模型的响应
|
||||
"""
|
||||
try:
|
||||
logger.info(f"{self.__class__.__name__}:开始:执行对话")
|
||||
self.current_chat_turns = 0 # 重置对话轮次计数器
|
||||
|
||||
# 更新对话历史
|
||||
self.append_chat_history({"role": "system", "content": system_prompt})
|
||||
self.append_chat_history({"role": "user", "content": prompt})
|
||||
|
||||
# 获取历史消息用于本次对话
|
||||
response = await self.model.chat(
|
||||
history=self.chat_history,
|
||||
agent_name=self.__class__.__name__,
|
||||
sub_title=sub_title,
|
||||
)
|
||||
response_content = response.choices[0].message.content
|
||||
self.chat_history.append({"role": "assistant", "content": response_content})
|
||||
logger.info(f"{self.__class__.__name__}:完成:执行对话")
|
||||
return response_content
|
||||
except Exception as e:
|
||||
error_msg = f"执行过程中遇到错误: {str(e)}"
|
||||
logger.error(f"Agent执行失败: {str(e)}")
|
||||
return error_msg
|
||||
|
||||
def append_chat_history(self, msg: dict) -> None:
|
||||
self.clear_memory()
|
||||
self.chat_history.append(msg)
|
||||
# self.user_output.data_recorder.append_chat_history(
|
||||
# msg, agent_name=self.__class__.__name__
|
||||
# )
|
||||
|
||||
def clear_memory(self):
|
||||
if len(self.chat_history) <= self.max_memory:
|
||||
return
|
||||
logger.info(f"{self.__class__.__name__}:清除记忆")
|
||||
|
||||
# 使用切片保留第一条和最后两条消息
|
||||
self.chat_history = self.chat_history[:2] + self.chat_history[-5:]
|
||||
|
||||
|
||||
class ModelerAgent(Agent): # 继承自Agent类
|
||||
def __init__(
|
||||
self,
|
||||
model: LLM,
|
||||
max_chat_turns: int = 30, # 添加最大对话轮次限制
|
||||
) -> None:
|
||||
super().__init__(model, max_chat_turns)
|
||||
self.system_prompt = MODELER_PROMPT
|
||||
from app.core.llm.llm import LLM
|
||||
from app.models.model import CoderToWriter
|
||||
from app.core.prompts import CODER_PROMPT
|
||||
from app.utils.common_utils import get_current_files
|
||||
import json
|
||||
from app.core.prompts import get_reflection_prompt, get_completion_check_prompt
|
||||
from app.core.functions import coder_tools
|
||||
|
||||
|
||||
# 代码强
|
||||
@@ -290,7 +206,11 @@ class CoderAgent(Agent): # 同样继承自Agent类
|
||||
):
|
||||
logger.info("没有调用工具,代表任务已完成")
|
||||
task_completed = True
|
||||
return completion_response.choices[0].message.content
|
||||
return CoderToWriter(
|
||||
coder_response=completion_response.choices[
|
||||
0
|
||||
].message.content
|
||||
)
|
||||
else:
|
||||
logger.info("没有工具,代表任务完成")
|
||||
|
||||
@@ -304,152 +224,4 @@ class CoderAgent(Agent): # 同样继承自Agent类
|
||||
|
||||
logger.info(f"{self.__class__.__name__}:完成:执行子任务: {subtask_title}")
|
||||
|
||||
return response.choices[0].message.content
|
||||
|
||||
|
||||
# 长文本
|
||||
# TODO: 并行 parallel
|
||||
# TODO: 获取当前文件下的文件
|
||||
# TODO: 引用cites tool
|
||||
class WriterAgent(Agent): # 同样继承自Agent类
|
||||
def __init__(
|
||||
self,
|
||||
task_id: str,
|
||||
model: LLM,
|
||||
max_chat_turns: int = 10, # 添加最大对话轮次限制
|
||||
comp_template: CompTemplate = CompTemplate,
|
||||
format_output: FormatOutPut = FormatOutPut.Markdown,
|
||||
user_output: UserOutput = None,
|
||||
scholar: OpenAlexScholar = None,
|
||||
) -> None:
|
||||
super().__init__(task_id, model, max_chat_turns, user_output)
|
||||
self.format_out_put = format_output
|
||||
self.comp_template = comp_template
|
||||
self.scholar = scholar
|
||||
self.system_prompt = get_writer_prompt(format_output)
|
||||
self.available_images: list[str] = []
|
||||
|
||||
async def run(
|
||||
self,
|
||||
prompt: str,
|
||||
available_images: list[str] = None,
|
||||
sub_title: str = None,
|
||||
) -> str:
|
||||
"""
|
||||
执行写作任务
|
||||
Args:
|
||||
prompt: 写作提示
|
||||
available_images: 可用的图片相对路径列表(如 20250420-173744-9f87792c/编号_分布.png)
|
||||
sub_title: 子任务标题
|
||||
"""
|
||||
logger.info(f"subtitle是:{sub_title}")
|
||||
|
||||
if available_images:
|
||||
self.available_images = available_images
|
||||
# 拼接成完整URL
|
||||
image_list = ",".join(available_images)
|
||||
image_prompt = f"\n可用的图片链接列表:\n{image_list}\n请在写作时适当引用这些图片链接。"
|
||||
ic(image_prompt)
|
||||
prompt = prompt + image_prompt
|
||||
|
||||
logger.info(f"{self.__class__.__name__}:开始:执行对话")
|
||||
self.current_chat_turns += 1 # 重置对话轮次计数器
|
||||
|
||||
# 更新对话历史
|
||||
self.append_chat_history({"role": "system", "content": self.system_prompt})
|
||||
self.append_chat_history({"role": "user", "content": prompt})
|
||||
|
||||
# 获取历史消息用于本次对话
|
||||
response = await self.model.chat(
|
||||
history=self.chat_history,
|
||||
tools=writer_tools,
|
||||
tool_choice="auto",
|
||||
agent_name=self.__class__.__name__,
|
||||
sub_title=sub_title,
|
||||
)
|
||||
|
||||
if (
|
||||
hasattr(response.choices[0].message, "tool_calls")
|
||||
and response.choices[0].message.tool_calls
|
||||
):
|
||||
logger.info("检测到工具调用")
|
||||
tool_call = response.choices[0].message.tool_calls[0]
|
||||
tool_id = tool_call.id
|
||||
tool_call.function.name
|
||||
if tool_call.function.name == "search_papers":
|
||||
logger.info("调用工具: search_papers")
|
||||
await redis_manager.publish_message(
|
||||
self.task_id,
|
||||
SystemMessage(content=f"写作手调用{tool_call.function.name}工具"),
|
||||
)
|
||||
|
||||
query = json.loads(tool_call.function.arguments)["query"]
|
||||
|
||||
full_content = response.choices[0].message.content
|
||||
# 更新对话历史 - 添加助手的响应
|
||||
self.append_chat_history(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": full_content,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": tool_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search_papers",
|
||||
"arguments": json.dumps({"query": query}),
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
papers = self.scholar.search_papers(query)
|
||||
except Exception as e:
|
||||
logger.error(f"搜索文献失败: {str(e)}")
|
||||
return f"搜索文献失败: {str(e)}"
|
||||
# TODO: pass to frontend
|
||||
self.scholar.print_papers(papers)
|
||||
self.append_chat_history(
|
||||
{
|
||||
"role": "tool",
|
||||
"content": papers,
|
||||
"tool_call_id": tool_id,
|
||||
"name": "search_papers",
|
||||
}
|
||||
)
|
||||
next_response = await self.model.chat(
|
||||
history=self.chat_history,
|
||||
tools=writer_tools,
|
||||
tool_choice="auto",
|
||||
agent_name=self.__class__.__name__,
|
||||
sub_title=sub_title,
|
||||
)
|
||||
response_content = next_response.choices[0].message.content
|
||||
else:
|
||||
response_content = response.choices[0].message.content
|
||||
self.chat_history.append({"role": "assistant", "content": response_content})
|
||||
logger.info(f"{self.__class__.__name__}:完成:执行对话")
|
||||
return response_content
|
||||
|
||||
async def summarize(self) -> str:
|
||||
"""
|
||||
总结对话内容
|
||||
"""
|
||||
try:
|
||||
self.append_chat_history(
|
||||
{"role": "user", "content": "请简单总结以上完成什么任务取得什么结果:"}
|
||||
)
|
||||
# 获取历史消息用于本次对话
|
||||
response = await self.model.chat(
|
||||
history=self.chat_history, agent_name=self.__class__.__name__
|
||||
)
|
||||
self.append_chat_history(
|
||||
{"role": "assistant", "content": response.choices[0].message.content}
|
||||
)
|
||||
return response.choices[0].message.content
|
||||
except Exception as e:
|
||||
logger.error(f"总结生成失败: {str(e)}")
|
||||
# 返回一个基础总结,避免完全失败
|
||||
return "由于网络原因无法生成详细总结,但已完成主要任务处理。"
|
||||
return CoderToWriter(coder_response=response.choices[0].message.content)
|
||||
@@ -0,0 +1,46 @@
|
||||
from app.core.agents.agent import Agent
|
||||
from app.core.llm.llm import LLM
|
||||
from app.core.prompts import COORDINATOR_PROMPT
|
||||
import json
|
||||
from app.utils.log_util import logger
|
||||
from app.models.model import CoordinatorToModeler
|
||||
|
||||
|
||||
class CoordinatorAgent(Agent):
|
||||
def __init__(
|
||||
self,
|
||||
task_id: str,
|
||||
model: LLM,
|
||||
max_chat_turns: int = 30,
|
||||
) -> None:
|
||||
super().__init__(task_id, model, max_chat_turns)
|
||||
self.system_prompt = COORDINATOR_PROMPT
|
||||
|
||||
async def run(self, ques_all: str) -> CoordinatorToModeler:
|
||||
"""用户输入问题 使用LLM 格式化 questions"""
|
||||
# TODO: "note": <补充说明,如果没有补充说明,请填 null>,
|
||||
self.append_chat_history({"role": "system", "content": self.system_prompt})
|
||||
self.append_chat_history({"role": "user", "content": ques_all})
|
||||
|
||||
response = await self.model.chat(
|
||||
history=self.chat_history,
|
||||
agent_name=self.__class__.__name__,
|
||||
)
|
||||
json_str = response.choices[0].message.content
|
||||
|
||||
if not json_str.startswith("```json"):
|
||||
logger.info(f"拒绝回答用户非数学建模请求:{json_str}")
|
||||
raise ValueError(f"拒绝回答用户非数学建模请求:{json_str}")
|
||||
|
||||
json_str = json_str.replace("```json", "").replace("```", "").strip()
|
||||
|
||||
if not json_str:
|
||||
raise ValueError("返回的 JSON 字符串为空,请检查输入内容。")
|
||||
|
||||
try:
|
||||
questions = json.loads(json_str)
|
||||
ques_count = questions["ques_count"]
|
||||
logger.info(f"questions:{questions}")
|
||||
return CoordinatorToModeler(questions=questions, ques_count=ques_count)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"JSON 解析错误: {e}")
|
||||
@@ -0,0 +1,44 @@
|
||||
from app.core.agents.agent import Agent
|
||||
from app.core.llm.llm import LLM
|
||||
from app.core.prompts import MODELER_PROMPT
|
||||
from app.models.model import CoordinatorToModeler, ModelerToCoder
|
||||
from app.utils.log_util import logger
|
||||
import json
|
||||
|
||||
|
||||
class ModelerAgent(Agent): # 继承自Agent类
|
||||
def __init__(
|
||||
self,
|
||||
task_id: str,
|
||||
model: LLM,
|
||||
max_chat_turns: int = 30, # 添加最大对话轮次限制
|
||||
) -> None:
|
||||
super().__init__(task_id, model, max_chat_turns)
|
||||
self.system_prompt = MODELER_PROMPT
|
||||
|
||||
async def run(self, coordinator_to_modeler: CoordinatorToModeler) -> ModelerToCoder:
|
||||
self.append_chat_history({"role": "system", "content": self.system_prompt})
|
||||
self.append_chat_history(
|
||||
{
|
||||
"role": "user",
|
||||
"content": coordinator_to_modeler.questions.model_dump_json(),
|
||||
}
|
||||
)
|
||||
|
||||
response = await self.model.chat(
|
||||
history=self.chat_history,
|
||||
agent_name=self.__class__.__name__,
|
||||
)
|
||||
|
||||
json_str = response.choices[0].message.content
|
||||
|
||||
json_str = json_str.replace("```json", "").replace("```", "").strip()
|
||||
|
||||
if not json_str:
|
||||
raise ValueError("返回的 JSON 字符串为空,请检查输入内容。")
|
||||
|
||||
try:
|
||||
questions_solution = json.loads(json_str)
|
||||
return ModelerToCoder(questions_solution=questions_solution)
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"JSON 解析错误: {e}")
|
||||
@@ -0,0 +1,164 @@
|
||||
from app.core.agents.agent import Agent
|
||||
from app.core.llm.llm import LLM
|
||||
from app.core.prompts import get_writer_prompt
|
||||
from app.utils.enums import CompTemplate, FormatOutPut
|
||||
from app.models.user_output import UserOutput
|
||||
from app.tools.openalex_scholar import OpenAlexScholar
|
||||
from app.utils.log_util import logger
|
||||
from app.utils.redis_manager import redis_manager
|
||||
from app.schemas.response import SystemMessage
|
||||
import json
|
||||
from app.core.functions import writer_tools
|
||||
from app.utils.common_utils import get_footnotes
|
||||
from icecream import ic
|
||||
from app.models.model import WriterResponse
|
||||
|
||||
|
||||
# 长文本
|
||||
# 长文本
|
||||
# TODO: 并行 parallel
|
||||
# TODO: 获取当前文件下的文件
|
||||
# TODO: 引用cites tool
|
||||
class WriterAgent(Agent): # 同样继承自Agent类
|
||||
def __init__(
|
||||
self,
|
||||
task_id: str,
|
||||
model: LLM,
|
||||
max_chat_turns: int = 10, # 添加最大对话轮次限制
|
||||
comp_template: CompTemplate = CompTemplate,
|
||||
format_output: FormatOutPut = FormatOutPut.Markdown,
|
||||
user_output: UserOutput = None,
|
||||
scholar: OpenAlexScholar = None,
|
||||
) -> None:
|
||||
super().__init__(task_id, model, max_chat_turns, user_output)
|
||||
self.format_out_put = format_output
|
||||
self.comp_template = comp_template
|
||||
self.scholar = scholar
|
||||
self.system_prompt = get_writer_prompt(format_output)
|
||||
self.available_images: list[str] = []
|
||||
|
||||
async def run(
|
||||
self,
|
||||
prompt: str,
|
||||
available_images: list[str] = None,
|
||||
sub_title: str = None,
|
||||
) -> WriterResponse:
|
||||
"""
|
||||
执行写作任务
|
||||
Args:
|
||||
prompt: 写作提示
|
||||
available_images: 可用的图片相对路径列表(如 20250420-173744-9f87792c/编号_分布.png)
|
||||
sub_title: 子任务标题
|
||||
"""
|
||||
logger.info(f"subtitle是:{sub_title}")
|
||||
|
||||
if available_images:
|
||||
self.available_images = available_images
|
||||
# 拼接成完整URL
|
||||
image_list = ",".join(available_images)
|
||||
image_prompt = f"\n可用的图片链接列表:\n{image_list}\n请在写作时适当引用这些图片链接。"
|
||||
ic(image_prompt)
|
||||
prompt = prompt + image_prompt
|
||||
|
||||
logger.info(f"{self.__class__.__name__}:开始:执行对话")
|
||||
self.current_chat_turns += 1 # 重置对话轮次计数器
|
||||
|
||||
# 更新对话历史
|
||||
self.append_chat_history({"role": "system", "content": self.system_prompt})
|
||||
self.append_chat_history({"role": "user", "content": prompt})
|
||||
|
||||
# 获取历史消息用于本次对话
|
||||
response = await self.model.chat(
|
||||
history=self.chat_history,
|
||||
tools=writer_tools,
|
||||
tool_choice="auto",
|
||||
agent_name=self.__class__.__name__,
|
||||
sub_title=sub_title,
|
||||
)
|
||||
|
||||
if (
|
||||
hasattr(response.choices[0].message, "tool_calls")
|
||||
and response.choices[0].message.tool_calls
|
||||
):
|
||||
logger.info("检测到工具调用")
|
||||
tool_call = response.choices[0].message.tool_calls[0]
|
||||
tool_id = tool_call.id
|
||||
tool_call.function.name
|
||||
if tool_call.function.name == "search_papers":
|
||||
logger.info("调用工具: search_papers")
|
||||
await redis_manager.publish_message(
|
||||
self.task_id,
|
||||
SystemMessage(content=f"写作手调用{tool_call.function.name}工具"),
|
||||
)
|
||||
|
||||
query = json.loads(tool_call.function.arguments)["query"]
|
||||
footnotes = get_footnotes(query)
|
||||
full_content = response.choices[0].message.content
|
||||
# 更新对话历史 - 添加助手的响应
|
||||
self.append_chat_history(
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": full_content,
|
||||
"tool_calls": [
|
||||
{
|
||||
"id": tool_id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "search_papers",
|
||||
"arguments": json.dumps({"query": query}),
|
||||
},
|
||||
}
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
try:
|
||||
papers = self.scholar.search_papers(query)
|
||||
except Exception as e:
|
||||
logger.error(f"搜索文献失败: {str(e)}")
|
||||
return f"搜索文献失败: {str(e)}"
|
||||
# TODO: pass to frontend
|
||||
papers_str = self.scholar.papers_to_str(papers)
|
||||
logger.info(f"搜索文献结果\n{papers_str}")
|
||||
self.append_chat_history(
|
||||
{
|
||||
"role": "tool",
|
||||
"content": papers_str,
|
||||
"tool_call_id": tool_id,
|
||||
"name": "search_papers",
|
||||
}
|
||||
)
|
||||
next_response = await self.model.chat(
|
||||
history=self.chat_history,
|
||||
tools=writer_tools,
|
||||
tool_choice="auto",
|
||||
agent_name=self.__class__.__name__,
|
||||
sub_title=sub_title,
|
||||
)
|
||||
response_content = next_response.choices[0].message.content
|
||||
else:
|
||||
response_content = response.choices[0].message.content
|
||||
self.chat_history.append({"role": "assistant", "content": response_content})
|
||||
logger.info(f"{self.__class__.__name__}:完成:执行对话")
|
||||
return WriterResponse(response_content=response_content, footnotes=footnotes)
|
||||
|
||||
async def summarize(self) -> str:
|
||||
"""
|
||||
总结对话内容
|
||||
"""
|
||||
try:
|
||||
self.append_chat_history(
|
||||
{"role": "user", "content": "请简单总结以上完成什么任务取得什么结果:"}
|
||||
)
|
||||
# 获取历史消息用于本次对话
|
||||
response = await self.model.chat(
|
||||
history=self.chat_history, agent_name=self.__class__.__name__
|
||||
)
|
||||
self.append_chat_history(
|
||||
{"role": "assistant", "content": response.choices[0].message.content}
|
||||
)
|
||||
return response.choices[0].message.content
|
||||
except Exception as e:
|
||||
logger.error(f"总结生成失败: {str(e)}")
|
||||
# 返回一个基础总结,避免完全失败
|
||||
return "由于网络原因无法生成详细总结,但已完成主要任务处理。"
|
||||
@@ -0,0 +1,144 @@
|
||||
from app.models.user_output import UserOutput
|
||||
from app.tools.base_interpreter import BaseCodeInterpreter
|
||||
from app.core.agents.modeler_agent import ModelerToCoder
|
||||
|
||||
|
||||
class Flows:
|
||||
def __init__(self):
|
||||
self.flows: dict[str, dict] = {}
|
||||
|
||||
def set_flows(self, ques_count: int):
|
||||
ques_str = [f"ques{i}" for i in range(1, ques_count + 1)]
|
||||
seq = [
|
||||
"firstPage",
|
||||
"RepeatQues",
|
||||
"analysisQues",
|
||||
"modelAssumption",
|
||||
"symbol",
|
||||
"eda",
|
||||
*ques_str,
|
||||
"sensitivity_analysis",
|
||||
"judge",
|
||||
]
|
||||
self.flows = {key: {} for key in seq}
|
||||
|
||||
def get_solution_flows(
|
||||
self, questions: dict[str, str | int], modeler_response: ModelerToCoder
|
||||
):
|
||||
questions_quesx = {
|
||||
key: value
|
||||
for key, value in questions.items()
|
||||
if key.startswith("ques") and key != "ques_count"
|
||||
}
|
||||
ques_flow = {
|
||||
key: {
|
||||
"coder_prompt": f"""
|
||||
参考建模手给出的解决方案{modeler_response.questions_solution[key]}
|
||||
完成如下问题{value}
|
||||
""",
|
||||
}
|
||||
for key, value in questions_quesx.items()
|
||||
}
|
||||
flows = {
|
||||
"eda": {
|
||||
# TODO : 获取当前路径下的所有数据集
|
||||
"coder_prompt": f"""
|
||||
参考建模手给出的解决方案{modeler_response.questions_solution["eda"]}
|
||||
对当前目录下数据进行EDA分析(数据清洗,可视化),清洗后的数据保存当前目录下,**不需要复杂的模型**
|
||||
""",
|
||||
},
|
||||
**ques_flow,
|
||||
"sensitivity_analysis": {
|
||||
"coder_prompt": f"""
|
||||
参考建模手给出的解决方案{modeler_response.questions_solution["sensitivity_analysis"]}
|
||||
完成敏感性分析
|
||||
""",
|
||||
},
|
||||
}
|
||||
return flows
|
||||
|
||||
def get_write_flows(
|
||||
self, user_output: UserOutput, config_template: dict, bg_ques_all: str
|
||||
):
|
||||
model_build_solve = user_output.get_model_build_solve()
|
||||
flows = {
|
||||
"firstPage": f"""问题背景{bg_ques_all},不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["firstPage"]},撰写标题,摘要,关键词""",
|
||||
"RepeatQues": f"""问题背景{bg_ques_all},不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["RepeatQues"]},撰写问题重述""",
|
||||
"analysisQues": f"""问题背景{bg_ques_all},不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["analysisQues"]},撰写问题分析""",
|
||||
"modelAssumption": f"""问题背景{bg_ques_all},不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["modelAssumption"]},撰写模型假设""",
|
||||
"symbol": f"""不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["symbol"]},撰写符号说明部分""",
|
||||
"judge": f"""不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["judge"]},撰写模型的评价部分""",
|
||||
}
|
||||
return flows
|
||||
|
||||
def get_writer_prompt(
|
||||
self,
|
||||
key: str,
|
||||
coder_response: str,
|
||||
code_interpreter: BaseCodeInterpreter,
|
||||
config_template: dict,
|
||||
) -> str:
|
||||
"""根据不同的key生成对应的writer_prompt
|
||||
|
||||
Args:
|
||||
key: 任务类型
|
||||
coder_response: 代码执行结果
|
||||
|
||||
Returns:
|
||||
str: 生成的writer_prompt
|
||||
"""
|
||||
code_output = code_interpreter.get_code_output(key)
|
||||
|
||||
questions_quesx_keys = self.get_questions_quesx_keys()
|
||||
bgc = self.questions["background"]
|
||||
quesx_writer_prompt = {
|
||||
key: f"""
|
||||
问题背景{bgc},不需要编写代码,代码手得到的结果{coder_response},{code_output},按照如下模板撰写:{config_template[key]}
|
||||
"""
|
||||
for key in questions_quesx_keys
|
||||
}
|
||||
|
||||
writer_prompt = {
|
||||
"eda": f"""
|
||||
问题背景{bgc},不需要编写代码,代码手得到的结果{coder_response},{code_output},按照如下模板撰写:{config_template["eda"]}
|
||||
""",
|
||||
**quesx_writer_prompt,
|
||||
"sensitivity_analysis": f"""
|
||||
问题背景{bgc},不需要编写代码,代码手得到的结果{coder_response},{code_output},按照如下模板撰写:{config_template["sensitivity_analysis"]}
|
||||
""",
|
||||
}
|
||||
|
||||
if key in writer_prompt:
|
||||
return writer_prompt[key]
|
||||
else:
|
||||
raise ValueError(f"未知的任务类型: {key}")
|
||||
|
||||
def get_questions_quesx_keys(self) -> list[str]:
|
||||
"""获取问题1,2...的键"""
|
||||
return list(self.get_questions_quesx().keys())
|
||||
|
||||
def get_questions_quesx(self) -> dict[str, str]:
|
||||
"""获取问题1,2,3...的键值对"""
|
||||
# 获取所有以 "ques" 开头的键值对
|
||||
questions_quesx = {
|
||||
key: value
|
||||
for key, value in self.questions.items()
|
||||
if key.startswith("ques") and key != "ques_count"
|
||||
}
|
||||
return questions_quesx
|
||||
|
||||
def get_seq(self, ques_count: int) -> dict[str, str]:
|
||||
ques_str = [f"ques{i}" for i in range(1, ques_count + 1)]
|
||||
seq = [
|
||||
"firstPage",
|
||||
"RepeatQues",
|
||||
"analysisQues",
|
||||
"modelAssumption",
|
||||
"symbol",
|
||||
"eda",
|
||||
*ques_str,
|
||||
"sensitivity_analysis",
|
||||
"judge",
|
||||
"reference",
|
||||
]
|
||||
return {key: "" for key in seq}
|
||||
@@ -0,0 +1,40 @@
|
||||
from app.config.setting import settings
|
||||
from app.core.llm.llm import LLM
|
||||
|
||||
|
||||
class LLMFactory:
|
||||
task_id: str
|
||||
|
||||
def __init__(self, task_id: str) -> None:
|
||||
self.task_id = task_id
|
||||
|
||||
def get_all_llms(self) -> tuple[LLM, LLM, LLM, LLM]:
|
||||
coordinator_llm = LLM(
|
||||
api_key=settings.COORDINATOR_API_KEY,
|
||||
model=settings.COORDINATOR_MODEL,
|
||||
base_url=settings.COORDINATOR_BASE_URL,
|
||||
task_id=self.task_id,
|
||||
)
|
||||
|
||||
modeler_llm = LLM(
|
||||
api_key=settings.MODELER_API_KEY,
|
||||
model=settings.MODELER_MODEL,
|
||||
base_url=settings.MODELER_BASE_URL,
|
||||
task_id=self.task_id,
|
||||
)
|
||||
|
||||
coder_llm = LLM(
|
||||
api_key=settings.CODER_API_KEY,
|
||||
model=settings.CODER_MODEL,
|
||||
base_url=settings.CODER_BASE_URL,
|
||||
task_id=self.task_id,
|
||||
)
|
||||
|
||||
writer_llm = LLM(
|
||||
api_key=settings.WRITER_API_KEY,
|
||||
model=settings.WRITER_MODEL,
|
||||
base_url=settings.WRITER_BASE_URL,
|
||||
task_id=self.task_id,
|
||||
)
|
||||
|
||||
return coordinator_llm, modeler_llm, coder_llm, writer_llm
|
||||
+36
-18
@@ -1,14 +1,44 @@
|
||||
from app.utils.enums import FormatOutPut
|
||||
|
||||
|
||||
COORDINATOR_PROMPT = """
|
||||
判断用户输入的信息是否是数学建模问题
|
||||
如果是关于数学建模的,你将按照如下要求
|
||||
整理问题,将其交给建模手 ModelerAgent 分析
|
||||
{FORMAT_QUESTIONS_PROMPT}
|
||||
如果不是关于数学建模的,你将按照如下要求
|
||||
你会拒绝用户请求,输出一段拒绝的文字
|
||||
"""
|
||||
|
||||
FORMAT_QUESTIONS_PROMPT = """
|
||||
用户将提供给你一段题目信息,**请你不要更改题目信息,完整将用户输入的内容**,以 JSON 的形式输出,输出的 JSON 需遵守以下的格式:
|
||||
|
||||
{
|
||||
"title": <题目标题>
|
||||
"background": <题目背景,用户输入的一切不在title,ques1,ques2,ques3...中的内容都视为问题背景信息background>,
|
||||
"ques_count": <问题数量,number,int>,
|
||||
"ques1": <问题1>,
|
||||
"ques2": <问题2>,
|
||||
"ques3": <问题3,用户输入的存在多少问题,就输出多少问题ques1,ques2,ques3...以此类推>,
|
||||
}
|
||||
"""
|
||||
|
||||
# TODO: 设计成一个类?
|
||||
|
||||
MODELER_PROMPT = """
|
||||
role:你是一名数学建模经验丰富的建模手,负责建模部分。
|
||||
task:你需要根据用户要求和数据建立数学模型求解问题。
|
||||
task:你需要根据用户要求和数据对应每个问题建立数学模型求解问题。
|
||||
skill:熟练掌握各种数学建模的模型和思路
|
||||
output:数学建模的思路和使用到的模型
|
||||
attention:不需要给出代码,只需要给出思路和模型
|
||||
**不需要建立复杂的模型,简单规划需要步骤**
|
||||
format:以 JSON 的形式输出输出的 JSON,需遵守以下的格式:
|
||||
{
|
||||
"eda": <数据分析EDA方案>,
|
||||
"ques1": <问题1的建模思路和模型方案>,
|
||||
"ques2": <问题2的建模思路和模型方案>,
|
||||
"ques3": <问题3的建模思路和模型方案,用户输入的存在多少问题,就输出多少问题ques1,ques2,ques3...以此类推>,
|
||||
"sensitivity_analysis": <敏感性分析方案>,
|
||||
}
|
||||
"""
|
||||
|
||||
# TODO : 对于特大 csv 读取
|
||||
@@ -89,25 +119,13 @@ def get_writer_prompt(
|
||||
4. 严格按照参考用户输入的格式模板以及**正确的编号顺序**
|
||||
5. 不需要询问用户
|
||||
6. 当提到图片时,请使用提供的图片列表中的文件名
|
||||
7. when you write,check if you need to use tools search_papers to cite.if you need, markdown Footnote e.g.[^1]
|
||||
8. 对于问题背景和模型介绍,需查询文献调用tools search_papers
|
||||
7. when you write,check if you need to use tools search_papers to cite. if you need, markdown Footnote e.g.[^1]
|
||||
8. List all references at the end in markdown footnote format.
|
||||
9. Include an empty line between each citation for better readability.
|
||||
10. 对于问题背景和模型介绍,需查询文献调用tools search_papers
|
||||
"""
|
||||
|
||||
|
||||
FORMAT_QUESTIONS_PROMPT = """
|
||||
用户将提供给你一段题目信息,**请你不要更改题目信息,完整将用户输入的内容**,以 JSON 的形式输出,输出的 JSON 需遵守以下的格式:
|
||||
|
||||
{
|
||||
"title": <题目标题>
|
||||
"background": <题目背景,用户输入的一切不在title,ques1,ques2,ques3...中的内容都视为问题背景信息background>,
|
||||
"ques_count": <问题数量,number,int>,
|
||||
"ques1": <问题1>,
|
||||
"ques2": <问题2>,
|
||||
"ques3": <问题3,用户输入的存在多少问题,就输出多少问题ques1,ques2,ques3...以此类推>,
|
||||
}
|
||||
"""
|
||||
|
||||
|
||||
def get_reflection_prompt(error_message, code) -> str:
|
||||
return f"""The code execution encountered an error:
|
||||
{error_message}
|
||||
|
||||
+36
-157
@@ -1,5 +1,4 @@
|
||||
from app.core.agents import WriterAgent, CoderAgent
|
||||
from app.core.llm import LLM, simple_chat
|
||||
from app.core.agents import WriterAgent, CoderAgent, CoordinatorAgent, ModelerAgent
|
||||
from app.schemas.request import Problem
|
||||
from app.schemas.response import SystemMessage
|
||||
from app.tools.openalex_scholar import OpenAlexScholar
|
||||
@@ -8,10 +7,11 @@ from app.utils.common_utils import create_work_dir, get_config_template
|
||||
from app.models.user_output import UserOutput
|
||||
from app.config.setting import settings
|
||||
from app.tools.interpreter_factory import create_interpreter
|
||||
import json
|
||||
from app.utils.redis_manager import redis_manager
|
||||
from app.utils.notebook_serializer import NotebookSerializer
|
||||
from app.tools.base_interpreter import BaseCodeInterpreter
|
||||
from app.core.flows import Flows
|
||||
from app.core.llm.llm_factory import LLMFactory
|
||||
|
||||
|
||||
class WorkFlow:
|
||||
@@ -34,29 +34,37 @@ class MathModelWorkFlow(WorkFlow):
|
||||
self.task_id = problem.task_id
|
||||
self.work_dir = create_work_dir(self.task_id)
|
||||
|
||||
llm_model = LLM(
|
||||
api_key=settings.API_KEY,
|
||||
model=settings.MODEL,
|
||||
base_url=settings.BASE_URL,
|
||||
task_id=self.task_id,
|
||||
)
|
||||
llm_factory = LLMFactory(self.task_id)
|
||||
coordinator_llm, modeler_llm, coder_llm, writer_llm = llm_factory.get_all_llms()
|
||||
|
||||
coordinator_agent = CoordinatorAgent(self.task_id, coordinator_llm)
|
||||
|
||||
try:
|
||||
coordinator_response = await coordinator_agent.run(problem.ques_all)
|
||||
self.questions = coordinator_response.questions
|
||||
self.ques_count = coordinator_response.ques_count
|
||||
except Exception as e:
|
||||
# 非数学建模问题
|
||||
logger.error(f"CoordinatorAgent 执行失败: {e}")
|
||||
raise e
|
||||
|
||||
modeler_agent = ModelerAgent(self.task_id, modeler_llm)
|
||||
|
||||
modeler_response = await modeler_agent.run(coordinator_response)
|
||||
|
||||
await redis_manager.publish_message(
|
||||
self.task_id,
|
||||
SystemMessage(content="正在拆解问题问题"),
|
||||
)
|
||||
|
||||
await self.format_questions(problem.ques_all, llm_model)
|
||||
|
||||
user_output = UserOutput(work_dir=self.work_dir)
|
||||
|
||||
notebook_serializer = NotebookSerializer(work_dir=self.work_dir)
|
||||
|
||||
await redis_manager.publish_message(
|
||||
self.task_id,
|
||||
SystemMessage(content="正在创建代码沙盒环境"),
|
||||
)
|
||||
|
||||
notebook_serializer = NotebookSerializer(work_dir=self.work_dir)
|
||||
code_interpreter = await create_interpreter(
|
||||
kind="local",
|
||||
task_id=self.task_id,
|
||||
@@ -65,8 +73,7 @@ class MathModelWorkFlow(WorkFlow):
|
||||
timeout=3000,
|
||||
)
|
||||
|
||||
# Example usage
|
||||
scholar = OpenAlexScholar(email=settings.OPENALEX_EMAIL) # 请替换为您的真实邮箱
|
||||
scholar = OpenAlexScholar(email=settings.OPENALEX_EMAIL)
|
||||
|
||||
await redis_manager.publish_message(
|
||||
self.task_id,
|
||||
@@ -78,30 +85,31 @@ class MathModelWorkFlow(WorkFlow):
|
||||
SystemMessage(content="初始化代码手"),
|
||||
)
|
||||
|
||||
# modeler_agent
|
||||
coder_agent = CoderAgent(
|
||||
task_id=problem.task_id,
|
||||
model=llm_model,
|
||||
model=coder_llm,
|
||||
work_dir=self.work_dir,
|
||||
max_chat_turns=settings.MAX_CHAT_TURNS,
|
||||
max_retries=settings.MAX_RETRIES,
|
||||
code_interpreter=code_interpreter,
|
||||
)
|
||||
|
||||
# TODO: 自定义 writer_agent mode llm
|
||||
writer_agent = WriterAgent(
|
||||
task_id=problem.task_id,
|
||||
model=llm_model,
|
||||
model=writer_llm,
|
||||
comp_template=problem.comp_template,
|
||||
format_output=problem.format_output,
|
||||
scholar=scholar,
|
||||
)
|
||||
|
||||
################################################ solution steps
|
||||
solution_steps = self.get_solution_steps()
|
||||
flows = Flows()
|
||||
|
||||
################################################ solution steps
|
||||
solution_flows = flows.get_solution_flows(self.questions, modeler_response)
|
||||
config_template = get_config_template(problem.comp_template)
|
||||
|
||||
for key, value in solution_steps.items():
|
||||
for key, value in solution_flows.items():
|
||||
await redis_manager.publish_message(
|
||||
self.task_id,
|
||||
SystemMessage(content=f"代码手开始求解{key}"),
|
||||
@@ -116,9 +124,8 @@ class MathModelWorkFlow(WorkFlow):
|
||||
SystemMessage(content=f"代码手求解成功{key}", type="success"),
|
||||
)
|
||||
|
||||
# TODO: 是否可以不需要coder_response
|
||||
writer_prompt = self.get_writer_prompt(
|
||||
key, coder_response, code_interpreter, config_template
|
||||
writer_prompt = flows.get_writer_prompt(
|
||||
key, coder_response.code_response, code_interpreter, config_template
|
||||
)
|
||||
|
||||
await redis_manager.publish_message(
|
||||
@@ -147,147 +154,19 @@ class MathModelWorkFlow(WorkFlow):
|
||||
|
||||
################################################ write steps
|
||||
|
||||
flows = self.get_write_flows(user_output, config_template, problem.ques_all)
|
||||
for key, value in flows.items():
|
||||
write_flows = flows.get_write_flows(
|
||||
user_output, config_template, problem.ques_all
|
||||
)
|
||||
for key, value in write_flows.items():
|
||||
await redis_manager.publish_message(
|
||||
self.task_id,
|
||||
SystemMessage(content=f"论文手开始写{key}部分"),
|
||||
)
|
||||
|
||||
writer_response = await writer_agent.run(prompt=value, sub_title=key)
|
||||
|
||||
user_output.set_res(key, writer_response)
|
||||
|
||||
logger.info(user_output.get_res())
|
||||
|
||||
user_output.save_result(ques_count=self.ques_count)
|
||||
|
||||
async def format_questions(self, ques_all: str, model: LLM) -> None:
|
||||
"""用户输入问题 使用LLM 格式化 questions"""
|
||||
# TODO: "note": <补充说明,如果没有补充说明,请填 null>,
|
||||
from app.core.prompts import FORMAT_QUESTIONS_PROMPT
|
||||
|
||||
history = [
|
||||
{
|
||||
"role": "system",
|
||||
"content": FORMAT_QUESTIONS_PROMPT,
|
||||
},
|
||||
{"role": "user", "content": ques_all},
|
||||
]
|
||||
json_str = await simple_chat(model, history)
|
||||
json_str = json_str.replace("```json", "").replace("```", "").strip()
|
||||
|
||||
if not json_str:
|
||||
raise ValueError("返回的 JSON 字符串为空,请检查输入内容。")
|
||||
|
||||
try:
|
||||
self.questions = json.loads(json_str)
|
||||
self.ques_count = self.questions["ques_count"]
|
||||
logger.info(f"questions:{self.questions}")
|
||||
except json.JSONDecodeError as e:
|
||||
raise ValueError(f"JSON 解析错误: {e}")
|
||||
|
||||
def get_solution_steps(self):
|
||||
questions_quesx = {
|
||||
key: value
|
||||
for key, value in self.questions.items()
|
||||
if key.startswith("ques") and key != "ques_count"
|
||||
}
|
||||
ques_flow = {
|
||||
key: {
|
||||
"coder_prompt": f"""
|
||||
完成如下问题{value}
|
||||
""",
|
||||
}
|
||||
for key, value in questions_quesx.items()
|
||||
}
|
||||
flows = {
|
||||
"eda": {
|
||||
# TODO : 获取当前路径下的所有数据集
|
||||
"coder_prompt": """
|
||||
对当前目录下数据进行EDA分析(数据清洗,可视化),清洗后的数据保存当前目录下,**不需要复杂的模型**
|
||||
""",
|
||||
},
|
||||
**ques_flow,
|
||||
"sensitivity_analysis": {
|
||||
"coder_prompt": """
|
||||
根据上面建立的模型,选择一个模型,完成敏感性分析
|
||||
""",
|
||||
},
|
||||
}
|
||||
return flows
|
||||
|
||||
def get_writer_prompt(
|
||||
self,
|
||||
key: str,
|
||||
coder_response: str,
|
||||
code_interpreter: BaseCodeInterpreter,
|
||||
config_template: dict,
|
||||
) -> str:
|
||||
"""根据不同的key生成对应的writer_prompt
|
||||
|
||||
Args:
|
||||
key: 任务类型
|
||||
coder_response: 代码执行结果
|
||||
|
||||
Returns:
|
||||
str: 生成的writer_prompt
|
||||
"""
|
||||
code_output = code_interpreter.get_code_output(key)
|
||||
|
||||
# TODO: 结果{coder_response} 是否需要
|
||||
# TODO: 将当前产生的文件,路径发送给 writer_agent
|
||||
questions_quesx_keys = self.get_questions_quesx_keys()
|
||||
# TODO: 小标题编号
|
||||
# 题号最多6题
|
||||
bgc = self.questions["background"]
|
||||
quesx_writer_prompt = {
|
||||
key: f"""
|
||||
问题背景{bgc},不需要编写代码,代码手得到的结果{coder_response},{code_output},按照如下模板撰写:{config_template[key]}
|
||||
"""
|
||||
for key in questions_quesx_keys
|
||||
}
|
||||
|
||||
writer_prompt = {
|
||||
"eda": f"""
|
||||
问题背景{bgc},不需要编写代码,代码手得到的结果{coder_response},{code_output},按照如下模板撰写:{config_template["eda"]}
|
||||
""",
|
||||
**quesx_writer_prompt,
|
||||
"sensitivity_analysis": f"""
|
||||
问题背景{bgc},不需要编写代码,代码手得到的结果{coder_response},{code_output},按照如下模板撰写:{config_template["sensitivity_analysis"]}
|
||||
""",
|
||||
}
|
||||
|
||||
if key in writer_prompt:
|
||||
return writer_prompt[key]
|
||||
else:
|
||||
raise ValueError(f"未知的任务类型: {key}")
|
||||
|
||||
def get_questions_quesx_keys(self) -> list[str]:
|
||||
"""获取问题1,2...的键"""
|
||||
return list(self.get_questions_quesx().keys())
|
||||
|
||||
def get_questions_quesx(self) -> dict[str, str]:
|
||||
"""获取问题1,2,3...的键值对"""
|
||||
# 获取所有以 "ques" 开头的键值对
|
||||
questions_quesx = {
|
||||
key: value
|
||||
for key, value in self.questions.items()
|
||||
if key.startswith("ques") and key != "ques_count"
|
||||
}
|
||||
return questions_quesx
|
||||
|
||||
def get_write_flows(
|
||||
self, user_output: UserOutput, config_template: dict, bg_ques_all: str
|
||||
):
|
||||
model_build_solve = user_output.get_model_build_solve()
|
||||
flows = {
|
||||
"firstPage": f"""问题背景{bg_ques_all},不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["firstPage"]},撰写标题,摘要,关键词""",
|
||||
"RepeatQues": f"""问题背景{bg_ques_all},不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["RepeatQues"]},撰写问题重述""",
|
||||
"analysisQues": f"""问题背景{bg_ques_all},不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["analysisQues"]},撰写问题分析""",
|
||||
"modelAssumption": f"""问题背景{bg_ques_all},不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["modelAssumption"]},撰写模型假设""",
|
||||
"symbol": f"""不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["symbol"]},撰写符号说明部分""",
|
||||
"judge": f"""不需要编写代码,根据模型的求解的信息{model_build_solve},按照如下模板撰写:{config_template["judge"]},撰写模型的评价部分""",
|
||||
# TODO: 修改参考文献插入方式
|
||||
"reference": f"""不需要编写代码,根据模型的求解的信息{model_build_solve},可以生成参考文献,按照如下模板撰写:{config_template["reference"]},撰写参考文献""",
|
||||
}
|
||||
return flows
|
||||
|
||||
@@ -1,7 +1,22 @@
|
||||
from pydantic import BaseModel
|
||||
from typing import Any
|
||||
|
||||
|
||||
class CoordinatorToModeler(BaseModel):
|
||||
questions: dict
|
||||
ques_count: int
|
||||
|
||||
|
||||
class ModelerToCoder(BaseModel):
|
||||
questions_solution: dict[str, str]
|
||||
|
||||
|
||||
class CoderToWriter(BaseModel):
|
||||
code_response: str
|
||||
code_execution_result: str
|
||||
code_output: str
|
||||
created_images: list[str]
|
||||
|
||||
|
||||
class WriterResponse(BaseModel):
|
||||
response_content: Any
|
||||
footnotes: list[str] | None = None
|
||||
|
||||
@@ -1,20 +1,30 @@
|
||||
import os
|
||||
from app.utils.data_recorder import DataRecorder
|
||||
from app.models.model import WriterResponse
|
||||
|
||||
|
||||
class UserOutput:
|
||||
def __init__(self, work_dir: str, data_recorder: DataRecorder | None = None):
|
||||
self.work_dir = work_dir
|
||||
self.res: dict[str, str] = {
|
||||
# "eda": "",
|
||||
# "ques1": "",
|
||||
self.res: dict[str, dict] = {
|
||||
# "eda": {
|
||||
# "response_content": "",
|
||||
# "footnotes": "",
|
||||
# },
|
||||
# "ques1": {
|
||||
# "response_content": "",
|
||||
# "footnotes": "",
|
||||
# },
|
||||
}
|
||||
self.data_recorder = data_recorder
|
||||
self.cost_time = 0.0
|
||||
self.initialized = True
|
||||
|
||||
def set_res(self, key: str, value: str):
|
||||
self.res[key] = value # TODO: 换种数据类型有顺序
|
||||
def set_res(self, key: str, writer_response: WriterResponse):
|
||||
self.res[key] = {
|
||||
"response_content": writer_response.response_content,
|
||||
"footnotes": writer_response.footnotes,
|
||||
}
|
||||
|
||||
def get_res(self):
|
||||
return self.res
|
||||
@@ -44,7 +54,49 @@ class UserOutput:
|
||||
"judge",
|
||||
"reference",
|
||||
]
|
||||
return "\n".join([self.res.get(key, "") for key in seq])
|
||||
|
||||
# 收集所有内容和脚注
|
||||
all_content = []
|
||||
all_footnotes = []
|
||||
footnote_counter = 1
|
||||
|
||||
for key in seq:
|
||||
if key not in self.res:
|
||||
continue
|
||||
|
||||
content = self.res[key]["response_content"]
|
||||
footnotes = self.res[key]["footnotes"]
|
||||
|
||||
# 更新内容中的脚注引用编号
|
||||
if footnotes:
|
||||
# 获取当前内容中的所有脚注引用
|
||||
current_footnotes = footnotes.split("\n")
|
||||
|
||||
# 更新内容中的脚注引用编号
|
||||
for i, _ in enumerate(current_footnotes, start=footnote_counter):
|
||||
content = content.replace(
|
||||
f"[^{i - footnote_counter + 1}]", f"[^{i}]"
|
||||
)
|
||||
|
||||
# 更新脚注编号
|
||||
updated_footnotes = []
|
||||
for i, footnote in enumerate(current_footnotes, start=footnote_counter):
|
||||
updated_footnote = footnote.replace(
|
||||
f"[^{i - footnote_counter + 1}]:", f"[^{i}]:"
|
||||
)
|
||||
updated_footnotes.append(updated_footnote)
|
||||
|
||||
footnote_counter += len(current_footnotes)
|
||||
all_footnotes.extend(updated_footnotes)
|
||||
|
||||
all_content.append(content)
|
||||
|
||||
# 合并所有内容和脚注
|
||||
final_content = "\n".join(all_content)
|
||||
if all_footnotes:
|
||||
final_content += "\n\n" + "\n".join(all_footnotes)
|
||||
|
||||
return final_content
|
||||
|
||||
def save_result(self, ques_count):
|
||||
res_path = os.path.join(self.work_dir, "res.md")
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
from typing import List, Literal, Union
|
||||
from typing import Literal, Union
|
||||
from app.utils.enums import AgentType
|
||||
from pydantic import BaseModel, Field
|
||||
from uuid import uuid4
|
||||
@@ -23,7 +23,7 @@ class UserMessage(Message):
|
||||
|
||||
class AgentMessage(Message):
|
||||
msg_type: str = "agent"
|
||||
agent_type: AgentType # CoderAgent | WriterAgent
|
||||
agent_type: AgentType # CoordinatorAgent | ModelerAgent | CoderAgent | WriterAgent
|
||||
|
||||
|
||||
class CodeExecution(BaseModel):
|
||||
|
||||
@@ -47,7 +47,7 @@ class OpenAlexScholar:
|
||||
# 拼接单词形成文本
|
||||
return " ".join(words).strip()
|
||||
|
||||
def search_papers(self, query: str, limit: int = 10) -> List[Dict[str, Any]]:
|
||||
def search_papers(self, query: str, limit: int = 8) -> List[Dict[str, Any]]:
|
||||
"""Search for papers using OpenAlex API.
|
||||
|
||||
Args:
|
||||
@@ -148,20 +148,21 @@ class OpenAlexScholar:
|
||||
|
||||
return papers
|
||||
|
||||
def print_papers(self, papers: List[Dict[str, Any]]):
|
||||
def papers_to_str(self, papers: List[Dict[str, Any]]) -> str:
|
||||
"""将文献列表转换为字符串"""
|
||||
result = ""
|
||||
for paper in papers:
|
||||
print("\n" + "=" * 80)
|
||||
print(f"标题: {paper['title']}")
|
||||
print(f"\n摘要: {paper['abstract']}")
|
||||
print("\n作者:")
|
||||
for author in paper["authors"]:
|
||||
print(f"- {author['name']}")
|
||||
if author["institution"]:
|
||||
print(f" 所属机构: {author['institution']}")
|
||||
print(f"\n引用次数: {paper['citations_count']}")
|
||||
print(f"发表年份: {paper['publication_year']}")
|
||||
print(f"\n引用格式:\n{paper['citation_format']}")
|
||||
print("=" * 80)
|
||||
result += "\n" + "=" * 80
|
||||
result += f"\n标题: {paper['title']}"
|
||||
result += f"\n摘要: {paper['abstract']}"
|
||||
result += "\n作者:"
|
||||
for author in paper["authors"]:
|
||||
result += f"- {author['name']}"
|
||||
result += f"\n引用次数: {paper['citations_count']}"
|
||||
result += f"\n发表年份: {paper['publication_year']}"
|
||||
result += f"\n引用格式:\n{paper['citation_format']}"
|
||||
result += "=" * 80
|
||||
return result
|
||||
|
||||
def _format_citation(self, work: Dict[str, Any]) -> str:
|
||||
"""Format citation in a readable format."""
|
||||
|
||||
@@ -7,6 +7,7 @@ from app.utils.log_util import logger
|
||||
import re
|
||||
import pypandoc
|
||||
from app.config.setting import settings
|
||||
from icecream import ic
|
||||
|
||||
|
||||
def create_task_id() -> str:
|
||||
@@ -96,7 +97,6 @@ def md_2_docx(task_id: str):
|
||||
str(work_dir),
|
||||
"--mathml", # MathML 格式公式
|
||||
"--standalone",
|
||||
# "--extract-media=" + str(md_dir / "generated_images") # 按需启用
|
||||
]
|
||||
|
||||
pypandoc.convert_file(
|
||||
@@ -108,3 +108,12 @@ def md_2_docx(task_id: str):
|
||||
)
|
||||
print(f"转换完成: {docx_path}")
|
||||
logger.info(f"转换完成: {docx_path}")
|
||||
|
||||
|
||||
def get_footnotes(text: str):
|
||||
# 匹配脚注定义
|
||||
footnotes = re.findall(r"\[\^(\d+)\]:\s*(.+?)(?=\n\[\^|\n\n|\Z)", text, re.DOTALL)
|
||||
|
||||
for num, content in footnotes:
|
||||
ic(f"[^{num}] {content.strip()}")
|
||||
return footnotes
|
||||
|
||||
@@ -12,6 +12,8 @@ class FormatOutPut(str, Enum):
|
||||
|
||||
|
||||
class AgentType(str, Enum):
|
||||
COORDINATOR = "CoordinatorAgent"
|
||||
MODELER = "ModelerAgent"
|
||||
CODER = "CoderAgent"
|
||||
WRITER = "WriterAgent"
|
||||
SYSTEM = "SystemAgent"
|
||||
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 140 KiB |
@@ -33,9 +33,9 @@ const props = defineProps({
|
||||
user: {
|
||||
type: Object,
|
||||
default: () => ({
|
||||
name: 'John Doe',
|
||||
email: 'john.doe@example.com',
|
||||
avatar: 'https://github.com/shadcn.png'
|
||||
name: 'San Jin',
|
||||
email: 'mathmodel@mathmodel.com',
|
||||
avatar: 'https://github.com/jihe520.png'
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
Reference in New Issue
Block a user