fix: preserve agent debug conversation context
This commit is contained in:
@@ -1,6 +1,7 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime
|
||||
from typing import Literal
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
@@ -77,12 +78,18 @@ class PromptSaveRequest(BaseModel):
|
||||
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):
|
||||
promptContent: str = Field(min_length=1)
|
||||
modelId: int = Field(gt=0)
|
||||
knowledgeIds: list[int] = Field(default_factory=list)
|
||||
knowledgeVersions: dict[int, int] = Field(default_factory=dict)
|
||||
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)
|
||||
topP: float | None = Field(default=None, ge=0, le=1)
|
||||
topK: int | None = Field(default=None, ge=1, le=1000)
|
||||
|
||||
@@ -2,6 +2,7 @@ from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator
|
||||
from types import SimpleNamespace
|
||||
|
||||
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.model_stream_service import ModelStreamService
|
||||
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:
|
||||
@staticmethod
|
||||
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(
|
||||
db,
|
||||
question=payload.question,
|
||||
history=history,
|
||||
version_overrides=payload.knowledgeVersions,
|
||||
preview_knowledge_ids=payload.knowledgeIds,
|
||||
prompt_override=payload.promptContent,
|
||||
)
|
||||
debug_messages = [
|
||||
{"role": "system", "content": payload.promptContent.strip()},
|
||||
*(rag_result.messages or []),
|
||||
]
|
||||
return RagResult(
|
||||
question=rag_result.question,
|
||||
knowledge_scopes=rag_result.knowledge_scopes,
|
||||
chunks=rag_result.chunks,
|
||||
prompt=PromptService.render_messages(debug_messages),
|
||||
prompt=rag_result.prompt,
|
||||
allow_general_knowledge=rag_result.allow_general_knowledge,
|
||||
retrieval_log_id=rag_result.retrieval_log_id,
|
||||
tool_trace=rag_result.tool_trace,
|
||||
messages=debug_messages,
|
||||
messages=rag_result.messages,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -83,6 +83,7 @@ class KnowledgeAgentService:
|
||||
version_overrides: dict[int, int] | None = None,
|
||||
preview_knowledge_ids: list[int] | None = None,
|
||||
context_trace: list[dict] | None = None,
|
||||
prompt_override: str | None = None,
|
||||
) -> RagResult:
|
||||
started = perf_counter()
|
||||
catalog = cls.get_knowledge_catalog(
|
||||
@@ -205,6 +206,7 @@ class KnowledgeAgentService:
|
||||
history,
|
||||
session_summary,
|
||||
summary_up_to_message_id,
|
||||
prompt_override=prompt_override,
|
||||
)
|
||||
return RagResult(
|
||||
question=question,
|
||||
@@ -309,15 +311,23 @@ class KnowledgeAgentService:
|
||||
text = question.strip()
|
||||
if len(text) <= 16:
|
||||
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))
|
||||
|
||||
@staticmethod
|
||||
def _fallback_rewrite(question: str, history) -> str:
|
||||
previous_user = next(
|
||||
(message.content.strip() for message in reversed(history) if message.role == "user" and message.content.strip()),
|
||||
"",
|
||||
)
|
||||
return f"关于“{previous_user}”的追问:{question.strip()}" if previous_user else question.strip()
|
||||
previous_users = [
|
||||
message.content.strip()
|
||||
for message in history
|
||||
if message.role == "user" and message.content.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
|
||||
def get_knowledge_catalog(
|
||||
|
||||
@@ -118,8 +118,9 @@ class PromptService:
|
||||
history: list[ChatMessage] | None = None,
|
||||
session_summary: str | None = None,
|
||||
summary_up_to_message_id: int | None = None,
|
||||
prompt_override: str | None = None,
|
||||
) -> 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))
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
import asyncio
|
||||
from decimal import Decimal
|
||||
import json
|
||||
from unittest.mock import patch
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
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
|
||||
|
||||
|
||||
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])
|
||||
def test_agent_debug_stream_respects_reasoning_visibility(reasoning_visible):
|
||||
async def chunks():
|
||||
|
||||
@@ -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")
|
||||
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"]
|
||||
|
||||
|
||||
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():
|
||||
terms = KnowledgeAgentService._query_terms("合一的作业是什么?")
|
||||
|
||||
|
||||
Reference in New Issue
Block a user