mirror of
https://github.com/infiniflow/ragflow.git
synced 2026-08-08 00:18:12 +08:00
Fix attachments not take effect in agentic chat (#17895)
This commit is contained in:
@@ -290,10 +290,6 @@ class DialogService(CommonService):
|
||||
|
||||
|
||||
async def async_chat_solo(dialog, messages, stream=True, session_id=None):
|
||||
attachments = ""
|
||||
image_attachments = []
|
||||
image_files = []
|
||||
|
||||
if dialog.llm_id:
|
||||
if dialog.tenant_llm_id:
|
||||
try:
|
||||
@@ -319,12 +315,8 @@ async def async_chat_solo(dialog, messages, stream=True, session_id=None):
|
||||
|
||||
chat_mdl = LLMBundle(dialog.tenant_id, model_config, langfuse_session_id=session_id)
|
||||
factory = model_config.get("llm_factory", "") if model_config else ""
|
||||
if "files" in messages[-1]:
|
||||
if model_config["model_type"] == "chat":
|
||||
text_attachments, image_attachments = split_file_attachments(messages[-1]["files"])
|
||||
else:
|
||||
text_attachments, image_files = split_file_attachments(messages[-1]["files"], raw=True)
|
||||
attachments = "\n\n".join(text_attachments)
|
||||
|
||||
text_attachments_content, image_attachments, image_files = get_files_content(messages[-1], model_config["model_type"])
|
||||
|
||||
prompt_config = dialog.prompt_config
|
||||
tts_mdl = None
|
||||
@@ -332,8 +324,8 @@ async def async_chat_solo(dialog, messages, stream=True, session_id=None):
|
||||
default_tts_model = get_tenant_default_model_by_type(dialog.tenant_id, LLMType.TTS)
|
||||
tts_mdl = LLMBundle(dialog.tenant_id, default_tts_model, trace_context=chat_mdl.trace_context, langfuse_session_id=session_id)
|
||||
msg = [{"role": m["role"], "content": re.sub(r"##\d+\$\$", "", m["content"])} for m in messages if m["role"] != "system"]
|
||||
if attachments and msg:
|
||||
msg[-1]["content"] += attachments
|
||||
if text_attachments_content and msg:
|
||||
msg[-1]["content"] += text_attachments_content
|
||||
if model_config["model_type"] == "chat" and image_attachments:
|
||||
convert_last_user_msg_to_multimodal(msg, image_attachments, factory)
|
||||
sys_date = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
|
||||
@@ -402,6 +394,19 @@ def get_models(dialog, trace_context=None, langfuse_session_id=None):
|
||||
return kbs, embd_mdl, rerank_mdl, chat_mdl, tts_mdl
|
||||
|
||||
|
||||
def get_files_content(last_message, model_type):
|
||||
text_attachments_content = ""
|
||||
image_attachments = []
|
||||
image_files = []
|
||||
if "files" in last_message:
|
||||
if model_type == "chat":
|
||||
text_attachments, image_attachments = split_file_attachments(last_message["files"])
|
||||
else:
|
||||
text_attachments, image_files = split_file_attachments(last_message["files"], raw=True)
|
||||
text_attachments_content = "\n\n".join(text_attachments)
|
||||
return text_attachments_content, image_attachments, image_files
|
||||
|
||||
|
||||
def split_file_attachments(files: list[dict] | None, raw: bool = False) -> tuple[list[str], list[str] | list[dict]]:
|
||||
if not files:
|
||||
return [], []
|
||||
@@ -636,40 +641,35 @@ async def async_chat(dialog, messages, stream=True, **kwargs):
|
||||
|
||||
retriever = settings.retriever
|
||||
questions = [m["content"] for m in messages if m["role"] == "user"][-3:]
|
||||
attachments = None
|
||||
if "doc_ids" in kwargs:
|
||||
attachments = [doc_id for doc_id in kwargs["doc_ids"].split(",") if doc_id]
|
||||
attachments_ = ""
|
||||
image_attachments = []
|
||||
image_files = []
|
||||
if "doc_ids" in messages[-1]:
|
||||
attachments = [doc_id for doc_id in messages[-1]["doc_ids"] if doc_id]
|
||||
if "files" in messages[-1]:
|
||||
if llm_model_config["model_type"] == "chat":
|
||||
text_attachments, image_attachments = split_file_attachments(messages[-1]["files"])
|
||||
else:
|
||||
text_attachments, image_files = split_file_attachments(messages[-1]["files"], raw=True)
|
||||
attachments_ = "\n\n".join(text_attachments)
|
||||
|
||||
prompt_config = dialog.prompt_config
|
||||
# Get scoped_doc_ids
|
||||
scoped_doc_ids = None
|
||||
if "doc_ids" in kwargs:
|
||||
scoped_doc_ids = [doc_id for doc_id in kwargs["doc_ids"].split(",") if doc_id]
|
||||
if "doc_ids" in messages[-1]:
|
||||
scoped_doc_ids = [doc_id for doc_id in messages[-1]["doc_ids"] if doc_id]
|
||||
if dialog.meta_data_filter:
|
||||
attachments = await apply_meta_data_filter(
|
||||
scoped_doc_ids = await apply_meta_data_filter(
|
||||
dialog.meta_data_filter,
|
||||
None,
|
||||
questions[-1],
|
||||
chat_mdl,
|
||||
attachments,
|
||||
scoped_doc_ids,
|
||||
kb_ids=dialog.kb_ids,
|
||||
metas_loader=lambda: DocMetadataService.get_flatted_meta_by_kbs(dialog.kb_ids),
|
||||
)
|
||||
|
||||
# Get chat attachments
|
||||
text_attachments_content, image_attachments, image_files = get_files_content(messages[-1], llm_model_config["model_type"])
|
||||
|
||||
prompt_config = dialog.prompt_config
|
||||
include_reference_metadata, metadata_fields = _resolve_reference_metadata(prompt_config, request_payload=kwargs)
|
||||
field_map = KnowledgebaseService.get_field_map(dialog.kb_ids)
|
||||
logging.debug(f"field_map retrieved: {field_map}")
|
||||
# try to use sql if field mapping is good to go
|
||||
if field_map:
|
||||
logging.debug("Use SQL to retrieval:{}".format(questions[-1]))
|
||||
ans = await use_sql(questions[-1], field_map, dialog.tenant_id, chat_mdl, prompt_config.get("quote", True), dialog.kb_ids, doc_ids=attachments)
|
||||
ans = await use_sql(questions[-1], field_map, dialog.tenant_id, chat_mdl, prompt_config.get("quote", True), dialog.kb_ids, doc_ids=scoped_doc_ids)
|
||||
# For aggregate queries (COUNT, SUM, etc.), chunks may be empty but answer is still valid
|
||||
if ans and (ans.get("reference", {}).get("chunks") or ans.get("answer")):
|
||||
if include_reference_metadata and ans.get("reference", {}).get("chunks"):
|
||||
@@ -689,7 +689,7 @@ async def async_chat(dialog, messages, stream=True, **kwargs):
|
||||
logging.warning("prompt_config['parameters'] is missing 'knowledge' entry despite kb_ids being set; auto-fixing.")
|
||||
prompt_config.setdefault("parameters", []).append({"key": "knowledge", "optional": False})
|
||||
param_keys.append("knowledge")
|
||||
logging.debug(f"attachments={attachments}, param_keys={param_keys}, embd_mdl={embd_mdl}")
|
||||
logging.debug(f"scoped_doc_ids={scoped_doc_ids}, param_keys={param_keys}, embd_mdl={embd_mdl}")
|
||||
|
||||
sys_date = datetime.now(timezone.utc).strftime("%Y-%m-%d %H:%M:%S")
|
||||
kwargs["date"] = sys_date
|
||||
@@ -735,7 +735,7 @@ async def async_chat(dialog, messages, stream=True, **kwargs):
|
||||
page_size=dialog.top_n,
|
||||
similarity_threshold=0.2,
|
||||
vector_similarity_weight=0.3,
|
||||
doc_ids=attachments,
|
||||
doc_ids=scoped_doc_ids,
|
||||
),
|
||||
internet_enabled=use_web_search,
|
||||
)
|
||||
@@ -770,7 +770,7 @@ async def async_chat(dialog, messages, stream=True, **kwargs):
|
||||
dialog.top_n,
|
||||
dialog.similarity_threshold,
|
||||
dialog.vector_similarity_weight,
|
||||
doc_ids=attachments,
|
||||
doc_ids=scoped_doc_ids,
|
||||
top=dialog.top_k,
|
||||
aggs=True,
|
||||
rerank_mdl=rerank_mdl,
|
||||
@@ -828,7 +828,7 @@ async def async_chat(dialog, messages, stream=True, **kwargs):
|
||||
kwargs.setdefault("knowledge", "")
|
||||
gen_conf = dialog.llm_setting
|
||||
|
||||
system_content = prompt_config["system"].format(**kwargs) + attachments_
|
||||
system_content = prompt_config["system"].format(**kwargs) + text_attachments_content
|
||||
# If knowledge was retrieved but the template has no {knowledge}
|
||||
# placeholder, auto-append it so the LLM still sees the context.
|
||||
if knowledges and "{knowledge}" not in prompt_config.get("system", ""):
|
||||
@@ -1878,6 +1878,12 @@ async def rag_agent(dialog, messages, stream=True, **kwargs):
|
||||
yield ans
|
||||
return
|
||||
kbs, embd_mdl, rerank_mdl, chat_mdl, tts_mdl = get_models(dialog)
|
||||
model_type = chat_mdl.model_config["model_type"]
|
||||
factory = chat_mdl.model_config.get("llm_factory", "") if chat_mdl.model_config else ""
|
||||
text_attachments_content, image_attachments, image_files = get_files_content(messages[-1], model_type)
|
||||
agent_messages = deepcopy(messages)
|
||||
if model_type == "chat" and image_attachments:
|
||||
convert_last_user_msg_to_multimodal(agent_messages, image_attachments, factory)
|
||||
use_web_search = _should_use_web_search(prompt_config, kwargs.get("internet"))
|
||||
logging.debug("web_search kb=%s configured=%s internet=%r enabled=%s", bool(dialog.kb_ids), has_web_search_provider(prompt_config), kwargs.get("internet"), use_web_search)
|
||||
tenant_ids = list(set([kb.tenant_id for kb in kbs]))
|
||||
@@ -1923,6 +1929,7 @@ async def rag_agent(dialog, messages, stream=True, **kwargs):
|
||||
empty_response=prompt_config.get("empty_response", ""),
|
||||
do_refer=False,
|
||||
thinking_mode=thinking_mode,
|
||||
text_attachments_content=text_attachments_content,
|
||||
)
|
||||
|
||||
async def decorate_answer(answer):
|
||||
@@ -2011,7 +2018,10 @@ async def rag_agent(dialog, messages, stream=True, **kwargs):
|
||||
|
||||
async def _drive_stream():
|
||||
try:
|
||||
stream_iter = chat_mdl.async_chat_streamly_delta(rag_tools.sys_prompt(), messages, gen_conf)
|
||||
if model_type == "chat":
|
||||
stream_iter = chat_mdl.async_chat_streamly_delta(rag_tools.sys_prompt(), agent_messages, gen_conf)
|
||||
else:
|
||||
stream_iter = chat_mdl.async_chat_streamly_delta(rag_tools.sys_prompt(), agent_messages, gen_conf, images=image_files)
|
||||
async for kind, value, state in _stream_with_think_delta(stream_iter):
|
||||
event_queue.put_nowait(("stream", kind, value, state))
|
||||
except Exception:
|
||||
@@ -2111,8 +2121,11 @@ async def rag_agent(dialog, messages, stream=True, **kwargs):
|
||||
final["audio_binary"] = None
|
||||
yield final
|
||||
else:
|
||||
answer = await chat_mdl.async_chat(rag_tools.sys_prompt(), messages, gen_conf)
|
||||
user_content = messages[-1].get("content", "[content not available]")
|
||||
if model_type == "chat":
|
||||
answer = await chat_mdl.async_chat(rag_tools.sys_prompt(), agent_messages, gen_conf)
|
||||
else:
|
||||
answer = await chat_mdl.async_chat(rag_tools.sys_prompt(), agent_messages, gen_conf, images=image_files)
|
||||
user_content = agent_messages[-1].get("content", "[content not available]")
|
||||
logging.debug("User: {}|Assistant: {}".format(user_content, answer))
|
||||
res = await decorate_answer(answer)
|
||||
res["audio_binary"] = tts(tts_mdl, answer)
|
||||
|
||||
@@ -83,6 +83,7 @@ class RAGTools:
|
||||
empty_response: str = "",
|
||||
do_refer: bool | None = True,
|
||||
thinking_mode: str = "medium",
|
||||
text_attachments_content: str = "",
|
||||
):
|
||||
self.tenant_ids = tenant_ids
|
||||
self.chat_mdl = chat_mdl.clone()
|
||||
@@ -114,6 +115,7 @@ class RAGTools:
|
||||
self.user_defined_prompts = user_defined_prompts or {}
|
||||
self.empty_response = empty_response
|
||||
self.do_refer = do_refer
|
||||
self.text_attachments_content = text_attachments_content or ""
|
||||
# Optional sink used by the outer agent stream to preserve the final
|
||||
# answer deltas produced by the inner research graph. The tool API
|
||||
# still returns the complete string to the caller, but the stream
|
||||
@@ -572,6 +574,19 @@ class RAGTools:
|
||||
|
||||
if self.tool_started_sink is not None:
|
||||
self.tool_started_sink()
|
||||
if self.text_attachments_content:
|
||||
self.kbinfos = {
|
||||
"chunks": [
|
||||
{
|
||||
"id": "chat_attachment",
|
||||
"chunk_id": "chat_attachment",
|
||||
"doc_id": "chat_attachment",
|
||||
"docnm_kwd": "Chat attachment",
|
||||
"content_with_weight": self.text_attachments_content,
|
||||
}
|
||||
],
|
||||
"doc_aggs": [{"doc_id": "chat_attachment", "doc_name": "Chat attachment", "count": 1}],
|
||||
}
|
||||
messages = [{"role": "user", "content": question}] if question else []
|
||||
final = ""
|
||||
async for kind, delta in _split_think_stream(run_agentic_rag(self, messages)):
|
||||
|
||||
@@ -143,6 +143,8 @@ async def agentic_research(state: dict, tools) -> dict:
|
||||
if action == "ANSWER_PARTIAL":
|
||||
return _finalize(ctx, tools, partial=True)
|
||||
if action == "ABSTAIN":
|
||||
if getattr(tools, "text_attachments_content", ""):
|
||||
return {"verdict": verdict.__dict__, "kbinfos": tools.kbinfos}
|
||||
tools.kbinfos["chunks"] = []
|
||||
return {"verdict": verdict.__dict__, "abstain": True}
|
||||
if action == "REPLAN":
|
||||
|
||||
@@ -78,6 +78,8 @@ async def decompose_and_search(state: dict, tools) -> dict:
|
||||
"kbinfos": tools.kbinfos,
|
||||
}
|
||||
if action == "ABSTAIN":
|
||||
if getattr(tools, "text_attachments_content", ""):
|
||||
return {"verdict": verdict.__dict__, "kbinfos": tools.kbinfos}
|
||||
tools.kbinfos["chunks"] = []
|
||||
return {"verdict": verdict.__dict__, "abstain": True}
|
||||
|
||||
|
||||
31
test/unit_test/rag/advanced_rag/test_agentic_rag.py
Normal file
31
test/unit_test/rag/advanced_rag/test_agentic_rag.py
Normal file
@@ -0,0 +1,31 @@
|
||||
from copy import deepcopy
|
||||
|
||||
import pytest
|
||||
|
||||
from rag.advanced_rag.agentic_rag import RAGTools
|
||||
|
||||
|
||||
class FakeChatModel:
|
||||
max_length = 8192
|
||||
|
||||
def clone(self):
|
||||
return self
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rag_tool_adds_text_attachment_as_evidence(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_run_agentic_rag(tools, messages):
|
||||
captured["kbinfos"] = deepcopy(tools.kbinfos)
|
||||
captured["messages"] = messages
|
||||
yield "answer"
|
||||
|
||||
monkeypatch.setattr("rag.advanced_rag.agentic_rag_graph.run_agentic_rag", fake_run_agentic_rag)
|
||||
|
||||
tools = RAGTools([], FakeChatModel(), text_attachments_content="attached facts")
|
||||
|
||||
assert await tools.rag("What is attached?") == "answer"
|
||||
assert captured["messages"] == [{"role": "user", "content": "What is attached?"}]
|
||||
assert captured["kbinfos"]["chunks"][0]["docnm_kwd"] == "Chat attachment"
|
||||
assert captured["kbinfos"]["chunks"][0]["content_with_weight"] == "attached facts"
|
||||
Reference in New Issue
Block a user