test: verify real kbqa workflow

This commit is contained in:
gqt
2026-07-13 14:23:06 +08:00
parent 52184aa00e
commit 2361d62867
32 changed files with 656 additions and 12 deletions
+1
View File
@@ -26,6 +26,7 @@ def create_agent_bundle(settings: Settings) -> AgentBundle:
system_prompt=RESEARCH_PROMPT,
context_schema=ResearchContext,
middleware=[
RequireResearchMiddleware(tool_name="search_knowledge_base"),
ModelCallLimitMiddleware(run_limit=4, exit_behavior="error"),
ToolCallLimitMiddleware(
tool_name="search_knowledge_base",
+1
View File
@@ -9,6 +9,7 @@ def create_chat_model(settings: Settings) -> ChatOpenAI:
api_key=settings.dashscope_api_key,
base_url=settings.dashscope_base_url,
model=settings.chat_model,
extra_body={"enable_thinking": False},
temperature=0.2,
streaming=True,
timeout=120.0,
+7 -3
View File
@@ -5,10 +5,13 @@ from langchain.messages import HumanMessage, ToolMessage
class RequireResearchMiddleware(AgentMiddleware):
def __init__(self, tool_name: str = "research") -> None:
self.tool_name = tool_name
@staticmethod
def _has_current_research(messages: list) -> bool:
def _has_current_tool(messages: list, tool_name: str) -> bool:
for message in reversed(messages):
if isinstance(message, ToolMessage) and message.name == "research":
if isinstance(message, ToolMessage) and message.name == tool_name:
return True
if isinstance(message, HumanMessage):
return False
@@ -19,5 +22,6 @@ class RequireResearchMiddleware(AgentMiddleware):
request: ModelRequest,
handler: Callable[[ModelRequest], Awaitable[ModelResponse]],
) -> ModelResponse:
tool_choice = "none" if self._has_current_research(request.messages) else "required"
has_tool_result = self._has_current_tool(request.messages, self.tool_name)
tool_choice = "none" if has_tool_result else "required"
return await handler(request.override(tool_choice=tool_choice))
+4 -2
View File
@@ -14,6 +14,7 @@ class MilvusVectorStore:
def ensure_collection(self) -> None:
if self.client.has_collection(self.collection_name):
self.client.load_collection(self.collection_name)
return
schema = self.client.create_schema(auto_id=False, enable_dynamic_field=False)
schema.add_field(
@@ -44,6 +45,7 @@ class MilvusVectorStore:
schema=schema,
index_params=index_params,
)
self.client.load_collection(self.collection_name)
def upsert(self, chunks: list[PreparedChunk], vectors: list[list[float]]) -> None:
if len(chunks) != len(vectors):
@@ -67,7 +69,7 @@ class MilvusVectorStore:
data=[vector],
anns_field="vector",
limit=limit,
output_fields=["document_id", "order_index"],
output_fields=["chunk_id", "document_id", "order_index"],
search_params={"metric_type": "COSINE"},
)
hits: list[VectorHit] = []
@@ -75,7 +77,7 @@ class MilvusVectorStore:
entity = item.get("entity", {})
hits.append(
VectorHit(
chunk_id=str(item["id"]),
chunk_id=str(entity["chunk_id"]),
document_id=str(entity["document_id"]),
order_index=int(entity["order_index"]),
score=float(item["distance"]),