Skip to content
Closed
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
1 change: 1 addition & 0 deletions python/packages/azure-ai-search/README.md
Original file line number Diff line number Diff line change
Expand Up @@ -19,5 +19,6 @@ See the [Azure AI Search context provider examples](../../samples/02-agents/cont

- Semantic search with hybrid (vector + keyword) queries
- Agentic mode with Knowledge Bases for complex multi-hop reasoning
- Per-source retrieval parameters (e.g. OData filters) in agentic mode
- Environment variable configuration with Settings class
- API key and managed identity authentication
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,7 @@
KnowledgeBaseRetrievalResponse,
KnowledgeRetrievalIntent,
KnowledgeRetrievalSemanticIntent,
KnowledgeSourceParams,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalLowReasoningEffort as KBRetrievalLowReasoningEffort,
Expand Down Expand Up @@ -239,6 +240,7 @@ def __init__(
azure_openai_api_key: str | None = None,
knowledge_base_output_mode: KnowledgeBaseOutputModeLiteral = "extractive_data",
retrieval_reasoning_effort: RetrievalReasoningEffortLiteral = "minimal",
knowledge_source_params: list[KnowledgeSourceParams] | None = None,
agentic_message_history_count: int = _DEFAULT_AGENTIC_MESSAGE_HISTORY_COUNT,
env_file_path: str | None = None,
env_file_encoding: str | None = None,
Expand All @@ -264,6 +266,8 @@ def __init__(
azure_openai_api_key: Optional Azure OpenAI API key for Knowledge Base creation.
knowledge_base_output_mode: Output mode for Knowledge Base retrieval.
retrieval_reasoning_effort: Reasoning effort for query planning.
knowledge_source_params: Optional per-source retrieval parameters (for example an
OData ``filter_add_on``) forwarded to the agentic retrieval request.
agentic_message_history_count: Number of recent messages included in retrieval.
env_file_path: Optional ``.env`` file checked before process environment variables.
env_file_encoding: Encoding for the ``.env`` file.
Expand Down Expand Up @@ -292,6 +296,7 @@ def __init__(
azure_openai_api_key: str | None = None,
knowledge_base_output_mode: KnowledgeBaseOutputModeLiteral = "extractive_data",
retrieval_reasoning_effort: RetrievalReasoningEffortLiteral = "minimal",
knowledge_source_params: list[KnowledgeSourceParams] | None = None,
agentic_message_history_count: int = _DEFAULT_AGENTIC_MESSAGE_HISTORY_COUNT,
env_file_path: str | None = None,
env_file_encoding: str | None = None,
Expand All @@ -317,6 +322,8 @@ def __init__(
azure_openai_api_key: Unused when connecting to an existing Knowledge Base.
knowledge_base_output_mode: Output mode for Knowledge Base retrieval.
retrieval_reasoning_effort: Reasoning effort for query planning.
knowledge_source_params: Optional per-source retrieval parameters (for example an
OData ``filter_add_on``) forwarded to the agentic retrieval request.
agentic_message_history_count: Number of recent messages included in retrieval.
env_file_path: Optional ``.env`` file checked before process environment variables.
env_file_encoding: Encoding for the ``.env`` file.
Expand Down Expand Up @@ -345,6 +352,7 @@ def __init__(
azure_openai_api_key: str | None = None,
knowledge_base_output_mode: KnowledgeBaseOutputModeLiteral = "extractive_data",
retrieval_reasoning_effort: RetrievalReasoningEffortLiteral = "minimal",
knowledge_source_params: list[KnowledgeSourceParams] | None = None,
agentic_message_history_count: int = _DEFAULT_AGENTIC_MESSAGE_HISTORY_COUNT,
env_file_path: str | None = None,
env_file_encoding: str | None = None,
Expand Down Expand Up @@ -374,6 +382,8 @@ def __init__(
azure_openai_api_key: Optional Azure OpenAI API key for Knowledge Base creation.
knowledge_base_output_mode: Output mode for Knowledge Base retrieval.
retrieval_reasoning_effort: Reasoning effort for query planning.
knowledge_source_params: Optional per-source retrieval parameters (for example an
OData ``filter_add_on``) forwarded to the agentic retrieval request.
agentic_message_history_count: Number of recent messages included in retrieval.
env_file_path: Optional ``.env`` file checked before process environment variables.
env_file_encoding: Encoding for the ``.env`` file.
Expand Down Expand Up @@ -401,6 +411,7 @@ def __init__(
azure_openai_api_key: str | None = None,
knowledge_base_output_mode: KnowledgeBaseOutputModeLiteral = "extractive_data",
retrieval_reasoning_effort: RetrievalReasoningEffortLiteral = "minimal",
knowledge_source_params: list[KnowledgeSourceParams] | None = None,
agentic_message_history_count: int = _DEFAULT_AGENTIC_MESSAGE_HISTORY_COUNT,
env_file_path: str | None = None,
env_file_encoding: str | None = None,
Expand Down Expand Up @@ -429,6 +440,8 @@ def __init__(
azure_openai_api_key: Azure OpenAI API key.
knowledge_base_output_mode: Output mode for Knowledge Base retrieval.
retrieval_reasoning_effort: Reasoning effort for Knowledge Base query planning.
knowledge_source_params: Optional per-source retrieval parameters (for example an
OData ``filter_add_on``) forwarded to the agentic retrieval request.
agentic_message_history_count: Number of recent messages for agentic mode.
env_file_path: Path to environment file for loading settings.
env_file_encoding: Encoding of the environment file.
Expand Down Expand Up @@ -474,6 +487,9 @@ def __init__(
if mode == "agentic" and settings.get("index_name") and not model:
raise ValueError("model is required for agentic mode when creating Knowledge Base from index.")

if knowledge_source_params is not None and mode != "agentic":
raise ValueError("knowledge_source_params is only supported in agentic mode.")

resolved_credential: AzureKeyCredential | AsyncTokenCredential
if credential:
resolved_credential = credential # type: ignore[assignment]
Expand Down Expand Up @@ -505,6 +521,7 @@ def __init__(
self.knowledge_base_output_mode = knowledge_base_output_mode
self.retrieval_reasoning_effort = retrieval_reasoning_effort
self.agentic_message_history_count = agentic_message_history_count
self._knowledge_source_params = knowledge_source_params

self._use_existing_knowledge_base = False
if mode == "agentic":
Expand Down Expand Up @@ -830,6 +847,7 @@ async def _agentic_search(self, messages: list[Message]) -> list[Message]:
retrieval_reasoning_effort=reasoning_effort,
output_mode=output_mode,
include_activity=True,
knowledge_source_params=self._knowledge_source_params,

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Should this preserve the source-data params that are now built on main before forwarding caller params? Current main resolves _knowledge_source_names and sends include_reference_source_data=True; this branch sends None on the default path, so existing agentic retrieval loses ref.source_data again after the conflict is resolved this way. Could we merge caller-supplied fields with the resolved-source defaults instead of replacing them?

)
else:
kb_messages = self._prepare_messages_for_kb_search(messages)
Expand All @@ -838,6 +856,7 @@ async def _agentic_search(self, messages: list[Message]) -> list[Message]:
retrieval_reasoning_effort=reasoning_effort,
output_mode=output_mode,
include_activity=True,
knowledge_source_params=self._knowledge_source_params,
)

if not self._retrieval_client:
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -3,14 +3,15 @@

import os
from types import SimpleNamespace
from typing import Any, cast
from typing import Any, Literal, cast
from unittest.mock import AsyncMock, Mock, patch

import pytest
from agent_framework import Content, Message
from agent_framework._sessions import AgentSession, SessionContext
from agent_framework.exceptions import SettingNotFoundError
from azure.core.credentials import AzureKeyCredential
from azure.search.documents.knowledgebases.models import KnowledgeSourceParams, SearchIndexKnowledgeSourceParams

from agent_framework_azure_ai_search._context_provider import AzureAISearchContextProvider

Expand Down Expand Up @@ -315,6 +316,17 @@ def test_agentic_explicit_index_ignores_env_kb_name(self) -> None:
assert provider.knowledge_base_name == "idx-kb"
assert provider._use_existing_knowledge_base is False

def test_knowledge_source_params_in_semantic_mode_raises(self) -> None:
with pytest.raises(ValueError, match="agentic mode"):
cast(Any, AzureAISearchContextProvider)(
source_id="s",
endpoint="https://test.search.windows.net",
index_name="idx",
api_key="key",
mode="semantic",
knowledge_source_params=[SearchIndexKnowledgeSourceParams(knowledge_source_name="src")],
)


# -- __aenter__ / __aexit__ ---------------------------------------------------

Expand Down Expand Up @@ -1335,6 +1347,47 @@ async def test_none_response_returns_default_message(self) -> None:
assert len(results) == 1
assert results[0].text == "No results found from Knowledge Base."

@pytest.mark.parametrize("effort", ["minimal", "medium"])
async def test_knowledge_source_params_reach_request(self, effort: Literal["minimal", "medium"]) -> None:
provider = _make_provider()
provider._knowledge_base_initialized = True
provider.knowledge_base_name = "kb"
provider.retrieval_reasoning_effort = effort
params: list[KnowledgeSourceParams] = [
SearchIndexKnowledgeSourceParams(knowledge_source_name="src", filter_add_on="category eq 'public'")
]
provider._knowledge_source_params = params

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

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

request = mock_retrieval.retrieve.call_args.kwargs["retrieval_request"]
assert request.knowledge_source_params is params

async def test_default_sends_no_knowledge_source_params(self) -> None:
provider = _make_provider()
provider._knowledge_base_initialized = True
provider.knowledge_base_name = "kb"
provider.retrieval_reasoning_effort = "minimal"

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

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

request = mock_retrieval.retrieve.call_args.kwargs["retrieval_request"]
assert request.knowledge_source_params is None


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

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -85,6 +85,13 @@ async def main() -> None:
# Optional: Configure retrieval behavior
knowledge_base_output_mode="extractive_data", # or "answer_synthesis"
retrieval_reasoning_effort="minimal", # or "medium", "low"
# Optional: per-source params, e.g. an OData filter for multi-tenant isolation
# (import SearchIndexKnowledgeSourceParams from azure.search.documents.knowledgebases.models):
# knowledge_source_params=[
# SearchIndexKnowledgeSourceParams(
# knowledge_source_name="my-source", filter_add_on="category eq 'public'"
# )
# ],
)
else:
# Auto-create Knowledge Base from index
Expand Down