merge: attachment prompt idempotent system message into main

This commit is contained in:
肖泽涛 2026-02-21 02:23:53 +08:00
commit d43164576f

View File

@ -10,9 +10,12 @@ from typing import NotRequired
from langchain.agents import AgentState from langchain.agents import AgentState
from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse from langchain.agents.middleware import AgentMiddleware, ModelRequest, ModelResponse
from langchain_core.messages import SystemMessage
from src.utils import logger from src.utils import logger
ATTACHMENT_PROMPT_MARKER = "<!-- attachment_context -->"
class AttachmentState(AgentState): class AttachmentState(AgentState):
"""扩展 AgentState 以支持附件""" """扩展 AgentState 以支持附件"""
@ -60,7 +63,7 @@ class AttachmentMiddleware(AgentMiddleware[AttachmentState]):
LangChain 标准中间件 State 中读取附件并注入提示词 LangChain 标准中间件 State 中读取附件并注入提示词
LangGraph 会自动从 checkpointer 恢复 state包括 attachments LangGraph 会自动从 checkpointer 恢复 state包括 attachments
request.state 中读取附件将其转换为 SystemMessage 并注入到消息列表开头 request.state 中读取附件将其转换为上下文块 并注入到系统提示词中
""" """
state_schema = AttachmentState state_schema = AttachmentState
@ -78,23 +81,21 @@ class AttachmentMiddleware(AgentMiddleware[AttachmentState]):
if attachment_prompt: if attachment_prompt:
logger.info("AttachmentMiddleware: injecting 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) if ATTACHMENT_PROMPT_MARKER in existing_text:
insert_idx = 0 logger.info("AttachmentMiddleware: attachment prompt already injected, skip")
for idx, msg in enumerate(messages): return await handler(request)
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 not is_system: merged_blocks = existing_blocks + [
break {"type": "text", "text": f"{ATTACHMENT_PROMPT_MARKER}\n{attachment_prompt}"}
insert_idx = idx + 1 ]
request = request.override(system_message=SystemMessage(content=merged_blocks))
messages.insert(insert_idx, {"role": "system", "content": attachment_prompt})
request = request.override(messages=messages)
return await handler(request) return await handler(request)