test: verify real kbqa workflow
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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"]),
|
||||
|
||||
Reference in New Issue
Block a user