diff --git a/src/agents/common/middlewares/attachment_middleware.py b/src/agents/common/middlewares/attachment_middleware.py index 663cb0d8..0e423508 100644 --- a/src/agents/common/middlewares/attachment_middleware.py +++ b/src/agents/common/middlewares/attachment_middleware.py @@ -10,9 +10,12 @@ from typing import NotRequired from langchain.agents import AgentState from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse +from langchain_core.messages import SystemMessage from src.utils import logger +ATTACHMENT_PROMPT_MARKER = "" + class AttachmentState(AgentState): """扩展 AgentState 以支持附件""" @@ -60,7 +63,7 @@ class AttachmentMiddleware(AgentMiddleware[AttachmentState]): LangChain 标准中间件:从 State 中读取附件并注入提示词。 LangGraph 会自动从 checkpointer 恢复 state,包括 attachments。 - 从 request.state 中读取附件,将其转换为 SystemMessage 并注入到消息列表开头。 + 从 request.state 中读取附件,将其转换为上下文块 并注入到系统提示词中。 """ state_schema = AttachmentState @@ -78,23 +81,21 @@ class AttachmentMiddleware(AgentMiddleware[AttachmentState]): if attachment_prompt: logger.info("AttachmentMiddleware: injecting attachment prompt") + existing_blocks = list(request.system_message.content_blocks) if request.system_message else [] + existing_text = "\n".join( + block.get("text", "") + for block in existing_blocks + if isinstance(block, dict) and block.get("type") == "text" + ) - messages = list(request.messages) - insert_idx = 0 - for idx, msg in enumerate(messages): - if isinstance(msg, dict): - role = msg.get("role") or msg.get("type") - is_system = role == "system" - else: - msg_type = getattr(msg, "type", None) or getattr(msg, "role", None) - is_system = msg_type == "system" + if ATTACHMENT_PROMPT_MARKER in existing_text: + logger.info("AttachmentMiddleware: attachment prompt already injected, skip") + return await handler(request) - if not is_system: - break - insert_idx = idx + 1 - - messages.insert(insert_idx, {"role": "system", "content": attachment_prompt}) - request = request.override(messages=messages) + merged_blocks = existing_blocks + [ + {"type": "text", "text": f"{ATTACHMENT_PROMPT_MARKER}\n{attachment_prompt}"} + ] + request = request.override(system_message=SystemMessage(content=merged_blocks)) return await handler(request)