fix: preserve agent debug conversation context

This commit is contained in:
2026-07-17 16:06:43 +08:00
parent cadf44e0b8
commit 537978082a
7 changed files with 92 additions and 15 deletions

View File

@@ -275,6 +275,10 @@ async function debugAgent() {
if (!agentForm.question.trim()) return ElMessage.error("请输入调试问题"); if (!agentForm.question.trim()) return ElMessage.error("请输入调试问题");
if (!promptContent.value.trim()) return ElMessage.error("主提示词不能为空"); if (!promptContent.value.trim()) return ElMessage.error("主提示词不能为空");
const question = agentForm.question.trim(); const question = agentForm.question.trim();
const conversationHistory = agentPreviewMessages.value
.slice(1)
.filter((message) => ["user", "assistant"].includes(message.role) && message.content.trim() && !message.streaming)
.map((message) => ({ role: message.role, content: message.content.trim() }));
agentPreviewMessages.value.push({ role: "user", content: question }); agentPreviewMessages.value.push({ role: "user", content: question });
agentPreviewMessages.value.push({ role: "assistant", content: "", streaming: true }); agentPreviewMessages.value.push({ role: "assistant", content: "", streaming: true });
const assistantIndex = agentPreviewMessages.value.length - 1; const assistantIndex = agentPreviewMessages.value.length - 1;
@@ -291,6 +295,7 @@ async function debugAgent() {
modelId: agentForm.modelId, modelId: agentForm.modelId,
knowledgeIds: agentForm.knowledgeIds, knowledgeIds: agentForm.knowledgeIds,
question, question,
history: conversationHistory,
temperature: agentForm.temperature, temperature: agentForm.temperature,
topP: agentForm.topP, topP: agentForm.topP,
topK: agentForm.topK, topK: agentForm.topK,

View File

@@ -1,6 +1,7 @@
from __future__ import annotations from __future__ import annotations
from datetime import datetime from datetime import datetime
from typing import Literal
from pydantic import BaseModel, Field from pydantic import BaseModel, Field
@@ -77,12 +78,18 @@ class PromptSaveRequest(BaseModel):
promptContent: str = Field(min_length=1) promptContent: str = Field(min_length=1)
class AgentDebugHistoryMessage(BaseModel):
role: Literal["user", "assistant"]
content: str = Field(min_length=1, max_length=20000)
class AgentDebugRequest(BaseModel): class AgentDebugRequest(BaseModel):
promptContent: str = Field(min_length=1) promptContent: str = Field(min_length=1)
modelId: int = Field(gt=0) modelId: int = Field(gt=0)
knowledgeIds: list[int] = Field(default_factory=list) knowledgeIds: list[int] = Field(default_factory=list)
knowledgeVersions: dict[int, int] = Field(default_factory=dict) knowledgeVersions: dict[int, int] = Field(default_factory=dict)
question: str = Field(min_length=1, max_length=2000) question: str = Field(min_length=1, max_length=2000)
history: list[AgentDebugHistoryMessage] = Field(default_factory=list, max_length=100)
temperature: float | None = Field(default=None, ge=0, le=2) temperature: float | None = Field(default=None, ge=0, le=2)
topP: float | None = Field(default=None, ge=0, le=1) topP: float | None = Field(default=None, ge=0, le=1)
topK: int | None = Field(default=None, ge=1, le=1000) topK: int | None = Field(default=None, ge=1, le=1000)

View File

@@ -2,6 +2,7 @@ from __future__ import annotations
import asyncio import asyncio
from collections.abc import AsyncIterator from collections.abc import AsyncIterator
from types import SimpleNamespace
from sqlalchemy.orm import Session from sqlalchemy.orm import Session
@@ -12,31 +13,33 @@ from app.services.admin_service import OperationLogService
from app.services.knowledge_agent_service import KnowledgeAgentService from app.services.knowledge_agent_service import KnowledgeAgentService
from app.services.model_stream_service import ModelStreamService from app.services.model_stream_service import ModelStreamService
from app.services.reasoning_policy_service import ReasoningPolicyService from app.services.reasoning_policy_service import ReasoningPolicyService
from app.services.rag_service import PromptService, RagResult from app.services.rag_service import RagResult
class AgentDebugService: class AgentDebugService:
@staticmethod @staticmethod
async def build_result(db: Session, payload: AgentDebugRequest) -> RagResult: async def build_result(db: Session, payload: AgentDebugRequest) -> RagResult:
history = [
SimpleNamespace(id=index + 1, role=item.role, content=item.content)
for index, item in enumerate(payload.history)
]
rag_result = await KnowledgeAgentService.build_result( rag_result = await KnowledgeAgentService.build_result(
db, db,
question=payload.question, question=payload.question,
history=history,
version_overrides=payload.knowledgeVersions, version_overrides=payload.knowledgeVersions,
preview_knowledge_ids=payload.knowledgeIds, preview_knowledge_ids=payload.knowledgeIds,
prompt_override=payload.promptContent,
) )
debug_messages = [
{"role": "system", "content": payload.promptContent.strip()},
*(rag_result.messages or []),
]
return RagResult( return RagResult(
question=rag_result.question, question=rag_result.question,
knowledge_scopes=rag_result.knowledge_scopes, knowledge_scopes=rag_result.knowledge_scopes,
chunks=rag_result.chunks, chunks=rag_result.chunks,
prompt=PromptService.render_messages(debug_messages), prompt=rag_result.prompt,
allow_general_knowledge=rag_result.allow_general_knowledge, allow_general_knowledge=rag_result.allow_general_knowledge,
retrieval_log_id=rag_result.retrieval_log_id, retrieval_log_id=rag_result.retrieval_log_id,
tool_trace=rag_result.tool_trace, tool_trace=rag_result.tool_trace,
messages=debug_messages, messages=rag_result.messages,
) )
@staticmethod @staticmethod

View File

@@ -83,6 +83,7 @@ class KnowledgeAgentService:
version_overrides: dict[int, int] | None = None, version_overrides: dict[int, int] | None = None,
preview_knowledge_ids: list[int] | None = None, preview_knowledge_ids: list[int] | None = None,
context_trace: list[dict] | None = None, context_trace: list[dict] | None = None,
prompt_override: str | None = None,
) -> RagResult: ) -> RagResult:
started = perf_counter() started = perf_counter()
catalog = cls.get_knowledge_catalog( catalog = cls.get_knowledge_catalog(
@@ -205,6 +206,7 @@ class KnowledgeAgentService:
history, history,
session_summary, session_summary,
summary_up_to_message_id, summary_up_to_message_id,
prompt_override=prompt_override,
) )
return RagResult( return RagResult(
question=question, question=question,
@@ -309,15 +311,23 @@ class KnowledgeAgentService:
text = question.strip() text = question.strip()
if len(text) <= 16: if len(text) <= 16:
return True return True
if re.search(r"(上面|前面|之前|刚才|此前|前一个|原来).{0,8}(问题|内容|回答|练习|流程)", text):
return True
if re.search(r"(再|重新|继续).{0,6}(回答|解释|说明|分析|看看)", text):
return True
return bool(re.search(r"(这个|那个|上述|上面|前面|其中|第二|继续|然后呢|怎么办|怎么做|什么意思|为什么)$", text)) return bool(re.search(r"(这个|那个|上述|上面|前面|其中|第二|继续|然后呢|怎么办|怎么做|什么意思|为什么)$", text))
@staticmethod @staticmethod
def _fallback_rewrite(question: str, history) -> str: def _fallback_rewrite(question: str, history) -> str:
previous_user = next( previous_users = [
(message.content.strip() for message in reversed(history) if message.role == "user" and message.content.strip()), message.content.strip()
"", for message in history
) if message.role == "user" and message.content.strip()
return f"关于“{previous_user}”的追问:{question.strip()}" if previous_user else question.strip() ][-3:]
if not previous_users:
return question.strip()
context = "\n".join(f"历史问题{index + 1}{content}" for index, content in enumerate(previous_users))
return f"结合以下历史问题回答当前追问:\n{context}\n当前追问:{question.strip()}"
@staticmethod @staticmethod
def get_knowledge_catalog( def get_knowledge_catalog(

View File

@@ -118,8 +118,9 @@ class PromptService:
history: list[ChatMessage] | None = None, history: list[ChatMessage] | None = None,
session_summary: str | None = None, session_summary: str | None = None,
summary_up_to_message_id: int | None = None, summary_up_to_message_id: int | None = None,
prompt_override: str | None = None,
) -> list[dict[str, str]]: ) -> list[dict[str, str]]:
prompt = cls._load_active_prompt(db) prompt = prompt_override.strip() if prompt_override and prompt_override.strip() else cls._load_active_prompt(db)
# 知识库上下文 # 知识库上下文
context = "\n\n".join(cls._format_chunk(index, chunk) for index, chunk in enumerate(chunks, start=1)) context = "\n\n".join(cls._format_chunk(index, chunk) for index, chunk in enumerate(chunks, start=1))

View File

@@ -1,7 +1,7 @@
import asyncio import asyncio
from decimal import Decimal from decimal import Decimal
import json import json
from unittest.mock import patch from unittest.mock import AsyncMock, patch
import pytest import pytest
from sqlalchemy import create_engine from sqlalchemy import create_engine
@@ -164,6 +164,37 @@ def test_debug_stream_setting_overrides_model_without_changing_it():
assert AgentDebugService.overrides(payload)["stream_enabled"] == 0 assert AgentDebugService.overrides(payload)["stream_enabled"] == 0
def test_debug_preview_passes_conversation_history_and_replaces_saved_prompt():
rag_result = RagResult(
question="再回答上面的问题",
knowledge_scopes=[],
chunks=[],
prompt="调试提示词渲染结果",
allow_general_knowledge=True,
messages=[{"role": "system", "content": "调试主提示词"}],
)
payload = AgentDebugRequest(
promptContent="调试主提示词",
modelId=1,
question="再回答上面的问题",
history=[
{"role": "user", "content": "最开始的问题"},
{"role": "assistant", "content": "第一次回答"},
],
)
with _database() as db:
build_result = AsyncMock(return_value=rag_result)
with patch("app.services.agent_debug_service.KnowledgeAgentService.build_result", build_result):
result = asyncio.run(AgentDebugService.build_result(db, payload))
kwargs = build_result.await_args.kwargs
assert [item.content for item in kwargs["history"]] == ["最开始的问题", "第一次回答"]
assert kwargs["prompt_override"] == "调试主提示词"
assert result.messages == rag_result.messages
assert result.prompt == "调试提示词渲染结果"
@pytest.mark.parametrize("reasoning_visible", [True, False]) @pytest.mark.parametrize("reasoning_visible", [True, False])
def test_agent_debug_stream_respects_reasoning_visibility(reasoning_visible): def test_agent_debug_stream_respects_reasoning_visibility(reasoning_visible):
async def chunks(): async def chunks():

View File

@@ -269,10 +269,30 @@ def test_contextual_follow_up_is_rewritten_before_agent_decision():
rewrite = next(item for item in result.tool_trace if item["tool"] == "rewrite_contextual_question") rewrite = next(item for item in result.tool_trace if item["tool"] == "rewrite_contextual_question")
decision = next(item for item in result.tool_trace if item["tool"] == "agent_decision") decision = next(item for item in result.tool_trace if item["tool"] == "agent_decision")
assert rewrite["rewrittenQuestion"] == "关于“课程退款条件有哪些?”的追问:那第二种情况呢?" assert rewrite["rewrittenQuestion"] == (
"结合以下历史问题回答当前追问:\n"
"历史问题1课程退款条件有哪些\n"
"当前追问:那第二种情况呢?"
)
assert decision["request"]["question"] == rewrite["rewrittenQuestion"] assert decision["request"]["question"] == rewrite["rewrittenQuestion"]
def test_long_follow_up_reference_is_detected_and_keeps_multiple_user_questions():
history = [
ChatMessage(id=1, session_id=1, user_id=1, role="user", content="觉式练习和慧氏练习分别是什么,区别是什么?"),
ChatMessage(id=2, session_id=1, user_id=1, role="assistant", content="暂时无法确认。"),
ChatMessage(id=3, session_id=1, user_id=1, role="user", content="阴影人格练习步骤是什么?"),
ChatMessage(id=4, session_id=1, user_id=1, role="assistant", content="先稳定状态,再进行觉察。"),
]
question = "你能根据这个练习的基本方向再来回答一下我上面的问题吗?"
assert KnowledgeAgentService._needs_contextual_rewrite(question) is True
rewritten = KnowledgeAgentService._fallback_rewrite(question, history)
assert "觉式练习和慧氏练习" in rewritten
assert "阴影人格练习步骤" in rewritten
assert question in rewritten
def test_homework_overview_expands_practice_terms_and_section_limit(): def test_homework_overview_expands_practice_terms_and_section_limit():
terms = KnowledgeAgentService._query_terms("合一的作业是什么?") terms = KnowledgeAgentService._query_terms("合一的作业是什么?")