Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -63,7 +63,7 @@
KnowledgeBaseRetrievalResponse,
KnowledgeRetrievalIntent,
KnowledgeRetrievalSemanticIntent,
SearchIndexKnowledgeSourceParams,
KnowledgeSourceParams,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalMinimalReasoningEffort as KBRetrievalMinimalReasoningEffort,
Expand All @@ -89,7 +89,7 @@
KnowledgeBaseRetrievalResponse,
KnowledgeRetrievalIntent,
KnowledgeRetrievalSemanticIntent,
SearchIndexKnowledgeSourceParams,
KnowledgeSourceParams,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalMinimalReasoningEffort as KBRetrievalMinimalReasoningEffort,
Expand Down Expand Up @@ -602,7 +602,7 @@ def __init__(
)

self._knowledge_base_initialized = False
self._knowledge_source_names: list[str] = []
self._knowledge_source_params: list[KnowledgeSourceParams] = []

def _common_client_kwargs(self) -> dict[str, Any]:
"""Build the keyword arguments shared by every Azure AI Search client.
Expand Down Expand Up @@ -814,13 +814,22 @@ async def _ensure_knowledge_base(self) -> None:
credential=self.credential,
**self._common_client_kwargs(),
)
# Resolve the existing KB's real knowledge source names so agentic
# retrieval can request reference source data per source. Without
# this, source names were left unset ("None-source").
self._knowledge_source_names = []
# The KB references expose names but not source kinds, while runtime
# parameter kinds must match the source definitions.
knowledge_source_params: list[KnowledgeSourceParams] = []
if self._index_client is not None:
kb = await self._index_client.get_knowledge_base(knowledge_base_name)
self._knowledge_source_names = [ks.name for ks in (kb.knowledge_sources or [])]
for knowledge_source_reference in kb.knowledge_sources or []:
knowledge_source = await self._index_client.get_knowledge_source(knowledge_source_reference.name)
# The base type preserves kinds unknown to this SDK instead of dropping or mismatching them.
knowledge_source_params.append(
KnowledgeSourceParams(
knowledge_source_name=knowledge_source_reference.name,
include_reference_source_data=True,
kind=knowledge_source.kind,
)
)
self._knowledge_source_params = knowledge_source_params
self._knowledge_base_initialized = True
return

Expand All @@ -834,7 +843,13 @@ async def _ensure_knowledge_base(self) -> None:
raise ValueError("index_name is required when creating Knowledge Base from index")

knowledge_source_name = f"{self.index_name}-source"
self._knowledge_source_names = [knowledge_source_name]
self._knowledge_source_params = [
KnowledgeSourceParams(
knowledge_source_name=knowledge_source_name,
include_reference_source_data=True,
kind="searchIndex",
)
]
try:
await self._index_client.get_knowledge_source(knowledge_source_name)
except ResourceNotFoundError:
Expand Down Expand Up @@ -935,14 +950,8 @@ async def _agentic_search(self, messages: list[Message]) -> list[Message]:

# Request reference source data per knowledge source so ref.source_data
# is populated when the source has source_data_fields configured (#5095).
if self._knowledge_source_names:
request_kwargs["knowledge_source_params"] = [
SearchIndexKnowledgeSourceParams(
knowledge_source_name=name,
include_reference_source_data=True,
)
for name in self._knowledge_source_names
]
if self._knowledge_source_params:
request_kwargs["knowledge_source_params"] = self._knowledge_source_params

retrieval_request = KnowledgeBaseRetrievalRequest(**request_kwargs)

Expand Down
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
# Copyright (c) Microsoft. All rights reserved.
# pyright: reportPrivateUsage=false

import asyncio
import os
from collections.abc import Iterator
from contextlib import contextmanager
Expand Down Expand Up @@ -1205,6 +1206,27 @@ async def test_existing_kb_sets_initialized(self) -> None:
await provider._ensure_knowledge_base()
assert provider._knowledge_base_initialized is True

async def test_close_then_existing_kb_reuse_clears_source_params(self) -> None:
provider = _make_provider(mode="agentic", index_name=None, knowledge_base_name="old-kb")
provider._knowledge_source_params = [
_context_provider.KnowledgeSourceParams(
knowledge_source_name="old-source",
include_reference_source_data=True,
kind="searchIndex",
)
]

provider._index_client = AsyncMock()
provider._retrieval_client = AsyncMock()
await provider.close()
provider.knowledge_base_name = "new-kb"

with patch("agent_framework_azure_ai_search._context_provider.KnowledgeBaseRetrievalClient") as mock_cls:
mock_cls.return_value = AsyncMock()
await provider._ensure_knowledge_base()

assert provider._knowledge_source_params == []

async def test_missing_index_client_raises(self) -> None:
provider = _make_provider()
provider._knowledge_base_initialized = False
Expand Down Expand Up @@ -1547,7 +1569,13 @@ async def test_requests_reference_source_data_per_source(self) -> None:
provider._knowledge_base_initialized = True
provider.knowledge_base_name = "kb"
provider.retrieval_reasoning_effort = "minimal"
provider._knowledge_source_names = ["test-index-source"]
provider._knowledge_source_params = [
_context_provider.KnowledgeSourceParams(
knowledge_source_name="test-index-source",
include_reference_source_data=True,
kind="searchIndex",
)
]

mock_result = Mock()
mock_result.response = []
Expand All @@ -1568,13 +1596,13 @@ async def test_requests_reference_source_data_per_source(self) -> None:
assert params[0].include_reference_source_data is True

async def test_no_source_params_when_no_sources(self) -> None:
# When no knowledge source names are resolved, the request must not
# When no knowledge source parameters are resolved, the request must not
# carry an empty/placeholder knowledge_source_params list.
provider = _make_provider()
provider._knowledge_base_initialized = True
provider.knowledge_base_name = "kb"
provider.retrieval_reasoning_effort = "minimal"
provider._knowledge_source_names = []
provider._knowledge_source_params = []

mock_result = Mock()
mock_result.response = []
Expand All @@ -1589,6 +1617,68 @@ async def test_no_source_params_when_no_sources(self) -> None:
request = mock_retrieval.retrieve.call_args.kwargs["retrieval_request"]
assert request.knowledge_source_params is None

async def test_existing_mixed_source_kb_preserves_source_kinds(self) -> None:
provider = _make_provider(mode="agentic", index_name=None, knowledge_base_name="kb")
provider.retrieval_reasoning_effort = "minimal"

provider._index_client = AsyncMock()
provider._index_client.get_knowledge_base.return_value = SimpleNamespace(
knowledge_sources=[
SimpleNamespace(name="web-source"),
SimpleNamespace(name="index-source"),
SimpleNamespace(name="unknown-source"),
]
)
provider._index_client.get_knowledge_source.side_effect = [
SimpleNamespace(kind="web"),
SimpleNamespace(kind="searchIndex"),
SimpleNamespace(kind="futureKind"),
]

mock_result = Mock(response=[], references=None)
provider._retrieval_client = AsyncMock()
provider._retrieval_client.retrieve.return_value = mock_result

await provider._agentic_search([Message(role="user", contents=["query"])])

request = provider._retrieval_client.retrieve.call_args.kwargs["retrieval_request"]
assert [(param.knowledge_source_name, param.kind) for param in request.knowledge_source_params] == [
("web-source", "web"),
("index-source", "searchIndex"),
("unknown-source", "futureKind"),
]

async def test_concurrent_existing_kb_initialization_publishes_complete_source_params(self) -> None:
provider = _make_provider(mode="agentic", index_name=None, knowledge_base_name="kb")
provider._retrieval_client = AsyncMock()

both_knowledge_base_reads_started = asyncio.Event()
knowledge_base_read_count = 0

async def get_knowledge_base(_: str) -> SimpleNamespace:
nonlocal knowledge_base_read_count
knowledge_base_read_count += 1
if knowledge_base_read_count == 2:
both_knowledge_base_reads_started.set()
await asyncio.wait_for(both_knowledge_base_reads_started.wait(), timeout=1)
return SimpleNamespace(
knowledge_sources=[SimpleNamespace(name="web-source"), SimpleNamespace(name="index-source")]
)

async def get_knowledge_source(name: str) -> SimpleNamespace:
return SimpleNamespace(kind="web" if name == "web-source" else "searchIndex")

provider._index_client = AsyncMock()
provider._index_client.get_knowledge_base.side_effect = get_knowledge_base
provider._index_client.get_knowledge_source.side_effect = get_knowledge_source

await asyncio.gather(provider._ensure_knowledge_base(), provider._ensure_knowledge_base())

assert [(param.knowledge_source_name, param.kind) for param in provider._knowledge_source_params] == [
("web-source", "web"),
("index-source", "searchIndex"),
]


# -- before_run: agentic mode --------------------------------------------------

Expand Down
Loading