mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Python: Enhance Azure AI Search Citations with Document URLs in Foundry V2 (#4028)
* Python: Enhance Azure AI Search citations with document URLs in Foundry V2 (Responses API) Override _parse_response_from_openai and _parse_chunk_from_openai in RawAzureAIClient to extract get_urls from azure_ai_search_call_output items and enrich url_citation annotations with document-specific URLs. - Non-streaming: first pass collects get_urls, post-processes annotations - Streaming: captures search output state, enriches url_citation events (also handles url_citation annotation type not handled by base class) - Updated V2 sample to demonstrate citation URL extraction - Added 14 unit tests covering extraction, enrichment, and edge cases Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * refactor: rework search citation enrichment to override _inner_get_response - Remove all direct openai/pydantic imports from _client.py - Override _inner_get_response instead of _parse_response_from_openai/_parse_chunk_from_openai - Use closure-local state for streaming instead of instance-level _streaming_search_get_urls - Add _build_url_citation_content helper for streaming url_citation handling - Fix mypy errors by using str(value or '') for Annotation TypedDict fields - Fix docstring to say 'citation' instead of 'url_citation' - Update tests to match new approach Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * fix: handle streaming search citations from output_item.done events The azure_ai_search_call_output item only has populated output data (including get_urls) in the response.output_item.done event, not in the response.output_item.added event. Also removed the search_get_urls guard on url_citation handling so annotations are always produced even if get_urls haven't been captured yet. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * addressed comments * refactor: address PR review - eliminate type: ignore[assignment] pattern Call super()._inner_get_response() independently in each branch instead of once at the top with union type reassignment. Non-streaming uses two-arg super() in the closure; streaming uses cast() for type narrowing. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * refactor: remove defensive patterns per PR review - Replace all getattr() with direct attribute access - Remove cast() for streaming branch, use type: ignore[assignment] - Simplify _build_url_citation_content to use dict access directly - Simplify _extract_azure_search_urls to use item.type/item.output - Handle empty list output from streaming 'added' events - Update tests to match actual runtime types (objects, not dicts) Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> * mypy fix * small fixes --------- Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
@@ -13,14 +13,18 @@ import pytest
|
||||
from agent_framework import (
|
||||
Agent,
|
||||
AgentResponse,
|
||||
Annotation,
|
||||
ChatOptions,
|
||||
ChatResponse,
|
||||
ChatResponseUpdate,
|
||||
Content,
|
||||
Message,
|
||||
ResponseStream,
|
||||
SupportsChatGetResponse,
|
||||
tool,
|
||||
)
|
||||
from agent_framework._settings import load_settings
|
||||
from agent_framework.openai._responses_client import RawOpenAIResponsesClient
|
||||
from azure.ai.projects.aio import AIProjectClient
|
||||
from azure.ai.projects.models import (
|
||||
ApproximateLocation,
|
||||
@@ -1774,3 +1778,370 @@ def test_get_image_generation_tool_with_options() -> None:
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
|
||||
# region Azure AI Search Citation Enhancement Tests
|
||||
|
||||
|
||||
def test_extract_azure_search_urls_with_dict_items(mock_project_client: MagicMock) -> None:
|
||||
"""Test _extract_azure_search_urls with dict-style output (after JSON parsing)."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
mock_output = {
|
||||
"documents": [{"id": "1", "url": "https://search.example.com/"}],
|
||||
"get_urls": [
|
||||
"https://search.example.com/indexes/idx/docs/1?api-version=2024-07-01",
|
||||
"https://search.example.com/indexes/idx/docs/2?api-version=2024-07-01",
|
||||
],
|
||||
}
|
||||
mock_search_item = MagicMock()
|
||||
mock_search_item.type = "azure_ai_search_call_output"
|
||||
mock_search_item.output = mock_output
|
||||
|
||||
mock_call_item = MagicMock()
|
||||
mock_call_item.type = "azure_ai_search_call"
|
||||
|
||||
mock_msg_item = MagicMock()
|
||||
mock_msg_item.type = "message"
|
||||
|
||||
urls = client._extract_azure_search_urls([mock_call_item, mock_search_item, mock_msg_item])
|
||||
assert len(urls) == 2
|
||||
assert urls[0] == "https://search.example.com/indexes/idx/docs/1?api-version=2024-07-01"
|
||||
assert urls[1] == "https://search.example.com/indexes/idx/docs/2?api-version=2024-07-01"
|
||||
|
||||
|
||||
def test_extract_azure_search_urls_with_object_items(mock_project_client: MagicMock) -> None:
|
||||
"""Test _extract_azure_search_urls with object-style output items."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
mock_output = MagicMock()
|
||||
mock_output.get_urls = ["https://example.com/doc/1", "https://example.com/doc/2"]
|
||||
mock_item = MagicMock()
|
||||
mock_item.type = "azure_ai_search_call_output"
|
||||
mock_item.output = mock_output
|
||||
|
||||
urls = client._extract_azure_search_urls([mock_item])
|
||||
assert urls == ["https://example.com/doc/1", "https://example.com/doc/2"]
|
||||
|
||||
|
||||
def test_extract_azure_search_urls_no_search_items(mock_project_client: MagicMock) -> None:
|
||||
"""Test _extract_azure_search_urls with no search output items."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
mock_item = MagicMock()
|
||||
mock_item.type = "message"
|
||||
urls = client._extract_azure_search_urls([mock_item])
|
||||
assert urls == []
|
||||
|
||||
|
||||
def test_extract_azure_search_urls_with_json_string_output(mock_project_client: MagicMock) -> None:
|
||||
"""Test _extract_azure_search_urls with JSON string output (non-streaming pydantic extra field)."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
json_output = json.dumps({
|
||||
"documents": [{"id": "1"}],
|
||||
"get_urls": [
|
||||
"https://search.example.com/indexes/idx/docs/1?api-version=2024-07-01",
|
||||
],
|
||||
})
|
||||
mock_item = MagicMock()
|
||||
mock_item.type = "azure_ai_search_call_output"
|
||||
mock_item.output = json_output
|
||||
|
||||
urls = client._extract_azure_search_urls([mock_item])
|
||||
assert len(urls) == 1
|
||||
assert urls[0] == "https://search.example.com/indexes/idx/docs/1?api-version=2024-07-01"
|
||||
|
||||
|
||||
def test_get_search_doc_url_valid(mock_project_client: MagicMock) -> None:
|
||||
"""Test _get_search_doc_url with valid doc_N title."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
get_urls = ["https://example.com/doc/0", "https://example.com/doc/1", "https://example.com/doc/2"]
|
||||
|
||||
assert client._get_search_doc_url("doc_0", get_urls) == "https://example.com/doc/0"
|
||||
assert client._get_search_doc_url("doc_1", get_urls) == "https://example.com/doc/1"
|
||||
assert client._get_search_doc_url("doc_2", get_urls) == "https://example.com/doc/2"
|
||||
|
||||
|
||||
def test_get_search_doc_url_out_of_range(mock_project_client: MagicMock) -> None:
|
||||
"""Test _get_search_doc_url with out-of-range index."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
get_urls = ["https://example.com/doc/0"]
|
||||
assert client._get_search_doc_url("doc_5", get_urls) is None
|
||||
|
||||
|
||||
def test_get_search_doc_url_no_match(mock_project_client: MagicMock) -> None:
|
||||
"""Test _get_search_doc_url with non-matching title."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
get_urls = ["https://example.com/doc/0"]
|
||||
assert client._get_search_doc_url("some_title", get_urls) is None
|
||||
assert client._get_search_doc_url(None, get_urls) is None
|
||||
assert client._get_search_doc_url("doc_0", []) is None
|
||||
|
||||
|
||||
def test_enrich_annotations_with_search_urls(mock_project_client: MagicMock) -> None:
|
||||
"""Test _enrich_annotations_with_search_urls enriches citation annotations."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
get_urls = [
|
||||
"https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01",
|
||||
"https://search.example.com/indexes/idx/docs/41?api-version=2024-07-01",
|
||||
]
|
||||
|
||||
content = Content.from_text(text="test response")
|
||||
content.annotations = [
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "doc_0",
|
||||
"url": "https://search.example.com/",
|
||||
},
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "doc_1",
|
||||
"url": "https://search.example.com/",
|
||||
},
|
||||
]
|
||||
|
||||
client._enrich_annotations_with_search_urls([content], get_urls)
|
||||
|
||||
assert content.annotations[0]["additional_properties"]["get_url"] == get_urls[0]
|
||||
assert content.annotations[1]["additional_properties"]["get_url"] == get_urls[1]
|
||||
|
||||
|
||||
def test_enrich_annotations_no_match(mock_project_client: MagicMock) -> None:
|
||||
"""Test _enrich_annotations_with_search_urls with non-matching titles."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
get_urls = ["https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01"]
|
||||
|
||||
content = Content.from_text(text="test response")
|
||||
content.annotations = [
|
||||
{
|
||||
"type": "citation",
|
||||
"title": "some_title",
|
||||
"url": "https://search.example.com/",
|
||||
},
|
||||
]
|
||||
|
||||
client._enrich_annotations_with_search_urls([content], get_urls)
|
||||
assert "additional_properties" not in content.annotations[0] or "get_url" not in content.annotations[0].get(
|
||||
"additional_properties", {}
|
||||
)
|
||||
|
||||
|
||||
def test_enrich_annotations_empty_get_urls(mock_project_client: MagicMock) -> None:
|
||||
"""Test _enrich_annotations_with_search_urls with empty get_urls."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
content = Content.from_text(text="test")
|
||||
content.annotations = [{"type": "citation", "title": "doc_0", "url": "https://example.com/"}]
|
||||
|
||||
# Should not raise or modify
|
||||
client._enrich_annotations_with_search_urls([content], [])
|
||||
assert "additional_properties" not in content.annotations[0]
|
||||
|
||||
|
||||
async def test_inner_get_response_enriches_non_streaming(mock_project_client: MagicMock) -> None:
|
||||
"""Test _inner_get_response enriches url_citation annotations for non-streaming responses."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
|
||||
# Build a ChatResponse with citation annotations and a raw_representation carrying search output
|
||||
content = Content.from_text(text="Here is the result【5:0†source】.")
|
||||
content.annotations = [
|
||||
Annotation(type="citation", title="doc_0", url="https://search.example.com/"),
|
||||
]
|
||||
msg = Message(role="assistant", contents=[content])
|
||||
mock_raw = MagicMock()
|
||||
mock_search_output = MagicMock()
|
||||
mock_search_output.type = "azure_ai_search_call_output"
|
||||
mock_search_output_data = MagicMock()
|
||||
mock_search_output_data.get_urls = [
|
||||
"https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01",
|
||||
]
|
||||
mock_search_output.output = mock_search_output_data
|
||||
mock_raw.output = [mock_search_output]
|
||||
|
||||
base_response = ChatResponse(messages=[msg], raw_representation=mock_raw)
|
||||
|
||||
async def _fake_awaitable() -> ChatResponse:
|
||||
return base_response
|
||||
|
||||
with patch.object(RawOpenAIResponsesClient, "_inner_get_response", return_value=_fake_awaitable()):
|
||||
result_awaitable = client._inner_get_response(messages=[], options={}, stream=False)
|
||||
result = await result_awaitable # type: ignore[misc]
|
||||
|
||||
ann = result.messages[0].contents[0].annotations[0]
|
||||
assert ann["additional_properties"]["get_url"] == (
|
||||
"https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01"
|
||||
)
|
||||
|
||||
|
||||
async def test_inner_get_response_no_search_output_non_streaming(mock_project_client: MagicMock) -> None:
|
||||
"""Test _inner_get_response passes through when no search output exists."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
|
||||
content = Content.from_text(text="Hello world")
|
||||
msg = Message(role="assistant", contents=[content])
|
||||
mock_raw = MagicMock()
|
||||
mock_raw.output = []
|
||||
base_response = ChatResponse(messages=[msg], raw_representation=mock_raw)
|
||||
|
||||
async def _fake_awaitable() -> ChatResponse:
|
||||
return base_response
|
||||
|
||||
with patch.object(RawOpenAIResponsesClient, "_inner_get_response", return_value=_fake_awaitable()):
|
||||
result_awaitable = client._inner_get_response(messages=[], options={}, stream=False)
|
||||
result = await result_awaitable # type: ignore[misc]
|
||||
|
||||
assert result.messages[0].contents[0].text == "Hello world"
|
||||
|
||||
|
||||
def _create_mock_stream() -> MagicMock:
|
||||
"""Create a mock ResponseStream with working with_transform_hook."""
|
||||
mock_stream = MagicMock(spec=ResponseStream)
|
||||
mock_stream._transform_hooks = []
|
||||
mock_stream.with_transform_hook.side_effect = lambda hook: mock_stream._transform_hooks.append(hook) or mock_stream
|
||||
return mock_stream
|
||||
|
||||
|
||||
def test_inner_get_response_streaming_registers_hook(mock_project_client: MagicMock) -> None:
|
||||
"""Test _inner_get_response appends a transform hook to the stream for streaming responses."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
|
||||
mock_stream = _create_mock_stream()
|
||||
|
||||
with patch.object(RawOpenAIResponsesClient, "_inner_get_response", return_value=mock_stream):
|
||||
result = client._inner_get_response(messages=[], options={}, stream=True)
|
||||
|
||||
assert result is mock_stream
|
||||
assert len(mock_stream._transform_hooks) == 1
|
||||
|
||||
|
||||
def test_streaming_hook_captures_search_urls(mock_project_client: MagicMock) -> None:
|
||||
"""Test the streaming transform hook captures get_urls from search output events."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
|
||||
mock_stream = _create_mock_stream()
|
||||
|
||||
with patch.object(RawOpenAIResponsesClient, "_inner_get_response", return_value=mock_stream):
|
||||
client._inner_get_response(messages=[], options={}, stream=True)
|
||||
|
||||
hook = mock_stream._transform_hooks[0]
|
||||
|
||||
# Simulate azure_ai_search_call_output event
|
||||
mock_item = MagicMock()
|
||||
mock_item.type = "azure_ai_search_call_output"
|
||||
mock_item.output = MagicMock()
|
||||
mock_item.output.get_urls = [
|
||||
"https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01",
|
||||
]
|
||||
|
||||
raw_event = MagicMock()
|
||||
raw_event.type = "response.output_item.added"
|
||||
raw_event.item = mock_item
|
||||
|
||||
update = ChatResponseUpdate(raw_representation=raw_event)
|
||||
result = hook(update)
|
||||
assert result is update # passes through (no annotations to enrich)
|
||||
|
||||
|
||||
def test_streaming_hook_enriches_url_citation(mock_project_client: MagicMock) -> None:
|
||||
"""Test the streaming transform hook enriches url_citation annotations with get_urls."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
|
||||
mock_stream = _create_mock_stream()
|
||||
|
||||
with patch.object(RawOpenAIResponsesClient, "_inner_get_response", return_value=mock_stream):
|
||||
client._inner_get_response(messages=[], options={}, stream=True)
|
||||
|
||||
hook = mock_stream._transform_hooks[0]
|
||||
|
||||
# Step 1: Feed search output event to capture URLs
|
||||
mock_item = MagicMock()
|
||||
mock_item.type = "azure_ai_search_call_output"
|
||||
mock_item.output = MagicMock()
|
||||
mock_item.output.get_urls = [
|
||||
"https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01",
|
||||
"https://search.example.com/indexes/idx/docs/41?api-version=2024-07-01",
|
||||
]
|
||||
raw_output_event = MagicMock()
|
||||
raw_output_event.type = "response.output_item.added"
|
||||
raw_output_event.item = mock_item
|
||||
hook(ChatResponseUpdate(raw_representation=raw_output_event))
|
||||
|
||||
# Step 2: Feed url_citation annotation event (annotation is always a dict in streaming)
|
||||
raw_ann_event = MagicMock()
|
||||
raw_ann_event.type = "response.output_text.annotation.added"
|
||||
raw_ann_event.annotation = {
|
||||
"type": "url_citation",
|
||||
"title": "doc_0",
|
||||
"url": "https://search.example.com/",
|
||||
"start_index": 100,
|
||||
"end_index": 112,
|
||||
}
|
||||
raw_ann_event.annotation_index = 0
|
||||
|
||||
result = hook(ChatResponseUpdate(raw_representation=raw_ann_event))
|
||||
|
||||
# Verify the result has enriched annotation
|
||||
assert result.contents is not None
|
||||
found = False
|
||||
for content_item in result.contents:
|
||||
if hasattr(content_item, "annotations") and content_item.annotations:
|
||||
for ann in content_item.annotations:
|
||||
if isinstance(ann, dict) and ann.get("title") == "doc_0":
|
||||
found = True
|
||||
assert ann["additional_properties"]["get_url"] == (
|
||||
"https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01"
|
||||
)
|
||||
assert found, "Expected url_citation annotation with enriched get_url"
|
||||
|
||||
|
||||
def test_build_url_citation_content(mock_project_client: MagicMock) -> None:
|
||||
"""Test _build_url_citation_content creates Content with enriched Annotation."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
get_urls = ["https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01"]
|
||||
|
||||
annotation_data = {
|
||||
"type": "url_citation",
|
||||
"title": "doc_0",
|
||||
"url": "https://search.example.com/",
|
||||
"start_index": 100,
|
||||
"end_index": 112,
|
||||
}
|
||||
|
||||
raw_event = MagicMock()
|
||||
raw_event.annotation_index = 0
|
||||
|
||||
content = client._build_url_citation_content(annotation_data, get_urls, raw_event)
|
||||
|
||||
assert content.annotations is not None
|
||||
ann = content.annotations[0]
|
||||
assert ann["type"] == "citation"
|
||||
assert ann["title"] == "doc_0"
|
||||
assert ann["url"] == "https://search.example.com/"
|
||||
assert ann["additional_properties"]["get_url"] == get_urls[0]
|
||||
assert ann["annotated_regions"][0]["start_index"] == 100
|
||||
assert ann["annotated_regions"][0]["end_index"] == 112
|
||||
|
||||
|
||||
def test_build_url_citation_content_with_dict(mock_project_client: MagicMock) -> None:
|
||||
"""Test _build_url_citation_content handles dict-style annotation data."""
|
||||
client = create_test_azure_ai_client(mock_project_client)
|
||||
get_urls = ["https://search.example.com/indexes/idx/docs/16?api-version=2024-07-01"]
|
||||
|
||||
annotation_data = {
|
||||
"type": "url_citation",
|
||||
"title": "doc_1",
|
||||
"url": "https://search.example.com/",
|
||||
"start_index": 200,
|
||||
"end_index": 215,
|
||||
}
|
||||
|
||||
raw_event = MagicMock()
|
||||
raw_event.annotation_index = 1
|
||||
|
||||
content = client._build_url_citation_content(annotation_data, get_urls, raw_event)
|
||||
|
||||
assert content.annotations is not None
|
||||
ann = content.annotations[0]
|
||||
assert ann["type"] == "citation"
|
||||
assert ann["title"] == "doc_1"
|
||||
# doc_1 is out of range for a 1-element get_urls, so no get_url
|
||||
assert "get_url" not in ann.get("additional_properties", {})
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
Reference in New Issue
Block a user