merge: attachment prompt idempotent system message into main
This commit is contained in:
commit
d43164576f
@ -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)
|
||||||
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user