diff --git a/api/db/services/dialog_service.py b/api/db/services/dialog_service.py index a17a01a5b3..b57196aea9 100644 --- a/api/db/services/dialog_service.py +++ b/api/db/services/dialog_service.py @@ -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) diff --git a/rag/advanced_rag/agentic_rag.py b/rag/advanced_rag/agentic_rag.py index ca438b7a66..b076807918 100644 --- a/rag/advanced_rag/agentic_rag.py +++ b/rag/advanced_rag/agentic_rag.py @@ -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)): diff --git a/rag/advanced_rag/harness/orchestrator/agentic.py b/rag/advanced_rag/harness/orchestrator/agentic.py index f857435629..49c7387023 100644 --- a/rag/advanced_rag/harness/orchestrator/agentic.py +++ b/rag/advanced_rag/harness/orchestrator/agentic.py @@ -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": diff --git a/rag/advanced_rag/harness/orchestrator/decompose.py b/rag/advanced_rag/harness/orchestrator/decompose.py index ea23a17cf7..07c85d5a16 100644 --- a/rag/advanced_rag/harness/orchestrator/decompose.py +++ b/rag/advanced_rag/harness/orchestrator/decompose.py @@ -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} diff --git a/test/unit_test/rag/advanced_rag/test_agentic_rag.py b/test/unit_test/rag/advanced_rag/test_agentic_rag.py new file mode 100644 index 0000000000..9e826d942b --- /dev/null +++ b/test/unit_test/rag/advanced_rag/test_agentic_rag.py @@ -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"