Files
agent-framework/python/packages/azure-ai-search/tests/test_aisearch_context_provider.py
T
Eduard van Valkenburg 1e350ea22f Python: [BREAKING] PR2 — Wire context provider pipeline, remove old types, update all consumers (#3850)
* PR2: Wire context provider pipeline and update all internal consumers

- Replace AgentThread with AgentSession across all packages
- Replace ContextProvider with BaseContextProvider across all packages
- Replace context_provider param with context_providers (Sequence)
- Replace thread= with session= in run() signatures
- Replace get_new_thread() with create_session()
- Add get_session(service_session_id) to agent interface
- DurableAgentThread -> DurableAgentSession
- Remove _notify_thread_of_new_messages from WorkflowAgent
- Wire before_run/after_run context provider pipeline in RawAgent
- Auto-inject InMemoryHistoryProvider when no providers configured

* fix: update all tests for context provider pipeline, fix lazy-loaders, remove old test files

* refactor: update all sample files for context provider pipeline (AgentThread→AgentSession, ContextProvider→BaseContextProvider)

* fix: update remaining ag-ui references (client docstring, getting_started sample)

* fix: make get_session service_session_id keyword-only to avoid confusion with session_id

* refactor: rename _RunContext.thread_messages to session_messages

* refactor: remove _threads.py, _memory.py, and old provider files; migrate devui to use plain message lists

* rename: remove _new_ prefix from test files

* refactor: rewrite SlidingWindowChatMessageStore as SlidingWindowHistoryProvider(InMemoryHistoryProvider)

* fix: read full history from session state directly instead of reaching into provider internals

* fix: update stale .pyi stubs, sample imports, and README references for new provider types

* fix: remove stale message_store, _notify_thread_of_new_messages, and session_id.key references in samples

* refactor: merge context_providers and sessions sample folders into sessions, remove aggregate_context_provider

* refactor: UserInfoMemory stores state in session.state instead of instance attributes

* feat: add Pydantic BaseModel support to session state serialization

Pydantic models stored in session.state are now automatically serialized
via model_dump() and restored via model_validate() during to_dict()/from_dict()
round-trips. Models are auto-registered on first serialization; use
register_state_type() for cold-start deserialization.

Also export register_state_type as a public API.

* fix mem0

* Update sample README links and descriptions for session terminology

- Replace 'thread' with 'session' in sample descriptions across all READMEs
- Update file links for renamed samples (mem0_sessions, redis_sessions, etc.)
- Fix Threads section → Sessions section in main samples/README.md
- Update tools, middleware, workflows, durabletask, azure_functions READMEs
- Update architecture diagrams in concepts/tools/README.md
- Update migration guides (autogen, semantic-kernel)

* Fix broken Redis README link to renamed sample

* Fix Mem0 OSS client search: pass scoping params as direct kwargs

AsyncMemory (OSS) expects user_id/agent_id/run_id as direct kwargs,
while AsyncMemoryClient (Platform) expects them in a filters dict.
Adds tests for both client types.

Port of fix from #3844 to new Mem0ContextProvider.

* Fix rebase issues: restore missing _conversation_state.py and checkpoint decode logic

- Add back _conversation_state.py (encode/decode_chat_messages) lost in rebase
- Fix on_checkpoint_restore to decode cache/conversation with decode_chat_messages
- Fix on_checkpoint_restore to use decode_checkpoint_value for pending requests
- Add tests/workflow/__init__.py for relative import support
- Fix test_agent_executor checkpoint selection (checkpoints[1] not superstep)

* Add STORES_BY_DEFAULT ClassVar to skip redundant InMemoryHistoryProvider injection

Chat clients that store history server-side by default (OpenAI Responses API,
Azure AI Agent) now declare STORES_BY_DEFAULT = True. The agent checks this
during auto-injection and skips InMemoryHistoryProvider unless the user
explicitly sets store=False.

* Fix broken markdown links in azure_ai and redis READMEs

* Fix getting-started samples to use session API instead of removed thread/ContextProvider API

* updates to workflow as agent

* fix group chat import

* Rename Thread→Session throughout, fix service_session_id propagation, remove stale AGUIThread

- Fix: Propagate conversation_id from ChatResponse back to session.service_session_id
  in both streaming and non-streaming paths in _agents.py
- Rename AgentThreadException → AgentSessionException
- Remove stale AGUIThread from ag_ui lazy-loader
- Rename use_service_thread → use_service_session in ag-ui package
- Rename test functions from *_thread_* to *_session_*
- Rename sample files from *_thread* to *_session*
- Update docstrings and comments: thread → session
- Update _mcp.py kwargs filter: add 'session' alongside 'thread'
- Fix ContinuationToken docstring example: thread=thread → session=session
- Fix _clients.py docstring: 'Agent threads' → 'Agent sessions'

* Fix broken markdown links after thread→session file renames

* fix azure ai test
2026-02-12 21:00:32 +00:00

294 lines
11 KiB
Python

# Copyright (c) Microsoft. All rights reserved.
# pyright: reportPrivateUsage=false
import os
from unittest.mock import AsyncMock, patch
import pytest
from agent_framework import Message
from agent_framework._sessions import AgentSession, SessionContext
from agent_framework.exceptions import ServiceInitializationError, SettingNotFoundError
from agent_framework_azure_ai_search._context_provider import AzureAISearchContextProvider
# -- Helpers -------------------------------------------------------------------
class MockSearchResults:
"""Async-iterable mock for Azure SearchClient.search() results."""
def __init__(self, docs: list[dict]):
self._docs = docs
self._index = 0
def __aiter__(self):
return self
async def __anext__(self):
if self._index >= len(self._docs):
raise StopAsyncIteration
doc = self._docs[self._index]
self._index += 1
return doc
@pytest.fixture
def mock_search_client() -> AsyncMock:
"""Create a mock SearchClient that returns one document."""
client = AsyncMock()
async def _search(**kwargs):
return MockSearchResults([{"id": "doc1", "content": "test document"}])
client.search = AsyncMock(side_effect=_search)
return client
@pytest.fixture
def mock_search_client_empty() -> AsyncMock:
"""Create a mock SearchClient that returns no results."""
client = AsyncMock()
async def _search(**kwargs):
return MockSearchResults([])
client.search = AsyncMock(side_effect=_search)
return client
def _make_provider(**overrides) -> AzureAISearchContextProvider:
"""Create a semantic-mode provider with mocked internals (skips auto-discovery)."""
defaults = {
"source_id": "aisearch",
"endpoint": "https://test.search.windows.net",
"index_name": "test-index",
"api_key": "test-key",
}
defaults.update(overrides)
provider = AzureAISearchContextProvider(**defaults)
provider._auto_discovered_vector_field = True # skip auto-discovery
return provider
# -- Initialization: semantic mode ---------------------------------------------
class TestInitSemantic:
"""Initialization tests for semantic mode."""
def test_valid_init(self) -> None:
provider = _make_provider()
assert provider.source_id == "aisearch"
assert provider.endpoint == "https://test.search.windows.net"
assert provider.index_name == "test-index"
assert provider.mode == "semantic"
def test_source_id_set(self) -> None:
provider = _make_provider(source_id="my-source")
assert provider.source_id == "my-source"
def test_missing_endpoint_raises(self) -> None:
with patch.dict(os.environ, {}, clear=True), pytest.raises(SettingNotFoundError, match="endpoint"):
AzureAISearchContextProvider(
source_id="s",
endpoint=None,
index_name="idx",
api_key="key",
)
def test_missing_index_name_semantic_raises(self) -> None:
with pytest.raises(SettingNotFoundError, match="index_name"):
AzureAISearchContextProvider(
source_id="s",
endpoint="https://test.search.windows.net",
index_name=None,
api_key="key",
)
def test_env_variable_fallback(self) -> None:
env = {
"AZURE_SEARCH_ENDPOINT": "https://env.search.windows.net",
"AZURE_SEARCH_INDEX_NAME": "env-index",
"AZURE_SEARCH_API_KEY": "env-key",
}
with patch.dict(os.environ, env, clear=False):
provider = AzureAISearchContextProvider(source_id="env-test")
assert provider.endpoint == "https://env.search.windows.net"
assert provider.index_name == "env-index"
# -- Initialization: agentic mode validation -----------------------------------
class TestInitAgenticValidation:
"""Initialization validation tests for agentic mode."""
def test_both_index_and_kb_raises(self) -> None:
with pytest.raises(SettingNotFoundError, match="multiple were set"):
AzureAISearchContextProvider(
source_id="s",
endpoint="https://test.search.windows.net",
index_name="idx",
knowledge_base_name="kb",
api_key="key",
mode="agentic",
model_deployment_name="deploy",
azure_openai_resource_url="https://aoai.openai.azure.com",
)
def test_neither_index_nor_kb_raises(self) -> None:
with pytest.raises(SettingNotFoundError, match="none was set"):
AzureAISearchContextProvider(
source_id="s",
endpoint="https://test.search.windows.net",
api_key="key",
mode="agentic",
)
def test_missing_model_deployment_name_raises(self) -> None:
with pytest.raises(ServiceInitializationError, match="model_deployment_name"):
AzureAISearchContextProvider(
source_id="s",
endpoint="https://test.search.windows.net",
index_name="idx",
api_key="key",
mode="agentic",
azure_openai_resource_url="https://aoai.openai.azure.com",
)
def test_vector_field_without_embedding_raises(self) -> None:
with pytest.raises(ValueError, match="embedding_function"):
AzureAISearchContextProvider(
source_id="s",
endpoint="https://test.search.windows.net",
index_name="idx",
api_key="key",
vector_field_name="embedding",
)
# -- before_run: semantic mode -------------------------------------------------
class TestBeforeRunSemantic:
"""Tests for before_run in semantic mode."""
async def test_results_added_to_context(self, mock_search_client: AsyncMock) -> None:
provider = _make_provider()
provider._search_client = mock_search_client
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[Message(role="user", contents=["test query"])],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_search_client.search.assert_awaited_once()
msgs = ctx.context_messages.get("aisearch", [])
assert len(msgs) >= 2 # context_prompt + at least one result
assert msgs[0].text == provider.context_prompt
async def test_empty_input_no_search(self, mock_search_client: AsyncMock) -> None:
provider = _make_provider()
provider._search_client = mock_search_client
session = AgentSession(session_id="test-session")
ctx = SessionContext(input_messages=[], session_id="s1")
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_search_client.search.assert_not_awaited()
assert ctx.context_messages.get("aisearch") is None
async def test_no_results_no_messages(self, mock_search_client_empty: AsyncMock) -> None:
provider = _make_provider()
provider._search_client = mock_search_client_empty
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[Message(role="user", contents=["test query"])],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_search_client_empty.search.assert_awaited_once()
assert ctx.context_messages.get("aisearch") is None
async def test_context_prompt_prepended(self, mock_search_client: AsyncMock) -> None:
custom_prompt = "Custom search context:"
provider = _make_provider(context_prompt=custom_prompt)
provider._search_client = mock_search_client
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[Message(role="user", contents=["test query"])],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
msgs = ctx.context_messages["aisearch"]
assert msgs[0].text == custom_prompt
# -- before_run: message filtering ---------------------------------------------
class TestBeforeRunFiltering:
"""Tests that only user/assistant messages are used for search."""
async def test_filters_non_user_assistant(self, mock_search_client: AsyncMock) -> None:
provider = _make_provider()
provider._search_client = mock_search_client
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[
Message(role="system", contents=["system prompt"]),
Message(role="user", contents=["actual question"]),
],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_search_client.search.assert_awaited_once()
call_kwargs = mock_search_client.search.call_args[1]
# The search text should contain only the user message, not the system message
assert "actual question" in call_kwargs["search_text"]
assert "system prompt" not in call_kwargs["search_text"]
async def test_only_system_messages_no_search(self, mock_search_client: AsyncMock) -> None:
provider = _make_provider()
provider._search_client = mock_search_client
session = AgentSession(session_id="test-session")
ctx = SessionContext(
input_messages=[Message(role="system", contents=["system prompt"])],
session_id="s1",
)
await provider.before_run(agent=None, session=session, context=ctx, state=session.state) # type: ignore[arg-type]
mock_search_client.search.assert_not_awaited()
# -- __aexit__ -----------------------------------------------------------------
class TestAexit:
"""Tests for async context manager cleanup."""
async def test_closes_retrieval_client(self) -> None:
provider = _make_provider()
mock_retrieval = AsyncMock()
provider._retrieval_client = mock_retrieval
await provider.__aexit__(None, None, None)
mock_retrieval.close.assert_awaited_once()
assert provider._retrieval_client is None
async def test_no_retrieval_client_no_error(self) -> None:
provider = _make_provider()
assert provider._retrieval_client is None
await provider.__aexit__(None, None, None) # should not raise