Compare commits

..
Author SHA1 Message Date
Korolev DmitryandGitHub 4072c66b5f Merge branch 'main' into dmkorolev/threads 2025-11-19 12:58:45 +01:00
Korolev Dmitry 0c02824853 rm 2025-11-17 21:53:35 +01:00
Korolev Dmitry 836d22b205 comments 2025-11-17 21:51:52 +01:00
Korolev Dmitry 19b6f3a5d9 unblock 2025-11-17 21:50:47 +01:00
Korolev Dmitry 3b80c9e50d via agent thread store 2025-11-17 19:03:21 +01:00
154 changed files with 4484 additions and 14014 deletions
-2
View File
@@ -204,8 +204,6 @@ agents.md
# AI
.claude/
WARP.md
**/memory-bank/
**/projectBrief.md
# Azurite storage emulator files
*/__azurite_db_blob__.json
-3
View File
@@ -1,3 +0,0 @@
# Declarative Agents
This folder contains sample agent definitions than be ran using the declarative agent support, for python see the [declarative agent python sample folder](../python/samples/getting_started/declarative/).
-25
View File
@@ -1,25 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions is the language specified by the user. You return your answers in a JSON format. You must include Chat as the type in your response.
model:
id: =Env.AZURE_OPENAI_DEPLOYMENT_NAME
provider: AzureOpenAI
apiType: Chat
options:
temperature: 0.9
topP: 0.95
outputSchema:
properties:
language:
kind: string
required: true
description: The language of the answer.
answer:
kind: string
required: true
description: The answer text.
type:
kind: string
required: true
description: The type of the response.
@@ -1,25 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions is the language specified by the user. You return your answers in a JSON format. You must include Assistants as the type in your response.
model:
id: =Env.AZURE_OPENAI_DEPLOYMENT_NAME
provider: AzureOpenAI
apiType: Assistants
options:
temperature: 0.9
topP: 0.95
outputSchema:
properties:
language:
kind: string
required: true
description: The language of the answer.
answer:
kind: string
required: true
description: The answer text.
type:
kind: string
required: true
description: The type of the response.
@@ -1,28 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions is the language specified by the user. You return your answers in a JSON format. You must include Responses as the type in your response.
model:
id: =Env.AZURE_OPENAI_DEPLOYMENT_NAME
provider: AzureOpenAI
apiType: Responses
options:
text:
verbosity: medium
connection:
kind: remote
endpoint: =Env.AZURE_OPENAI_ENDPOINT
outputSchema:
properties:
language:
kind: string
required: true
description: The language of the answer.
answer:
kind: string
required: true
description: The answer text.
type:
kind: string
required: true
description: The type of the response.
-18
View File
@@ -1,18 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions is the language specified by the user. You return your answers in a JSON format.
model:
options:
temperature: 0.9
topP: 0.95
outputSchema:
properties:
language:
kind: string
required: true
description: The language of the answer.
answer:
kind: string
required: true
description: The answer text.
-27
View File
@@ -1,27 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions using the tools provided.
model:
options:
allowMultipleToolCalls: true
chatToolMode: auto
tools:
- kind: function
name: GetWeather
description: Get the weather for a given location.
bindings:
get_weather: get_weather
parameters:
properties:
location:
kind: string
description: The city and state, e.g. San Francisco, CA
required: true
unit:
kind: string
description: The unit of temperature. Possible values are 'celsius' and 'fahrenheit'.
required: false
enum:
- celsius
- fahrenheit
@@ -1,21 +0,0 @@
kind: Prompt
name: MicrosoftLearnAgent
description: Microsoft Learn Agent
instructions: You answer questions by searching the Microsoft Learn content only.
model:
id: =Env.AZURE_FOUNDRY_PROJECT_MODEL_ID
options:
temperature: 0.9
topP: 0.95
connection:
kind: remote
endpoint: =Env.AZURE_FOUNDRY_PROJECT_ENDPOINT
tools:
- kind: mcp
name: microsoft_learn
description: Get information from Microsoft Learn.
url: https://learn.microsoft.com/api/mcp
approvalMode:
kind: never
allowedTools:
- microsoft_docs_search
@@ -1,22 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions is the language specified by the user. You return your answers in a JSON format.
model:
id: =Env.AZURE_FOUNDRY_PROJECT_MODEL_ID
options:
temperature: 0.9
topP: 0.95
connection:
kind: remote
endpoint: =Env.AZURE_FOUNDRY_PROJECT_ENDPOINT
outputSchema:
properties:
language:
kind: string
required: true
description: The language of the answer.
answer:
kind: string
required: true
description: The answer text.
-28
View File
@@ -1,28 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions is the language specified by the user. You return your answers in a JSON format. You must include Chat as the type in your response.
model:
id: =Env.OPENAI_MODEL
provider: OpenAI
apiType: Chat
options:
temperature: 0.9
topP: 0.95
connection:
kind: key
key: =Env.OPENAI_API_KEY
outputSchema:
properties:
language:
kind: string
required: true
description: The language of the answer.
answer:
kind: string
required: true
description: The answer text.
type:
kind: string
required: true
description: The type of the response.
@@ -1,30 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions is the language specified by the user. You return your answers in a JSON format. You must include Assistants as the type in your response.
model:
id: =Env.OPENAI_MODEL
provider: OpenAI
apiType: Assistants
options:
temperature: 0.9
topP: 0.95
connection:
kind: key
key: =Env.OPENAI_APIKEY
outputSchema:
name: AssistantResponse
description: The response from the assistant.
properties:
language:
kind: string
required: true
description: The language of the answer.
answer:
kind: string
required: true
description: The answer text.
type:
kind: string
required: true
description: The type of the response.
-28
View File
@@ -1,28 +0,0 @@
kind: Prompt
name: Assistant
description: Helpful assistant
instructions: You are a helpful assistant. You answer questions is the language specified by the user. You return your answers in a JSON format. You must include Responses as the type in your response.
model:
id: =Env.OPENAI_MODEL
provider: OpenAI
apiType: Responses
options:
text:
verbosity: medium
connection:
kind: key
key: =Env.OPENAI_APIKEY
outputSchema:
properties:
language:
kind: string
required: true
description: The language of the answer.
answer:
kind: string
required: true
description: The answer text.
type:
kind: string
required: true
description: The type of the response.
-1
View File
@@ -101,7 +101,6 @@
<Project Path="samples/GettingStarted/AgentWithRAG/AgentWithRAG_Step01_BasicTextRAG/AgentWithRAG_Step01_BasicTextRAG.csproj" />
<Project Path="samples/GettingStarted/AgentWithRAG/AgentWithRAG_Step02_CustomVectorStoreRAG/AgentWithRAG_Step02_CustomVectorStoreRAG.csproj" />
<Project Path="samples/GettingStarted/AgentWithRAG/AgentWithRAG_Step03_CustomRAGDataSource/AgentWithRAG_Step03_CustomRAGDataSource.csproj" />
<Project Path="samples/GettingStarted/AgentWithRAG/AgentWithRAG_Step04_FoundryServiceRAG/AgentWithRAG_Step04_FoundryServiceRAG.csproj" />
</Folder>
<Folder Name="/Samples/GettingStarted/FoundryAgents/">
<File Path="samples/GettingStarted/FoundryAgents/README.md" />
@@ -1,26 +0,0 @@
<Project Sdk="Microsoft.NET.Sdk">
<PropertyGroup>
<OutputType>Exe</OutputType>
<TargetFramework>net9.0</TargetFramework>
<Nullable>enable</Nullable>
<ImplicitUsings>enable</ImplicitUsings>
</PropertyGroup>
<ItemGroup>
<PackageReference Include="Azure.AI.Projects" />
<PackageReference Include="Azure.Identity" />
</ItemGroup>
<ItemGroup>
<ProjectReference Include="..\..\..\..\src\Microsoft.Agents.AI.AzureAI\Microsoft.Agents.AI.AzureAI.csproj" />
</ItemGroup>
<ItemGroup>
<None Update="contoso-outdoors-knowledge-base.md">
<CopyToOutputDirectory>Always</CopyToOutputDirectory>
</None>
</ItemGroup>
</Project>
@@ -1,60 +0,0 @@
// Copyright (c) Microsoft. All rights reserved.
// This sample shows how to use the built in RAG capabilities that the Foundry service provides when using AI Agents provided by Foundry.
using System.ClientModel;
using Azure.AI.Projects;
using Azure.Identity;
using Microsoft.Agents.AI;
using Microsoft.Extensions.AI;
using OpenAI;
using OpenAI.Files;
using OpenAI.VectorStores;
var endpoint = Environment.GetEnvironmentVariable("AZURE_FOUNDRY_PROJECT_ENDPOINT") ?? throw new InvalidOperationException("AZURE_FOUNDRY_PROJECT_ENDPOINT is not set.");
var deploymentName = Environment.GetEnvironmentVariable("AZURE_FOUNDRY_PROJECT_DEPLOYMENT_NAME") ?? "gpt-4o-mini";
// Create an AI Project client and get an OpenAI client that works with the foundry service.
AIProjectClient aiProjectClient = new(
new Uri(endpoint),
new AzureCliCredential());
OpenAIClient openAIClient = aiProjectClient.GetProjectOpenAIClient();
// Upload the file that contains the data to be used for RAG to the Foundry service.
OpenAIFileClient fileClient = openAIClient.GetOpenAIFileClient();
ClientResult<OpenAIFile> uploadResult = await fileClient.UploadFileAsync(
filePath: "contoso-outdoors-knowledge-base.md",
purpose: FileUploadPurpose.Assistants);
// Create a vector store in the Foundry service using the uploaded file.
VectorStoreClient vectorStoreClient = openAIClient.GetVectorStoreClient();
ClientResult<VectorStore> vectorStoreCreate = await vectorStoreClient.CreateVectorStoreAsync(options: new VectorStoreCreationOptions()
{
Name = "contoso-outdoors-knowledge-base",
FileIds = { uploadResult.Value.Id }
});
var fileSearchTool = new HostedFileSearchTool() { Inputs = [new HostedVectorStoreContent(vectorStoreCreate.Value.Id)] };
AIAgent agent = await aiProjectClient
.CreateAIAgentAsync(
model: deploymentName,
name: "AskContoso",
instructions: "You are a helpful support specialist for Contoso Outdoors. Answer questions using the provided context and cite the source document when available.",
tools: [fileSearchTool]);
AgentThread thread = agent.GetNewThread();
Console.WriteLine(">> Asking about returns\n");
Console.WriteLine(await agent.RunAsync("Hi! I need help understanding the return policy.", thread));
Console.WriteLine("\n>> Asking about shipping\n");
Console.WriteLine(await agent.RunAsync("How long does standard shipping usually take?", thread));
Console.WriteLine("\n>> Asking about product care\n");
Console.WriteLine(await agent.RunAsync("What is the best way to maintain the TrailRunner tent fabric?", thread));
// Cleanup
await fileClient.DeleteFileAsync(uploadResult.Value.Id);
await vectorStoreClient.DeleteVectorStoreAsync(vectorStoreCreate.Value.Id);
await aiProjectClient.Agents.DeleteAgentAsync(agent.Name);
@@ -1,19 +0,0 @@
# Contoso Outdoors Knowledge Base
## Contoso Outdoors Return Policy
Customers may return any item within 30 days of delivery. Items should be unused and include original packaging. Refunds are issued to the original payment method within 5 business days of inspection.
## Contoso Outdoors Shipping Guide
Standard shipping is free on orders over $50 and typically arrives in 3-5 business days within the continental United States. Expedited options are available at checkout.
## Product Information
### TrailRunner Tent
The TrailRunner Tent is a lightweight, 2-person tent designed for easy setup and durability. It features waterproof materials, ventilation windows, and a compact carry bag.
#### Care Instructions
Clean the tent fabric with lukewarm water and a non-detergent soap. Allow it to air dry completely before storage and avoid prolonged UV exposure to extend the lifespan of the waterproof coating.
@@ -7,4 +7,3 @@ These samples show how to create an agent with the Agent Framework that uses Ret
|[Basic Text RAG](./AgentWithRAG_Step01_BasicTextRAG/)|This sample demonstrates how to create and run a basic agent with simple text Retrieval Augmented Generation (RAG).|
|[RAG with Vector Store and custom schema](./AgentWithRAG_Step02_CustomVectorStoreRAG/)|This sample demonstrates how to create and run an agent that uses Retrieval Augmented Generation (RAG) with a vector store. It also uses a custom schema for the documents stored in the vector store.|
|[RAG with custom RAG data source](./AgentWithRAG_Step03_CustomRAGDataSource/)|This sample demonstrates how to create and run an agent that uses Retrieval Augmented Generation (RAG) with a custom RAG data source.|
|[RAG with Foundry VectorStore service](./AgentWithRAG_Step04_FoundryServiceRAG/)|This sample demonstrates how to create and run an agent that uses Retrieval Augmented Generation (RAG) with the Foundry VectorStore service.|
@@ -34,6 +34,7 @@ internal sealed partial class IdGenerator
this._random = randomSeed.HasValue ? new Random(randomSeed.Value) : null;
this.ResponseId = responseId ?? NewId("resp", random: this._random);
this.ConversationId = conversationId ?? NewId("conv", random: this._random);
this.IsNewConversation = conversationId is null;
this._partitionId = GetPartitionIdOrDefault(this.ConversationId) ?? string.Empty;
}
@@ -59,6 +60,11 @@ internal sealed partial class IdGenerator
/// </summary>
public string ConversationId { get; }
/// <summary>
/// Gets a value indicating whether this is a new conversation.
/// </summary>
public bool IsNewConversation { get; }
/// <summary>
/// Generates a new ID.
/// </summary>
@@ -26,6 +26,11 @@ internal sealed class AgentInvocationContext(IdGenerator idGenerator, JsonSerial
/// </summary>
public string ConversationId => this.IdGenerator.ConversationId;
/// <summary>
/// Returns true, if conversation is new.
/// </summary>
public bool IsNewConversation => this.IdGenerator.IsNewConversation;
/// <summary>
/// Gets the JSON serializer options.
/// </summary>
@@ -81,10 +81,13 @@ internal sealed class HostedAgentResponseExecutor : IResponseExecutor
[EnumeratorCancellation] CancellationToken cancellationToken = default)
{
string agentName = GetAgentName(request)!;
AIAgent agent = this._serviceProvider.GetRequiredKeyedService<AIAgent>(agentName);
string conversationId = context.ConversationId;
var agent = this._serviceProvider.GetRequiredKeyedService<AIAgent>(agentName);
var threadStore = this._serviceProvider.GetKeyedService<AgentThreadStore>(agent.Name);
var chatOptions = new ChatOptions
{
ConversationId = request.Conversation?.Id,
Temperature = (float?)request.Temperature,
TopP = (float?)request.TopP,
MaxOutputTokens = request.MaxOutputTokens,
@@ -94,16 +97,25 @@ internal sealed class HostedAgentResponseExecutor : IResponseExecutor
var options = new ChatClientAgentRunOptions(chatOptions);
var messages = new List<ChatMessage>();
AgentThread thread = !context.IsNewConversation && threadStore is not null
? await threadStore.GetThreadAsync(agent, conversationId, cancellationToken).ConfigureAwait(false)
: agent.GetNewThread();
foreach (var inputMessage in request.Input.GetInputMessages())
{
messages.Add(inputMessage.ToChatMessage());
}
await foreach (var streamingEvent in agent.RunStreamingAsync(messages, options: options, cancellationToken: cancellationToken)
await foreach (var streamingEvent in agent.RunStreamingAsync(messages, thread, options: options, cancellationToken: cancellationToken)
.ToStreamingResponseAsync(request, context, cancellationToken).ConfigureAwait(false))
{
yield return streamingEvent;
}
if (threadStore is not null && thread is not null)
{
await threadStore.SaveThreadAsync(agent, conversationId, thread, cancellationToken).ConfigureAwait(false);
}
}
/// <summary>
@@ -4,7 +4,6 @@ using System;
using System.Collections.Generic;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Extensions.AI;
using Moq;
namespace Microsoft.Agents.AI.Workflows.Declarative.UnitTests;
@@ -16,18 +15,10 @@ internal sealed class MockAgentProvider : Mock<WorkflowAgentProvider>
{
public IList<string> ExistingConversationIds { get; } = [];
public ChatMessage? TestChatMessage { get; set; }
public MockAgentProvider()
{
this.Setup(provider => provider.CreateConversationAsync(It.IsAny<CancellationToken>()))
.Returns(() => Task.FromResult(this.CreateConversationId()));
this.Setup(provider => provider.GetMessageAsync(
It.IsAny<string>(),
It.IsAny<string>(),
It.IsAny<CancellationToken>()))
.Returns(Task.FromResult(this.CreateChatMessage()));
}
private string CreateConversationId()
@@ -37,13 +28,4 @@ internal sealed class MockAgentProvider : Mock<WorkflowAgentProvider>
return newConversationId;
}
private ChatMessage CreateChatMessage()
{
this.TestChatMessage = new ChatMessage(ChatRole.User, Guid.NewGuid().ToString("N"))
{
MessageId = Guid.NewGuid().ToString("N"),
};
return this.TestChatMessage;
}
}
@@ -1,67 +0,0 @@
// Copyright (c) Microsoft. All rights reserved.
using System.Threading.Tasks;
using Microsoft.Agents.AI.Workflows.Declarative.Extensions;
using Microsoft.Agents.AI.Workflows.Declarative.ObjectModel;
using Microsoft.Bot.ObjectModel;
using Microsoft.Extensions.AI;
using Xunit.Abstractions;
namespace Microsoft.Agents.AI.Workflows.Declarative.UnitTests.ObjectModel;
/// <summary>
/// Tests for <see cref="RetrieveConversationMessageExecutor"/>.
/// </summary>
public sealed class RetrieveConversationMessageExecutorTest(ITestOutputHelper output) : WorkflowActionExecutorTest(output)
{
[Fact]
public async Task RetrieveMessageSuccessfullyAsync()
{
// Arrange, Act, Assert
await this.ExecuteTestAsync(nameof(RetrieveMessageSuccessfullyAsync),
"TestMessage");
}
private async Task ExecuteTestAsync(
string displayName,
string variableName)
{
// Arrange
MockAgentProvider mockAgentProvider = new();
RetrieveConversationMessage model = this.CreateModel(
this.FormatDisplayName(displayName),
FormatVariablePath(variableName),
"TestConversationId",
"DefaultMessageId");
RetrieveConversationMessageExecutor action = new(model, mockAgentProvider.Object, this.State);
// Act
await this.ExecuteAsync(action);
// Assert
ChatMessage testMessage = mockAgentProvider.TestChatMessage ?? new ChatMessage();
VerifyModel(model, action);
this.VerifyState(variableName, testMessage.ToRecord());
}
private RetrieveConversationMessage CreateModel(
string displayName,
string messageVariable,
string conversationId,
string messageId)
{
RetrieveConversationMessage.Builder actionBuilder =
new()
{
Id = this.CreateActionId(),
DisplayName = this.FormatDisplayName(displayName),
Message = PropertyPath.Create(messageVariable),
ConversationId = StringExpression.Literal(conversationId),
MessageId = StringExpression.Literal(messageId)
};
return AssignParent<RetrieveConversationMessage>(actionBuilder);
}
}
-1
View File
@@ -59,7 +59,6 @@
"OPENAI",
"opentelemetry",
"OTEL",
"powerfx",
"protos",
"pydantic",
"pytestmark",
-8
View File
@@ -3,14 +3,6 @@ AZURE_AI_PROJECT_ENDPOINT=""
AZURE_AI_MODEL_DEPLOYMENT_NAME=""
# Bing connection for web search (optional, used by samples with web search)
BING_CONNECTION_ID=""
# Azure AI Search (optional, used by AzureAISearchContextProvider samples)
AZURE_SEARCH_ENDPOINT=""
AZURE_SEARCH_API_KEY=""
AZURE_SEARCH_INDEX_NAME=""
AZURE_SEARCH_SEMANTIC_CONFIG=""
AZURE_SEARCH_KNOWLEDGE_BASE_NAME=""
# Note: For agentic mode Knowledge Bases, also set AZURE_OPENAI_ENDPOINT below
# (different from AZURE_AI_PROJECT_ENDPOINT - Knowledge Base needs OpenAI endpoint for model calls)
# OpenAI
OPENAI_API_KEY=""
OPENAI_CHAT_MODEL_ID=""
+1 -33
View File
@@ -7,36 +7,6 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
## [Unreleased]
## [1.0.0b251120] - 2025-11-20
### Added
- **agent-framework-core**: Introducing support for declarative YAML spec ([#2002](https://github.com/microsoft/agent-framework/pull/2002))
- **agent-framework-core**: Use AI Foundry evaluators for self-reflection ([#2250](https://github.com/microsoft/agent-framework/pull/2250))
- **agent-framework-core**: Propagate `as_tool()` kwargs and add runtime context + middleware sample ([#2311](https://github.com/microsoft/agent-framework/pull/2311))
- **agent-framework-anthropic**: Anthropic Foundry integration ([#2302](https://github.com/microsoft/agent-framework/pull/2302))
- **samples**: M365 Agent SDK Hosting sample ([#2292](https://github.com/microsoft/agent-framework/pull/2292))
- **samples**: Foundry Sample for A2A + SharePoint Samples ([#2313](https://github.com/microsoft/agent-framework/pull/2313))
### Changed
- **agent-framework-azurefunctions**: [BREAKING] Schema changes for Azure Functions package ([#2151](https://github.com/microsoft/agent-framework/pull/2151))
- **agent-framework-core**: Move evaluation folders under `evaluations` ([#2355](https://github.com/microsoft/agent-framework/pull/2355))
- **agent-framework-core**: Move red teaming files to their own folder ([#2333](https://github.com/microsoft/agent-framework/pull/2333))
- **agent-framework-core**: "fix all" task now single source of truth ([#2303](https://github.com/microsoft/agent-framework/pull/2303))
- **agent-framework-core**: Improve and clean up exception handling ([#2337](https://github.com/microsoft/agent-framework/pull/2337), [#2319](https://github.com/microsoft/agent-framework/pull/2319))
- **agent-framework-core**: Clean up imports ([#2318](https://github.com/microsoft/agent-framework/pull/2318))
### Fixed
- **agent-framework-azure-ai**: Fix for Azure AI client ([#2358](https://github.com/microsoft/agent-framework/pull/2358))
- **agent-framework-core**: Fix tool execution bleed-over in aiohttp/Bot Framework scenarios ([#2314](https://github.com/microsoft/agent-framework/pull/2314))
- **agent-framework-core**: `@ai_function` now correctly handles `self` parameter ([#2266](https://github.com/microsoft/agent-framework/pull/2266))
- **agent-framework-core**: Resolve string annotations in `FunctionExecutor` ([#2308](https://github.com/microsoft/agent-framework/pull/2308))
- **agent-framework-core**: Langfuse observability captures ChatAgent system instructions ([#2316](https://github.com/microsoft/agent-framework/pull/2316))
- **agent-framework-core**: Incomplete URL substring sanitization fix ([#2274](https://github.com/microsoft/agent-framework/pull/2274))
- **observability**: Handle datetime serialization in tool results ([#2248](https://github.com/microsoft/agent-framework/pull/2248))
## [1.0.0b251117] - 2025-11-17
### Fixed
@@ -290,9 +260,7 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0
For more information, see the [announcement blog post](https://devblogs.microsoft.com/foundry/introducing-microsoft-agent-framework-the-open-source-engine-for-agentic-ai-apps/).
[Unreleased]: https://github.com/microsoft/agent-framework/compare/python-1.0.0b251120...HEAD
[1.0.0b251120]: https://github.com/microsoft/agent-framework/compare/python-1.0.0b251117...python-1.0.0b251120
[1.0.0b251117]: https://github.com/microsoft/agent-framework/compare/python-1.0.0b251114...python-1.0.0b251117
[Unreleased]: https://github.com/microsoft/agent-framework/compare/python-1.0.0b251114...HEAD
[1.0.0b251114]: https://github.com/microsoft/agent-framework/compare/python-1.0.0b251112.post1...python-1.0.0b251114
[1.0.0b251112.post1]: https://github.com/microsoft/agent-framework/compare/python-1.0.0b251112...python-1.0.0b251112.post1
[1.0.0b251112]: https://github.com/microsoft/agent-framework/compare/python-1.0.0b251111...python-1.0.0b251112
+1 -1
View File
@@ -4,7 +4,7 @@ description = "A2A integration for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251120"
version = "1.0.0b251114"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
+2 -2
View File
@@ -16,7 +16,7 @@ pip install agent-framework-ag-ui
from fastapi import FastAPI
from agent_framework import ChatAgent
from agent_framework.azure import AzureOpenAIChatClient
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework_ag_ui import add_agent_framework_fastapi_endpoint
# Create your agent
agent = ChatAgent(
@@ -41,7 +41,7 @@ add_agent_framework_fastapi_endpoint(app, agent, "/")
```python
import asyncio
from agent_framework import TextContent
from agent_framework.ag_ui import AGUIChatClient
from agent_framework_ag_ui import AGUIChatClient
async def main():
async with AGUIChatClient(endpoint="http://localhost:8000/") as client:
@@ -91,4 +91,4 @@ def add_agent_framework_fastapi_endpoint(
)
except Exception as e:
logger.error(f"Error in agent endpoint: {e}", exc_info=True)
return {"error": "An internal error has occurred."}
return {"error": str(e)}
@@ -18,7 +18,7 @@ All example agents are factory functions that accept any `ChatClientProtocol`-co
from fastapi import FastAPI
from agent_framework.azure import AzureOpenAIChatClient
from agent_framework.openai import OpenAIChatClient
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework_ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework_ag_ui_examples.agents import simple_agent, weather_agent
app = FastAPI()
@@ -40,7 +40,7 @@ add_agent_framework_fastapi_endpoint(app, weather_agent(openai_client), "/weathe
from fastapi import FastAPI
from agent_framework import ChatAgent
from agent_framework.azure import AzureOpenAIChatClient
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework_ag_ui import add_agent_framework_fastapi_endpoint
# Create your agent
agent = ChatAgent(
@@ -136,7 +136,7 @@ The server exposes endpoints at:
```python
from fastapi import FastAPI
from agent_framework.azure import AzureOpenAIChatClient
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework_ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework_ag_ui_examples.agents import (
simple_agent,
weather_agent,
@@ -188,8 +188,8 @@ You can create your own agent factories following the same pattern as the exampl
```python
from agent_framework import ChatAgent, ai_function
from agent_framework import ChatClientProtocol
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework._clients import ChatClientProtocol
from agent_framework_ag_ui import AgentFrameworkAgent
@ai_function
def my_tool(param: str) -> str:
@@ -2,8 +2,10 @@
"""Example agent demonstrating predictive state updates with document writing."""
from agent_framework import ChatAgent, ChatClientProtocol, ai_function
from agent_framework.ag_ui import AgentFrameworkAgent, DocumentWriterConfirmationStrategy
from agent_framework import ChatAgent, ai_function
from agent_framework._clients import ChatClientProtocol
from agent_framework_ag_ui import AgentFrameworkAgent, DocumentWriterConfirmationStrategy
@ai_function
@@ -4,7 +4,8 @@
from enum import Enum
from agent_framework import ChatAgent, ChatClientProtocol, ai_function
from agent_framework import ChatAgent, ai_function
from agent_framework._clients import ChatClientProtocol
from pydantic import BaseModel, Field
@@ -4,10 +4,12 @@
from enum import Enum
from agent_framework import ChatAgent, ChatClientProtocol, ai_function
from agent_framework.ag_ui import AgentFrameworkAgent, RecipeConfirmationStrategy
from agent_framework import ChatAgent, ai_function
from agent_framework._clients import ChatClientProtocol
from pydantic import BaseModel, Field
from agent_framework_ag_ui import AgentFrameworkAgent, RecipeConfirmationStrategy
class SkillLevel(str, Enum):
"""The skill level required for the recipe."""
@@ -4,8 +4,10 @@
import asyncio
from agent_framework import ChatAgent, ChatClientProtocol, ai_function
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework import ChatAgent, ai_function
from agent_framework._clients import ChatClientProtocol
from agent_framework_ag_ui import AgentFrameworkAgent
@ai_function
@@ -2,7 +2,8 @@
"""Simple agentic chat example (Feature 1: Agentic Chat)."""
from agent_framework import ChatAgent, ChatClientProtocol
from agent_framework import ChatAgent
from agent_framework._clients import ChatClientProtocol
def simple_agent(chat_client: ChatClientProtocol) -> ChatAgent:
@@ -2,8 +2,10 @@
"""Example agent demonstrating human-in-the-loop with function approvals."""
from agent_framework import ChatAgent, ChatClientProtocol, ai_function
from agent_framework.ag_ui import AgentFrameworkAgent, TaskPlannerConfirmationStrategy
from agent_framework import ChatAgent, ai_function
from agent_framework._clients import ChatClientProtocol
from agent_framework_ag_ui import AgentFrameworkAgent, TaskPlannerConfirmationStrategy
@ai_function(approval_mode="always_require")
@@ -18,10 +18,12 @@ from ag_ui.core import (
TextMessageStartEvent,
ToolCallStartEvent,
)
from agent_framework import ChatAgent, ChatClientProtocol, ai_function
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework import ChatAgent, ai_function
from agent_framework._clients import ChatClientProtocol
from pydantic import BaseModel, Field
from agent_framework_ag_ui import AgentFrameworkAgent
class StepStatus(str, Enum):
"""Status of a task step."""
@@ -4,8 +4,10 @@
from typing import Any
from agent_framework import AIFunction, ChatAgent, ChatClientProtocol
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework import AIFunction, ChatAgent
from agent_framework._clients import ChatClientProtocol
from agent_framework_ag_ui import AgentFrameworkAgent
# Declaration-only tools (func=None) - actual rendering happens on the client side
generate_haiku = AIFunction[Any, str](
@@ -4,7 +4,8 @@
from typing import Any
from agent_framework import ChatAgent, ChatClientProtocol, ai_function
from agent_framework import ChatAgent, ai_function
from agent_framework._clients import ChatClientProtocol
@ai_function
@@ -2,10 +2,11 @@
"""Backend tool rendering endpoint."""
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework.azure import AzureOpenAIChatClient
from fastapi import FastAPI
from agent_framework_ag_ui import add_agent_framework_fastapi_endpoint
from ...agents.weather_agent import weather_agent
@@ -6,11 +6,12 @@ import logging
import os
import uvicorn
from agent_framework.ag_ui import add_agent_framework_fastapi_endpoint
from agent_framework.azure import AzureOpenAIChatClient
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from agent_framework_ag_ui import add_agent_framework_fastapi_endpoint
from ..agents.document_writer_agent import document_writer_agent
from ..agents.human_in_the_loop_agent import human_in_the_loop_agent
from ..agents.recipe_agent import recipe_agent
@@ -10,7 +10,7 @@ standard chat interface.
import asyncio
import os
from agent_framework.ag_ui import AGUIChatClient
from agent_framework_ag_ui import AGUIChatClient
async def main():
@@ -13,7 +13,8 @@ import asyncio
import os
from agent_framework import ai_function
from agent_framework.ag_ui import AGUIChatClient
from agent_framework_ag_ui import AGUIChatClient
@ai_function
@@ -23,7 +23,8 @@ import logging
import os
from agent_framework import ChatAgent, FunctionCallContent, FunctionResultContent, TextContent, ai_function
from agent_framework.ag_ui import AGUIChatClient
from agent_framework_ag_ui import AGUIChatClient
# Enable debug logging
logging.basicConfig(
+1 -1
View File
@@ -1,6 +1,6 @@
[project]
name = "agent-framework-ag-ui"
version = "1.0.0b251120"
version = "1.0.0b251117"
description = "AG-UI protocol integration for Agent Framework"
readme = "README.md"
license-files = ["LICENSE"]
+1
View File
@@ -0,0 +1 @@
# Copyright (c) Microsoft. All rights reserved.
@@ -11,7 +11,7 @@ from agent_framework._types import ChatResponseUpdate
async def test_agent_initialization_basic():
"""Test basic agent initialization without state schema."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -28,7 +28,7 @@ async def test_agent_initialization_basic():
async def test_agent_initialization_with_state_schema():
"""Test agent initialization with state_schema."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -43,7 +43,7 @@ async def test_agent_initialization_with_state_schema():
async def test_agent_initialization_with_predict_state_config():
"""Test agent initialization with predict_state_config."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -58,7 +58,7 @@ async def test_agent_initialization_with_predict_state_config():
async def test_run_started_event_emission():
"""Test RunStartedEvent is emitted at start of run."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -81,7 +81,7 @@ async def test_run_started_event_emission():
async def test_predict_state_custom_event_emission():
"""Test PredictState CustomEvent is emitted when predict_state_config is present."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -112,7 +112,7 @@ async def test_predict_state_custom_event_emission():
async def test_initial_state_snapshot_with_schema():
"""Test initial StateSnapshotEvent emission when state_schema present."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -141,7 +141,7 @@ async def test_initial_state_snapshot_with_schema():
async def test_state_initialization_object_type():
"""Test state initialization with object type in schema."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -167,7 +167,7 @@ async def test_state_initialization_object_type():
async def test_state_initialization_array_type():
"""Test state initialization with array type in schema."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -193,7 +193,7 @@ async def test_state_initialization_array_type():
async def test_run_finished_event_emission():
"""Test RunFinishedEvent is emitted at end of run."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -214,7 +214,7 @@ async def test_run_finished_event_emission():
async def test_tool_result_confirm_changes_accepted():
"""Test confirm_changes tool result handling when accepted."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -260,7 +260,7 @@ async def test_tool_result_confirm_changes_accepted():
async def test_tool_result_confirm_changes_rejected():
"""Test confirm_changes tool result handling when rejected."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -293,7 +293,7 @@ async def test_tool_result_confirm_changes_rejected():
async def test_tool_result_function_approval_accepted():
"""Test function approval tool result when steps are accepted."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -338,7 +338,7 @@ async def test_tool_result_function_approval_accepted():
async def test_tool_result_function_approval_rejected():
"""Test function approval tool result when rejected."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -374,7 +374,7 @@ async def test_tool_result_function_approval_rejected():
async def test_thread_metadata_tracking():
"""Test that thread metadata includes ag_ui_thread_id and ag_ui_run_id."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
thread_metadata = {}
@@ -405,7 +405,7 @@ async def test_thread_metadata_tracking():
async def test_state_context_injection():
"""Test that current state is injected into thread metadata."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
thread_metadata = {}
@@ -436,7 +436,7 @@ async def test_state_context_injection():
async def test_no_messages_provided():
"""Test handling when no messages are provided."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -459,7 +459,7 @@ async def test_no_messages_provided():
async def test_message_end_event_emission():
"""Test TextMessageEndEvent is emitted for assistant messages."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -486,7 +486,7 @@ async def test_message_end_event_emission():
async def test_error_handling_with_exception():
"""Test that exceptions during agent execution are re-raised."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class FailingChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -506,7 +506,7 @@ async def test_error_handling_with_exception():
async def test_json_decode_error_in_tool_result():
"""Test handling of orphaned tool result - should be sanitized out."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -543,7 +543,7 @@ async def test_json_decode_error_in_tool_result():
async def test_suppressed_summary_with_document_state():
"""Test suppressed summary uses document state for confirmation message."""
from agent_framework.ag_ui import AgentFrameworkAgent, DocumentWriterConfirmationStrategy
from agent_framework_ag_ui import AgentFrameworkAgent, DocumentWriterConfirmationStrategy
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
+1 -1
View File
@@ -154,7 +154,7 @@ async def test_endpoint_error_handling():
assert response.status_code == 200
content = json.loads(response.content)
assert "error" in content
assert content["error"] == "An internal error has occurred."
assert "Expecting value" in content["error"]
async def test_endpoint_multiple_paths():
@@ -32,7 +32,7 @@ class GenericOutput(BaseModel):
async def test_structured_output_with_recipe():
"""Test structured output processing with recipe state."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -70,7 +70,7 @@ async def test_structured_output_with_recipe():
async def test_structured_output_with_steps():
"""Test structured output processing with steps state."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -109,7 +109,7 @@ async def test_structured_output_with_steps():
async def test_structured_output_with_no_schema_match():
"""Test structured output when response fields don't match state_schema keys."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -138,7 +138,7 @@ async def test_structured_output_with_no_schema_match():
async def test_structured_output_without_schema():
"""Test structured output without state_schema treats all fields as state."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class DataOutput(BaseModel):
"""Output with data and info fields."""
@@ -175,7 +175,7 @@ async def test_structured_output_without_schema():
async def test_no_structured_output_when_no_response_format():
"""Test that structured output path is skipped when no response_format."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -200,7 +200,7 @@ async def test_no_structured_output_when_no_response_format():
async def test_structured_output_with_message_field():
"""Test structured output that includes a message field."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
@@ -234,7 +234,7 @@ async def test_structured_output_with_message_field():
async def test_empty_updates_no_structured_processing():
"""Test that empty updates don't trigger structured output processing."""
from agent_framework.ag_ui import AgentFrameworkAgent
from agent_framework_ag_ui import AgentFrameworkAgent
class MockChatClient:
async def get_streaming_response(self, messages, chat_options, **kwargs):
-21
View File
@@ -1,21 +0,0 @@
MIT License
Copyright (c) Microsoft Corporation.
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE
-23
View File
@@ -1,23 +0,0 @@
# Get Started with Microsoft Agent Framework Azure AI Search
Please install this package via pip:
```bash
pip install agent-framework-aisearch --pre
```
## Azure AI Search Integration
The Azure AI Search integration provides context providers for RAG (Retrieval Augmented Generation) capabilities with two modes:
- **Semantic Mode**: Fast hybrid search (vector + keyword) with semantic ranking
- **Agentic Mode**: Multi-hop reasoning using Knowledge Bases for complex queries
### Basic Usage Example
See the [Azure AI Search context provider examples](https://github.com/microsoft/agent-framework/tree/main/python/samples/getting_started/agents/azure_ai/) which demonstrate:
- Semantic search with hybrid (vector + keyword) queries
- Agentic mode with Knowledge Bases for complex multi-hop reasoning
- Environment variable configuration with Settings class
- API key and managed identity authentication
@@ -1,16 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
import importlib.metadata
from ._search_provider import AzureAISearchContextProvider, AzureAISearchSettings
try:
__version__ = importlib.metadata.version(__name__)
except importlib.metadata.PackageNotFoundError:
__version__ = "0.0.0" # Fallback for development mode
__all__ = [
"AzureAISearchContextProvider",
"AzureAISearchSettings",
"__version__",
]
@@ -1,914 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
"""Azure AI Search Context Provider for Agent Framework.
This module provides context providers for Azure AI Search integration with two modes:
- Agentic: Recommended for most scenarios. Uses Knowledge Bases for query planning and
multi-hop reasoning. Slightly slower with more token consumption, but more accurate.
- Semantic: Fast hybrid search (vector + keyword) with semantic ranker. Best for simple
queries where speed is critical.
See: https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/foundry-iq-boost-response-relevance-by-36-with-agentic-retrieval/4470720
"""
import sys
from collections.abc import Awaitable, Callable, MutableSequence
from typing import TYPE_CHECKING, Any, ClassVar, Literal
from agent_framework import ChatMessage, Context, ContextProvider, Role
from agent_framework._logging import get_logger
from agent_framework._pydantic import AFBaseSettings
from agent_framework.exceptions import ServiceInitializationError
from azure.core.credentials import AzureKeyCredential
from azure.core.credentials_async import AsyncTokenCredential
from azure.core.exceptions import ResourceNotFoundError
from azure.search.documents.aio import SearchClient
from azure.search.documents.indexes.aio import SearchIndexClient
from azure.search.documents.indexes.models import (
AzureOpenAIVectorizerParameters,
KnowledgeBase,
KnowledgeBaseAzureOpenAIModel,
KnowledgeRetrievalLowReasoningEffort,
KnowledgeRetrievalMediumReasoningEffort,
KnowledgeRetrievalMinimalReasoningEffort,
KnowledgeRetrievalOutputMode,
KnowledgeRetrievalReasoningEffort,
KnowledgeSourceReference,
SearchIndexKnowledgeSource,
SearchIndexKnowledgeSourceParameters,
)
from azure.search.documents.models import (
QueryCaptionType,
QueryType,
VectorizableTextQuery,
VectorizedQuery,
)
from pydantic import SecretStr, ValidationError
# Type checking imports for optional agentic mode dependencies
if TYPE_CHECKING:
from azure.search.documents.knowledgebases.aio import KnowledgeBaseRetrievalClient
from azure.search.documents.knowledgebases.models import (
KnowledgeBaseMessage,
KnowledgeBaseMessageTextContent,
KnowledgeBaseRetrievalRequest,
KnowledgeRetrievalIntent,
KnowledgeRetrievalSemanticIntent,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalLowReasoningEffort as KBRetrievalLowReasoningEffort,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalMediumReasoningEffort as KBRetrievalMediumReasoningEffort,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalMinimalReasoningEffort as KBRetrievalMinimalReasoningEffort,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalOutputMode as KBRetrievalOutputMode,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalReasoningEffort as KBRetrievalReasoningEffort,
)
# Runtime imports for agentic mode (optional dependency)
try:
from azure.search.documents.knowledgebases.aio import KnowledgeBaseRetrievalClient
from azure.search.documents.knowledgebases.models import (
KnowledgeBaseMessage,
KnowledgeBaseMessageTextContent,
KnowledgeBaseRetrievalRequest,
KnowledgeRetrievalIntent,
KnowledgeRetrievalSemanticIntent,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalLowReasoningEffort as KBRetrievalLowReasoningEffort,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalMediumReasoningEffort as KBRetrievalMediumReasoningEffort,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalMinimalReasoningEffort as KBRetrievalMinimalReasoningEffort,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalOutputMode as KBRetrievalOutputMode,
)
from azure.search.documents.knowledgebases.models import (
KnowledgeRetrievalReasoningEffort as KBRetrievalReasoningEffort,
)
_agentic_retrieval_available = True
except ImportError:
_agentic_retrieval_available = False
if sys.version_info >= (3, 11):
from typing import Self # pragma: no cover
else:
from typing_extensions import Self # pragma: no cover
if sys.version_info >= (3, 12):
from typing import override # type: ignore # pragma: no cover
else:
from typing_extensions import override # type: ignore[import] # pragma: no cover
# Module-level constants
logger = get_logger("agent_framework.azure")
_DEFAULT_AGENTIC_MESSAGE_HISTORY_COUNT = 10
class AzureAISearchSettings(AFBaseSettings):
"""Settings for Azure AI Search Context Provider with auto-loading from environment.
The settings are first loaded from environment variables with the prefix 'AZURE_SEARCH_'.
If the environment variables are not found, the settings can be loaded from a .env file.
Keyword Args:
endpoint: Azure AI Search endpoint URL.
Can be set via environment variable AZURE_SEARCH_ENDPOINT.
index_name: Name of the search index.
Can be set via environment variable AZURE_SEARCH_INDEX_NAME.
api_key: API key for authentication (optional, use managed identity if not provided).
Can be set via environment variable AZURE_SEARCH_API_KEY.
env_file_path: If provided, the .env settings are read from this file path location.
env_file_encoding: The encoding of the .env file, defaults to 'utf-8'.
Examples:
.. code-block:: python
from agent_framework_aisearch import AzureAISearchSettings
# Using environment variables
# Set AZURE_SEARCH_ENDPOINT=https://mysearch.search.windows.net
# Set AZURE_SEARCH_INDEX_NAME=my-index
settings = AzureAISearchSettings()
# Or passing parameters directly
settings = AzureAISearchSettings(
endpoint="https://mysearch.search.windows.net",
index_name="my-index",
)
# Or loading from a .env file
settings = AzureAISearchSettings(env_file_path="path/to/.env")
"""
env_prefix: ClassVar[str] = "AZURE_SEARCH_"
endpoint: str | None = None
index_name: str | None = None
api_key: SecretStr | None = None
class AzureAISearchContextProvider(ContextProvider):
"""Azure AI Search Context Provider with hybrid search and semantic ranking.
This provider retrieves relevant documents from Azure AI Search to provide context
to the AI agent. It supports two modes:
- **agentic**: Recommended for most scenarios. Uses Knowledge Bases for query planning
and multi-hop reasoning. Slightly slower with more token consumption, but provides
more accurate results (up to 36% improvement in response relevance).
- **semantic** (default): Fast hybrid search combining vector and keyword search
with semantic reranking. Best for simple queries where speed is critical.
Examples:
Using environment variables (recommended):
.. code-block:: python
from agent_framework_aisearch import AzureAISearchContextProvider
from azure.identity.aio import DefaultAzureCredential
# Set AZURE_SEARCH_ENDPOINT and AZURE_SEARCH_INDEX_NAME in environment
search_provider = AzureAISearchContextProvider(credential=DefaultAzureCredential())
Semantic hybrid search with API key:
.. code-block:: python
# Direct API key string
search_provider = AzureAISearchContextProvider(
endpoint="https://mysearch.search.windows.net",
index_name="my-index",
api_key="my-api-key",
mode="semantic",
)
Loading from .env file:
.. code-block:: python
# Load settings from a .env file
search_provider = AzureAISearchContextProvider(
credential=DefaultAzureCredential(), env_file_path="path/to/.env"
)
Agentic retrieval for complex queries:
.. code-block:: python
# Use agentic mode for multi-hop reasoning
# Note: azure_openai_resource_url is the OpenAI endpoint for Knowledge Base model calls,
# which is different from azure_ai_project_endpoint (the AI Foundry project endpoint)
search_provider = AzureAISearchContextProvider(
endpoint="https://mysearch.search.windows.net",
index_name="my-index",
credential=DefaultAzureCredential(),
mode="agentic",
azure_openai_resource_url="https://myresource.openai.azure.com",
model_deployment_name="gpt-4o",
knowledge_base_name="my-knowledge-base",
)
"""
_DEFAULT_SEARCH_CONTEXT_PROMPT = "Use the following context to answer the question:"
def __init__(
self,
endpoint: str | None = None,
index_name: str | None = None,
api_key: str | AzureKeyCredential | None = None,
credential: AsyncTokenCredential | None = None,
*,
mode: Literal["semantic", "agentic"] = "semantic",
top_k: int = 5,
semantic_configuration_name: str | None = None,
vector_field_name: str | None = None,
embedding_function: Callable[[str], Awaitable[list[float]]] | None = None,
context_prompt: str | None = None,
# Agentic mode parameters (Knowledge Base)
azure_ai_project_endpoint: str | None = None,
azure_openai_resource_url: str | None = None,
model_deployment_name: str | None = None,
model_name: str | None = None,
knowledge_base_name: str | None = None,
retrieval_instructions: str | None = None,
azure_openai_api_key: str | None = None,
knowledge_base_output_mode: Literal["extractive_data", "answer_synthesis"] = "extractive_data",
retrieval_reasoning_effort: Literal["minimal", "medium", "low"] = "minimal",
agentic_message_history_count: int = _DEFAULT_AGENTIC_MESSAGE_HISTORY_COUNT,
env_file_path: str | None = None,
env_file_encoding: str | None = None,
) -> None:
"""Initialize Azure AI Search Context Provider.
Args:
endpoint: Azure AI Search endpoint URL.
Can also be set via environment variable AZURE_SEARCH_ENDPOINT.
index_name: Name of the search index to query.
Can also be set via environment variable AZURE_SEARCH_INDEX_NAME.
api_key: API key for authentication (string or AzureKeyCredential).
Can also be set via environment variable AZURE_SEARCH_API_KEY.
credential: AsyncTokenCredential for managed identity authentication.
Use this for Entra ID authentication instead of api_key.
mode: Search mode - "semantic" for hybrid search with semantic ranking (fast)
or "agentic" for multi-hop reasoning (slower). Default: "semantic".
top_k: Maximum number of documents to retrieve. Only applies to semantic mode.
In agentic mode, the server-side Knowledge Base determines retrieval based on
query complexity and reasoning effort. Default: 5.
semantic_configuration_name: Name of semantic configuration in the index.
Required for semantic ranking. If None, uses index default.
vector_field_name: Name of the vector field in the index for hybrid search.
Required if using vector search. Default: None (keyword search only).
embedding_function: Async function to generate embeddings for vector search.
Signature: async def embed(text: str) -> list[float]
Required if vector_field_name is specified and no server-side vectorization.
context_prompt: Custom prompt to prepend to retrieved context.
Default: "Use the following context to answer the question:"
azure_ai_project_endpoint: Azure AI Foundry project endpoint URL.
This is NOT the same as azure_openai_resource_url - the project endpoint is used
for Azure AI Foundry services, while the OpenAI endpoint is used by the Knowledge
Base to call the model for query planning. Required for agentic mode.
Example: "https://myproject.services.ai.azure.com/api/projects/myproject"
azure_openai_resource_url: Azure OpenAI resource URL for Knowledge Base model calls.
This is the OpenAI endpoint used by the Knowledge Base to call the LLM for
query planning and reasoning. This is separate from the project endpoint because
the Knowledge Base directly calls Azure OpenAI for its internal operations.
Required for agentic mode. Example: "https://myresource.openai.azure.com"
model_deployment_name: Model deployment name in Azure OpenAI for Knowledge Base.
This is the deployment name the Knowledge Base uses to call the LLM.
Required for agentic mode.
model_name: The underlying model name (e.g., "gpt-4o", "gpt-4o-mini").
If not provided, defaults to model_deployment_name. Used for Knowledge Base configuration.
knowledge_base_name: Name for the Knowledge Base. Required for agentic mode.
retrieval_instructions: Custom instructions for the Knowledge Base's
retrieval planning. Only used in agentic mode.
azure_openai_api_key: Azure OpenAI API key for Knowledge Base to call the model.
Only needed when using API key authentication instead of managed identity.
knowledge_base_output_mode: Output mode for Knowledge Base retrieval. Only used in agentic mode.
"extractive_data": Returns raw chunks without synthesis (default, recommended for agent integration).
"answer_synthesis": Returns synthesized answer from the LLM.
Some knowledge sources require answer_synthesis mode. Default: "extractive_data".
retrieval_reasoning_effort: Reasoning effort for Knowledge Base query planning. Only used in agentic mode.
"minimal": Fastest, basic query planning.
"medium": Moderate reasoning with some query decomposition.
"low": Lower reasoning effort than medium.
Default: "minimal".
agentic_message_history_count: Number of recent messages from conversation history to send to
the Knowledge Base. This context helps with query planning in agentic mode, allowing the
Knowledge Base to understand the conversation flow and generate better retrieval queries.
There is no technical limit - adjust based on your use case. Default: 10.
env_file_path: Path to environment file for loading settings.
env_file_encoding: Encoding of the environment file.
Examples:
.. code-block:: python
from agent_framework_aisearch import AzureAISearchContextProvider
from azure.identity.aio import DefaultAzureCredential
# Using environment variables
# Set AZURE_SEARCH_ENDPOINT=https://mysearch.search.windows.net
# Set AZURE_SEARCH_INDEX_NAME=my-index
credential = DefaultAzureCredential()
provider = AzureAISearchContextProvider(credential=credential)
# Or passing parameters directly
provider = AzureAISearchContextProvider(
endpoint="https://mysearch.search.windows.net",
index_name="my-index",
credential=credential,
)
# Or loading from a .env file
provider = AzureAISearchContextProvider(credential=credential, env_file_path="path/to/.env")
"""
# Load settings from environment/file
try:
settings = AzureAISearchSettings(
endpoint=endpoint,
index_name=index_name,
api_key=api_key if isinstance(api_key, str) else None,
env_file_path=env_file_path,
env_file_encoding=env_file_encoding,
)
except ValidationError as ex:
raise ServiceInitializationError("Failed to create Azure AI Search settings.", ex) from ex
# Validate required parameters
if not settings.endpoint:
raise ServiceInitializationError(
"Azure AI Search endpoint is required. Set via 'endpoint' parameter "
"or 'AZURE_SEARCH_ENDPOINT' environment variable."
)
if not settings.index_name:
raise ServiceInitializationError(
"Azure AI Search index name is required. Set via 'index_name' parameter "
"or 'AZURE_SEARCH_INDEX_NAME' environment variable."
)
# Determine the credential to use
resolved_credential: AzureKeyCredential | AsyncTokenCredential
if credential:
# AsyncTokenCredential takes precedence
resolved_credential = credential
elif isinstance(api_key, AzureKeyCredential):
resolved_credential = api_key
elif settings.api_key:
resolved_credential = AzureKeyCredential(settings.api_key.get_secret_value())
else:
raise ServiceInitializationError(
"Azure credential is required. Provide 'api_key' or 'credential' parameter "
"or set 'AZURE_SEARCH_API_KEY' environment variable."
)
self.endpoint = settings.endpoint
self.index_name = settings.index_name
self.credential = resolved_credential
self.mode = mode
self.top_k = top_k
self.semantic_configuration_name = semantic_configuration_name
self.vector_field_name = vector_field_name
self.embedding_function = embedding_function
self.context_prompt = context_prompt or self._DEFAULT_SEARCH_CONTEXT_PROMPT
# Agentic mode parameters (Knowledge Base)
self.azure_openai_resource_url = azure_openai_resource_url
self.azure_openai_deployment_name = model_deployment_name
# If model_name not provided, default to deployment name
self.model_name = model_name or model_deployment_name
self.knowledge_base_name = knowledge_base_name
self.retrieval_instructions = retrieval_instructions
self.azure_openai_api_key = azure_openai_api_key
self.azure_ai_project_endpoint = azure_ai_project_endpoint
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
# Auto-discover vector field if not specified
self._auto_discovered_vector_field = False
self._use_vectorizable_query = False # Will be set to True if server-side vectorization detected
if not vector_field_name and mode == "semantic":
# Attempt to auto-discover vector field from index schema
# This will be done lazily on first search to avoid blocking initialization
pass
# Validation
if vector_field_name and not embedding_function:
raise ValueError("embedding_function is required when vector_field_name is specified")
if mode == "agentic":
if not _agentic_retrieval_available:
raise ImportError(
"Agentic retrieval requires azure-search-documents >= 11.7.0b1 with Knowledge Base support. "
"Please upgrade: pip install azure-search-documents>=11.7.0b1"
)
if not self.azure_openai_resource_url:
raise ValueError(
"azure_openai_resource_url is required for agentic mode. "
"This should be your Azure OpenAI endpoint (e.g., 'https://myresource.openai.azure.com')"
)
if not self.azure_openai_deployment_name:
raise ValueError("model_deployment_name is required for agentic mode")
if not knowledge_base_name:
raise ValueError("knowledge_base_name is required for agentic mode")
# Create search client for semantic mode
self._search_client = SearchClient(
endpoint=self.endpoint,
index_name=self.index_name,
credential=self.credential,
)
# Create index client and retrieval client for agentic mode (Knowledge Base)
self._index_client: SearchIndexClient | None = None
self._retrieval_client: KnowledgeBaseRetrievalClient | None = None
if mode == "agentic":
self._index_client = SearchIndexClient(
endpoint=self.endpoint,
credential=self.credential,
)
# Retrieval client will be created after Knowledge Base initialization
self._knowledge_base_initialized = False
async def __aenter__(self) -> Self:
"""Async context manager entry."""
return self
async def __aexit__(
self,
exc_type: type[BaseException] | None,
exc_val: BaseException | None,
exc_tb: Any,
) -> None:
"""Async context manager exit - cleanup clients.
Args:
exc_type: Exception type if an error occurred.
exc_val: Exception value if an error occurred.
exc_tb: Exception traceback if an error occurred.
"""
# Close retrieval client if it was created
if self._retrieval_client is not None:
await self._retrieval_client.close()
self._retrieval_client = None
@override
async def invoking(
self,
messages: ChatMessage | MutableSequence[ChatMessage],
**kwargs: Any,
) -> Context:
"""Retrieve relevant context from Azure AI Search before model invocation.
Args:
messages: User messages to use for context retrieval.
**kwargs: Additional arguments (unused).
Returns:
Context object with retrieved documents as messages.
"""
# Convert to list and filter to USER/ASSISTANT messages with text only
messages_list = [messages] if isinstance(messages, ChatMessage) else list(messages)
filtered_messages = [
msg
for msg in messages_list
if msg and msg.text and msg.text.strip() and msg.role in [Role.USER, Role.ASSISTANT]
]
if not filtered_messages:
return Context()
# Perform search based on mode
if self.mode == "semantic":
# Semantic mode: flatten messages to single query
query = "\n".join(msg.text for msg in filtered_messages)
search_result_parts = await self._semantic_search(query)
else: # agentic
# Agentic mode: pass recent messages as conversation history
recent_messages = filtered_messages[-self.agentic_message_history_count :]
search_result_parts = await self._agentic_search(recent_messages)
# Format results as context - return multiple messages for each result part
if not search_result_parts:
return Context()
# Create context messages: first message with prompt, then one message per result part
context_messages = [ChatMessage(role=Role.USER, text=self.context_prompt)]
context_messages.extend([ChatMessage(role=Role.USER, text=part) for part in search_result_parts])
return Context(messages=context_messages)
def _find_vector_fields(self, index: Any) -> list[str]:
"""Find all fields that can store vectors (have dimensions defined).
Args:
index: SearchIndex object from Azure Search.
Returns:
List of vector field names.
"""
return [
field.name
for field in index.fields
if field.vector_search_dimensions is not None and field.vector_search_dimensions > 0
]
def _find_vectorizable_fields(self, index: Any, vector_fields: list[str]) -> list[str]:
"""Find vector fields that have auto-vectorization configured.
These are fields that have a vectorizer in their profile, meaning the index
can automatically vectorize text queries without needing a client-side embedding function.
Args:
index: SearchIndex object from Azure Search.
vector_fields: List of vector field names.
Returns:
List of vectorizable field names (subset of vector_fields).
"""
vectorizable_fields: list[str] = []
# Check if index has vector search configuration
if not index.vector_search or not index.vector_search.profiles:
return vectorizable_fields
# For each vector field, check if it has a vectorizer configured
for field in index.fields:
if field.name in vector_fields and field.vector_search_profile_name:
# Find the profile for this field
profile = next(
(p for p in index.vector_search.profiles if p.name == field.vector_search_profile_name), None
)
if profile and hasattr(profile, "vectorizer_name") and profile.vectorizer_name:
# This field has server-side vectorization configured
vectorizable_fields.append(field.name)
return vectorizable_fields
async def _auto_discover_vector_field(self) -> None:
"""Auto-discover vector field from index schema.
Attempts to find vector fields in the index and detect which have server-side
vectorization configured. Prioritizes vectorizable fields (which can auto-embed text)
over regular vector fields (which require client-side embedding).
"""
if self._auto_discovered_vector_field or self.vector_field_name:
return # Already discovered or manually specified
try:
# Use existing index client or create temporary one
if not self._index_client:
self._index_client = SearchIndexClient(endpoint=self.endpoint, credential=self.credential)
index_client = self._index_client
# Get index schema
index = await index_client.get_index(self.index_name)
# Step 1: Find all vector fields
vector_fields = self._find_vector_fields(index)
if not vector_fields:
# No vector fields found - keyword search only
logger.info(f"No vector fields found in index '{self.index_name}'. Using keyword-only search.")
self._auto_discovered_vector_field = True
return
# Step 2: Find which vector fields have server-side vectorization
vectorizable_fields = self._find_vectorizable_fields(index, vector_fields)
# Step 3: Decide which field to use
if vectorizable_fields:
# Prefer vectorizable fields (server-side embedding)
if len(vectorizable_fields) == 1:
self.vector_field_name = vectorizable_fields[0]
self._auto_discovered_vector_field = True
self._use_vectorizable_query = True # Use VectorizableTextQuery
logger.info(
f"Auto-discovered vectorizable field '{self.vector_field_name}' "
f"with server-side vectorization. No embedding_function needed."
)
else:
# Multiple vectorizable fields
logger.warning(
f"Multiple vectorizable fields found: {vectorizable_fields}. "
f"Please specify vector_field_name explicitly. Using keyword-only search."
)
elif len(vector_fields) == 1:
# Single vector field without vectorizer - needs client-side embedding
self.vector_field_name = vector_fields[0]
self._auto_discovered_vector_field = True
self._use_vectorizable_query = False
if not self.embedding_function:
logger.warning(
f"Auto-discovered vector field '{self.vector_field_name}' without server-side vectorization. "
f"Provide embedding_function for vector search, or it will fall back to keyword-only search."
)
self.vector_field_name = None
else:
# Multiple vector fields without vectorizers
logger.warning(
f"Multiple vector fields found: {vector_fields}. "
f"Please specify vector_field_name explicitly. Using keyword-only search."
)
except Exception as e:
# Log warning but continue with keyword search
logger.warning(f"Failed to auto-discover vector field: {e}. Using keyword-only search.")
self._auto_discovered_vector_field = True # Mark as attempted
async def _semantic_search(self, query: str) -> list[str]:
"""Perform semantic hybrid search with semantic ranking.
This is the recommended mode for most use cases. It combines:
- Vector search (if embedding_function provided)
- Keyword search (BM25)
- Semantic reranking (if semantic_configuration_name provided)
Args:
query: Search query text.
Returns:
List of formatted search result strings, one per document.
"""
# Auto-discover vector field if not already done
await self._auto_discover_vector_field()
vector_queries: list[VectorizableTextQuery | VectorizedQuery] = []
# Build vector query based on server-side vectorization or client-side embedding
if self.vector_field_name:
# Use larger k for vector query when semantic reranker is enabled for better ranking quality
vector_k = max(self.top_k, 50) if self.semantic_configuration_name else self.top_k
if self._use_vectorizable_query:
# Server-side vectorization: Index will auto-embed the text query
vector_queries = [
VectorizableTextQuery(
text=query,
k_nearest_neighbors=vector_k,
fields=self.vector_field_name,
)
]
elif self.embedding_function:
# Client-side embedding: We provide the vector
query_vector = await self.embedding_function(query)
vector_queries = [
VectorizedQuery(
vector=query_vector,
k_nearest_neighbors=vector_k,
fields=self.vector_field_name,
)
]
# else: vector_field_name is set but no vectorization available - skip vector search
# Build search parameters
search_params: dict[str, Any] = {
"search_text": query,
"top": self.top_k,
}
if vector_queries:
search_params["vector_queries"] = vector_queries
# Add semantic ranking if configured
if self.semantic_configuration_name:
search_params["query_type"] = QueryType.SEMANTIC
search_params["semantic_configuration_name"] = self.semantic_configuration_name
search_params["query_caption"] = QueryCaptionType.EXTRACTIVE
# Execute search
results = await self._search_client.search(**search_params) # type: ignore[reportUnknownVariableType]
# Format results with citations
formatted_results: list[str] = []
async for doc in results: # type: ignore[reportUnknownVariableType]
# Extract document ID for citation
doc_id = doc.get("id") or doc.get("@search.id") # type: ignore[reportUnknownVariableType]
# Use full document chunks with citation
doc_text: str = self._extract_document_text(doc, doc_id=doc_id) # type: ignore[reportUnknownArgumentType]
if doc_text:
formatted_results.append(doc_text) # type: ignore[reportUnknownArgumentType]
return formatted_results
async def _ensure_knowledge_base(self) -> None:
"""Ensure Knowledge Base and knowledge source are created.
This method is idempotent - it will only create resources if they don't exist.
Note: Azure SDK uses KnowledgeAgent classes internally, but the feature
is marketed as "Knowledge Bases" in Azure AI Search.
"""
if self._knowledge_base_initialized or not self._index_client:
return
# Runtime validation for agentic mode parameters
if not self.knowledge_base_name:
raise ValueError("knowledge_base_name is required for agentic mode")
if not self.azure_openai_resource_url:
raise ValueError("azure_openai_resource_url is required for agentic mode")
if not self.azure_openai_deployment_name:
raise ValueError("model_deployment_name is required for agentic mode")
knowledge_base_name = self.knowledge_base_name
# Step 1: Create or get knowledge source
knowledge_source_name = f"{self.index_name}-source"
try:
# Try to get existing knowledge source
await self._index_client.get_knowledge_source(knowledge_source_name)
except ResourceNotFoundError:
# Create new knowledge source if it doesn't exist
knowledge_source = SearchIndexKnowledgeSource(
name=knowledge_source_name,
description=f"Knowledge source for {self.index_name} search index",
search_index_parameters=SearchIndexKnowledgeSourceParameters(
search_index_name=self.index_name,
),
)
await self._index_client.create_knowledge_source(knowledge_source)
# Step 2: Create or update Knowledge Base
# Always create/update to ensure configuration is current
aoai_params = AzureOpenAIVectorizerParameters(
resource_url=self.azure_openai_resource_url,
deployment_name=self.azure_openai_deployment_name,
model_name=self.model_name,
api_key=self.azure_openai_api_key,
)
# Map output mode string to SDK enum
output_mode = (
KnowledgeRetrievalOutputMode.EXTRACTIVE_DATA
if self.knowledge_base_output_mode == "extractive_data"
else KnowledgeRetrievalOutputMode.ANSWER_SYNTHESIS
)
# Map reasoning effort string to SDK class
reasoning_effort_map: dict[str, KnowledgeRetrievalReasoningEffort] = {
"minimal": KnowledgeRetrievalMinimalReasoningEffort(),
"medium": KnowledgeRetrievalMediumReasoningEffort(),
"low": KnowledgeRetrievalLowReasoningEffort(),
}
reasoning_effort = reasoning_effort_map[self.retrieval_reasoning_effort]
knowledge_base = KnowledgeBase(
name=knowledge_base_name,
description=f"Knowledge Base for multi-hop retrieval across {self.index_name}",
knowledge_sources=[
KnowledgeSourceReference(
name=knowledge_source_name,
)
],
models=[KnowledgeBaseAzureOpenAIModel(azure_open_ai_parameters=aoai_params)],
output_mode=output_mode,
retrieval_reasoning_effort=reasoning_effort,
)
await self._index_client.create_or_update_knowledge_base(knowledge_base)
self._knowledge_base_initialized = True
# Create retrieval client now that Knowledge Base is initialized
if _agentic_retrieval_available and self._retrieval_client is None:
self._retrieval_client = KnowledgeBaseRetrievalClient(
endpoint=self.endpoint,
knowledge_base_name=knowledge_base_name,
credential=self.credential,
)
async def _agentic_search(self, messages: list[ChatMessage]) -> list[str]:
"""Perform agentic retrieval with multi-hop reasoning using Knowledge Bases.
This mode uses query planning and is slightly slower than semantic search,
but provides more accurate results through intelligent retrieval.
This method uses Azure AI Search Knowledge Bases which:
1. Analyze the query and plan sub-queries
2. Retrieve relevant documents across multiple sources
3. Perform multi-hop reasoning with an LLM
4. Synthesize a comprehensive answer with references
Args:
messages: Conversation history to use for retrieval context.
Returns:
List of answer parts from the Knowledge Base, one per content item.
"""
# Ensure Knowledge Base is initialized
await self._ensure_knowledge_base()
# Map reasoning effort string to SDK class (for retrieval requests)
reasoning_effort_map: dict[str, KBRetrievalReasoningEffort] = {
"minimal": KBRetrievalMinimalReasoningEffort(),
"medium": KBRetrievalMediumReasoningEffort(),
"low": KBRetrievalLowReasoningEffort(),
}
reasoning_effort = reasoning_effort_map[self.retrieval_reasoning_effort]
# Map output mode string to SDK enum (for retrieval requests)
output_mode = (
KBRetrievalOutputMode.EXTRACTIVE_DATA
if self.knowledge_base_output_mode == "extractive_data"
else KBRetrievalOutputMode.ANSWER_SYNTHESIS
)
# For minimal reasoning, use intents API; for medium/low, use messages API
if self.retrieval_reasoning_effort == "minimal":
# Minimal reasoning uses intents with a single search query
query = "\n".join(msg.text for msg in messages if msg.text)
intents: list[KnowledgeRetrievalIntent] = [KnowledgeRetrievalSemanticIntent(search=query)]
retrieval_request = KnowledgeBaseRetrievalRequest(
intents=intents,
retrieval_reasoning_effort=reasoning_effort,
output_mode=output_mode,
include_activity=True,
)
else:
# Medium/low reasoning uses messages with conversation history
kb_messages = [
KnowledgeBaseMessage(
role=msg.role.value if hasattr(msg.role, "value") else str(msg.role),
content=[KnowledgeBaseMessageTextContent(text=msg.text)],
)
for msg in messages
if msg.text
]
retrieval_request = KnowledgeBaseRetrievalRequest(
messages=kb_messages,
retrieval_reasoning_effort=reasoning_effort,
output_mode=output_mode,
include_activity=True,
)
# Use reusable retrieval client
if not self._retrieval_client:
raise RuntimeError("Retrieval client not initialized. Ensure Knowledge Base is set up correctly.")
# Perform retrieval via Knowledge Base
retrieval_result = await self._retrieval_client.retrieve(retrieval_request=retrieval_request)
# Extract answer parts from response
if retrieval_result.response and len(retrieval_result.response) > 0:
# Get the assistant's response (last message)
assistant_message = retrieval_result.response[-1]
if assistant_message.content:
# Extract all text content items as separate parts
answer_parts: list[str] = []
for content_item in assistant_message.content:
# Check if this is a text content item
if isinstance(content_item, KnowledgeBaseMessageTextContent) and content_item.text:
answer_parts.append(content_item.text)
if answer_parts:
return answer_parts
# Fallback if no answer generated
return ["No results found from Knowledge Base."]
def _extract_document_text(self, doc: dict[str, Any], doc_id: str | None = None) -> str:
"""Extract readable text from a search document with optional citation.
Args:
doc: Search result document.
doc_id: Optional document ID for citation.
Returns:
Formatted document text with citation if doc_id provided.
"""
# Try common text field names
text = ""
for field in ["content", "text", "description", "body", "chunk"]:
if doc.get(field):
text = str(doc[field])
break
# Fallback: concatenate all string fields
if not text:
text_parts: list[str] = []
for key, value in doc.items():
if isinstance(value, str) and not key.startswith("@") and key != "id":
text_parts.append(f"{key}: {value}")
text = " | ".join(text_parts) if text_parts else ""
# Add citation if document ID provided
if doc_id and text:
return f"[Source: {doc_id}] {text}"
return text
-91
View File
@@ -1,91 +0,0 @@
[project]
name = "agent-framework-aisearch"
description = "Azure AI Search integration for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251118"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
urls.release_notes = "https://github.com/microsoft/agent-framework/releases?q=tag%3Apython-1&expanded=true"
urls.issues = "https://github.com/microsoft/agent-framework/issues"
classifiers = [
"License :: OSI Approved :: MIT License",
"Development Status :: 4 - Beta",
"Intended Audience :: Developers",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
"Programming Language :: Python :: 3.14",
"Typing :: Typed",
]
dependencies = [
"agent-framework-core",
"azure-search-documents==11.7.0b2",
]
[tool.uv]
prerelease = "if-necessary-or-explicit"
environments = [
"sys_platform == 'darwin'",
"sys_platform == 'linux'",
"sys_platform == 'win32'"
]
[tool.uv-dynamic-versioning]
fallback-version = "0.0.0"
[tool.pytest.ini_options]
testpaths = 'tests'
addopts = "-ra -q -r fEX"
asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "function"
filterwarnings = [
"ignore:Support for class-based `config` is deprecated:DeprecationWarning:pydantic.*"
]
timeout = 120
[tool.ruff]
extend = "../../pyproject.toml"
exclude = ["examples"]
[tool.coverage.run]
omit = [
"**/__init__.py"
]
[tool.pyright]
extends = "../../pyproject.toml"
exclude = ['tests']
[tool.mypy]
plugins = ['pydantic.mypy']
strict = true
python_version = "3.10"
ignore_missing_imports = true
disallow_untyped_defs = true
no_implicit_optional = true
check_untyped_defs = true
warn_return_any = true
show_error_codes = true
warn_unused_ignores = false
disallow_incomplete_defs = true
disallow_untyped_decorators = true
[tool.bandit]
targets = ["agent_framework_aisearch"]
exclude_dirs = ["tests"]
[tool.poe]
executor.type = "uv"
include = "../../shared_tasks.toml"
[tool.poe.tasks]
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_aisearch"
test = "pytest --cov=agent_framework_aisearch --cov-report=term-missing:skip-covered tests"
[build-system]
requires = ["flit-core >= 3.11,<4.0"]
build-backend = "flit_core.buildapi"
@@ -1,992 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
# pyright: reportPrivateUsage=false
import os
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
from agent_framework import ChatMessage, Context, Role
from agent_framework.azure import AzureAISearchContextProvider
from agent_framework.exceptions import ServiceInitializationError
from azure.core.credentials import AzureKeyCredential
from azure.core.exceptions import ResourceNotFoundError
from agent_framework_aisearch import AzureAISearchSettings
@pytest.fixture
def mock_search_client() -> AsyncMock:
"""Create a mock SearchClient."""
mock_client = AsyncMock()
mock_client.search = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock()
return mock_client
@pytest.fixture
def mock_index_client() -> AsyncMock:
"""Create a mock SearchIndexClient."""
mock_client = AsyncMock()
mock_client.get_knowledge_source = AsyncMock()
mock_client.create_knowledge_source = AsyncMock()
mock_client.get_agent = AsyncMock()
mock_client.create_agent = AsyncMock()
mock_client.__aenter__ = AsyncMock(return_value=mock_client)
mock_client.__aexit__ = AsyncMock()
return mock_client
@pytest.fixture
def sample_messages() -> list[ChatMessage]:
"""Create sample chat messages for testing."""
return [
ChatMessage(role=Role.USER, text="What is in the documents?"),
]
class TestAzureAISearchSettings:
"""Test AzureAISearchSettings configuration."""
def test_settings_with_direct_values(self) -> None:
"""Test settings with direct values."""
settings = AzureAISearchSettings(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
)
assert settings.endpoint == "https://test.search.windows.net"
assert settings.index_name == "test-index"
# api_key is now SecretStr
assert settings.api_key.get_secret_value() == "test-key"
def test_settings_with_env_file_path(self) -> None:
"""Test settings with env_file_path parameter."""
settings = AzureAISearchSettings(
endpoint="https://test.search.windows.net",
index_name="test-index",
env_file_path="test.env",
)
assert settings.endpoint == "https://test.search.windows.net"
assert settings.index_name == "test-index"
def test_provider_uses_settings_from_env(self) -> None:
"""Test that provider creates settings internally from env."""
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
)
assert provider.endpoint == "https://test.search.windows.net"
assert provider.index_name == "test-index"
def test_provider_missing_endpoint_raises_error(self) -> None:
"""Test that provider raises ServiceInitializationError without endpoint."""
# Use patch.dict to clear environment and pass env_file_path="" to prevent .env file loading
clean_env = {k: v for k, v in os.environ.items() if not k.startswith("AZURE_SEARCH_")}
with (
patch.dict(os.environ, clean_env, clear=True),
pytest.raises(ServiceInitializationError, match="endpoint is required"),
):
AzureAISearchContextProvider(
index_name="test-index",
api_key="test-key",
env_file_path="", # Disable .env file loading
)
def test_provider_missing_index_name_raises_error(self) -> None:
"""Test that provider raises ServiceInitializationError without index_name."""
# Use patch.dict to clear environment and pass env_file_path="" to prevent .env file loading
clean_env = {k: v for k, v in os.environ.items() if not k.startswith("AZURE_SEARCH_")}
with (
patch.dict(os.environ, clean_env, clear=True),
pytest.raises(ServiceInitializationError, match="index name is required"),
):
AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
api_key="test-key",
env_file_path="", # Disable .env file loading
)
def test_provider_missing_credential_raises_error(self) -> None:
"""Test that provider raises ServiceInitializationError without credential."""
# Use patch.dict to clear environment and pass env_file_path="" to prevent .env file loading
clean_env = {k: v for k, v in os.environ.items() if not k.startswith("AZURE_SEARCH_")}
with (
patch.dict(os.environ, clean_env, clear=True),
pytest.raises(ServiceInitializationError, match="credential is required"),
):
AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
env_file_path="", # Disable .env file loading
)
class TestSearchProviderInitialization:
"""Test initialization and configuration of AzureAISearchContextProvider."""
def test_init_semantic_mode_minimal(self) -> None:
"""Test initialization with minimal semantic mode parameters."""
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
assert provider.endpoint == "https://test.search.windows.net"
assert provider.index_name == "test-index"
assert provider.mode == "semantic"
assert provider.top_k == 5
def test_init_semantic_mode_with_vector_field_requires_embedding_function(self) -> None:
"""Test that vector_field_name requires embedding_function."""
with pytest.raises(ValueError, match="embedding_function is required"):
AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
vector_field_name="embedding",
)
def test_init_agentic_mode_requires_azure_openai_resource_url(self) -> None:
"""Test that agentic mode requires azure_openai_resource_url."""
with pytest.raises(ValueError, match="azure_openai_resource_url"):
AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
)
def test_init_agentic_mode_requires_model_deployment_name(self) -> None:
"""Test that agentic mode requires model_deployment_name."""
with pytest.raises(ValueError, match="model_deployment_name"):
AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
azure_openai_resource_url="https://test.openai.azure.com",
)
def test_init_agentic_mode_requires_knowledge_base_name(self) -> None:
"""Test that agentic mode requires knowledge_base_name."""
with pytest.raises(ValueError, match="knowledge_base_name"):
AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="gpt-4o",
azure_openai_resource_url="https://test.openai.azure.com",
)
def test_init_agentic_mode_with_all_params(self) -> None:
"""Test initialization with all agentic mode parameters."""
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="my-gpt-4o-deployment",
model_name="gpt-4o",
knowledge_base_name="test-kb",
azure_openai_resource_url="https://test.openai.azure.com",
)
assert provider.mode == "agentic"
assert provider.azure_ai_project_endpoint == "https://test.services.ai.azure.com"
assert provider.azure_openai_resource_url == "https://test.openai.azure.com"
assert provider.azure_openai_deployment_name == "my-gpt-4o-deployment"
assert provider.model_name == "gpt-4o"
assert provider.knowledge_base_name == "test-kb"
def test_init_model_name_defaults_to_deployment_name(self) -> None:
"""Test that model_name defaults to deployment_name if not provided."""
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="gpt-4o",
knowledge_base_name="test-kb",
azure_openai_resource_url="https://test.openai.azure.com",
)
assert provider.model_name == "gpt-4o"
def test_init_with_custom_context_prompt(self) -> None:
"""Test initialization with custom context prompt."""
custom_prompt = "Use the following information:"
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
context_prompt=custom_prompt,
)
assert provider.context_prompt == custom_prompt
def test_init_uses_default_context_prompt(self) -> None:
"""Test that default context prompt is used when not provided."""
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
assert provider.context_prompt == provider._DEFAULT_SEARCH_CONTEXT_PROMPT
class TestSemanticSearch:
"""Test semantic search functionality."""
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_semantic_search_basic(
self, mock_search_class: MagicMock, sample_messages: list[ChatMessage]
) -> None:
"""Test basic semantic search without vector search."""
# Setup mock
mock_search_client = AsyncMock()
mock_results = AsyncMock()
mock_results.__aiter__.return_value = iter([{"content": "Test document content"}])
mock_search_client.search.return_value = mock_results
mock_search_class.return_value = mock_search_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
context = await provider.invoking(sample_messages)
assert isinstance(context, Context)
assert len(context.messages) > 1 # First message is prompt, rest are results
# First message should be the context prompt
assert "Use the following context" in context.messages[0].text
# Second message should contain the search result
assert "Test document content" in context.messages[1].text
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_semantic_search_empty_query(self, mock_search_class: MagicMock) -> None:
"""Test that empty queries return empty context."""
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
# Empty message
context = await provider.invoking([ChatMessage(role=Role.USER, text="")])
assert isinstance(context, Context)
assert len(context.messages) == 0
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_semantic_search_with_vector_query(
self, mock_search_class: MagicMock, sample_messages: list[ChatMessage]
) -> None:
"""Test semantic search with vector query."""
# Setup mock
mock_search_client = AsyncMock()
mock_results = AsyncMock()
mock_results.__aiter__.return_value = iter([{"content": "Vector search result"}])
mock_search_client.search.return_value = mock_results
mock_search_class.return_value = mock_search_client
# Mock embedding function
async def mock_embed(text: str) -> list[float]:
return [0.1, 0.2, 0.3]
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
vector_field_name="embedding",
embedding_function=mock_embed,
)
context = await provider.invoking(sample_messages)
assert isinstance(context, Context)
assert len(context.messages) > 0
# Verify that search was called
mock_search_client.search.assert_called_once()
class TestKnowledgeBaseSetup:
"""Test Knowledge Base setup for agentic mode."""
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_ensure_knowledge_base_creates_when_not_exists(
self, mock_search_class: MagicMock, mock_index_class: MagicMock
) -> None:
"""Test that Knowledge Base is created when it doesn't exist."""
# Setup mocks
mock_index_client = AsyncMock()
mock_index_client.get_knowledge_source.side_effect = ResourceNotFoundError("Not found")
mock_index_client.create_knowledge_source = AsyncMock()
mock_index_client.get_knowledge_base.side_effect = ResourceNotFoundError("Not found")
mock_index_client.create_or_update_knowledge_base = AsyncMock()
mock_index_class.return_value = mock_index_client
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="gpt-4o",
model_name="gpt-4o",
knowledge_base_name="test-kb",
azure_openai_resource_url="https://test.openai.azure.com",
)
await provider._ensure_knowledge_base()
# Verify knowledge source was created
mock_index_client.create_knowledge_source.assert_called_once()
# Verify Knowledge Base was created
mock_index_client.create_or_update_knowledge_base.assert_called_once()
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_ensure_knowledge_base_skips_when_exists(
self, mock_search_class: MagicMock, mock_index_class: MagicMock
) -> None:
"""Test that Knowledge Base setup is skipped when already exists."""
# Setup mocks
mock_index_client = AsyncMock()
mock_index_client.get_knowledge_source.return_value = MagicMock() # Exists
mock_index_client.get_knowledge_base.return_value = MagicMock() # Exists
mock_index_class.return_value = mock_index_client
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="gpt-4o",
knowledge_base_name="test-kb",
azure_openai_resource_url="https://test.openai.azure.com",
)
await provider._ensure_knowledge_base()
# Verify nothing was created
mock_index_client.create_knowledge_source.assert_not_called()
mock_index_client.create_agent.assert_not_called()
class TestContextProviderLifecycle:
"""Test context provider lifecycle methods."""
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_context_manager(self, mock_search_class: MagicMock) -> None:
"""Test that provider can be used as async context manager."""
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
async with AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
) as provider:
assert provider is not None
assert isinstance(provider, AzureAISearchContextProvider)
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.KnowledgeBaseRetrievalClient")
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_context_manager_agentic_cleanup(
self, mock_search_class: MagicMock, mock_index_class: MagicMock, mock_retrieval_class: MagicMock
) -> None:
"""Test that agentic mode provider cleans up retrieval client."""
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
mock_index_client = AsyncMock()
mock_index_class.return_value = mock_index_client
mock_retrieval_client = AsyncMock()
mock_retrieval_client.close = AsyncMock()
mock_retrieval_class.return_value = mock_retrieval_client
async with AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="gpt-4o",
knowledge_base_name="test-kb",
azure_openai_resource_url="https://test.openai.azure.com",
) as provider:
# Simulate retrieval client being created
provider._retrieval_client = mock_retrieval_client
# Verify cleanup was called
mock_retrieval_client.close.assert_called_once()
def test_string_api_key_conversion(self) -> None:
"""Test that string api_key is converted to AzureKeyCredential."""
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="my-api-key", # String api_key
mode="semantic",
)
assert isinstance(provider.credential, AzureKeyCredential)
class TestMessageFiltering:
"""Test message filtering functionality."""
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_filters_non_user_assistant_messages(self, mock_search_class: MagicMock) -> None:
"""Test that only USER and ASSISTANT messages are processed."""
# Setup mock
mock_search_client = AsyncMock()
mock_results = AsyncMock()
mock_results.__aiter__.return_value = iter([{"content": "Test result"}])
mock_search_client.search.return_value = mock_results
mock_search_class.return_value = mock_search_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
# Mix of message types
messages = [
ChatMessage(role=Role.SYSTEM, text="System message"),
ChatMessage(role=Role.USER, text="User message"),
ChatMessage(role=Role.ASSISTANT, text="Assistant message"),
ChatMessage(role=Role.TOOL, text="Tool message"),
]
context = await provider.invoking(messages)
# Should have processed only USER and ASSISTANT messages
assert isinstance(context, Context)
mock_search_client.search.assert_called_once()
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_filters_empty_messages(self, mock_search_class: MagicMock) -> None:
"""Test that empty/whitespace messages are filtered out."""
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
# Messages with empty/whitespace text
messages = [
ChatMessage(role=Role.USER, text=""),
ChatMessage(role=Role.USER, text=" "),
ChatMessage(role=Role.USER, text=None),
]
context = await provider.invoking(messages)
# Should return empty context
assert len(context.messages) == 0
class TestCitations:
"""Test citation functionality."""
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_citations_included_in_semantic_search(self, mock_search_class: MagicMock) -> None:
"""Test that citations are included in semantic search results."""
# Setup mock with document ID
mock_search_client = AsyncMock()
mock_results = AsyncMock()
mock_doc = {"id": "doc123", "content": "Test document content"}
mock_results.__aiter__.return_value = iter([mock_doc])
mock_search_client.search.return_value = mock_results
mock_search_class.return_value = mock_search_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
context = await provider.invoking([ChatMessage(role=Role.USER, text="test query")])
# Check that citation is included
assert isinstance(context, Context)
assert len(context.messages) > 1 # First message is prompt, rest are results
# Citation should be in the result message (second message)
assert "[Source: doc123]" in context.messages[1].text
assert "Test document content" in context.messages[1].text
class TestAgenticSearch:
"""Test agentic search functionality."""
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.KnowledgeBaseRetrievalClient")
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_agentic_search_basic(
self,
mock_search_class: MagicMock,
mock_index_class: MagicMock,
mock_retrieval_class: MagicMock,
sample_messages: list[ChatMessage],
) -> None:
"""Test basic agentic search with Knowledge Base retrieval."""
# Setup search client mock
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
# Setup index client mock
mock_index_client = AsyncMock()
mock_index_client.get_knowledge_source.side_effect = ResourceNotFoundError("Not found")
mock_index_client.create_knowledge_source = AsyncMock()
mock_index_client.create_or_update_knowledge_base = AsyncMock()
mock_index_class.return_value = mock_index_client
# Setup retrieval client mock with response
mock_retrieval_client = AsyncMock()
mock_response = MagicMock()
mock_message = MagicMock()
mock_content = MagicMock()
mock_content.text = "Agentic search result"
# Make it pass isinstance check
from agent_framework_aisearch._search_provider import _agentic_retrieval_available
if _agentic_retrieval_available:
from azure.search.documents.knowledgebases.models import KnowledgeBaseMessageTextContent
mock_content.__class__ = KnowledgeBaseMessageTextContent
mock_message.content = [mock_content]
mock_response.response = [mock_message]
mock_retrieval_client.retrieve.return_value = mock_response
mock_retrieval_client.close = AsyncMock()
mock_retrieval_class.return_value = mock_retrieval_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="gpt-4o",
knowledge_base_name="test-kb",
azure_openai_resource_url="https://test.openai.azure.com",
)
context = await provider.invoking(sample_messages)
assert isinstance(context, Context)
# Should have at least the prompt message
assert len(context.messages) >= 1
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.KnowledgeBaseRetrievalClient")
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_agentic_search_no_results(
self,
mock_search_class: MagicMock,
mock_index_class: MagicMock,
mock_retrieval_class: MagicMock,
sample_messages: list[ChatMessage],
) -> None:
"""Test agentic search when no results are returned."""
# Setup mocks
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
mock_index_client = AsyncMock()
mock_index_client.get_knowledge_source.side_effect = ResourceNotFoundError("Not found")
mock_index_client.create_knowledge_source = AsyncMock()
mock_index_client.create_or_update_knowledge_base = AsyncMock()
mock_index_class.return_value = mock_index_client
# Empty response
mock_retrieval_client = AsyncMock()
mock_response = MagicMock()
mock_response.response = []
mock_retrieval_client.retrieve.return_value = mock_response
mock_retrieval_client.close = AsyncMock()
mock_retrieval_class.return_value = mock_retrieval_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="gpt-4o",
knowledge_base_name="test-kb",
azure_openai_resource_url="https://test.openai.azure.com",
)
context = await provider.invoking(sample_messages)
assert isinstance(context, Context)
# Should have fallback message
assert len(context.messages) >= 1
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.KnowledgeBaseRetrievalClient")
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_agentic_search_with_medium_reasoning(
self,
mock_search_class: MagicMock,
mock_index_class: MagicMock,
mock_retrieval_class: MagicMock,
sample_messages: list[ChatMessage],
) -> None:
"""Test agentic search with medium reasoning effort."""
# Setup mocks
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
mock_index_client = AsyncMock()
mock_index_client.get_knowledge_source.side_effect = ResourceNotFoundError("Not found")
mock_index_client.create_knowledge_source = AsyncMock()
mock_index_client.create_or_update_knowledge_base = AsyncMock()
mock_index_class.return_value = mock_index_client
mock_retrieval_client = AsyncMock()
mock_response = MagicMock()
mock_message = MagicMock()
mock_content = MagicMock()
mock_content.text = "Medium reasoning result"
from agent_framework_aisearch._search_provider import _agentic_retrieval_available
if _agentic_retrieval_available:
from azure.search.documents.knowledgebases.models import KnowledgeBaseMessageTextContent
mock_content.__class__ = KnowledgeBaseMessageTextContent
mock_message.content = [mock_content]
mock_response.response = [mock_message]
mock_retrieval_client.retrieve.return_value = mock_response
mock_retrieval_client.close = AsyncMock()
mock_retrieval_class.return_value = mock_retrieval_client
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="agentic",
azure_ai_project_endpoint="https://test.services.ai.azure.com",
model_deployment_name="gpt-4o",
knowledge_base_name="test-kb",
azure_openai_resource_url="https://test.openai.azure.com",
retrieval_reasoning_effort="medium", # Test medium reasoning
)
context = await provider.invoking(sample_messages)
assert isinstance(context, Context)
assert len(context.messages) >= 1
class TestVectorFieldAutoDiscovery:
"""Test vector field auto-discovery functionality."""
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_auto_discovers_single_vector_field(
self, mock_search_class: MagicMock, mock_index_class: MagicMock
) -> None:
"""Test that single vector field is auto-discovered."""
# Setup search client mock
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
# Setup index client mock
mock_index_client = AsyncMock()
mock_index = MagicMock()
# Create mock field with vector_search_dimensions attribute
mock_vector_field = MagicMock()
mock_vector_field.name = "embedding_vector"
mock_vector_field.vector_search_dimensions = 1536
mock_index.fields = [mock_vector_field]
mock_index_client.get_index.return_value = mock_index
mock_index_client.close = AsyncMock()
mock_index_class.return_value = mock_index_client
# Create provider without specifying vector_field_name
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
# Trigger auto-discovery
await provider._auto_discover_vector_field()
# Vector field should be auto-discovered but not used without embedding function
assert provider._auto_discovered_vector_field is True
# Should be cleared since no embedding function
assert provider.vector_field_name is None
@pytest.mark.asyncio
async def test_vector_detection_accuracy(self) -> None:
"""Test that vector field detection logic correctly identifies vector fields."""
from azure.search.documents.indexes.models import SearchField
# Create real SearchField objects to test the detection logic
vector_field = SearchField(
name="embedding_vector", type="Collection(Edm.Single)", vector_search_dimensions=1536, searchable=True
)
string_field = SearchField(name="content", type="Edm.String", searchable=True)
number_field = SearchField(name="price", type="Edm.Double", filterable=True)
# Test detection logic directly
is_vector_1 = vector_field.vector_search_dimensions is not None and vector_field.vector_search_dimensions > 0
is_vector_2 = string_field.vector_search_dimensions is not None and string_field.vector_search_dimensions > 0
is_vector_3 = number_field.vector_search_dimensions is not None and number_field.vector_search_dimensions > 0
# Only the vector field should be detected
assert is_vector_1 is True
assert is_vector_2 is False
assert is_vector_3 is False
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_no_false_positives_on_string_fields(
self, mock_search_class: MagicMock, mock_index_class: MagicMock
) -> None:
"""Test that regular string fields are not detected as vector fields."""
# Setup search client mock
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
# Setup index with only string fields (no vectors)
mock_index_client = AsyncMock()
mock_index = MagicMock()
# All fields have vector_search_dimensions = None
mock_fields = []
for name in ["id", "title", "content", "category"]:
field = MagicMock()
field.name = name
field.vector_search_dimensions = None
field.vector_search_profile_name = None
mock_fields.append(field)
mock_index.fields = mock_fields
mock_index_client.get_index.return_value = mock_index
mock_index_client.close = AsyncMock()
mock_index_class.return_value = mock_index_client
# Create provider
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
# Trigger auto-discovery
await provider._auto_discover_vector_field()
# Should NOT detect any vector fields
assert provider.vector_field_name is None
assert provider._auto_discovered_vector_field is True
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_multiple_vector_fields_without_vectorizer(
self, mock_search_class: MagicMock, mock_index_class: MagicMock
) -> None:
"""Test that multiple vector fields without vectorizer logs warning and uses keyword search."""
# Setup search client mock
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
# Setup index with multiple vector fields (no vectorizers)
mock_index_client = AsyncMock()
mock_index = MagicMock()
# Multiple vector fields
mock_fields = []
for name in ["embedding1", "embedding2"]:
field = MagicMock()
field.name = name
field.vector_search_dimensions = 1536
field.vector_search_profile_name = None # No vectorizer
mock_fields.append(field)
mock_index.fields = mock_fields
mock_index.vector_search = None # No vector search config
mock_index_client.get_index.return_value = mock_index
mock_index_client.close = AsyncMock()
mock_index_class.return_value = mock_index_client
# Create provider
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
# Trigger auto-discovery
await provider._auto_discover_vector_field()
# Should NOT use any vector field (multiple fields, can't choose)
assert provider.vector_field_name is None
assert provider._auto_discovered_vector_field is True
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_multiple_vectorizable_fields(
self, mock_search_class: MagicMock, mock_index_class: MagicMock
) -> None:
"""Test that multiple vectorizable fields logs warning and uses keyword search."""
# Setup search client mock
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
# Setup index with multiple vectorizable fields
mock_index_client = AsyncMock()
mock_index = MagicMock()
# Multiple vector fields with vectorizers
mock_fields = []
for name in ["embedding1", "embedding2"]:
field = MagicMock()
field.name = name
field.vector_search_dimensions = 1536
field.vector_search_profile_name = f"{name}-profile"
mock_fields.append(field)
mock_index.fields = mock_fields
# Setup vector search config with profiles that have vectorizers
mock_profile1 = MagicMock()
mock_profile1.name = "embedding1-profile"
mock_profile1.vectorizer_name = "vectorizer1"
mock_profile2 = MagicMock()
mock_profile2.name = "embedding2-profile"
mock_profile2.vectorizer_name = "vectorizer2"
mock_index.vector_search = MagicMock()
mock_index.vector_search.profiles = [mock_profile1, mock_profile2]
mock_index_client.get_index.return_value = mock_index
mock_index_client.close = AsyncMock()
mock_index_class.return_value = mock_index_client
# Create provider
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
# Trigger auto-discovery
await provider._auto_discover_vector_field()
# Should NOT use any vector field (multiple vectorizable fields, can't choose)
assert provider.vector_field_name is None
assert provider._auto_discovered_vector_field is True
@pytest.mark.asyncio
@patch("agent_framework_aisearch._search_provider.SearchIndexClient")
@patch("agent_framework_aisearch._search_provider.SearchClient")
async def test_single_vectorizable_field_detected(
self, mock_search_class: MagicMock, mock_index_class: MagicMock
) -> None:
"""Test that single vectorizable field is auto-detected for server-side vectorization."""
# Setup search client mock
mock_search_client = AsyncMock()
mock_search_class.return_value = mock_search_client
# Setup index with single vectorizable field
mock_index_client = AsyncMock()
mock_index = MagicMock()
# Single vector field with vectorizer
mock_field = MagicMock()
mock_field.name = "embedding"
mock_field.vector_search_dimensions = 1536
mock_field.vector_search_profile_name = "embedding-profile"
mock_index.fields = [mock_field]
# Setup vector search config with profile that has vectorizer
mock_profile = MagicMock()
mock_profile.name = "embedding-profile"
mock_profile.vectorizer_name = "openai-vectorizer"
mock_index.vector_search = MagicMock()
mock_index.vector_search.profiles = [mock_profile]
mock_index_client.get_index.return_value = mock_index
mock_index_client.close = AsyncMock()
mock_index_class.return_value = mock_index_client
# Create provider
provider = AzureAISearchContextProvider(
endpoint="https://test.search.windows.net",
index_name="test-index",
api_key="test-key",
mode="semantic",
)
# Trigger auto-discovery
await provider._auto_discover_vector_field()
# Should detect the vectorizable field
assert provider.vector_field_name == "embedding"
assert provider._auto_discovered_vector_field is True
assert provider._use_vectorizable_query is True # Server-side vectorization
+1 -1
View File
@@ -4,7 +4,7 @@ description = "Anthropic integration for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251120"
version = "1.0.0b251114"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
@@ -310,8 +310,8 @@ class AzureAIClient(OpenAIBaseResponsesClient):
return run_options
async def initialize_client(self) -> None:
"""Initialize OpenAI client."""
self.client = self.project_client.get_openai_client() # type: ignore
"""Initialize OpenAI client asynchronously."""
self.client = await self.project_client.get_openai_client() # type: ignore
def _update_agent_name(self, agent_name: str | None) -> None:
"""Update the agent name in the chat client.
+2 -2
View File
@@ -4,7 +4,7 @@ description = "Azure AI Foundry integration for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251120"
version = "1.0.0b251114"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
@@ -24,7 +24,7 @@ classifiers = [
]
dependencies = [
"agent-framework-core",
"azure-ai-projects >= 2.0.0b2",
"azure-ai-projects >= 2.0.0b1",
"azure-ai-agents == 1.2.0b5",
"aiohttp",
]
@@ -268,7 +268,7 @@ async def test_azure_ai_client_initialize_client(mock_project_client: MagicMock)
client = create_test_azure_ai_client(mock_project_client)
mock_openai_client = MagicMock()
mock_project_client.get_openai_client = MagicMock(return_value=mock_openai_client)
mock_project_client.get_openai_client = AsyncMock(return_value=mock_openai_client)
await client.initialize_client()
@@ -16,29 +16,27 @@ import azure.functions as func
from agent_framework import AgentProtocol, get_logger
from ._callbacks import AgentResponseCallbackProtocol
from ._constants import (
DEFAULT_MAX_POLL_RETRIES,
DEFAULT_POLL_INTERVAL_SECONDS,
MIMETYPE_APPLICATION_JSON,
MIMETYPE_TEXT_PLAIN,
REQUEST_RESPONSE_FORMAT_JSON,
REQUEST_RESPONSE_FORMAT_TEXT,
THREAD_ID_FIELD,
THREAD_ID_HEADER,
WAIT_FOR_RESPONSE_FIELD,
WAIT_FOR_RESPONSE_HEADER,
)
from ._durable_agent_state import DurableAgentState
from ._entities import create_agent_entity
from ._errors import IncomingRequestError
from ._models import AgentSessionId, RunRequest
from ._orchestration import AgentOrchestrationContextType, DurableAIAgent
from ._state import AgentState
logger = get_logger("agent_framework.azurefunctions")
THREAD_ID_FIELD: str = "thread_id"
RESPONSE_FORMAT_JSON: str = "json"
RESPONSE_FORMAT_TEXT: str = "text"
WAIT_FOR_RESPONSE_FIELD: str = "wait_for_response"
WAIT_FOR_RESPONSE_HEADER: str = "x-ms-wait-for-response"
EntityHandler = Callable[[df.DurableEntityContext], None]
HandlerT = TypeVar("HandlerT", bound=Callable[..., Any])
DEFAULT_MAX_POLL_RETRIES: int = 30
DEFAULT_POLL_INTERVAL_SECONDS: float = 1.0
if TYPE_CHECKING:
class DFAppBase:
@@ -319,11 +317,11 @@ class AgentFunctionApp(DFAppBase):
"""
logger.debug(f"[HTTP Trigger] Received request on route: /api/agents/{agent_name}/run")
request_response_format: str = REQUEST_RESPONSE_FORMAT_JSON
response_format: str = RESPONSE_FORMAT_JSON
thread_id: str | None = None
try:
req_body, message, request_response_format = self._parse_incoming_request(req)
req_body, message, response_format = self._parse_incoming_request(req)
thread_id = self._resolve_thread_id(req=req, req_body=req_body)
wait_for_response = self._should_wait_for_response(req=req, req_body=req_body)
@@ -336,7 +334,7 @@ class AgentFunctionApp(DFAppBase):
return self._create_http_response(
payload={"error": "Message is required"},
status_code=400,
request_response_format=request_response_format,
response_format=response_format,
thread_id=thread_id,
)
@@ -353,7 +351,6 @@ class AgentFunctionApp(DFAppBase):
message,
thread_id,
correlation_id,
request_response_format,
)
logger.debug("Signalling entity %s with request: %s", entity_instance_id, run_request)
await client.signal_entity(entity_instance_id, "run_agent", run_request)
@@ -373,7 +370,7 @@ class AgentFunctionApp(DFAppBase):
return self._create_http_response(
payload=result,
status_code=200 if result.get("status") == "success" else 500,
request_response_format=request_response_format,
response_format=response_format,
thread_id=thread_id,
)
@@ -386,7 +383,7 @@ class AgentFunctionApp(DFAppBase):
return self._create_http_response(
payload=accepted_response,
status_code=202,
request_response_format=request_response_format,
response_format=response_format,
thread_id=thread_id,
)
@@ -395,7 +392,7 @@ class AgentFunctionApp(DFAppBase):
return self._create_http_response(
payload={"error": str(exc)},
status_code=exc.status_code,
request_response_format=request_response_format,
response_format=response_format,
thread_id=thread_id,
)
except ValueError as exc:
@@ -403,7 +400,7 @@ class AgentFunctionApp(DFAppBase):
return self._create_http_response(
payload={"error": "Invalid JSON"},
status_code=400,
request_response_format=request_response_format,
response_format=response_format,
thread_id=thread_id,
)
except Exception as exc:
@@ -411,7 +408,7 @@ class AgentFunctionApp(DFAppBase):
return self._create_http_response(
payload={"error": str(exc)},
status_code=500,
request_response_format=request_response_format,
response_format=response_format,
thread_id=thread_id,
)
@@ -469,7 +466,7 @@ class AgentFunctionApp(DFAppBase):
return func.HttpResponse(
json.dumps({"status": "healthy", "agents": agent_info, "agent_count": len(self.agents)}),
status_code=200,
mimetype=MIMETYPE_APPLICATION_JSON,
mimetype="application/json",
)
_ = health_check
@@ -494,7 +491,7 @@ class AgentFunctionApp(DFAppBase):
self,
client: df.DurableOrchestrationClient,
entity_instance_id: df.EntityId,
) -> DurableAgentState | None:
) -> AgentState | None:
state_response = await client.read_entity_state(entity_instance_id)
if not state_response or not state_response.entity_exists:
return None
@@ -505,7 +502,9 @@ class AgentFunctionApp(DFAppBase):
typed_state_payload = cast(dict[str, Any], state_payload)
return DurableAgentState.from_dict(typed_state_payload)
agent_state = AgentState()
agent_state.restore_state(typed_state_payload)
return agent_state
async def _get_response_from_entity(
self,
@@ -581,58 +580,31 @@ class AgentFunctionApp(DFAppBase):
return result
def _build_response_payload(
self,
*,
response: str | None,
message: str,
thread_id: str,
status: str,
correlation_id: str,
extra_fields: dict[str, Any] | None = None,
) -> dict[str, Any]:
"""Create a consistent response structure and allow optional extra fields."""
payload = {
"response": response,
"message": message,
THREAD_ID_FIELD: thread_id,
"status": status,
"correlation_id": correlation_id,
}
if extra_fields:
payload.update(extra_fields)
return payload
async def _build_timeout_result(self, message: str, thread_id: str, correlation_id: str) -> dict[str, Any]:
"""Create the timeout response."""
return self._build_response_payload(
response="Agent is still processing or timed out...",
message=message,
thread_id=thread_id,
status="timeout",
correlation_id=correlation_id,
)
return {
"response": "Agent is still processing or timed out...",
"message": message,
THREAD_ID_FIELD: thread_id,
"status": "timeout",
"correlation_id": correlation_id,
}
def _build_success_result(
self, response_data: dict[str, Any], message: str, thread_id: str, correlation_id: str, state: DurableAgentState
self, response_data: dict[str, Any], message: str, thread_id: str, correlation_id: str, state: AgentState
) -> dict[str, Any]:
"""Build the success result returned to the HTTP caller."""
return self._build_response_payload(
response=response_data.get("content"),
message=message,
thread_id=thread_id,
status="success",
correlation_id=correlation_id,
extra_fields={"message_count": response_data.get("message_count", state.message_count)},
)
return {
"response": response_data.get("content"),
"message": message,
THREAD_ID_FIELD: thread_id,
"status": "success",
"message_count": response_data.get("message_count", state.message_count),
"correlation_id": correlation_id,
}
def _build_request_data(
self,
req_body: dict[str, Any],
message: str,
thread_id: str,
correlation_id: str,
request_response_format: str,
self, req_body: dict[str, Any], message: str, thread_id: str, correlation_id: str
) -> dict[str, Any]:
"""Create the durable entity request payload."""
enable_tool_calls_value = req_body.get("enable_tool_calls")
@@ -641,7 +613,6 @@ class AgentFunctionApp(DFAppBase):
return RunRequest(
message=message,
role=req_body.get("role"),
request_response_format=request_response_format,
response_format=req_body.get("response_format"),
enable_tool_calls=enable_tool_calls,
thread_id=thread_id,
@@ -650,23 +621,23 @@ class AgentFunctionApp(DFAppBase):
def _build_accepted_response(self, message: str, thread_id: str, correlation_id: str) -> dict[str, Any]:
"""Build the response returned when not waiting for completion."""
return self._build_response_payload(
response="Agent request accepted",
message=message,
thread_id=thread_id,
status="accepted",
correlation_id=correlation_id,
)
return {
"response": "Agent request accepted",
"message": message,
THREAD_ID_FIELD: thread_id,
"status": "accepted",
"correlation_id": correlation_id,
}
def _create_http_response(
self,
payload: dict[str, Any] | str,
status_code: int,
request_response_format: str,
response_format: str,
thread_id: str | None,
) -> func.HttpResponse:
"""Create the HTTP response using helper serializers for clarity."""
if request_response_format == REQUEST_RESPONSE_FORMAT_TEXT:
if response_format == RESPONSE_FORMAT_TEXT:
return self._build_plain_text_response(payload=payload, status_code=status_code, thread_id=thread_id)
return self._build_json_response(payload=payload, status_code=status_code)
@@ -679,13 +650,13 @@ class AgentFunctionApp(DFAppBase):
) -> func.HttpResponse:
"""Return a plain-text response with optional thread identifier header."""
body_text = payload if isinstance(payload, str) else self._convert_payload_to_text(payload)
headers = {THREAD_ID_HEADER: thread_id} if thread_id is not None else None
return func.HttpResponse(body_text, status_code=status_code, mimetype=MIMETYPE_TEXT_PLAIN, headers=headers)
headers = {"x-ms-thread-id": thread_id} if thread_id is not None else None
return func.HttpResponse(body_text, status_code=status_code, mimetype="text/plain", headers=headers)
def _build_json_response(self, payload: dict[str, Any] | str, status_code: int) -> func.HttpResponse:
"""Return the JSON response, serializing dictionaries as needed."""
body_json = payload if isinstance(payload, str) else json.dumps(payload)
return func.HttpResponse(body_json, status_code=status_code, mimetype=MIMETYPE_APPLICATION_JSON)
return func.HttpResponse(body_json, status_code=status_code, mimetype="application/json")
def _convert_payload_to_text(self, payload: dict[str, Any]) -> str:
"""Convert a structured payload into a human-readable text response."""
@@ -731,19 +702,18 @@ class AgentFunctionApp(DFAppBase):
normalized_content_type = self._extract_content_type(headers)
body_parser, body_format = self._select_body_parser(normalized_content_type)
prefers_json = self._accepts_json_response(headers)
request_response_format = self._select_request_response_format(
body_format=body_format, prefers_json=prefers_json
)
response_format = self._select_response_format(body_format=body_format, prefers_json=prefers_json)
req_body, message = body_parser(req)
return req_body, message, request_response_format
return req_body, message, response_format
def _extract_normalized_headers(self, req: func.HttpRequest) -> dict[str, str]:
"""Create a lowercase header mapping from the incoming request."""
headers: dict[str, str] = {}
raw_headers = req.headers
if isinstance(raw_headers, Mapping):
for key, value in raw_headers.items():
header_mapping: Mapping[str, Any] = cast(Mapping[str, Any], raw_headers)
for key, value in header_mapping.items():
if value is not None:
headers[str(key).lower()] = str(value)
return headers
@@ -759,9 +729,9 @@ class AgentFunctionApp(DFAppBase):
normalized_content_type: str,
) -> tuple[Callable[[func.HttpRequest], tuple[dict[str, Any], str]], str]:
"""Choose the body parser and declared body format."""
if normalized_content_type in {MIMETYPE_APPLICATION_JSON} or normalized_content_type.endswith("+json"):
return self._parse_json_body, REQUEST_RESPONSE_FORMAT_JSON
return self._parse_text_body, REQUEST_RESPONSE_FORMAT_TEXT
if normalized_content_type in {"application/json"} or normalized_content_type.endswith("+json"):
return self._parse_json_body, RESPONSE_FORMAT_JSON
return self._parse_text_body, RESPONSE_FORMAT_TEXT
@staticmethod
def _accepts_json_response(headers: dict[str, str]) -> bool:
@@ -772,16 +742,16 @@ class AgentFunctionApp(DFAppBase):
for value in accept_header.split(","):
media_type = value.split(";")[0].strip().lower()
if media_type == MIMETYPE_APPLICATION_JSON:
if media_type == "application/json":
return True
return False
@staticmethod
def _select_request_response_format(body_format: str, prefers_json: bool) -> str:
def _select_response_format(body_format: str, prefers_json: bool) -> str:
"""Combine body format and accept preference to determine response format."""
if body_format == REQUEST_RESPONSE_FORMAT_JSON or prefers_json:
return REQUEST_RESPONSE_FORMAT_JSON
return REQUEST_RESPONSE_FORMAT_TEXT
if body_format == RESPONSE_FORMAT_JSON or prefers_json:
return RESPONSE_FORMAT_JSON
return RESPONSE_FORMAT_TEXT
@staticmethod
def _parse_json_body(req: func.HttpRequest) -> tuple[dict[str, Any], str]:
@@ -6,6 +6,8 @@ This module enables callers of AgentFunctionApp to supply streaming and final-re
invoked during durable entity execution.
"""
from __future__ import annotations
from dataclasses import dataclass
from typing import Protocol
@@ -1,19 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
"""Constants for Azure Functions Agent Framework integration."""
# Supported request/response formats and MIME types
REQUEST_RESPONSE_FORMAT_JSON: str = "json"
REQUEST_RESPONSE_FORMAT_TEXT: str = "text"
MIMETYPE_APPLICATION_JSON: str = "application/json"
MIMETYPE_TEXT_PLAIN: str = "text/plain"
# Field and header names
THREAD_ID_FIELD: str = "thread_id"
THREAD_ID_HEADER: str = "x-ms-thread-id"
WAIT_FOR_RESPONSE_FIELD: str = "wait_for_response"
WAIT_FOR_RESPONSE_HEADER: str = "x-ms-wait-for-response"
# Polling configuration
DEFAULT_MAX_POLL_RETRIES: int = 30
DEFAULT_POLL_INTERVAL_SECONDS: float = 1.0
@@ -11,30 +11,14 @@ import asyncio
import inspect
import json
from collections.abc import AsyncIterable, Callable
from datetime import datetime, timezone
from typing import Any, cast
import azure.durable_functions as df
from agent_framework import (
AgentProtocol,
AgentRunResponse,
AgentRunResponseUpdate,
ChatMessage,
ErrorContent,
Role,
get_logger,
)
from agent_framework import AgentProtocol, AgentRunResponse, AgentRunResponseUpdate, Role, get_logger
from ._callbacks import AgentCallbackContext, AgentResponseCallbackProtocol
from ._durable_agent_state import (
DurableAgentState,
DurableAgentStateData,
DurableAgentStateEntry,
DurableAgentStateMessage,
DurableAgentStateRequest,
DurableAgentStateResponse,
)
from ._models import AgentResponse, RunRequest
from ._state import AgentState
logger = get_logger("agent_framework.azurefunctions.entities")
@@ -54,11 +38,11 @@ class AgentEntity:
Attributes:
agent: The AgentProtocol instance
state: The DurableAgentState managing conversation history
state: The AgentState managing conversation history
"""
agent: AgentProtocol
state: DurableAgentState
state: AgentState
def __init__(
self,
@@ -72,27 +56,11 @@ class AgentEntity:
callback: Optional callback invoked during streaming updates and final responses
"""
self.agent = agent
self.state = DurableAgentState()
self.state = AgentState()
self.callback = callback
logger.debug(f"[AgentEntity] Initialized with agent type: {type(agent).__name__}")
def _is_error_response(self, entry: DurableAgentStateEntry) -> bool:
"""Check if a conversation history entry is an error response.
Error responses should be kept in history for tracking but not sent to the agent
since Azure OpenAI doesn't support 'error' content type.
Args:
entry: A conversation history entry (DurableAgentStateEntry or dict)
Returns:
True if the entry is a response containing error content, False otherwise
"""
if isinstance(entry, DurableAgentStateResponse):
return entry.is_error
return False
async def run_agent(
self,
context: df.DurableEntityContext,
@@ -126,27 +94,26 @@ class AgentEntity:
raise ValueError("RunRequest must include a thread_id")
if not correlation_id:
raise ValueError("RunRequest must include a correlation_id")
role = run_request.role or Role.USER
response_format = run_request.response_format
enable_tool_calls = run_request.enable_tool_calls
state_request = DurableAgentStateRequest.from_run_request(run_request)
self.state.data.conversation_history.append(state_request)
logger.debug(f"[AgentEntity.run_agent] Received message: {message}")
logger.debug(f"[AgentEntity.run_agent] Thread ID: {thread_id}")
logger.debug(f"[AgentEntity.run_agent] Correlation ID: {correlation_id}")
logger.debug(f"[AgentEntity.run_agent] Role: {role.value}")
logger.debug(f"[AgentEntity.run_agent] Enable tool calls: {enable_tool_calls}")
logger.debug(f"[AgentEntity.run_agent] Response format: {'provided' if response_format else 'none'}")
logger.debug(f"[AgentEntity.run_agent] Received Message: {state_request}")
# Store message in history with role
self.state.add_user_message(message, role=role, correlation_id=correlation_id)
logger.debug("[AgentEntity.run_agent] Executing agent...")
try:
logger.debug("[AgentEntity.run_agent] Starting agent invocation")
# Build messages from conversation history, excluding error responses
# Error responses are kept in history for tracking but not sent to the agent
chat_messages: list[ChatMessage] = [
m.to_chat_message()
for entry in self.state.data.conversation_history
if not self._is_error_response(entry)
for m in entry.messages
]
run_kwargs: dict[str, Any] = {"messages": chat_messages}
run_kwargs: dict[str, Any] = {"messages": self.state.get_chat_messages()}
if not enable_tool_calls:
run_kwargs["tools"] = None
if response_format:
@@ -166,8 +133,8 @@ class AgentEntity:
response_text = None
structured_response = None
response_str: str | None = None
response_str: str | None = None
try:
if response_format:
try:
@@ -189,19 +156,18 @@ class AgentEntity:
)
response_text = "Error extracting response"
state_response = DurableAgentStateResponse.from_run_response(correlation_id, agent_run_response)
self.state.data.conversation_history.append(state_response)
agent_response = AgentResponse(
response=response_text,
message=str(message),
thread_id=str(thread_id),
status="success",
message_count=len(self.state.data.conversation_history),
message_count=self.state.message_count,
structured_response=structured_response,
)
result = agent_response.to_dict()
content = json.dumps(structured_response) if structured_response else (response_text or "")
self.state.add_assistant_message(content, agent_run_response, correlation_id)
logger.debug("[AgentEntity.run_agent] AgentRunResponse stored in conversation history")
return result
@@ -215,28 +181,12 @@ class AgentEntity:
logger.error(f"Error type: {type(exc).__name__}")
logger.error(f"Full traceback:\n{error_traceback}")
# Create error message
error_message = DurableAgentStateMessage.from_chat_message(
ChatMessage(
role=Role.ASSISTANT, contents=[ErrorContent(message=str(exc), error_code=type(exc).__name__)]
)
)
# Create and store error response in conversation history
error_state_response = DurableAgentStateResponse(
correlation_id=correlation_id,
created_at=datetime.now(tz=timezone.utc),
messages=[error_message],
is_error=True,
)
self.state.data.conversation_history.append(error_state_response)
error_response = AgentResponse(
response=f"Error: {exc!s}",
message=str(message),
thread_id=str(thread_id),
status="error",
message_count=len(self.state.data.conversation_history),
message_count=self.state.message_count,
error=str(exc),
error_type=type(exc).__name__,
)
@@ -383,7 +333,7 @@ class AgentEntity:
def reset(self, context: df.DurableEntityContext) -> None:
"""Reset the entity state (clear conversation history)."""
logger.debug("[AgentEntity.reset] Resetting entity state")
self.state.data = DurableAgentStateData(conversation_history=[])
self.state.reset()
logger.debug("[AgentEntity.reset] State reset complete")
@@ -412,7 +362,7 @@ def create_agent_entity(
entity = AgentEntity(agent, callback)
if current_state is not None:
entity.state = DurableAgentState.from_dict(current_state)
entity.state.restore_state(current_state)
logger.debug(
"[entity_function] Restored entity from state (message_count: %s)", entity.state.message_count
)
@@ -442,9 +392,8 @@ def create_agent_entity(
logger.error("[entity_function] Unknown operation: %s", operation)
context.set_result({"error": f"Unknown operation: {operation}"})
logger.debug("State dict: %s", entity.state.to_dict())
context.set_state(entity.state.to_dict())
logger.info(f"[entity_function] Operation {operation} completed successfully")
logger.debug(f"[entity_function] Operation {operation} completed successfully")
except Exception as exc:
import traceback
@@ -2,6 +2,8 @@
"""Custom exception types for the durable agent framework."""
from __future__ import annotations
class IncomingRequestError(ValueError):
"""Raised when an incoming HTTP request cannot be parsed or validated."""
@@ -17,8 +17,6 @@ from typing import TYPE_CHECKING, Any, cast
import azure.durable_functions as df
from agent_framework import AgentThread, Role
from ._constants import REQUEST_RESPONSE_FORMAT_TEXT
if TYPE_CHECKING: # pragma: no cover - type checking imports only
from pydantic import BaseModel
@@ -280,43 +278,35 @@ class RunRequest:
Attributes:
message: The message to send to the agent
request_response_format: The desired response format (e.g., "text" or "json")
role: The role of the message sender (user, system, or assistant)
response_format: Optional Pydantic BaseModel type describing the structured response format
enable_tool_calls: Whether to enable tool calls for this request
thread_id: Optional thread ID for tracking
correlation_id: Optional correlation ID for tracking the response to this specific request
created_at: Optional timestamp when the request was created
"""
message: str
request_response_format: str
role: Role = Role.USER
response_format: type[BaseModel] | None = None
enable_tool_calls: bool = True
thread_id: str | None = None
correlation_id: str | None = None
created_at: str | None = None
def __init__(
self,
message: str,
request_response_format: str = REQUEST_RESPONSE_FORMAT_TEXT,
role: Role | str | None = Role.USER,
response_format: type[BaseModel] | None = None,
enable_tool_calls: bool = True,
thread_id: str | None = None,
correlation_id: str | None = None,
created_at: str | None = None,
) -> None:
self.message = message
self.role = self.coerce_role(role)
self.response_format = response_format
self.request_response_format = request_response_format
self.enable_tool_calls = enable_tool_calls
self.thread_id = thread_id
self.correlation_id = correlation_id
self.created_at = created_at
@staticmethod
def coerce_role(value: Role | str | None) -> Role:
@@ -336,17 +326,13 @@ class RunRequest:
"message": self.message,
"enable_tool_calls": self.enable_tool_calls,
"role": self.role.value,
"request_response_format": self.request_response_format,
}
if self.response_format:
result["response_format"] = _serialize_response_format(self.response_format)
if self.thread_id:
result["thread_id"] = self.thread_id
if self.correlation_id:
result["correlationId"] = self.correlation_id
if self.created_at:
result["created_at"] = self.created_at
result["correlation_id"] = self.correlation_id
return result
@classmethod
@@ -354,13 +340,11 @@ class RunRequest:
"""Create RunRequest from dictionary."""
return cls(
message=data.get("message", ""),
request_response_format=data.get("request_response_format", REQUEST_RESPONSE_FORMAT_TEXT),
role=cls.coerce_role(data.get("role")),
response_format=_deserialize_response_format(data.get("response_format")),
enable_tool_calls=data.get("enable_tool_calls", True),
thread_id=data.get("thread_id"),
correlation_id=data.get("correlationId"),
created_at=data.get("created_at"),
correlation_id=data.get("correlation_id"),
)
@@ -0,0 +1,179 @@
# Copyright (c) Microsoft. All rights reserved.
"""Agent State Management.
This module defines the AgentState class for managing conversation state and
serializing agent framework responses.
"""
from collections.abc import MutableMapping
from datetime import datetime, timezone
from typing import Any, cast
from agent_framework import AgentRunResponse, ChatMessage, Role, get_logger
logger = get_logger("agent_framework.azurefunctions.state")
class AgentState:
"""Manages agent conversation state using agent_framework types (ChatMessage, AgentRunResponse).
This class handles:
- Conversation history tracking using ChatMessage objects
- Agent response storage using AgentRunResponse objects with correlation IDs
- State persistence and restoration
- Message counting
"""
def __init__(self) -> None:
"""Initialize empty agent state."""
self.conversation_history: list[ChatMessage] = []
self.last_response: str | None = None
self.message_count: int = 0
def _current_timestamp(self) -> str:
"""Return an ISO 8601 UTC timestamp."""
return datetime.now(timezone.utc).isoformat()
def add_user_message(
self,
content: str,
role: Role = Role.USER,
correlation_id: str | None = None,
) -> None:
"""Add a user message to the conversation history as a ChatMessage object.
Args:
content: The message content
role: The message role (user, system, etc.)
correlation_id: Optional correlation identifier associated with the user message
"""
self.message_count += 1
timestamp = self._current_timestamp()
additional_props: MutableMapping[str, Any] = {"timestamp": timestamp}
if correlation_id is not None:
additional_props["correlation_id"] = correlation_id
chat_message = ChatMessage(role=role, text=content, additional_properties=additional_props)
self.conversation_history.append(chat_message)
logger.debug(f"Added {role} ChatMessage to history (message #{self.message_count})")
def add_assistant_message(
self, content: str, agent_response: AgentRunResponse, correlation_id: str | None = None
) -> None:
"""Add an assistant message to the conversation history with full agent response.
Args:
content: The text content of the response
agent_response: The AgentRunResponse object from the agent framework
correlation_id: Optional correlation ID for tracking this response
"""
self.last_response = content
timestamp = self._current_timestamp()
serialized_response = self.serialize_response(agent_response)
# Create a ChatMessage for the assistant response
# The agent_response already contains messages, but we store it as a custom ChatMessage
# with the agent_response stored in additional_properties for full metadata preservation
additional_props: dict[str, Any] = {
"agent_response": serialized_response,
"correlation_id": correlation_id,
"timestamp": timestamp,
"message_count": self.message_count,
}
chat_message = ChatMessage(role="assistant", text=content, additional_properties=additional_props)
self.conversation_history.append(chat_message)
logger.debug(
f"Added assistant ChatMessage to history with AgentRunResponse metadata (correlation_id: {correlation_id})"
)
def get_chat_messages(self) -> list[ChatMessage]:
"""Return a copy of the full conversation history."""
return list(self.conversation_history)
def try_get_agent_response(self, correlation_id: str) -> dict[str, Any] | None:
"""Get an agent response by correlation ID.
Args:
correlation_id: The correlation ID to look up
Returns:
The agent response data if found, None otherwise
"""
for message in reversed(self.conversation_history):
metadata = getattr(message, "additional_properties", {}) or {}
if metadata.get("correlation_id") == correlation_id:
return self._build_agent_response_payload(message, metadata)
return None
def serialize_response(self, response: AgentRunResponse) -> dict[str, Any]:
"""Serialize an ``AgentRunResponse`` to a dictionary.
Args:
response: The agent framework response object
Returns:
Dictionary containing all response fields
"""
try:
return response.to_dict()
except Exception as exc: # pragma: no cover - defensive logging path
logger.warning(f"Error serializing response: {exc}")
return {"response": str(response), "serialization_error": str(exc)}
def to_dict(self) -> dict[str, Any]:
"""Get the current state as a dictionary for persistence.
Returns:
Dictionary containing conversation_history (as serialized ChatMessages),
last_response, and message_count
"""
return {
"conversation_history": [msg.to_dict() for msg in self.conversation_history],
"last_response": self.last_response,
"message_count": self.message_count,
}
def restore_state(self, state: dict[str, Any]) -> None:
"""Restore state from a dictionary, reconstructing ChatMessage objects.
Args:
state: Dictionary containing conversation_history, last_response, and message_count
"""
# Restore conversation history as ChatMessage objects
history_data = state.get("conversation_history", [])
restored_history: list[ChatMessage] = []
for raw_message in history_data:
if isinstance(raw_message, dict):
restored_history.append(ChatMessage.from_dict(cast(dict[str, Any], raw_message)))
else:
restored_history.append(cast(ChatMessage, raw_message))
self.conversation_history = restored_history
self.last_response = state.get("last_response")
self.message_count = state.get("message_count", 0)
logger.debug("Restored state: %s ChatMessages in history", len(self.conversation_history))
def reset(self) -> None:
"""Reset the state to empty."""
self.conversation_history = []
self.last_response = None
self.message_count = 0
logger.debug("State reset to empty")
def __repr__(self) -> str:
"""String representation of the state."""
return f"AgentState(messages={self.message_count}, history_length={len(self.conversation_history)})"
def _build_agent_response_payload(self, message: ChatMessage, metadata: dict[str, Any]) -> dict[str, Any]:
"""Construct the agent response payload returned to callers."""
return {
"content": message.text,
"agent_response": metadata.get("agent_response"),
"message_count": metadata.get("message_count", self.message_count),
"timestamp": metadata.get("timestamp"),
"correlation_id": metadata.get("correlation_id"),
}
@@ -4,7 +4,7 @@ description = "Azure Functions integration for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251120"
version = "1.0.0b251114"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
@@ -27,11 +27,6 @@ dependencies = [
"azure-functions-durable",
]
[dependency-groups]
dev = [
"types-python-dateutil>=2.9.0",
]
[tool.uv]
prerelease = "if-necessary-or-explicit"
environments = [
@@ -16,8 +16,6 @@ Usage:
import pytest
from agent_framework_azurefunctions._constants import THREAD_ID_HEADER
from .testutils import SampleTestHelper, skip_if_azure_functions_integration_tests_disabled
# Module-level markers - applied to all tests in this file
@@ -69,7 +67,7 @@ class TestSampleSingleAgent:
# Agent responded with plain text when the request body was text/plain.
assert response.text.strip()
assert response.headers.get(THREAD_ID_HEADER) is not None
assert response.headers.get("x-ms-thread-id") is not None
def test_thread_id_in_query(self) -> None:
"""Test using thread_id in query parameter."""
@@ -79,7 +77,7 @@ class TestSampleSingleAgent:
assert response.status_code in [200, 202]
assert response.text.strip()
assert response.headers.get(THREAD_ID_HEADER) == "test-query-thread"
assert response.headers.get("x-ms-thread-id") == "test-query-thread"
def test_conversation_continuity(self) -> None:
"""Test conversation context is maintained across requests."""
@@ -94,7 +92,7 @@ class TestSampleSingleAgent:
if response1.status_code == 200:
data1 = response1.json()
assert data1["message_count"] == 2 # Initial + reply
assert data1["message_count"] == 1
# Second message in same session
response2 = SampleTestHelper.post_json(
@@ -102,7 +100,7 @@ class TestSampleSingleAgent:
)
assert response2.status_code == 200
data2 = response2.json()
assert data2["message_count"] == 4
assert data2["message_count"] == 2
else:
# In async mode, we can't easily test message count
# Just verify we can make multiple calls
+38 -101
View File
@@ -13,13 +13,7 @@ from agent_framework import AgentRunResponse, ChatMessage
from agent_framework_azurefunctions import AgentFunctionApp
from agent_framework_azurefunctions._app import WAIT_FOR_RESPONSE_FIELD, WAIT_FOR_RESPONSE_HEADER
from agent_framework_azurefunctions._constants import (
MIMETYPE_APPLICATION_JSON,
MIMETYPE_TEXT_PLAIN,
THREAD_ID_HEADER,
)
from agent_framework_azurefunctions._durable_agent_state import DurableAgentState
from agent_framework_azurefunctions._entities import AgentEntity, create_agent_entity
from agent_framework_azurefunctions._entities import AgentEntity, AgentState, create_agent_entity
TFunc = TypeVar("TFunc", bound=Callable[..., Any])
@@ -339,14 +333,14 @@ class TestAgentEntityOperations:
result = await entity.run_agent(
mock_context,
{"message": "Test message", "thread_id": "test-conv-123", "correlationId": "corr-app-entity-1"},
{"message": "Test message", "thread_id": "test-conv-123", "correlation_id": "corr-app-entity-1"},
)
assert result["status"] == "success"
assert result["response"] == "Test response"
assert result["message"] == "Test message"
assert result["thread_id"] == "test-conv-123"
assert entity.state.message_count == 2
assert entity.state.message_count == 1
async def test_entity_stores_conversation_history(self) -> None:
"""Test that the entity stores conversation history."""
@@ -360,29 +354,18 @@ class TestAgentEntityOperations:
# Send first message
await entity.run_agent(
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlationId": "corr-app-entity-2"}
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlation_id": "corr-app-entity-2"}
)
# Each conversation turn creates 2 entries: request and response
history = entity.state.data.conversation_history[0].messages # Request entry
assert len(history) == 1 # Just the user message
# Send second message
await entity.run_agent(
mock_context, {"message": "Message 2", "thread_id": "conv-2", "correlationId": "corr-app-entity-2b"}
)
# Now we have 4 entries total (2 requests + 2 responses)
# Access the first request entry
history2 = entity.state.data.conversation_history[2].messages # Second request entry
assert len(history2) == 1 # Just the user message
history = entity.state.conversation_history
assert len(history) == 2 # User + assistant
user_msg = history[0]
user_role = getattr(user_msg.role, "value", user_msg.role)
assert user_role == "user"
assert user_msg.text == "Message 1"
assistant_msg = entity.state.data.conversation_history[1].messages[0]
assistant_msg = history[1]
assistant_role = getattr(assistant_msg.role, "value", assistant_msg.role)
assert assistant_role == "assistant"
assert assistant_msg.text == "Response 1"
@@ -397,17 +380,17 @@ class TestAgentEntityOperations:
entity = AgentEntity(mock_agent)
mock_context = Mock()
assert len(entity.state.data.conversation_history) == 0
assert entity.state.message_count == 0
await entity.run_agent(
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlationId": "corr-app-entity-3a"}
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlation_id": "corr-app-entity-3a"}
)
assert len(entity.state.data.conversation_history) == 2
assert entity.state.message_count == 1
await entity.run_agent(
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlationId": "corr-app-entity-3b"}
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlation_id": "corr-app-entity-3b"}
)
assert len(entity.state.data.conversation_history) == 4
assert entity.state.message_count == 2
def test_entity_reset(self) -> None:
"""Test that entity reset clears state."""
@@ -415,13 +398,19 @@ class TestAgentEntityOperations:
entity = AgentEntity(mock_agent)
# Set some state
entity.state = DurableAgentState()
entity.state.message_count = 10
entity.state.last_response = "Some response"
entity.state.conversation_history = [
ChatMessage(role="user", text="test", additional_properties={"timestamp": "2024-01-01T00:00:00Z"})
]
# Reset
mock_context = Mock()
entity.reset(mock_context)
assert len(entity.state.data.conversation_history) == 0
assert entity.state.message_count == 0
assert entity.state.last_response is None
assert len(entity.state.conversation_history) == 0
class TestAgentEntityFactory:
@@ -449,7 +438,7 @@ class TestAgentEntityFactory:
mock_context.get_input.return_value = {
"message": "Test message",
"thread_id": "conv-123",
"correlationId": "corr-app-factory-1",
"correlation_id": "corr-app-factory-1",
}
mock_context.get_state.return_value = None
@@ -469,27 +458,9 @@ class TestAgentEntityFactory:
mock_context = Mock()
mock_context.operation_name = "reset"
mock_context.get_state.return_value = {
"schemaVersion": "1.0.0",
"data": {
"conversationHistory": [
{
"$type": "request",
"correlationId": "corr-reset-test",
"createdAt": "2024-01-01T00:00:00Z",
"messages": [
{
"role": "user",
"contents": [
{
"$type": "text",
"text": "test",
}
],
}
],
}
],
},
"message_count": 5,
"conversation_history": [{"role": "user", "content": "test"}],
"last_response": "Test",
}
# Execute entity function
@@ -526,53 +497,19 @@ class TestAgentEntityFactory:
# Mock context with existing state
existing_state = {
"schemaVersion": "1.0.0",
"data": {
"conversationHistory": [
{
"$type": "request",
"correlationId": "corr-existing-1",
"createdAt": "2024-01-01T00:00:00Z",
"messages": [
{
"role": "user",
"contents": [
{
"$type": "text",
"text": "msg1",
}
],
}
],
},
{
"$type": "response",
"correlationId": "corr-existing-1",
"createdAt": "2024-01-01T00:05:00Z",
"messages": [
{
"role": "assistant",
"contents": [
{
"$type": "text",
"text": "resp1",
}
],
}
],
},
],
},
"message_count": 3,
"conversation_history": [{"role": "user", "content": "msg1"}, {"role": "assistant", "content": "resp1"}],
"last_response": "resp1",
}
mock_context = Mock()
mock_context.operation_name = "reset"
mock_context.get_state.return_value = existing_state
with patch.object(DurableAgentState, "from_dict", wraps=DurableAgentState.from_dict) as from_dict_mock:
with patch.object(AgentState, "restore_state") as restore_state_mock:
entity_function(mock_context)
from_dict_mock.assert_called_once_with(existing_state)
restore_state_mock.assert_called_once_with(existing_state)
class TestErrorHandling:
@@ -587,7 +524,7 @@ class TestErrorHandling:
mock_context = Mock()
result = await entity.run_agent(
mock_context, {"message": "Test message", "thread_id": "conv-1", "correlationId": "corr-app-error-1"}
mock_context, {"message": "Test message", "thread_id": "conv-1", "correlation_id": "corr-app-error-1"}
)
assert result["status"] == "error"
@@ -663,7 +600,7 @@ class TestIncomingRequestParsing:
app = self._create_app()
request = Mock()
request.headers = {"accept": MIMETYPE_APPLICATION_JSON}
request.headers = {"accept": "application/json"}
request.params = {}
request.get_json.side_effect = ValueError("Invalid JSON")
request.get_body.return_value = b"Plain text message"
@@ -737,8 +674,8 @@ class TestHttpRunRoute:
response = await handler(request, client)
assert response.status_code == 202
assert response.mimetype == MIMETYPE_TEXT_PLAIN
assert response.headers.get(THREAD_ID_HEADER) is not None
assert response.mimetype == "text/plain"
assert response.headers.get("x-ms-thread-id") is not None
assert response.get_body().decode("utf-8") == "Agent request accepted"
signal_args = client.signal_entity.call_args[0]
@@ -756,7 +693,7 @@ class TestHttpRunRoute:
handler = self._get_run_handler(mock_agent)
request = Mock()
request.headers = {WAIT_FOR_RESPONSE_HEADER: "false", "Accept": MIMETYPE_APPLICATION_JSON}
request.headers = {WAIT_FOR_RESPONSE_HEADER: "false", "Accept": "application/json"}
request.params = {}
request.route_params = {}
request.get_json.side_effect = ValueError("Invalid JSON")
@@ -767,8 +704,8 @@ class TestHttpRunRoute:
response = await handler(request, client)
assert response.status_code == 202
assert response.mimetype == MIMETYPE_APPLICATION_JSON
assert response.headers.get(THREAD_ID_HEADER) is None
assert response.mimetype == "application/json"
assert response.headers.get("x-ms-thread-id") is None
body = response.get_body().decode("utf-8")
assert '"status": "accepted"' in body
@@ -791,8 +728,8 @@ class TestHttpRunRoute:
response = await handler(request, client)
assert response.status_code == 400
assert response.mimetype == MIMETYPE_TEXT_PLAIN
assert response.headers.get(THREAD_ID_HEADER) is not None
assert response.mimetype == "text/plain"
assert response.headers.get("x-ms-thread-id") is not None
assert response.get_body().decode("utf-8") == "Message is required"
client.signal_entity.assert_not_called()
@@ -15,20 +15,14 @@ import pytest
from agent_framework import AgentRunResponse, AgentRunResponseUpdate, ChatMessage, Role
from pydantic import BaseModel
from agent_framework_azurefunctions._durable_agent_state import (
DurableAgentState,
DurableAgentStateData,
DurableAgentStateMessage,
DurableAgentStateRequest,
DurableAgentStateTextContent,
)
from agent_framework_azurefunctions._entities import AgentEntity, create_agent_entity
from agent_framework_azurefunctions._models import RunRequest
from agent_framework_azurefunctions._state import AgentState
TFunc = TypeVar("TFunc", bound=Callable[..., Any])
def _role_value(chat_message: DurableAgentStateMessage) -> str:
def _role_value(chat_message: ChatMessage) -> str:
"""Helper to extract the string role from a ChatMessage."""
role = getattr(chat_message, "role", None)
role_value = getattr(role, "value", role)
@@ -77,9 +71,9 @@ class TestAgentEntityInit:
entity = AgentEntity(mock_agent)
assert entity.agent == mock_agent
assert len(entity.state.data.conversation_history) == 0
assert entity.state.data.extension_data is None
assert entity.state.schema_version == "1.0.0"
assert entity.state.conversation_history == []
assert entity.state.last_response is None
assert entity.state.message_count == 0
def test_init_stores_agent_reference(self) -> None:
"""Test that the agent reference is stored correctly."""
@@ -118,7 +112,7 @@ class TestAgentEntityRunAgent:
mock_context = Mock()
result = await entity.run_agent(
mock_context, {"message": "Test message", "thread_id": "conv-123", "correlationId": "corr-entity-1"}
mock_context, {"message": "Test message", "thread_id": "conv-123", "correlation_id": "corr-entity-1"}
)
# Verify agent.run was called
@@ -129,8 +123,8 @@ class TestAgentEntityRunAgent:
assert len(sent_messages) == 1
sent_message = sent_messages[0]
assert isinstance(sent_message, ChatMessage)
assert getattr(sent_message, "text", None) == "Test message"
assert getattr(sent_message.role, "value", sent_message.role) == "user"
assert sent_message.text == "Test message"
assert _role_value(sent_message) == "user"
# Verify result
assert result["status"] == "success"
@@ -164,7 +158,7 @@ class TestAgentEntityRunAgent:
{
"message": "Tell me something",
"thread_id": "session-1",
"correlationId": "corr-stream-1",
"correlation_id": "corr-stream-1",
},
)
@@ -211,7 +205,7 @@ class TestAgentEntityRunAgent:
{
"message": "Hi",
"thread_id": "session-2",
"correlationId": "corr-final-1",
"correlation_id": "corr-final-1",
},
)
@@ -239,20 +233,19 @@ class TestAgentEntityRunAgent:
mock_context = Mock()
await entity.run_agent(
mock_context, {"message": "User message", "thread_id": "conv-1", "correlationId": "corr-entity-2"}
mock_context, {"message": "User message", "thread_id": "conv-1", "correlation_id": "corr-entity-2"}
)
# Should have 1 entry: user message + assistant response
user_history = entity.state.data.conversation_history[0].messages
assistant_history = entity.state.data.conversation_history[1].messages
# Should have 2 entries: user message + assistant response
history = entity.state.conversation_history
assert len(user_history) == 1
assert len(history) == 2
user_msg = user_history[0]
user_msg = history[0]
assert _role_value(user_msg) == "user"
assert user_msg.text == "User message"
assistant_msg = assistant_history[0]
assistant_msg = history[1]
assert _role_value(assistant_msg) == "assistant"
assert assistant_msg.text == "Agent response"
@@ -264,22 +257,41 @@ class TestAgentEntityRunAgent:
entity = AgentEntity(mock_agent)
mock_context = Mock()
assert len(entity.state.data.conversation_history) == 0
assert entity.state.message_count == 0
await entity.run_agent(
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlationId": "corr-entity-3a"}
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlation_id": "corr-entity-3a"}
)
assert len(entity.state.data.conversation_history) == 2
assert entity.state.message_count == 1
await entity.run_agent(
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlationId": "corr-entity-3b"}
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlation_id": "corr-entity-3b"}
)
assert len(entity.state.data.conversation_history) == 4
assert entity.state.message_count == 2
await entity.run_agent(
mock_context, {"message": "Message 3", "thread_id": "conv-1", "correlationId": "corr-entity-3c"}
mock_context, {"message": "Message 3", "thread_id": "conv-1", "correlation_id": "corr-entity-3c"}
)
assert len(entity.state.data.conversation_history) == 6
assert entity.state.message_count == 3
async def test_run_agent_stores_last_response(self) -> None:
"""Test that run_agent stores the last response."""
mock_agent = Mock()
mock_agent.run = AsyncMock(return_value=_agent_response("Response 1"))
entity = AgentEntity(mock_agent)
mock_context = Mock()
await entity.run_agent(
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlation_id": "corr-entity-4a"}
)
assert entity.state.last_response == "Response 1"
mock_agent.run = AsyncMock(return_value=_agent_response("Response 2"))
await entity.run_agent(
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlation_id": "corr-entity-4b"}
)
assert entity.state.last_response == "Response 2"
async def test_run_agent_with_none_thread_id(self) -> None:
"""Test run_agent with a None thread identifier."""
@@ -291,7 +303,7 @@ class TestAgentEntityRunAgent:
with pytest.raises(ValueError, match="thread_id"):
await entity.run_agent(
mock_context, {"message": "Message", "thread_id": None, "correlationId": "corr-entity-5"}
mock_context, {"message": "Message", "thread_id": None, "correlation_id": "corr-entity-5"}
)
async def test_run_agent_handles_response_without_text_attribute(self) -> None:
@@ -310,7 +322,7 @@ class TestAgentEntityRunAgent:
mock_context = Mock()
result = await entity.run_agent(
mock_context, {"message": "Message", "thread_id": "conv-1", "correlationId": "corr-entity-6"}
mock_context, {"message": "Message", "thread_id": "conv-1", "correlation_id": "corr-entity-6"}
)
# Should handle gracefully
@@ -326,7 +338,7 @@ class TestAgentEntityRunAgent:
mock_context = Mock()
result = await entity.run_agent(
mock_context, {"message": "Message", "thread_id": "conv-1", "correlationId": "corr-entity-7"}
mock_context, {"message": "Message", "thread_id": "conv-1", "correlation_id": "corr-entity-7"}
)
assert result["status"] == "success"
@@ -342,18 +354,18 @@ class TestAgentEntityRunAgent:
# Send multiple messages
await entity.run_agent(
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlationId": "corr-entity-8a"}
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlation_id": "corr-entity-8a"}
)
await entity.run_agent(
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlationId": "corr-entity-8b"}
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlation_id": "corr-entity-8b"}
)
await entity.run_agent(
mock_context, {"message": "Message 3", "thread_id": "conv-1", "correlationId": "corr-entity-8c"}
mock_context, {"message": "Message 3", "thread_id": "conv-1", "correlation_id": "corr-entity-8c"}
)
history = entity.state.data.conversation_history
history = entity.state.conversation_history
assert len(history) == 6
assert entity.state.message_count == 6
assert entity.state.message_count == 3
class TestAgentEntityReset:
@@ -364,47 +376,40 @@ class TestAgentEntityReset:
mock_agent = Mock()
entity = AgentEntity(mock_agent)
# Add some history with proper DurableAgentStateEntry objects
entity.state.data.conversation_history = [
DurableAgentStateRequest(
correlation_id="test-1",
created_at=datetime.now(),
messages=[
DurableAgentStateMessage(
role="user",
contents=[DurableAgentStateTextContent(text="msg1")],
)
],
),
# Add some history
entity.state.conversation_history = [
ChatMessage(role="user", text="msg1"),
ChatMessage(role="assistant", text="resp1"),
]
mock_context = Mock()
entity.reset(mock_context)
assert entity.state.data.conversation_history == []
assert entity.state.conversation_history == []
def test_reset_with_extension_data(self) -> None:
"""Test that reset works when entity has extension data."""
def test_reset_clears_last_response(self) -> None:
"""Test that reset clears the last response."""
mock_agent = Mock()
entity = AgentEntity(mock_agent)
# Set up some initial state with conversation history
entity.state.data = DurableAgentStateData(conversation_history=[], extension_data={"some_key": "some_value"})
entity.state.last_response = "Some response"
mock_context = Mock()
entity.reset(mock_context)
assert len(entity.state.data.conversation_history) == 0
assert entity.state.last_response is None
def test_reset_clears_message_count(self) -> None:
"""Test that reset clears the message count."""
mock_agent = Mock()
entity = AgentEntity(mock_agent)
entity.state.message_count = 10
mock_context = Mock()
entity.reset(mock_context)
assert len(entity.state.data.conversation_history) == 0
assert entity.state.message_count == 0
async def test_reset_after_conversation(self) -> None:
"""Test reset after a full conversation."""
@@ -416,22 +421,23 @@ class TestAgentEntityReset:
# Have a conversation
await entity.run_agent(
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlationId": "corr-entity-10a"}
mock_context, {"message": "Message 1", "thread_id": "conv-1", "correlation_id": "corr-entity-10a"}
)
await entity.run_agent(
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlationId": "corr-entity-10b"}
mock_context, {"message": "Message 2", "thread_id": "conv-1", "correlation_id": "corr-entity-10b"}
)
# Verify state before reset
assert entity.state.message_count == 4
assert len(entity.state.data.conversation_history) == 4
assert entity.state.message_count == 2
assert len(entity.state.conversation_history) == 4
# Reset
entity.reset(mock_context)
# Verify state after reset
assert entity.state.message_count == 0
assert len(entity.state.data.conversation_history) == 0
assert len(entity.state.conversation_history) == 0
assert entity.state.last_response is None
class TestCreateAgentEntity:
@@ -458,7 +464,7 @@ class TestCreateAgentEntity:
mock_context.get_input.return_value = {
"message": "Test message",
"thread_id": "conv-123",
"correlationId": "corr-entity-factory",
"correlation_id": "corr-entity-factory",
}
mock_context.get_state.return_value = None
@@ -479,22 +485,13 @@ class TestCreateAgentEntity:
mock_context = Mock()
mock_context.operation_name = "reset"
mock_context.get_state.return_value = {
"schemaVersion": "1.0.0",
"data": {
"conversationHistory": [
{
"$type": "request",
"correlationId": "test-correlation-id",
"createdAt": "2024-01-01T00:00:00Z",
"messages": [
{
"role": "user",
"contents": [{"$type": "text", "text": "test"}],
}
],
}
]
},
"message_count": 5,
"conversation_history": [
ChatMessage(
role="user", text="test", additional_properties={"timestamp": "2024-01-01T00:00:00Z"}
).to_dict()
],
"last_response": "Test",
}
# Execute
@@ -508,7 +505,9 @@ class TestCreateAgentEntity:
# Verify state was cleared
assert mock_context.set_state.called
state = mock_context.set_state.call_args[0][0]
assert state["data"]["conversationHistory"] == []
assert state["message_count"] == 0
assert state["conversation_history"] == []
assert state["last_response"] is None
def test_entity_function_handles_unknown_operation(self) -> None:
"""Test that the entity function handles unknown operations."""
@@ -548,7 +547,8 @@ class TestCreateAgentEntity:
assert result["status"] == "reset"
assert mock_context.set_state.called
state = mock_context.set_state.call_args[0][0]
assert state["data"] == {"conversationHistory": []}
assert state["message_count"] == 0
assert state["conversation_history"] == []
def test_entity_function_restores_existing_state(self) -> None:
"""Test that the entity function restores existing state."""
@@ -557,53 +557,26 @@ class TestCreateAgentEntity:
entity_function = create_agent_entity(mock_agent)
existing_state = {
"schemaVersion": "1.0.0",
"data": {
"conversationHistory": [
{
"$type": "request",
"correlationId": "corr-existing-1",
"createdAt": "2024-01-01T00:00:00Z",
"messages": [
{
"role": "user",
"contents": [
{
"$type": "text",
"text": "msg1",
}
],
}
],
},
{
"$type": "response",
"correlationId": "corr-existing-1",
"createdAt": "2024-01-01T00:05:00Z",
"messages": [
{
"role": "assistant",
"contents": [
{
"$type": "text",
"text": "resp1",
}
],
}
],
},
],
},
"message_count": 5,
"conversation_history": [
ChatMessage(
role="user", text="msg1", additional_properties={"timestamp": "2024-01-01T00:00:00Z"}
).to_dict(),
ChatMessage(
role="assistant", text="resp1", additional_properties={"timestamp": "2024-01-01T00:05:00Z"}
).to_dict(),
],
"last_response": "resp1",
}
mock_context = Mock()
mock_context.operation_name = "reset"
mock_context.get_state.return_value = existing_state
with patch.object(DurableAgentState, "from_dict", wraps=DurableAgentState.from_dict) as from_dict_mock:
with patch.object(AgentState, "restore_state") as restore_state_mock:
entity_function(mock_context)
from_dict_mock.assert_called_once_with(existing_state)
restore_state_mock.assert_called_once_with(existing_state)
class TestErrorHandling:
@@ -618,7 +591,7 @@ class TestErrorHandling:
mock_context = Mock()
result = await entity.run_agent(
mock_context, {"message": "Message", "thread_id": "conv-1", "correlationId": "corr-entity-error-1"}
mock_context, {"message": "Message", "thread_id": "conv-1", "correlation_id": "corr-entity-error-1"}
)
assert result["status"] == "error"
@@ -635,7 +608,7 @@ class TestErrorHandling:
mock_context = Mock()
result = await entity.run_agent(
mock_context, {"message": "Message", "thread_id": "conv-1", "correlationId": "corr-entity-error-2"}
mock_context, {"message": "Message", "thread_id": "conv-1", "correlation_id": "corr-entity-error-2"}
)
assert result["status"] == "error"
@@ -651,7 +624,7 @@ class TestErrorHandling:
mock_context = Mock()
result = await entity.run_agent(
mock_context, {"message": "Message", "thread_id": "conv-1", "correlationId": "corr-entity-error-3"}
mock_context, {"message": "Message", "thread_id": "conv-1", "correlation_id": "corr-entity-error-3"}
)
assert result["status"] == "error"
@@ -686,7 +659,7 @@ class TestErrorHandling:
result = await entity.run_agent(
mock_context,
{"message": "Test message", "thread_id": "conv-123", "correlationId": "corr-entity-error-4"},
{"message": "Test message", "thread_id": "conv-123", "correlation_id": "corr-entity-error-4"},
)
# Even on error, message info should be preserved
@@ -707,15 +680,15 @@ class TestConversationHistory:
mock_context = Mock()
await entity.run_agent(
mock_context, {"message": "Message", "thread_id": "conv-1", "correlationId": "corr-entity-history-1"}
mock_context, {"message": "Message", "thread_id": "conv-1", "correlation_id": "corr-entity-history-1"}
)
# Check both user and assistant messages have timestamps
for entry in entity.state.data.conversation_history:
timestamp = entry.created_at
for entry in entity.state.conversation_history:
timestamp = entry.additional_properties.get("timestamp")
assert timestamp is not None
# Verify timestamp is in ISO format
datetime.fromisoformat(str(timestamp))
datetime.fromisoformat(timestamp)
async def test_conversation_history_ordering(self) -> None:
"""Test that conversation history maintains the correct order."""
@@ -728,30 +701,29 @@ class TestConversationHistory:
mock_agent.run = AsyncMock(return_value=_agent_response("Response 1"))
await entity.run_agent(
mock_context,
{"message": "Message 1", "thread_id": "conv-1", "correlationId": "corr-entity-history-2a"},
{"message": "Message 1", "thread_id": "conv-1", "correlation_id": "corr-entity-history-2a"},
)
mock_agent.run = AsyncMock(return_value=_agent_response("Response 2"))
await entity.run_agent(
mock_context,
{"message": "Message 2", "thread_id": "conv-1", "correlationId": "corr-entity-history-2b"},
{"message": "Message 2", "thread_id": "conv-1", "correlation_id": "corr-entity-history-2b"},
)
mock_agent.run = AsyncMock(return_value=_agent_response("Response 3"))
await entity.run_agent(
mock_context,
{"message": "Message 3", "thread_id": "conv-1", "correlationId": "corr-entity-history-2c"},
{"message": "Message 3", "thread_id": "conv-1", "correlation_id": "corr-entity-history-2c"},
)
# Verify order
history = entity.state.data.conversation_history
# Each conversation turn creates 2 entries: request and response
assert history[0].messages[0].text == "Message 1" # Request 1
assert history[1].messages[0].text == "Response 1" # Response 1
assert history[2].messages[0].text == "Message 2" # Request 2
assert history[3].messages[0].text == "Response 2" # Response 2
assert history[4].messages[0].text == "Message 3" # Request 3
assert history[5].messages[0].text == "Response 3" # Response 3
history = entity.state.conversation_history
assert history[0].text == "Message 1"
assert history[1].text == "Response 1"
assert history[2].text == "Message 2"
assert history[3].text == "Response 2"
assert history[4].text == "Message 3"
assert history[5].text == "Response 3"
async def test_conversation_history_role_alternation(self) -> None:
"""Test that conversation history alternates between user and assistant roles."""
@@ -763,20 +735,19 @@ class TestConversationHistory:
await entity.run_agent(
mock_context,
{"message": "Message 1", "thread_id": "conv-1", "correlationId": "corr-entity-history-3a"},
{"message": "Message 1", "thread_id": "conv-1", "correlation_id": "corr-entity-history-3a"},
)
await entity.run_agent(
mock_context,
{"message": "Message 2", "thread_id": "conv-1", "correlationId": "corr-entity-history-3b"},
{"message": "Message 2", "thread_id": "conv-1", "correlation_id": "corr-entity-history-3b"},
)
# Check role alternation
history = entity.state.data.conversation_history
# Each conversation turn creates 2 entries: request and response
assert history[0].messages[0].role == "user" # Request 1
assert history[1].messages[0].role == "assistant" # Response 1
assert history[2].messages[0].role == "user" # Request 2
assert history[3].messages[0].role == "assistant" # Response 2
history = entity.state.conversation_history
assert _role_value(history[0]) == "user"
assert _role_value(history[1]) == "assistant"
assert _role_value(history[2]) == "user"
assert _role_value(history[3]) == "assistant"
class TestRunRequestSupport:
@@ -818,7 +789,7 @@ class TestRunRequestSupport:
"thread_id": "conv-456",
"role": "system",
"enable_tool_calls": False,
"correlationId": "corr-runreq-2",
"correlation_id": "corr-runreq-2",
}
result = await entity.run_agent(mock_context, request_dict)
@@ -857,9 +828,9 @@ class TestRunRequestSupport:
await entity.run_agent(mock_context, request)
# Check that system role was stored
history = entity.state.data.conversation_history
assert history[0].messages[0].role == "system"
assert history[0].messages[0].text == "System message"
history = entity.state.conversation_history
assert _role_value(history[0]) == "system"
assert history[0].text == "System message"
async def test_run_agent_with_response_format(self) -> None:
"""Test run_agent with a JSON response format."""
@@ -916,7 +887,7 @@ class TestRunRequestSupport:
"thread_id": "conv-789",
"role": "user",
"enable_tool_calls": True,
"correlationId": "corr-runreq-6",
"correlation_id": "corr-runreq-6",
}
mock_context.get_state.return_value = None
@@ -295,32 +295,32 @@ class TestRunRequest:
restored = RunRequest.from_dict(data)
assert restored.response_format is ModuleStructuredResponse
def test_init_with_correlationId(self) -> None:
"""Test RunRequest initialization with correlationId."""
def test_init_with_correlation_id(self) -> None:
"""Test RunRequest initialization with correlation_id."""
request = RunRequest(message="Test message", thread_id="thread-corr-init", correlation_id="corr-123")
assert request.message == "Test message"
assert request.correlation_id == "corr-123"
def test_to_dict_with_correlationId(self) -> None:
"""Test to_dict includes correlationId."""
def test_to_dict_with_correlation_id(self) -> None:
"""Test to_dict includes correlation_id."""
request = RunRequest(message="Test", thread_id="thread-corr-to-dict", correlation_id="corr-456")
data = request.to_dict()
assert data["message"] == "Test"
assert data["correlationId"] == "corr-456"
assert data["correlation_id"] == "corr-456"
def test_from_dict_with_correlationId(self) -> None:
"""Test from_dict with correlationId."""
data = {"message": "Test", "correlationId": "corr-789", "thread_id": "thread-corr-from-dict"}
def test_from_dict_with_correlation_id(self) -> None:
"""Test from_dict with correlation_id."""
data = {"message": "Test", "correlation_id": "corr-789", "thread_id": "thread-corr-from-dict"}
request = RunRequest.from_dict(data)
assert request.message == "Test"
assert request.correlation_id == "corr-789"
assert request.thread_id == "thread-corr-from-dict"
def test_round_trip_with_correlationId(self) -> None:
"""Test round-trip to_dict and from_dict with correlationId."""
def test_round_trip_with_correlation_id(self) -> None:
"""Test round-trip to_dict and from_dict with correlation_id."""
original = RunRequest(
message="Test message",
thread_id="thread-123",
@@ -136,8 +136,8 @@ class TestDurableAIAgent:
assert operation == "run_agent"
assert request["message"] == "Test message"
assert request["enable_tool_calls"] is True
assert "correlationId" in request
assert request["correlationId"] == "correlation-guid"
assert "correlation_id" in request
assert request["correlation_id"] == "correlation-guid"
assert "thread_id" in request
assert request["thread_id"] == "thread-guid"
@@ -145,7 +145,7 @@ class TestDurableAIAgent:
"""Test that run() works without explicit thread (creates unique session key)."""
mock_context = Mock()
mock_context.instance_id = "test-instance-002"
# Two calls to new_uuid: one for session_key, one for correlationId
# Two calls to new_uuid: one for session_key, one for correlation_id
mock_context.new_uuid = Mock(side_effect=["auto-generated-guid", "correlation-guid"])
mock_task = Mock()
@@ -164,7 +164,7 @@ class TestDurableAIAgent:
entity_id = call_args[0][0]
assert entity_id.name == "dafx-TestAgent"
assert entity_id.key == "auto-generated-guid"
# Should be called twice: once for session_key, once for correlationId
# Should be called twice: once for session_key, once for correlation_id
assert mock_context.new_uuid.call_count == 2
def test_run_with_response_format(self) -> None:
@@ -307,8 +307,8 @@ class TestOrchestrationIntegration:
mock_context.instance_id = "test-orchestration-001"
# new_uuid will be called 3 times:
# 1. thread creation
# 2. correlationId for first call
# 3. correlationId for second call
# 2. correlation_id for first call
# 3. correlation_id for second call
mock_context.new_uuid = Mock(side_effect=["deterministic-guid-001", "corr-1", "corr-2"])
# Track entity calls
@@ -0,0 +1,110 @@
# Copyright (c) Microsoft. All rights reserved.
"""Unit tests for AgentState correlation ID tracking."""
from unittest.mock import Mock
import pytest
from agent_framework import AgentRunResponse
from agent_framework_azurefunctions._state import AgentState
class TestAgentStateCorrelationId:
"""Test suite for AgentState correlation ID tracking."""
def _create_mock_response(self, text: str = "Response") -> Mock:
"""Create a mock AgentRunResponse with the provided text."""
mock_response = Mock(spec=AgentRunResponse)
mock_response.to_dict.return_value = {"text": text, "messages": []}
return mock_response
def test_add_assistant_message_with_correlation_id(self) -> None:
state = AgentState()
state.add_user_message("Hello", correlation_id="corr-123-request")
state.add_assistant_message("Response", self._create_mock_response(), correlation_id="corr-123")
message_metadata = state.conversation_history[-1].additional_properties or {}
assert message_metadata.get("correlation_id") == "corr-123"
response_data = state.try_get_agent_response("corr-123")
assert response_data is not None
assert response_data["content"] == "Response"
assert response_data["agent_response"] == {"text": "Response", "messages": []}
def test_try_get_agent_response_returns_response(self) -> None:
state = AgentState()
state.add_user_message("Hello", correlation_id="corr-200-request")
state.add_assistant_message("Response", self._create_mock_response(), correlation_id="corr-456")
response_data = state.try_get_agent_response("corr-456")
assert response_data is not None
assert response_data["content"] == "Response"
def test_try_get_agent_response_returns_none_for_missing_id(self) -> None:
state = AgentState()
state.add_user_message("Hello", correlation_id="corr-300-request")
state.add_assistant_message("Response", self._create_mock_response(), correlation_id="corr-123")
assert state.try_get_agent_response("non-existent") is None
def test_multiple_responses_tracked_separately(self) -> None:
state = AgentState()
for index in range(3):
state.add_user_message(f"Message {index}", correlation_id=f"corr-{index}-request")
state.add_assistant_message(
f"Response {index}",
self._create_mock_response(text=f"Response {index}"),
correlation_id=f"corr-{index}",
)
for index in range(3):
payload = state.try_get_agent_response(f"corr-{index}")
assert payload is not None
assert payload["content"] == f"Response {index}"
def test_add_assistant_message_without_correlation_id(self) -> None:
state = AgentState()
state.add_user_message("Hello", correlation_id="corr-400-request")
state.add_assistant_message("Response", self._create_mock_response())
assert state.try_get_agent_response("missing") is None
assert state.last_response == "Response"
def test_to_dict_does_not_duplicate_agent_responses(self) -> None:
state = AgentState()
state.add_user_message("Hello", correlation_id="corr-500-request")
state.add_assistant_message("Response", self._create_mock_response(), correlation_id="corr-123")
state_snapshot = state.to_dict()
assert "agent_responses" not in state_snapshot
metadata = state_snapshot["conversation_history"][-1]["additional_properties"]
assert metadata["correlation_id"] == "corr-123"
def test_restore_state_preserves_agent_response_lookup(self) -> None:
state = AgentState()
state.add_user_message("Hello", correlation_id="corr-600-request")
state.add_assistant_message("Response", self._create_mock_response(), correlation_id="corr-123")
restored_state = AgentState()
restored_state.restore_state(state.to_dict())
payload = restored_state.try_get_agent_response("corr-123")
assert payload is not None
assert payload["content"] == "Response"
def test_reset_clears_conversation_history(self) -> None:
state = AgentState()
state.add_user_message("Hello", correlation_id="corr-700-request")
state.add_assistant_message("Response", self._create_mock_response(), correlation_id="corr-123")
state.reset()
assert len(state.conversation_history) == 0
assert state.try_get_agent_response("corr-123") is None
if __name__ == "__main__":
pytest.main([__file__, "-v", "--tb=short"])
+1 -1
View File
@@ -4,7 +4,7 @@ description = "OpenAI ChatKit integration for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251120"
version = "1.0.0b251114"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
@@ -0,0 +1 @@
# Copyright (c) Microsoft. All rights reserved.
+1 -1
View File
@@ -4,7 +4,7 @@ description = "Copilot Studio integration for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251120"
version = "1.0.0b251114"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
@@ -5,7 +5,7 @@ import re
import sys
from collections.abc import AsyncIterable, Awaitable, Callable, MutableMapping, Sequence
from contextlib import AbstractAsyncContextManager, AsyncExitStack
from copy import deepcopy
from copy import copy
from itertools import chain
from typing import Any, ClassVar, Literal, Protocol, TypeVar, cast, runtime_checkable
from uuid import uuid4
@@ -454,16 +454,13 @@ class BaseAgent(SerializationMixin):
# Extract the input from kwargs using the specified arg_name
input_text = kwargs.get(arg_name, "")
# Forward all kwargs except the arg_name to support runtime context propagation
forwarded_kwargs = {k: v for k, v in kwargs.items() if k != arg_name}
if stream_callback is None:
# Use non-streaming mode
return (await self.run(input_text, **forwarded_kwargs)).text
return (await self.run(input_text)).text
# Use streaming mode - accumulate updates and create final response
response_updates: list[AgentRunResponseUpdate] = []
async for update in self.run_stream(input_text, **forwarded_kwargs):
async for update in self.run_stream(input_text):
response_updates.append(update)
if is_async_callback:
await stream_callback(update) # type: ignore[misc]
@@ -473,14 +470,12 @@ class BaseAgent(SerializationMixin):
# Create final text from accumulated updates
return AgentRunResponse.from_agent_run_response_updates(response_updates).text
agent_tool: AIFunction[BaseModel, str] = AIFunction(
return AIFunction(
name=tool_name,
description=tool_description,
func=agent_wrapper,
input_model=input_model, # type: ignore
)
agent_tool._forward_runtime_kwargs = True # type: ignore
return agent_tool
def _normalize_messages(
self,
@@ -594,7 +589,7 @@ class ChatAgent(BaseAgent):
chat_message_store_factory: Callable[[], ChatMessageStoreProtocol] | None = None,
context_providers: ContextProvider | list[ContextProvider] | AggregateContextProvider | None = None,
middleware: Middleware | list[Middleware] | None = None,
# chat options
# chat option params
allow_multiple_tool_calls: bool | None = None,
conversation_id: str | None = None,
frequency_penalty: float | None = None,
@@ -873,9 +868,7 @@ class ChatAgent(BaseAgent):
user=user,
**(additional_chat_options or {}),
)
# Filter chat_options from kwargs to prevent duplicate keyword argument
filtered_kwargs = {k: v for k, v in kwargs.items() if k != "chat_options"}
response = await self.chat_client.get_response(messages=thread_messages, chat_options=co, **filtered_kwargs)
response = await self.chat_client.get_response(messages=thread_messages, chat_options=co, **kwargs)
await self._update_thread_with_type_and_conversation_id(thread, response.conversation_id)
@@ -1007,11 +1000,9 @@ class ChatAgent(BaseAgent):
**(additional_chat_options or {}),
)
# Filter chat_options from kwargs to prevent duplicate keyword argument
filtered_kwargs = {k: v for k, v in kwargs.items() if k != "chat_options"}
response_updates: list[ChatResponseUpdate] = []
async for update in self.chat_client.get_streaming_response(
messages=thread_messages, chat_options=co, **filtered_kwargs
messages=thread_messages, chat_options=co, **kwargs
):
response_updates.append(update)
@@ -1245,7 +1236,7 @@ class ChatAgent(BaseAgent):
Raises:
AgentExecutionException: If the conversation IDs on the thread and agent don't match.
"""
chat_options = deepcopy(self.chat_options) if self.chat_options else ChatOptions()
chat_options = copy(self.chat_options) if self.chat_options else ChatOptions()
thread = thread or self.get_new_thread()
if thread.service_thread_id and thread.context_provider:
await thread.context_provider.thread_created(thread.service_thread_id)
@@ -214,7 +214,6 @@ def _merge_chat_options(
*,
base_chat_options: ChatOptions | Any | None,
model_id: str | None = None,
allow_multiple_tool_calls: bool | None = None,
frequency_penalty: float | None = None,
logit_bias: dict[str | int, float] | None = None,
max_tokens: int | None = None,
@@ -240,7 +239,6 @@ def _merge_chat_options(
Keyword Args:
base_chat_options: Optional base ChatOptions to merge with direct parameters.
model_id: The model_id to use for the agent.
allow_multiple_tool_calls: Whether to allow multiple tool calls in a single response.
frequency_penalty: The frequency penalty to use.
logit_bias: The logit bias to use.
max_tokens: The maximum number of tokens to generate.
@@ -272,7 +270,6 @@ def _merge_chat_options(
return base_chat_options & ChatOptions(
model_id=model_id,
allow_multiple_tool_calls=allow_multiple_tool_calls,
frequency_penalty=frequency_penalty,
logit_bias=logit_bias,
max_tokens=max_tokens,
@@ -488,7 +485,6 @@ class BaseChatClient(SerializationMixin, ABC):
self,
messages: str | ChatMessage | list[str] | list[ChatMessage],
*,
allow_multiple_tool_calls: bool | None = None,
frequency_penalty: float | None = None,
logit_bias: dict[str | int, float] | None = None,
max_tokens: int | None = None,
@@ -521,7 +517,6 @@ class BaseChatClient(SerializationMixin, ABC):
messages: The message or messages to send to the model.
Keyword Args:
allow_multiple_tool_calls: Whether to allow multiple tool calls in a single response.
frequency_penalty: The frequency penalty to use.
logit_bias: The logit bias to use.
max_tokens: The maximum number of tokens to generate.
@@ -550,7 +545,6 @@ class BaseChatClient(SerializationMixin, ABC):
chat_options = _merge_chat_options(
base_chat_options=kwargs.pop("chat_options", None),
model_id=model_id,
allow_multiple_tool_calls=allow_multiple_tool_calls,
frequency_penalty=frequency_penalty,
logit_bias=logit_bias,
max_tokens=max_tokens,
@@ -586,7 +580,6 @@ class BaseChatClient(SerializationMixin, ABC):
self,
messages: str | ChatMessage | list[str] | list[ChatMessage],
*,
allow_multiple_tool_calls: bool | None = None,
frequency_penalty: float | None = None,
logit_bias: dict[str | int, float] | None = None,
max_tokens: int | None = None,
@@ -619,7 +612,6 @@ class BaseChatClient(SerializationMixin, ABC):
messages: The message or messages to send to the model.
Keyword Args:
allow_multiple_tool_calls: Whether to allow multiple tool calls in a single response.
frequency_penalty: The frequency penalty to use.
logit_bias: The logit bias to use.
max_tokens: The maximum number of tokens to generate.
@@ -648,7 +640,6 @@ class BaseChatClient(SerializationMixin, ABC):
chat_options = _merge_chat_options(
base_chat_options=kwargs.pop("chat_options", None),
model_id=model_id,
allow_multiple_tool_calls=allow_multiple_tool_calls,
frequency_penalty=frequency_penalty,
logit_bias=logit_bias,
max_tokens=max_tokens,
+8 -65
View File
@@ -614,7 +614,6 @@ class AIFunction(BaseTool, Generic[ArgsT, ReturnT]):
**kwargs,
)
self.func = func
self._instance = None # Store the instance for bound methods
self.input_model = self._resolve_input_model(input_model)
self.approval_mode = approval_mode or "never_require"
if max_invocations is not None and max_invocations < 1:
@@ -627,47 +626,12 @@ class AIFunction(BaseTool, Generic[ArgsT, ReturnT]):
self.invocation_exception_count = 0
self._invocation_duration_histogram = _default_histogram()
self.type: Literal["ai_function"] = "ai_function"
self._forward_runtime_kwargs: bool = False
@property
def declaration_only(self) -> bool:
"""Indicate whether the function is declaration only (i.e., has no implementation)."""
# Check for explicit _declaration_only attribute first (used in tests)
if hasattr(self, "_declaration_only") and self._declaration_only:
return True
return self.func is None
def __get__(self, obj: Any, objtype: type | None = None) -> "AIFunction[ArgsT, ReturnT]":
"""Implement the descriptor protocol to support bound methods.
When an AIFunction is accessed as an attribute of a class instance,
this method is called to bind the instance to the function.
Args:
obj: The instance that owns the descriptor, or None for class access.
objtype: The type that owns the descriptor.
Returns:
A new AIFunction with the instance bound to the wrapped function.
"""
if obj is None:
# Accessed from the class, not an instance
return self
# Check if the wrapped function is a method (has 'self' parameter)
if self.func is not None:
sig = inspect.signature(self.func)
params = list(sig.parameters.keys())
if params and params[0] in {"self", "cls"}:
# Create a new AIFunction with the bound method
import copy
bound_func = copy.copy(self)
bound_func._instance = obj
return bound_func
return self
def _resolve_input_model(self, input_model: type[ArgsT] | Mapping[str, Any] | None) -> type[ArgsT]:
"""Resolve the input model for the function."""
if input_model is None:
@@ -682,7 +646,7 @@ class AIFunction(BaseTool, Generic[ArgsT, ReturnT]):
def __call__(self, *args: Any, **kwargs: Any) -> ReturnT | Awaitable[ReturnT]:
"""Call the wrapped function with the provided arguments."""
if self.declaration_only:
if self.func is None:
raise ToolException(f"Function '{self.name}' is declaration only and cannot be invoked.")
if self.max_invocations is not None and self.invocation_count >= self.max_invocations:
raise ToolException(
@@ -698,10 +662,7 @@ class AIFunction(BaseTool, Generic[ArgsT, ReturnT]):
)
self.invocation_count += 1
try:
# If we have a bound instance, call the function with self
if self._instance is not None:
return self.func(self._instance, *args, **kwargs)
return self.func(*args, **kwargs) # type:ignore[misc]
return self.func(*args, **kwargs)
except Exception:
self.invocation_exception_count += 1
raise
@@ -729,16 +690,11 @@ class AIFunction(BaseTool, Generic[ArgsT, ReturnT]):
global OBSERVABILITY_SETTINGS
from .observability import OBSERVABILITY_SETTINGS
original_kwargs = dict(kwargs)
tool_call_id = original_kwargs.pop("tool_call_id", None)
tool_call_id = kwargs.pop("tool_call_id", None)
if arguments is not None:
if not isinstance(arguments, self.input_model):
raise TypeError(f"Expected {self.input_model.__name__}, got {type(arguments).__name__}")
kwargs = arguments.model_dump(exclude_none=True)
if getattr(self, "_forward_runtime_kwargs", False) and original_kwargs:
kwargs.update(original_kwargs)
else:
kwargs = original_kwargs
if not OBSERVABILITY_SETTINGS.ENABLED: # type: ignore[name-defined]
logger.info(f"Function name: {self.name}")
logger.debug(f"Function arguments: {kwargs}")
@@ -902,12 +858,6 @@ def _parse_annotation(annotation: Any) -> Any:
def _create_input_model_from_func(func: Callable[..., Any], name: str) -> type[BaseModel]:
"""Create a Pydantic model from a function's signature."""
# Unwrap AIFunction objects to get the underlying function
from agent_framework._tools import AIFunction
if isinstance(func, AIFunction):
func = func.func # type: ignore[assignment]
sig = inspect.signature(func)
fields = {
pname: (
@@ -1278,20 +1228,15 @@ async def _auto_invoke_function(
parsed_args: dict[str, Any] = dict(function_call_content.parse_arguments() or {})
# Filter out internal framework kwargs before passing to tools.
runtime_kwargs: dict[str, Any] = {
key: value
for key, value in (custom_args or {}).items()
if key not in {"_function_middleware_pipeline", "middleware"}
}
# Merge with user-supplied args; right-hand side dominates, so parsed args win on conflicts.
merged_args: dict[str, Any] = (custom_args or {}) | parsed_args
try:
args = tool.input_model.model_validate(parsed_args)
args = tool.input_model.model_validate(merged_args)
except ValidationError as exc:
message = "Error: Argument parsing failed."
if config.include_detailed_errors:
message = f"{message} Exception: {exc}"
return FunctionResultContent(call_id=function_call_content.call_id, result=message, exception=exc)
if not middleware_pipeline or (
not hasattr(middleware_pipeline, "has_middlewares") and not middleware_pipeline.has_middlewares
):
@@ -1300,8 +1245,7 @@ async def _auto_invoke_function(
function_result = await tool.invoke(
arguments=args,
tool_call_id=function_call_content.call_id,
**runtime_kwargs if getattr(tool, "_forward_runtime_kwargs", False) else {},
)
) # type: ignore[arg-type]
return FunctionResultContent(
call_id=function_call_content.call_id,
result=function_result,
@@ -1317,14 +1261,13 @@ async def _auto_invoke_function(
middleware_context = FunctionInvocationContext(
function=tool,
arguments=args,
kwargs=runtime_kwargs.copy(),
kwargs=custom_args or {},
)
async def final_function_handler(context_obj: Any) -> Any:
return await tool.invoke(
arguments=context_obj.arguments,
tool_call_id=function_call_content.call_id,
**context_obj.kwargs if getattr(tool, "_forward_runtime_kwargs", False) else {},
)
try:
@@ -3173,40 +3173,6 @@ class ChatOptions(SerializationMixin):
self.top_p = top_p
self.user = user
def __deepcopy__(self, memo: dict[int, Any]) -> "ChatOptions":
"""Create a runtime-safe copy without deep-copying tool instances."""
clone = type(self).__new__(type(self))
memo[id(self)] = clone
for key, value in self.__dict__.items():
if key == "_tools":
setattr(clone, key, list(value) if value is not None else None)
continue
if key in {"logit_bias", "metadata", "additional_properties"}:
setattr(clone, key, self._safe_deepcopy_mapping(value, memo))
continue
setattr(clone, key, self._safe_deepcopy_value(value, memo))
return clone
@staticmethod
def _safe_deepcopy_mapping(
value: MutableMapping[str, Any] | None, memo: dict[int, Any]
) -> MutableMapping[str, Any] | None:
"""Deep copy helper that falls back to a shallow copy for problematic mappings."""
if value is None:
return None
try:
return deepcopy(value, memo) # type: ignore[arg-type]
except Exception:
return dict(value)
@staticmethod
def _safe_deepcopy_value(value: Any, memo: dict[int, Any]) -> Any:
"""Deep copy helper that avoids failing on non-copyable instances."""
try:
return deepcopy(value, memo)
except Exception:
return value
@property
def tools(self) -> list[ToolProtocol | MutableMapping[str, Any]] | None:
"""Return the tools that are specified."""
@@ -10,8 +10,6 @@ _IMPORTS: dict[str, tuple[str, str]] = {
"AgentResponseCallbackProtocol": ("agent_framework_azurefunctions", "azurefunctions"),
"AzureAIAgentClient": ("agent_framework_azure_ai", "azure-ai"),
"AzureAIClient": ("agent_framework_azure_ai", "azure-ai"),
"AzureAISearchContextProvider": ("agent_framework_aisearch", "aisearch"),
"AzureAISearchSettings": ("agent_framework_aisearch", "aisearch"),
"AzureOpenAIAssistantsClient": ("agent_framework.azure._assistants_client", "core"),
"AzureOpenAIChatClient": ("agent_framework.azure._chat_client", "core"),
"AzureAISettings": ("agent_framework_azure_ai", "azure-ai"),
@@ -1,23 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
import importlib
from typing import Any
IMPORT_PATH = "agent_framework_declarative"
PACKAGE_NAME = "agent-framework-declarative"
_IMPORTS = ["__version__", "AgentFactory", "DeclarativeLoaderError", "ProviderLookupError", "ProviderTypeMapping"]
def __getattr__(name: str) -> Any:
if name in _IMPORTS:
try:
return getattr(importlib.import_module(IMPORT_PATH), name)
except ModuleNotFoundError as exc:
raise ModuleNotFoundError(
f"The '{PACKAGE_NAME}' package is not installed, please do `pip install {PACKAGE_NAME}`"
) from exc
raise AttributeError(f"Module {IMPORT_PATH} has no attribute {name}.")
def __dir__() -> list[str]:
return _IMPORTS
@@ -1,17 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
from agent_framework_declarative import (
AgentFactory,
DeclarativeLoaderError,
ProviderLookupError,
ProviderTypeMapping,
__version__,
)
__all__ = [
"AgentFactory",
"DeclarativeLoaderError",
"ProviderLookupError",
"ProviderTypeMapping",
"__version__",
]
@@ -1021,7 +1021,7 @@ def use_observability(
.. code-block:: python
from agent_framework import use_observability, setup_observability
from agent_framework import ChatClientProtocol
from agent_framework._clients import ChatClientProtocol
# Decorate a custom chat client class
@@ -1104,7 +1104,6 @@ def _trace_agent_run(
if not OBSERVABILITY_SETTINGS.ENABLED:
# If model diagnostics are not enabled, just return the completion
return await run_func(self, messages=messages, thread=thread, **kwargs)
filtered_kwargs = {k: v for k, v in kwargs.items() if k != "chat_options"}
attributes = _get_span_attributes(
operation_name=OtelAttr.AGENT_INVOKE_OPERATION,
provider_name=provider_name,
@@ -1113,7 +1112,7 @@ def _trace_agent_run(
agent_description=self.description,
thread_id=thread.service_thread_id if thread else None,
chat_options=getattr(self, "chat_options", None),
**filtered_kwargs,
**kwargs,
)
with _get_span(attributes=attributes, span_name_attribute=OtelAttr.AGENT_NAME) as span:
if OBSERVABILITY_SETTINGS.SENSITIVE_DATA_ENABLED and messages:
@@ -1174,7 +1173,6 @@ def _trace_agent_run_stream(
all_updates: list["AgentRunResponseUpdate"] = []
filtered_kwargs = {k: v for k, v in kwargs.items() if k != "chat_options"}
attributes = _get_span_attributes(
operation_name=OtelAttr.AGENT_INVOKE_OPERATION,
provider_name=provider_name,
@@ -1183,7 +1181,7 @@ def _trace_agent_run_stream(
agent_description=self.description,
thread_id=thread.service_thread_id if thread else None,
chat_options=getattr(self, "chat_options", None),
**filtered_kwargs,
**kwargs,
)
with _get_span(attributes=attributes, span_name_attribute=OtelAttr.AGENT_NAME) as span:
if OBSERVABILITY_SETTINGS.SENSITIVE_DATA_ENABLED and messages:
@@ -1474,10 +1472,10 @@ def _to_otel_part(content: "Contents") -> dict[str, Any] | None:
elif isinstance(item, BaseModel):
res.append(item.model_dump(exclude_none=True))
else:
res.append(json.dumps(item, default=str))
response = json.dumps(res, default=str)
res.append(json.dumps(item))
response = json.dumps(res)
else:
response = json.dumps(content.result, default=str)
response = json.dumps(content.result)
return {"type": "tool_call_response", "id": content.call_id, "response": response}
case _:
# GenericPart in otel output messages json spec.
+1 -6
View File
@@ -4,7 +4,7 @@ description = "Microsoft Agent Framework for building AI Agents with Python. Thi
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251120"
version = "1.0.0b251114"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
@@ -43,15 +43,10 @@ dependencies = [
all = [
"agent-framework-a2a",
"agent-framework-ag-ui",
"agent-framework-aisearch",
"agent-framework-anthropic",
"agent-framework-azure-ai",
"agent-framework-azurefunctions",
"agent-framework-chatkit",
"agent-framework-copilotstudio",
"agent-framework-declarative",
"agent-framework-devui",
"agent-framework-lab",
"agent-framework-mem0",
"agent-framework-purview",
"agent-framework-redis",
@@ -115,26 +115,6 @@ async def test_chat_client_agent_prepare_thread_and_messages(chat_client: ChatCl
assert result_messages[1].text == "Test"
async def test_prepare_thread_does_not_mutate_agent_chat_options(chat_client: ChatClientProtocol) -> None:
tool = HostedCodeInterpreterTool()
agent = ChatAgent(chat_client=chat_client, tools=[tool])
assert agent.chat_options.tools is not None
base_tools = agent.chat_options.tools
thread = agent.get_new_thread()
_, prepared_chat_options, _ = await agent._prepare_thread_and_messages( # type: ignore[reportPrivateUsage]
thread=thread,
input_messages=[ChatMessage(role=Role.USER, text="Test")],
)
assert prepared_chat_options.tools is not None
assert base_tools is not prepared_chat_options.tools
prepared_chat_options.tools.append(HostedCodeInterpreterTool()) # type: ignore[arg-type]
assert len(agent.chat_options.tools) == 1
async def test_chat_client_agent_update_thread_id(chat_client_base: ChatClientProtocol) -> None:
mock_response = ChatResponse(
messages=[ChatMessage(role=Role.ASSISTANT, contents=[TextContent("test response")])],
@@ -1,315 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
"""Tests for kwargs propagation through as_tool() method."""
from collections.abc import Awaitable, Callable
from typing import Any
from agent_framework import ChatAgent, ChatMessage, ChatResponse, FunctionCallContent, agent_middleware
from agent_framework._middleware import AgentRunContext
from .conftest import MockChatClient
class TestAsToolKwargsPropagation:
"""Test cases for kwargs propagation through as_tool() delegation."""
async def test_as_tool_forwards_runtime_kwargs(self, chat_client: MockChatClient) -> None:
"""Test that runtime kwargs are forwarded through as_tool() to sub-agent."""
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
# Capture kwargs passed to the sub-agent
captured_kwargs.update(context.kwargs)
await next(context)
# Setup mock response
chat_client.responses = [
ChatResponse(messages=[ChatMessage(role="assistant", text="Response from sub-agent")]),
]
# Create sub-agent with middleware
sub_agent = ChatAgent(
chat_client=chat_client,
name="sub_agent",
middleware=[capture_middleware],
)
# Create tool from sub-agent
tool = sub_agent.as_tool(name="delegate", arg_name="task")
# Directly invoke the tool with kwargs (simulating what happens during agent execution)
_ = await tool.invoke(
arguments=tool.input_model(task="Test delegation"),
api_token="secret-xyz-123",
user_id="user-456",
session_id="session-789",
)
# Verify kwargs were forwarded to sub-agent
assert "api_token" in captured_kwargs, f"Expected 'api_token' in {captured_kwargs}"
assert captured_kwargs["api_token"] == "secret-xyz-123"
assert "user_id" in captured_kwargs
assert captured_kwargs["user_id"] == "user-456"
assert "session_id" in captured_kwargs
assert captured_kwargs["session_id"] == "session-789"
async def test_as_tool_excludes_arg_name_from_forwarded_kwargs(self, chat_client: MockChatClient) -> None:
"""Test that the arg_name parameter is not forwarded as a kwarg."""
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
captured_kwargs.update(context.kwargs)
await next(context)
# Setup mock response
chat_client.responses = [
ChatResponse(messages=[ChatMessage(role="assistant", text="Response from sub-agent")]),
]
sub_agent = ChatAgent(
chat_client=chat_client,
name="sub_agent",
middleware=[capture_middleware],
)
tool = sub_agent.as_tool(arg_name="custom_task")
# Invoke tool with both the arg_name field and additional kwargs
await tool.invoke(
arguments=tool.input_model(custom_task="Test task"),
api_token="token-123",
custom_task="should_be_excluded", # This should be filtered out
)
# The arg_name ("custom_task") should NOT be in the forwarded kwargs
assert "custom_task" not in captured_kwargs
# But other kwargs should be present
assert "api_token" in captured_kwargs
assert captured_kwargs["api_token"] == "token-123"
async def test_as_tool_nested_delegation_propagates_kwargs(self, chat_client: MockChatClient) -> None:
"""Test that kwargs propagate through multiple levels of delegation (A → B → C)."""
captured_kwargs_list: list[dict[str, Any]] = []
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
# Capture kwargs at each level
captured_kwargs_list.append(dict(context.kwargs))
await next(context)
# Setup mock responses to trigger nested tool invocation: B calls tool C, then completes.
chat_client.responses = [
ChatResponse(
messages=[
ChatMessage(
role="assistant",
contents=[
FunctionCallContent(
call_id="call_c_1",
name="call_c",
arguments='{"task": "Please execute agent_c"}',
)
],
)
]
),
ChatResponse(messages=[ChatMessage(role="assistant", text="Response from agent_c")]),
ChatResponse(messages=[ChatMessage(role="assistant", text="Response from agent_b")]),
]
# Create agent C (bottom level)
agent_c = ChatAgent(
chat_client=chat_client,
name="agent_c",
middleware=[capture_middleware],
)
# Create agent B (middle level) - delegates to C
agent_b = ChatAgent(
chat_client=chat_client,
name="agent_b",
tools=[agent_c.as_tool(name="call_c")],
middleware=[capture_middleware],
)
# Create tool from B for direct invocation
tool_b = agent_b.as_tool(name="call_b")
# Invoke tool B with kwargs - should propagate to both B and C
await tool_b.invoke(
arguments=tool_b.input_model(task="Test cascade"),
trace_id="trace-abc-123",
tenant_id="tenant-xyz",
)
# Verify both levels received the kwargs
# We should have 2 captures: one from B, one from C
assert len(captured_kwargs_list) >= 2
for kwargs_dict in captured_kwargs_list:
assert kwargs_dict.get("trace_id") == "trace-abc-123"
assert kwargs_dict.get("tenant_id") == "tenant-xyz"
async def test_as_tool_streaming_mode_forwards_kwargs(self, chat_client: MockChatClient) -> None:
"""Test that kwargs are forwarded in streaming mode."""
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
captured_kwargs.update(context.kwargs)
await next(context)
# Setup mock streaming responses
from agent_framework import ChatResponseUpdate, TextContent
chat_client.streaming_responses = [
[ChatResponseUpdate(text=TextContent(text="Streaming response"), role="assistant")],
]
sub_agent = ChatAgent(
chat_client=chat_client,
name="sub_agent",
middleware=[capture_middleware],
)
captured_updates: list[Any] = []
async def stream_callback(update: Any) -> None:
captured_updates.append(update)
tool = sub_agent.as_tool(stream_callback=stream_callback)
# Invoke tool with kwargs while streaming callback is active
await tool.invoke(
arguments=tool.input_model(task="Test streaming"),
api_key="streaming-key-999",
)
# Verify kwargs were forwarded even in streaming mode
assert "api_key" in captured_kwargs
assert captured_kwargs["api_key"] == "streaming-key-999"
assert len(captured_updates) == 1
async def test_as_tool_empty_kwargs_still_works(self, chat_client: MockChatClient) -> None:
"""Test that as_tool works correctly when no extra kwargs are provided."""
# Setup mock response
chat_client.responses = [
ChatResponse(messages=[ChatMessage(role="assistant", text="Response from agent")]),
]
sub_agent = ChatAgent(
chat_client=chat_client,
name="sub_agent",
)
tool = sub_agent.as_tool()
# Invoke without any extra kwargs - should work without errors
result = await tool.invoke(arguments=tool.input_model(task="Simple task"))
# Verify tool executed successfully
assert result is not None
async def test_as_tool_kwargs_with_chat_options(self, chat_client: MockChatClient) -> None:
"""Test that kwargs including chat_options are properly forwarded."""
captured_kwargs: dict[str, Any] = {}
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
captured_kwargs.update(context.kwargs)
await next(context)
# Setup mock response
chat_client.responses = [
ChatResponse(messages=[ChatMessage(role="assistant", text="Response with options")]),
]
sub_agent = ChatAgent(
chat_client=chat_client,
name="sub_agent",
middleware=[capture_middleware],
)
tool = sub_agent.as_tool()
# Invoke with various kwargs
await tool.invoke(
arguments=tool.input_model(task="Test with options"),
temperature=0.8,
max_tokens=500,
custom_param="custom_value",
)
# Verify all kwargs were forwarded
assert "temperature" in captured_kwargs
assert captured_kwargs["temperature"] == 0.8
assert "max_tokens" in captured_kwargs
assert captured_kwargs["max_tokens"] == 500
assert "custom_param" in captured_kwargs
assert captured_kwargs["custom_param"] == "custom_value"
async def test_as_tool_kwargs_isolated_per_invocation(self, chat_client: MockChatClient) -> None:
"""Test that kwargs are isolated per invocation and don't leak between calls."""
first_call_kwargs: dict[str, Any] = {}
second_call_kwargs: dict[str, Any] = {}
call_count = 0
@agent_middleware
async def capture_middleware(
context: AgentRunContext, next: Callable[[AgentRunContext], Awaitable[None]]
) -> None:
nonlocal call_count
call_count += 1
if call_count == 1:
first_call_kwargs.update(context.kwargs)
elif call_count == 2:
second_call_kwargs.update(context.kwargs)
await next(context)
# Setup mock responses for both calls
chat_client.responses = [
ChatResponse(messages=[ChatMessage(role="assistant", text="First response")]),
ChatResponse(messages=[ChatMessage(role="assistant", text="Second response")]),
]
sub_agent = ChatAgent(
chat_client=chat_client,
name="sub_agent",
middleware=[capture_middleware],
)
tool = sub_agent.as_tool()
# First call with specific kwargs
await tool.invoke(
arguments=tool.input_model(task="First task"),
session_id="session-1",
api_token="token-1",
)
# Second call with different kwargs
await tool.invoke(
arguments=tool.input_model(task="Second task"),
session_id="session-2",
api_token="token-2",
)
# Verify first call had its own kwargs
assert first_call_kwargs.get("session_id") == "session-1"
assert first_call_kwargs.get("api_token") == "token-1"
# Verify second call had its own kwargs (not leaked from first)
assert second_call_kwargs.get("session_id") == "session-2"
assert second_call_kwargs.get("api_token") == "token-2"
@@ -3,7 +3,6 @@
import pytest
from agent_framework import (
ChatAgent,
ChatClientProtocol,
ChatMessage,
ChatOptions,
@@ -128,148 +127,6 @@ async def test_base_client_with_streaming_function_calling(chat_client_base: Cha
assert exec_counter == 1
async def test_function_invocation_inside_aiohttp_server(chat_client_base: ChatClientProtocol):
import aiohttp
from aiohttp import web
exec_counter = 0
@ai_function(name="start_todo_investigation")
def ai_func(user_query: str) -> str:
nonlocal exec_counter
exec_counter += 1
return f"Investigated {user_query}"
chat_client_base.run_responses = [
ChatResponse(
messages=ChatMessage(
role="assistant",
contents=[
FunctionCallContent(
call_id="1",
name="start_todo_investigation",
arguments='{"user_query": "issue"}',
)
],
)
),
ChatResponse(messages=ChatMessage(role="assistant", text="done")),
]
agent = ChatAgent(chat_client=chat_client_base, tools=[ai_func])
async def handler(request: web.Request) -> web.Response:
thread = agent.get_new_thread()
result = await agent.run("Fix issue", thread=thread)
return web.Response(text=result.text or "")
app = web.Application()
app.add_routes([web.post("/run", handler)])
runner = web.AppRunner(app)
await runner.setup()
site = web.TCPSite(runner, "127.0.0.1", 0)
await site.start()
try:
port = site._server.sockets[0].getsockname()[1]
async with aiohttp.ClientSession() as session, session.post(f"http://127.0.0.1:{port}/run") as response:
assert response.status == 200
await response.text()
finally:
await runner.cleanup()
assert exec_counter == 1
async def test_function_invocation_in_threaded_aiohttp_app(chat_client_base: ChatClientProtocol):
import asyncio
import threading
from queue import Queue
import aiohttp
from aiohttp import web
exec_counter = 0
@ai_function(name="start_threaded_investigation")
def ai_func(user_query: str) -> str:
nonlocal exec_counter
exec_counter += 1
return f"Threaded {user_query}"
chat_client_base.run_responses = [
ChatResponse(
messages=ChatMessage(
role="assistant",
contents=[
FunctionCallContent(
call_id="thread-1",
name="start_threaded_investigation",
arguments='{"user_query": "issue"}',
)
],
)
),
ChatResponse(messages=ChatMessage(role="assistant", text="done")),
]
agent = ChatAgent(chat_client=chat_client_base, tools=[ai_func])
ready_event = threading.Event()
port_queue: Queue[int] = Queue()
shutdown_queue: Queue[tuple[asyncio.AbstractEventLoop, asyncio.Event]] = Queue()
async def init_app() -> web.Application:
async def handler(request: web.Request) -> web.Response:
thread = agent.get_new_thread()
result = await agent.run("Fix issue", thread=thread)
return web.Response(text=result.text or "")
app = web.Application()
app.add_routes([web.post("/run", handler)])
return app
def server_thread() -> None:
loop = asyncio.new_event_loop()
asyncio.set_event_loop(loop)
async def runner_main() -> None:
app = await init_app()
runner = web.AppRunner(app)
await runner.setup()
site = web.TCPSite(runner, "127.0.0.1", 0)
await site.start()
shutdown_event = asyncio.Event()
shutdown_queue.put((loop, shutdown_event))
port = site._server.sockets[0].getsockname()[1]
port_queue.put(port)
ready_event.set()
try:
await shutdown_event.wait()
finally:
await runner.cleanup()
try:
loop.run_until_complete(runner_main())
finally:
loop.close()
thread = threading.Thread(target=server_thread, daemon=True)
thread.start()
ready_event.wait(timeout=5)
assert ready_event.is_set()
loop_ref, shutdown_event = shutdown_queue.get(timeout=2)
port = port_queue.get(timeout=2)
async with aiohttp.ClientSession() as session, session.post(f"http://127.0.0.1:{port}/run") as response:
assert response.status == 200
await response.text()
loop_ref.call_soon_threadsafe(shutdown_event.set)
thread.join(timeout=5)
assert exec_counter == 1
@pytest.mark.parametrize(
"approval_required,num_functions",
[
@@ -1448,20 +1305,26 @@ async def test_approved_function_call_successful_execution(chat_client_base: Cha
assert success_result.result == "Success value1"
async def test_declaration_only_tool(chat_client_base: ChatClientProtocol):
"""Test that declaration_only tools without implementation (func=None) are not executed."""
async def test_declaration_only_tool_not_executed(chat_client_base: ChatClientProtocol):
"""Test that declaration_only tools are not executed."""
exec_counter = 0
@ai_function(name="declaration_func")
def declaration_func_inner(arg1: str) -> str:
nonlocal exec_counter
exec_counter += 1
return f"Result {arg1}"
# Create a new AIFunction with declaration_only set
from agent_framework import AIFunction
# Create a truly declaration-only function with no implementation
declaration_func = AIFunction(
name="declaration_func",
func=None,
description="A declaration-only function for testing",
input_model={"type": "object", "properties": {"arg1": {"type": "string"}}, "required": ["arg1"]},
func=declaration_func_inner,
additional_properties={"declaration_only": True},
)
# Verify it's marked as declaration_only
assert declaration_func.declaration_only is True
# Set declaration_only on the instance
object.__setattr__(declaration_func, "_declaration_only", True)
chat_client_base.run_responses = [
ChatResponse(
@@ -1475,6 +1338,8 @@ async def test_declaration_only_tool(chat_client_base: ChatClientProtocol):
response = await chat_client_base.get_response("hello", tool_choice="auto", tools=[declaration_func])
# Function should NOT be executed
assert exec_counter == 0
# Should have the function call in messages but not a result
function_calls = [
content
@@ -1484,15 +1349,6 @@ async def test_declaration_only_tool(chat_client_base: ChatClientProtocol):
]
assert len(function_calls) >= 1
# Should not have a function result
function_results = [
content
for msg in response.messages
for content in msg.contents
if isinstance(content, FunctionResultContent) and content.call_id == "1"
]
assert len(function_results) == 0
async def test_multiple_function_calls_parallel_execution(chat_client_base: ChatClientProtocol):
"""Test that multiple function calls are executed in parallel."""
@@ -1693,7 +1693,7 @@ def mock_function() -> AIFunction[Any, Any]:
@pytest.fixture
def mock_chat_client() -> Any:
"""Mock chat client for testing."""
from agent_framework import ChatClientProtocol
from agent_framework._clients import ChatClientProtocol
client = MagicMock(spec=ChatClientProtocol)
client.service_url = MagicMock(return_value="mock://test")
@@ -3,7 +3,7 @@
import logging
from collections.abc import MutableSequence
from typing import Any
from unittest.mock import Mock
from unittest.mock import MagicMock, Mock, patch
import pytest
from opentelemetry.sdk.trace.export.in_memory_span_exporter import InMemorySpanExporter
@@ -22,7 +22,6 @@ from agent_framework import (
ChatResponseUpdate,
Role,
UsageDetails,
ai_function,
prepend_agent_framework_to_user_agent,
)
from agent_framework.exceptions import AgentInitializationError, ChatClientInitializationError
@@ -479,46 +478,32 @@ async def test_agent_streaming_response_with_diagnostics_enabled_via_decorator(
assert span.attributes.get(OtelAttr.OUTPUT_MESSAGES) is not None # Streaming, so no usage yet
async def test_function_call_with_error_handling(span_exporter: InMemorySpanExporter):
"""Test that function call errors are properly captured in telemetry."""
async def test_agent_run_with_exception_handling(mock_chat_agent: AgentProtocol):
"""Test agent run with exception handling."""
# Create a function that raises an error using the decorator
@ai_function(name="failing_function", description="A function that fails")
async def failing_function(param: str) -> str:
raise ValueError("Function execution failed")
async def run_with_error(self, messages=None, *, thread=None, **kwargs):
raise RuntimeError("Agent run error")
span_exporter.clear()
mock_chat_agent.run = run_with_error
# Execute function and expect it to raise an error
with pytest.raises(ValueError, match="Function execution failed"):
await failing_function.invoke(param="test_value", tool_call_id="test_call_456")
agent = use_agent_observability(mock_chat_agent)()
# Verify span was created and error was captured
spans = span_exporter.get_finished_spans()
assert len(spans) == 1
span = spans[0]
from opentelemetry.trace import Span
# Verify span name and basic attributes
assert span.name == "execute_tool failing_function"
assert span.attributes is not None
assert span.attributes[OtelAttr.OPERATION.value] == OtelAttr.TOOL_EXECUTION_OPERATION
assert span.attributes[OtelAttr.TOOL_NAME] == "failing_function"
assert span.attributes[OtelAttr.TOOL_CALL_ID] == "test_call_456"
with (
patch("agent_framework.observability._get_span") as mock_get_span,
):
mock_span = MagicMock(spec=Span)
# Ensure the patched context manager returns mock_span when entered
mock_get_span.return_value.__enter__.return_value = mock_span
# Should raise the exception and call error handler
with pytest.raises(RuntimeError, match="Agent run error"):
await agent.run("Test message")
# Verify error status was set
assert span.status.status_code == StatusCode.ERROR
assert span.status.description is not None
assert "Function execution failed" in span.status.description
# Verify error type attribute was set
assert span.attributes[OtelAttr.ERROR_TYPE] == "ValueError"
# Verify exception event was recorded
assert len(span.events) > 0
exception_event = next((e for e in span.events if e.name == "exception"), None)
assert exception_event is not None
assert exception_event.attributes is not None
assert exception_event.attributes["exception.type"] == "ValueError"
exception_message = exception_event.attributes["exception.message"]
assert isinstance(exception_message, str)
assert "Function execution failed" in exception_message
# Verify error was recorded
# Check that both error attributes were set on the span
mock_span.set_attribute.assert_called_with(OtelAttr.ERROR_TYPE, "RuntimeError")
mock_span.record_exception.assert_called_once()
mock_span.set_status.assert_called_once_with(
status=StatusCode.ERROR, description=repr(RuntimeError("Agent run error"))
)
@@ -104,136 +104,6 @@ async def test_ai_function_decorator_with_async():
assert (await async_test_tool(1, 2)) == 3
def test_ai_function_decorator_in_class():
"""Test the ai_function decorator."""
class my_tools:
@ai_function(name="test_tool", description="A test tool")
def test_tool(self, x: int, y: int) -> int:
"""A simple function that adds two numbers."""
return x + y
test_tool = my_tools().test_tool
assert isinstance(test_tool, ToolProtocol)
assert isinstance(test_tool, AIFunction)
assert test_tool.name == "test_tool"
assert test_tool.description == "A test tool"
assert test_tool.parameters() == {
"properties": {"x": {"title": "X", "type": "integer"}, "y": {"title": "Y", "type": "integer"}},
"required": ["x", "y"],
"title": "test_tool_input",
"type": "object",
}
assert test_tool(1, 2) == 3
async def test_ai_function_decorator_shared_state():
"""Test that decorated methods maintain shared state across multiple calls and tool usage."""
class StatefulCounter:
"""A class that maintains a counter and provides decorated methods to interact with it."""
def __init__(self, initial_value: int = 0):
self.counter = initial_value
self.operation_log: list[str] = []
@ai_function(name="increment", description="Increment the counter")
def increment(self, amount: int) -> str:
"""Increment the counter by the given amount."""
self.counter += amount
self.operation_log.append(f"increment({amount})")
return f"Counter incremented by {amount}. New value: {self.counter}"
@ai_function(name="get_value", description="Get the current counter value")
def get_value(self) -> str:
"""Get the current counter value."""
self.operation_log.append("get_value()")
return f"Current counter value: {self.counter}"
@ai_function(name="multiply", description="Multiply the counter")
def multiply(self, factor: int) -> str:
"""Multiply the counter by the given factor."""
self.counter *= factor
self.operation_log.append(f"multiply({factor})")
return f"Counter multiplied by {factor}. New value: {self.counter}"
# Create a single instance with shared state
counter_instance = StatefulCounter(initial_value=10)
# Get the decorated methods - these will be used by different "agents" or tools
increment_tool = counter_instance.increment
get_value_tool = counter_instance.get_value
multiply_tool = counter_instance.multiply
# Verify they are AIFunction instances
assert isinstance(increment_tool, AIFunction)
assert isinstance(get_value_tool, AIFunction)
assert isinstance(multiply_tool, AIFunction)
# Tool 1 (increment) is used
result1 = increment_tool(5)
assert result1 == "Counter incremented by 5. New value: 15"
assert counter_instance.counter == 15
# Tool 2 (get_value) sees the state change from tool 1
result2 = get_value_tool()
assert result2 == "Current counter value: 15"
assert counter_instance.counter == 15
# Tool 3 (multiply) modifies the shared state
result3 = multiply_tool(3)
assert result3 == "Counter multiplied by 3. New value: 45"
assert counter_instance.counter == 45
# Tool 2 (get_value) sees the state change from tool 3
result4 = get_value_tool()
assert result4 == "Current counter value: 45"
assert counter_instance.counter == 45
# Tool 1 (increment) sees the current state and modifies it
result5 = increment_tool(10)
assert result5 == "Counter incremented by 10. New value: 55"
assert counter_instance.counter == 55
# Verify the operation log shows all operations in order
assert counter_instance.operation_log == [
"increment(5)",
"get_value()",
"multiply(3)",
"get_value()",
"increment(10)",
]
# Verify the parameters don't include 'self'
assert increment_tool.parameters() == {
"properties": {"amount": {"title": "Amount", "type": "integer"}},
"required": ["amount"],
"title": "increment_input",
"type": "object",
}
assert multiply_tool.parameters() == {
"properties": {"factor": {"title": "Factor", "type": "integer"}},
"required": ["factor"],
"title": "multiply_input",
"type": "object",
}
assert get_value_tool.parameters() == {
"properties": {},
"title": "get_value_input",
"type": "object",
}
# Test with invoke method as well (simulating agent execution)
result6 = await increment_tool.invoke(amount=5)
assert result6 == "Counter incremented by 5. New value: 60"
assert counter_instance.counter == 60
result7 = await get_value_tool.invoke()
assert result7 == "Current counter value: 60"
assert counter_instance.counter == 60
async def test_ai_function_invoke_telemetry_enabled(span_exporter: InMemorySpanExporter):
"""Test the ai_function invoke method with telemetry enabled."""
@@ -321,26 +191,6 @@ async def test_ai_function_invoke_telemetry_sensitive_disabled(span_exporter: In
assert attributes[OtelAttr.TOOL_CALL_ID] == "test_call_id"
async def test_ai_function_invoke_ignores_additional_kwargs() -> None:
"""Ensure ai_function tools drop unknown kwargs when invoked with validated arguments."""
@ai_function
async def simple_tool(message: str) -> str:
"""Echo tool."""
return message.upper()
args = simple_tool.input_model(message="hello world")
# These kwargs simulate runtime context passed through function invocation.
result = await simple_tool.invoke(
arguments=args,
api_token="secret-token",
chat_options={"model_id": "dummy"},
)
assert result == "HELLO WORLD"
async def test_ai_function_invoke_telemetry_with_pydantic_args(span_exporter: InMemorySpanExporter):
"""Test the ai_function invoke method with Pydantic model arguments."""
@@ -1,26 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
"""Test datetime serialization in observability telemetry."""
import json
from datetime import datetime
from agent_framework._types import FunctionResultContent
from agent_framework.observability import _to_otel_part
def test_datetime_in_tool_results() -> None:
"""Test that tool results with datetime values are serialized.
Reproduces issue #2219 where datetime objects caused TypeError.
"""
content = FunctionResultContent(
call_id="test-call",
result={"timestamp": datetime(2025, 11, 16, 10, 30, 0)},
)
result = _to_otel_part(content)
parsed = json.loads(result["response"])
# Datetime should be converted to string
assert isinstance(parsed["timestamp"], str)
-21
View File
@@ -1,21 +0,0 @@
MIT License
Copyright (c) Microsoft Corporation.
Permission is hereby granted, free of charge, to any person obtaining a copy
of this software and associated documentation files (the "Software"), to deal
in the Software without restriction, including without limitation the rights
to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
copies of the Software, and to permit persons to whom the Software is
furnished to do so, subject to the following conditions:
The above copyright notice and this permission notice shall be included in all
copies or substantial portions of the Software.
THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
SOFTWARE
-11
View File
@@ -1,11 +0,0 @@
# Get Started with Microsoft Agent Framework Declarative
Please install this package via pip:
```bash
pip install agent-framework-declarative --pre
```
## Declarative features
The declarative packages provides support for building agents based on a declarative yaml specification.
@@ -1,12 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
from importlib import metadata
from ._loader import AgentFactory, DeclarativeLoaderError, ProviderLookupError, ProviderTypeMapping
try:
__version__ = metadata.version(__name__)
except metadata.PackageNotFoundError:
__version__ = "0.0.0" # Fallback for development mode
__all__ = ["AgentFactory", "DeclarativeLoaderError", "ProviderLookupError", "ProviderTypeMapping", "__version__"]
@@ -1,422 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
from collections.abc import Callable, Mapping
from pathlib import Path
from typing import Any, Literal, TypedDict
import yaml
from agent_framework import (
AIFunction,
ChatAgent,
ChatClientProtocol,
HostedCodeInterpreterTool,
HostedFileContent,
HostedFileSearchTool,
HostedMCPSpecificApproval,
HostedMCPTool,
HostedVectorStoreContent,
HostedWebSearchTool,
ToolProtocol,
)
from agent_framework._tools import _create_model_from_json_schema # type: ignore
from agent_framework.exceptions import AgentFrameworkException
from dotenv import load_dotenv
from ._models import (
AnonymousConnection,
ApiKeyConnection,
CodeInterpreterTool,
FileSearchTool,
FunctionTool,
McpServerToolSpecifyApprovalMode,
McpTool,
Model,
ModelOptions,
PromptAgent,
ReferenceConnection,
RemoteConnection,
Tool,
WebSearchTool,
agent_schema_dispatch,
)
class ProviderTypeMapping(TypedDict, total=True):
package: str
name: str
model_id_field: str
PROVIDER_TYPE_OBJECT_MAPPING: dict[str, ProviderTypeMapping] = {
"AzureOpenAI.Chat": {
"package": "agent_framework.azure",
"name": "AzureOpenAIChatClient",
"model_id_field": "deployment_name",
},
"AzureOpenAI.Assistants": {
"package": "agent_framework.azure",
"name": "AzureOpenAIAssistantsClient",
"model_id_field": "deployment_name",
},
"AzureOpenAI.Responses": {
"package": "agent_framework.azure",
"name": "AzureOpenAIResponsesClient",
"model_id_field": "deployment_name",
},
"OpenAI.Chat": {
"package": "agent_framework.openai",
"name": "OpenAIChatClient",
"model_id_field": "model_id",
},
"OpenAI.Assistants": {
"package": "agent_framework.openai",
"name": "OpenAIAssistantsClient",
"model_id_field": "model_id",
},
"OpenAI.Responses": {
"package": "agent_framework.openai",
"name": "OpenAIResponsesClient",
"model_id_field": "model_id",
},
"AzureAIAgentClient": {
"package": "agent_framework.azure",
"name": "AzureAIAgentClient",
"model_id_field": "model_deployment_name",
},
"AzureAIClient": {
"package": "agent_framework.azure",
"name": "AzureAIClient",
"model_id_field": "model_deployment_name",
},
"Anthropic.Chat": {
"package": "agent_framework.anthropic",
"name": "AnthropicChatClient",
"model_id_field": "model_id",
},
}
class DeclarativeLoaderError(AgentFrameworkException):
"""Exception raised for errors in the declarative loader."""
pass
class ProviderLookupError(DeclarativeLoaderError):
"""Exception raised for errors in provider type lookup."""
pass
class AgentFactory:
def __init__(
self,
*,
chat_client: ChatClientProtocol | None = None,
bindings: Mapping[str, Any] | None = None,
connections: Mapping[str, Any] | None = None,
client_kwargs: Mapping[str, Any] | None = None,
additional_mappings: Mapping[str, ProviderTypeMapping] | None = None,
default_provider: str = "AzureAIClient",
env_file: str | None = None,
) -> None:
"""Create the agent factory, with bindings.
Args:
chat_client: An optional ChatClientProtocol instance to use as a dependency,
this will be passed to the ChatAgent that get's created.
If you need to create multiple agents with different chat clients,
do not pass this and instead provide the chat client in the YAML definition.
bindings: An optional dictionary of bindings to use when creating agents.
connections: An optional dictionary of connections to resolve ReferenceConnections.
client_kwargs: An optional dictionary of keyword arguments to pass to chat client constructor.
additional_mappings: An optional dictionary to extend the provider type to object mapping.
Should have the structure:
..code-block:: python
additional_mappings = {
"Provider.ApiType": {
"package": "package.name",
"name": "ClassName",
"model_id_field": "field_name_in_constructor",
},
...
}
Here, "Provider.ApiType" is the lookup key used when both provider and apiType are specified in the
model, "Provider" is also allowed.
Package refers to which model needs to be imported, Name is the class name of the ChatClientProtocol
implementation, and model_id_field is the name of the field in the constructor
that accepts the model.id value.
default_provider: The default provider used when model.provider is not specified,
default is "AzureAIClient".
env_file: An optional path to a .env file to load environment variables from.
"""
self.chat_client = chat_client
self.bindings = bindings
self.connections = connections
self.client_kwargs = client_kwargs or {}
self.additional_mappings = additional_mappings or {}
self.default_provider: str = default_provider
load_dotenv(dotenv_path=env_file)
def create_agent_from_yaml_path(self, yaml_path: str | Path) -> ChatAgent:
"""Create a ChatAgent from a YAML file path.
This method does the following things:
1. Loads the YAML file into a AgentSchema object using open and agent_schema_dispatch.
2. Validates that the loaded object is a PromptAgent.
3. Creates the appropriate ChatClient based on the model provider and apiType.
4. Parses the tools, options, and response format from the PromptAgent.
5. Creates and returns a ChatAgent instance with the configured properties.
Args:
yaml_path: Path to the YAML file representation of a AgentSchema object
Returns:
The ``ChatAgent`` instance created from the YAML file.
Raises:
DeclarativeLoaderError: If the YAML does not represent a PromptAgent.
ProviderLookupError: If the provider type is unknown or unsupported.
ValueError: If a ReferenceConnection cannot be resolved.
ModuleNotFoundError: If the required module for the provider type cannot be imported.
AttributeError: If the required class for the provider type cannot be found in the module.
"""
if not isinstance(yaml_path, Path):
yaml_path = Path(yaml_path)
if not yaml_path.exists():
raise DeclarativeLoaderError(f"YAML file not found at path: {yaml_path}")
with open(yaml_path) as f:
yaml_str = f.read()
return self.create_agent_from_yaml(yaml_str)
def create_agent_from_yaml(self, yaml_str: str) -> ChatAgent:
"""Create a ChatAgent from a YAML string.
This method does the following things:
1. Loads the YAML string into a AgentSchema object using agent_schema_dispatch.
2. Validates that the loaded object is a PromptAgent.
3. Creates the appropriate ChatClient based on the model provider and apiType.
4. Parses the tools, options, and response format from the PromptAgent.
5. Creates and returns a ChatAgent instance with the configured properties.
Args:
yaml_str: YAML string representation of a AgentSchema object
Returns:
The ``ChatAgent`` instance created from the YAML string.
Raises:
DeclarativeLoaderError: If the YAML does not represent a PromptAgent.
ProviderLookupError: If the provider type is unknown or unsupported.
ValueError: If a ReferenceConnection cannot be resolved.
ModuleNotFoundError: If the required module for the provider type cannot be imported.
AttributeError: If the required class for the provider type cannot be found in the module.
"""
prompt_agent = agent_schema_dispatch(yaml.safe_load(yaml_str))
if not isinstance(prompt_agent, PromptAgent):
raise DeclarativeLoaderError("Only yaml definitions for a PromptAgent are supported for agent creation.")
# Step 1: Create the ChatClient
client = self._get_client(prompt_agent)
# Step 2: Get the chat options
chat_options = self._parse_chat_options(prompt_agent.model)
if tools := self._parse_tools(prompt_agent.tools):
chat_options["tools"] = tools
if output_schema := prompt_agent.outputSchema:
chat_options["response_format"] = _create_model_from_json_schema("agent", output_schema.to_json_schema())
# Step 3: Create the agent instance
return ChatAgent(
chat_client=client,
name=prompt_agent.name,
description=prompt_agent.description,
instructions=prompt_agent.instructions,
**chat_options,
)
def _get_client(self, prompt_agent: PromptAgent) -> ChatClientProtocol:
"""Create the ChatClientProtocol instance based on the PromptAgent model."""
if not prompt_agent.model:
# if no model is defined, use the supplied chat_client
if self.chat_client:
return self.chat_client
raise DeclarativeLoaderError(
"ChatClient must be provided to create agent from PromptAgent, "
"alternatively define a model in the PromptAgent."
)
setup_dict: dict[str, Any] = {}
setup_dict.update(self.client_kwargs)
# parse connections
if prompt_agent.model.connection:
match prompt_agent.model.connection:
case ApiKeyConnection():
setup_dict["api_key"] = prompt_agent.model.connection.apiKey
if prompt_agent.model.connection.endpoint:
setup_dict["endpoint"] = prompt_agent.model.connection.endpoint
case RemoteConnection() | AnonymousConnection():
setup_dict["endpoint"] = prompt_agent.model.connection.endpoint
case ReferenceConnection():
if not self.connections:
raise ValueError("Connections must be provided to resolve ReferenceConnection")
# find the referenced connection
if prompt_agent.model.connection.name and (
value := self.connections.get(prompt_agent.model.connection.name)
):
setup_dict[prompt_agent.model.connection.name] = value
else:
raise ValueError(
f"ReferenceConnection with name {prompt_agent.model.connection.name} not found in provided "
"connections."
)
# Any client we create, needs a model.id
if not prompt_agent.model.id:
# if prompt_agent.model is defined, but no id, use the supplied chat_client
if self.chat_client:
return self.chat_client
# or raise, since we cannot create a client without model id
raise DeclarativeLoaderError(
"ChatClient must be provided to create agent from PromptAgent, or define model.id in the PromptAgent."
)
# if provider is defined, use that, if possible with apiType, fallback to default_provider
mapping = self._retrieve_provider_configuration(prompt_agent.model)
module_name = mapping["package"]
class_name = mapping["name"]
module = __import__(module_name, fromlist=[class_name])
agent_class = getattr(module, class_name)
setup_dict[mapping["model_id_field"]] = prompt_agent.model.id
return agent_class(**setup_dict) # type: ignore[no-any-return]
def _parse_chat_options(self, model: Model | None) -> dict[str, Any]:
"""Parse ModelOptions into chat options dictionary."""
chat_options: dict[str, Any] = {}
if not model or not model.options or not isinstance(model.options, ModelOptions):
return chat_options
options = model.options
if options.frequencyPenalty is not None:
chat_options["frequency_penalty"] = options.frequencyPenalty
if options.presencePenalty is not None:
chat_options["presence_penalty"] = options.presencePenalty
if options.maxOutputTokens is not None:
chat_options["max_tokens"] = options.maxOutputTokens
if options.temperature is not None:
chat_options["temperature"] = options.temperature
if options.topP is not None:
chat_options["top_p"] = options.topP
if options.seed is not None:
chat_options["seed"] = options.seed
if options.stopSequences:
chat_options["stop"] = options.stopSequences
if options.allowMultipleToolCalls is not None:
chat_options["allow_multiple_tool_calls"] = options.allowMultipleToolCalls
if (chat_tool_mode := options.additionalProperties.pop("chatToolMode", None)) is not None:
chat_options["tool_choice"] = chat_tool_mode
if options.additionalProperties:
chat_options["additional_chat_options"] = options.additionalProperties
return chat_options
def _parse_tools(self, tools: list[Tool] | None) -> list[ToolProtocol] | None:
"""Parse tool resources into ToolProtocol instances."""
if not tools:
return None
return [self._parse_tool(tool_resource) for tool_resource in tools]
def _parse_tool(self, tool_resource: Tool) -> ToolProtocol:
"""Parse a single tool resource into a ToolProtocol instance."""
match tool_resource:
case FunctionTool():
func: Callable[..., Any] | None = None
if self.bindings and tool_resource.bindings:
for binding in tool_resource.bindings:
if binding.name and (func := self.bindings.get(binding.name)):
break
return AIFunction( # type: ignore
name=tool_resource.name, # type: ignore
description=tool_resource.description, # type: ignore
input_model=tool_resource.parameters.to_json_schema() if tool_resource.parameters else None,
func=func,
)
case WebSearchTool():
return HostedWebSearchTool(
description=tool_resource.description, additional_properties=tool_resource.options
)
case FileSearchTool():
add_props: dict[str, Any] = {}
if tool_resource.ranker is not None:
add_props["ranker"] = tool_resource.ranker
if tool_resource.scoreThreshold is not None:
add_props["score_threshold"] = tool_resource.scoreThreshold
if tool_resource.filters:
add_props["filters"] = tool_resource.filters
return HostedFileSearchTool(
inputs=[HostedVectorStoreContent(id) for id in tool_resource.vectorStoreIds or []],
description=tool_resource.description,
max_results=tool_resource.maximumResultCount,
additional_properties=add_props,
)
case CodeInterpreterTool():
return HostedCodeInterpreterTool(
inputs=[HostedFileContent(file_id=file) for file in tool_resource.fileIds or []],
description=tool_resource.description,
)
case McpTool():
approval_mode: HostedMCPSpecificApproval | Literal["always_require", "never_require"] | None = None
if tool_resource.approvalMode is not None:
if tool_resource.approvalMode.kind == "always":
approval_mode = "always_require"
elif tool_resource.approvalMode.kind == "never":
approval_mode = "never_require"
elif isinstance(tool_resource.approvalMode, McpServerToolSpecifyApprovalMode):
approval_mode = {}
if tool_resource.approvalMode.alwaysRequireApprovalTools:
approval_mode["always_require_approval"] = (
tool_resource.approvalMode.alwaysRequireApprovalTools
)
if tool_resource.approvalMode.neverRequireApprovalTools:
approval_mode["never_require_approval"] = (
tool_resource.approvalMode.neverRequireApprovalTools
)
if not approval_mode:
approval_mode = None
return HostedMCPTool(
name=tool_resource.name, # type: ignore
description=tool_resource.description,
url=tool_resource.url, # type: ignore
allowed_tools=tool_resource.allowedTools,
approval_mode=approval_mode,
)
case _:
raise ValueError(f"Unsupported tool kind: {tool_resource.kind}")
def _retrieve_provider_configuration(self, model: Model) -> ProviderTypeMapping:
"""Retrieve the provider configuration based on the model's provider and apiType.
If only provider is specified, it will be used.
If both provider and apiType are specified, both will be used.
If neither is specified, the default_provider will be used.
Args:
model: The Model instance containing provider and apiType information.
Returns:
A dictionary containing the package, name, and model_id_field for the provider.
Raises:
ProviderLookupError: If the provider type is not supported or can't be found.
"""
class_lookup = (
f"{model.provider}.{model.apiType}"
if model.apiType
else f"{model.provider}"
if model.provider
else self.default_provider
)
if class_lookup in self.additional_mappings:
return self.additional_mappings[class_lookup]
if class_lookup not in PROVIDER_TYPE_OBJECT_MAPPING:
raise ProviderLookupError(f"Unsupported provider type: {class_lookup}")
return PROVIDER_TYPE_OBJECT_MAPPING[class_lookup]
File diff suppressed because it is too large Load Diff
@@ -1,97 +0,0 @@
[project]
name = "agent-framework-declarative"
description = "Declarative specification support for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "1.0.0b251120"
license-files = ["LICENSE"]
urls.homepage = "https://aka.ms/agent-framework"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
urls.release_notes = "https://github.com/microsoft/agent-framework/releases?q=tag%3Apython-1&expanded=true"
urls.issues = "https://github.com/microsoft/agent-framework/issues"
classifiers = [
"License :: OSI Approved :: MIT License",
"Development Status :: 4 - Beta",
"Intended Audience :: Developers",
"Programming Language :: Python :: 3",
"Programming Language :: Python :: 3.10",
"Programming Language :: Python :: 3.11",
"Programming Language :: Python :: 3.12",
"Programming Language :: Python :: 3.13",
"Typing :: Typed",
]
dependencies = [
"agent-framework-core",
"powerfx>=0.0.31; python_version < '3.14'",
"pyyaml>=6.0,<7.0",
]
[dependency-groups]
dev = [
"types-PyYaml"
]
[tool.uv]
prerelease = "if-necessary-or-explicit"
environments = [
"sys_platform == 'darwin'",
"sys_platform == 'linux'",
"sys_platform == 'win32'"
]
[tool.uv-dynamic-versioning]
fallback-version = "0.0.0"
[tool.pytest.ini_options]
testpaths = 'tests'
addopts = "-ra -q -r fEX"
asyncio_mode = "auto"
asyncio_default_fixture_loop_scope = "function"
filterwarnings = [
"ignore:Support for class-based `config` is deprecated:DeprecationWarning:pydantic.*"
]
timeout = 120
[tool.ruff]
extend = "../../pyproject.toml"
[tool.coverage.run]
omit = [
"**/__init__.py"
]
[tool.pyright]
extends = "../../pyproject.toml"
exclude = ['tests']
[tool.mypy]
plugins = ['pydantic.mypy']
strict = true
python_version = "3.10"
ignore_missing_imports = true
disallow_untyped_defs = true
no_implicit_optional = true
check_untyped_defs = true
warn_return_any = true
show_error_codes = true
warn_unused_ignores = false
disallow_incomplete_defs = true
disallow_untyped_decorators = true
exclude = [
'_models.py$',
]
[tool.bandit]
targets = ["agent_framework_declarative"]
exclude_dirs = ["tests"]
[tool.poe]
executor.type = "uv"
include = "../../shared_tasks.toml"
[tool.poe.tasks]
mypy = "mypy --config-file $POE_ROOT/pyproject.toml agent_framework_declarative"
test = "pytest --cov=agent_framework_declarative --cov-report=term-missing:skip-covered tests"
[build-system]
requires = ["flit-core >= 3.11,<4.0"]
build-backend = "flit_core.buildapi"
@@ -1,456 +0,0 @@
# Copyright (c) Microsoft. All rights reserved.
import sys
from pathlib import Path
from typing import Any
import pytest
import yaml
from agent_framework_declarative._models import (
AgentDefinition,
AgentManifest,
AnonymousConnection,
ApiKeyConnection,
ArrayProperty,
CodeInterpreterTool,
Connection,
CustomTool,
FileSearchTool,
FunctionTool,
McpServerApprovalMode,
McpServerToolAlwaysRequireApprovalMode,
McpServerToolNeverRequireApprovalMode,
McpServerToolSpecifyApprovalMode,
McpTool,
ModelResource,
ObjectProperty,
OpenApiTool,
PromptAgent,
Property,
PropertySchema,
ReferenceConnection,
RemoteConnection,
Resource,
ToolResource,
WebSearchTool,
agent_schema_dispatch,
)
pytestmark = pytest.mark.skipif(sys.version_info >= (3, 14), reason="Skipping on Python 3.14+")
@pytest.mark.parametrize(
"yaml_content,expected_type,expected_attributes",
[
# Agent Manifest (no kind field)
(
"""
name: my-manifest
description: A test manifest
""",
AgentManifest,
{"name": "my-manifest", "description": "A test manifest"},
),
# PromptAgent
(
"""
kind: Prompt
name: assistant
description: A helpful assistant
model:
id: gpt-4
""",
PromptAgent,
{"name": "assistant", "description": "A helpful assistant"},
),
# AgentDefinition
(
"""
kind: Agent
name: base-agent
description: A base agent
""",
AgentDefinition,
{"name": "base-agent", "description": "A base agent"},
),
# ModelResource
(
"""
kind: Model
name: my-model
id: gpt-4
""",
ModelResource,
{"name": "my-model", "id": "gpt-4"},
),
# ToolResource
(
"""
kind: Tool
name: my-tool
id: search-tool
""",
ToolResource,
{"name": "my-tool", "id": "search-tool"},
),
# Resource (base)
(
"""
kind: Resource
name: generic-resource
""",
Resource,
{"name": "generic-resource"},
),
# FunctionTool
(
"""
kind: function
name: get_weather
description: Get the weather
""",
FunctionTool,
{"name": "get_weather", "description": "Get the weather"},
),
# CustomTool
(
"""
kind: custom
name: custom_tool
description: A custom tool
""",
CustomTool,
{"name": "custom_tool", "description": "A custom tool"},
),
# WebSearchTool
(
"""
kind: web_search
name: search
description: Search the web
""",
WebSearchTool,
{"name": "search", "description": "Search the web"},
),
# FileSearchTool
(
"""
kind: file_search
name: file_search
description: Search files
""",
FileSearchTool,
{"name": "file_search", "description": "Search files"},
),
# McpTool
(
"""
kind: mcp
name: mcp_tool
description: An MCP tool
serverName: my-server
""",
McpTool,
{"name": "mcp_tool", "serverName": "my-server"},
),
# OpenApiTool
(
"""
kind: openapi
name: api_tool
description: An OpenAPI tool
specification: https://api.example.com/openapi.json
""",
OpenApiTool,
{"name": "api_tool", "specification": "https://api.example.com/openapi.json"},
),
# CodeInterpreterTool
(
"""
kind: code_interpreter
name: code_tool
description: A code interpreter tool
""",
CodeInterpreterTool,
{"name": "code_tool", "description": "A code interpreter tool"},
),
# ReferenceConnection
(
"""
kind: reference
name: my-connection
target: target-connection
""",
ReferenceConnection,
{"name": "my-connection", "target": "target-connection"},
),
# RemoteConnection
(
"""
kind: remote
endpoint: https://api.example.com
""",
RemoteConnection,
{"endpoint": "https://api.example.com"},
),
# ApiKeyConnection
(
"""
kind: key
apiKey: secret-key
endpoint: https://api.example.com
""",
ApiKeyConnection,
{"apiKey": "secret-key", "endpoint": "https://api.example.com"},
),
# AnonymousConnection
(
"""
kind: anonymous
endpoint: https://api.example.com
""",
AnonymousConnection,
{"endpoint": "https://api.example.com"},
),
# Connection (base)
(
"""
kind: connection
authenticationMode: oauth
""",
Connection,
{"authenticationMode": "oauth"},
),
# ArrayProperty
(
"""
kind: array
name: items
description: An array of items
""",
ArrayProperty,
{"name": "items", "description": "An array of items"},
),
# ObjectProperty
(
"""
kind: object
name: config
description: Configuration object
""",
ObjectProperty,
{"name": "config", "description": "Configuration object"},
),
# Property (base)
(
"""
kind: property
name: field
description: A property field
""",
Property,
{"name": "field", "description": "A property field"},
),
# McpServerToolAlwaysRequireApprovalMode
(
"""
kind: always
""",
McpServerToolAlwaysRequireApprovalMode,
{},
),
# McpServerToolNeverRequireApprovalMode
(
"""
kind: never
""",
McpServerToolNeverRequireApprovalMode,
{},
),
# McpServerToolSpecifyApprovalMode
(
"""
kind: specify
alwaysRequireApprovalTools: []
neverRequireApprovalTools: []
""",
McpServerToolSpecifyApprovalMode,
{},
),
# McpServerApprovalMode (base)
(
"""
kind: approval_mode
""",
McpServerApprovalMode,
{},
),
],
)
def test_agent_schema_dispatch_all_types(yaml_content: str, expected_type: type, expected_attributes: dict[str, Any]):
"""Test that agent_schema_dispatch correctly loads all MAML object types."""
result = agent_schema_dispatch(yaml.safe_load(yaml_content))
# Check the type is correct
assert isinstance(result, expected_type), f"Expected {expected_type.__name__}, got {type(result).__name__}"
# Check expected attributes
for attr_name, attr_value in expected_attributes.items():
assert hasattr(result, attr_name), f"Result missing attribute '{attr_name}'"
assert getattr(result, attr_name) == attr_value, (
f"Attribute '{attr_name}' has value {getattr(result, attr_name)}, expected {attr_value}"
)
def test_agent_schema_dispatch_unknown_kind():
"""Test that agent_schema_dispatch returns None for unknown kind."""
yaml_content = """
kind: unknown_type
name: test
"""
result = agent_schema_dispatch(yaml.safe_load(yaml_content))
assert result is None
def test_agent_schema_dispatch_complex_agent_manifest():
"""Test loading a complex agent manifest with nested objects."""
yaml_content = """
name: complex-manifest
description: A complete manifest
template:
kind: Prompt
name: assistant
description: A helpful assistant
model:
id: gpt-4
provider: openai
tools:
- kind: web_search
name: search
description: Search the web
- kind: function
name: calculator
description: Calculate math
resources:
- kind: model
name: model1
id: gpt-4
- kind: tool
name: tool1
id: search
"""
result = agent_schema_dispatch(yaml.safe_load(yaml_content))
assert isinstance(result, AgentManifest)
assert result.name == "complex-manifest"
assert result.description == "A complete manifest"
assert isinstance(result.template, PromptAgent)
assert result.template.name == "assistant"
assert len(result.resources) == 2
assert isinstance(result.resources[0], ModelResource)
assert isinstance(result.resources[1], ToolResource)
def test_agent_schema_dispatch_prompt_agent_with_tools():
"""Test loading a prompt agent with multiple tools."""
yaml_content = """
kind: Prompt
name: multi-tool-agent
description: Agent with multiple tools
model:
id: gpt-4
tools:
- kind: web_search
name: search
description: Search the web
- kind: function
name: get_weather
description: Get weather information
- kind: code_interpreter
name: code
description: Execute code
"""
result = agent_schema_dispatch(yaml.safe_load(yaml_content))
assert isinstance(result, PromptAgent)
assert result.name == "multi-tool-agent"
assert len(result.tools) == 3
# Tools are polymorphically created based on their kind
assert result.tools[0].kind == "web_search"
assert result.tools[1].kind == "function"
assert result.tools[2].kind == "code_interpreter"
def test_agent_schema_dispatch_model_resource():
"""Test loading a model resource."""
yaml_content = """
kind: Model
name: my-model
id: gpt-4
"""
result = agent_schema_dispatch(yaml.safe_load(yaml_content))
assert isinstance(result, ModelResource)
assert result.id == "gpt-4"
def test_agent_schema_dispatch_property_schema_with_nested_properties():
"""Test loading a property schema with nested properties."""
yaml_content = """
kind: property_schema
strict: true
properties:
- kind: property
name: name
description: User name
- kind: object
name: address
description: User address
properties:
- kind: property
name: street
description: Street address
- kind: property
name: city
description: City name
- kind: array
name: tags
description: User tags
"""
result = agent_schema_dispatch(yaml.safe_load(yaml_content))
assert isinstance(result, PropertySchema)
assert result.strict is True
assert len(result.properties) == 3
# Properties are polymorphically created based on their kind
assert result.properties[0].kind == "property"
assert result.properties[1].kind == "object"
assert result.properties[2].kind == "array"
def _get_agent_sample_yaml_files() -> list[tuple[Path, Path]]:
"""Helper function to collect all YAML files from agent-samples directory."""
current_file = Path(__file__)
repo_root = current_file.parent.parent.parent.parent # tests -> declarative -> packages -> python
agent_samples_dir = repo_root.parent / "agent-samples"
if not agent_samples_dir.exists():
return []
yaml_files = list(agent_samples_dir.rglob("*.yaml")) + list(agent_samples_dir.rglob("*.yml"))
return [(yaml_file, agent_samples_dir) for yaml_file in yaml_files]
@pytest.mark.parametrize(
"yaml_file,agent_samples_dir",
_get_agent_sample_yaml_files(),
ids=lambda x: x[0].name if isinstance(x, tuple) else str(x),
)
def test_agent_schema_dispatch_agent_samples(yaml_file: Path, agent_samples_dir: Path):
"""Test that agent_schema_dispatch successfully loads a YAML file from agent-samples directory."""
with open(yaml_file) as f:
content = f.read()
result = agent_schema_dispatch(yaml.safe_load(content))
# Result can be None for unknown kinds, but should not raise exceptions
assert result is not None, f"agent_schema_dispatch returned None for {yaml_file.relative_to(agent_samples_dir)}"
File diff suppressed because it is too large Load Diff

Some files were not shown because too many files have changed in this diff Show More