Compare commits

...
Author SHA1 Message Date
Evan MattsonandGitHub b0a7a1fcb8 Python: Fix WorkflowAgent event handling and kwargs forwarding (#2946)
* Fix kwargs propagation through workflow.as_agent()

* Fix WorkflowAgent to respect AgentExecutor output_response setting
2025-12-18 19:35:07 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>Chris
a841bdd1cc Bump Azure.AI.AgentServer.AgentFramework from 1.0.0-beta.4 to 1.0.0-beta.5 (#2854)
---
updated-dependencies:
- dependency-name: Azure.AI.AgentServer.AgentFramework
  dependency-version: 1.0.0-beta.5
  dependency-type: direct:production
  update-type: version-update:semver-patch
- dependency-name: Azure.AI.AgentServer.AgentFramework
  dependency-version: 1.0.0-beta.5
  dependency-type: direct:production
  update-type: version-update:semver-patch
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
Co-authored-by: Chris <66376200+crickman@users.noreply.github.com>
2025-12-18 18:36:13 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>Mark Wallace
d46adffe6c Bump AWSSDK.Extensions.Bedrock.MEAI from 4.0.4.11 to 4.0.5 (#2853)
---
updated-dependencies:
- dependency-name: AWSSDK.Extensions.Bedrock.MEAI
  dependency-version: 4.0.5
  dependency-type: direct:production
  update-type: version-update:semver-patch
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
Co-authored-by: Mark Wallace <127216156+markwallace-microsoft@users.noreply.github.com>
2025-12-18 17:25:54 +00:00
dependabot[bot]GitHubdependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
b0b5777363 Bump CommunityToolkit.Aspire.OllamaSharp from 13.0.0-beta.440 to 13.0.0 (#2856)
---
updated-dependencies:
- dependency-name: CommunityToolkit.Aspire.OllamaSharp
  dependency-version: 13.0.0
  dependency-type: direct:production
  update-type: version-update:semver-patch
...

Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2025-12-18 17:25:36 +00:00
Giles OdigweandGitHub 37b4cfd024 Python: Add Azure Managed Redis Support with Credential Provider (#2887)
* azure redis support

* small fixes

* azure managed redis sample

* fixes
2025-12-18 17:10:55 +00:00
CopilotGitHubstephentoubcopilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
ff9343d7cc .NET: Update Anthropic package to version 12.0.0 (#2914)
* Initial plan

* Update Anthropic package to version 12.0.0

Co-authored-by: stephentoub <2642209+stephentoub@users.noreply.github.com>

---------

Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
Co-authored-by: stephentoub <2642209+stephentoub@users.noreply.github.com>
2025-12-18 16:02:20 +00:00
Victor DibiaandGitHub 8ff34f9a43 Python: Add workflow cancellation sample (#2732)
* Add workflow cancellation sample

Add sample demonstrating how to cancel a running workflow using asyncio
tasks. Shows both cancellation mid-execution and normal completion paths.
Useful for implementing timeouts, graceful shutdown, or A2A executors.

* update docstring
2025-12-18 14:12:42 +00:00
Hao LuoandGitHub e3f8bfc645 Python: Fixes Run ID and Thread ID casing to align with AG-UI Typescript SDK (#2948)
* added camelCase input to run id and thread id aligning with @ag-ui/core

* fixed per copilot suggestions
2025-12-18 14:10:16 +00:00
Tao ChenandGitHub b4f2709b6d Python: Workflow add option to visualize internal executors (#2917)
* Workflow add option to visualize internal executors

* Address Copilot comments
2025-12-18 14:04:03 +00:00
Eduard van ValkenburgandGitHub e5c11d38d6 Python: cleanup and refactoring of chat clients (#2937)
* refactoring and unifying naming schemes of internal methods of chat clients

* set tool_choice to auto

* fix for mypy

* added note on naming and fix #2951

* fix responses

* fixes in azure ai agents client
2025-12-18 12:02:23 +00:00
a71f768331 .NET: [Breaking] Delete display name property (#2758)
* delete the AIAgent.DisplayName property

* use agent name as a first value for activity display name

* Update dotnet/src/Microsoft.Agents.AI.Workflows/Specialized/HandoffAgentExecutor.cs

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>

---------

Co-authored-by: Copilot <175728472+Copilot@users.noreply.github.com>
2025-12-18 09:22:45 +00:00
0298e0a401 Python: fix: correct BadRequestError when using Pydantic model in response_fo… (#1843)
* fix: correct BadRequestError when using Pydantic model in response_format

* Fix lint

---------

Co-authored-by: Evan Mattson <evan.mattson@microsoft.com>
2025-12-18 08:42:00 +00:00
Evan MattsonandGitHub ca1532cf22 Python: Move ollama samples to samples getting started dir (#2921)
* Move ollama samples to samples getting started dir

* Address feedback
2025-12-18 08:37:05 +00:00
Evan MattsonandGitHub 360839782c Pass kwargs into subworkflows (#2923) 2025-12-18 04:34:33 +00:00
Ege Ozan Ă–zyedekandGitHub ee53fe4666 Python: Correction of MCP image type conversion in _mcp.py (#2901)
* Correction of MCP image type conversion in  _mcp.py

* Added a new overload to the init function of the DataContent() type of the Agent Framework, edited the test case to correctly test the usage of the data and uri fields while using DataContent()

* Fixed tests related to the changes of the DataContent type, added testing for both string and byte representations
2025-12-17 16:11:39 +00:00
Dmytro StrukandGitHub 3cd805f0bf Added additional arguments for Azure AI agent (#2922) 2025-12-17 08:08:01 +00:00
CopilotGitHubSergeyMenshykhcopilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
c7ddb8aa14 .NET: Make DelegatingAIAgent abstract (#2797)
* Initial plan

* Make DelegatingAIAgent abstract

Co-authored-by: SergeyMenshykh <68852919+SergeyMenshykh@users.noreply.github.com>

---------

Co-authored-by: copilot-swe-agent[bot] <198982749+Copilot@users.noreply.github.com>
Co-authored-by: SergeyMenshykh <68852919+SergeyMenshykh@users.noreply.github.com>
2025-12-17 07:31:10 +00:00
Giles OdigweandGitHub d5527982b6 Python: Azure AI Agent with Bing Grounding Citations Sample (#2892)
* bing grounding sample with citations

* small fix

* fix
2025-12-17 00:43:38 +00:00
Dmytro StrukandGitHub ec1c5e9c11 Updated Ollama package version (#2920) 2025-12-17 00:42:27 +00:00
Evan MattsonandGitHub 06cdcb93f0 Fix Pydantic error when using Literal type for tool params (#2893) 2025-12-17 00:27:01 +00:00
81 changed files with 3035 additions and 1341 deletions
+3 -3
View File
@@ -11,13 +11,13 @@
</PropertyGroup>
<ItemGroup>
<!-- Aspire.* -->
<PackageVersion Include="Anthropic" Version="11.0.0" />
<PackageVersion Include="Anthropic" Version="12.0.0" />
<PackageVersion Include="Anthropic.Foundry" Version="0.1.0" />
<PackageVersion Include="Aspire.Azure.AI.OpenAI" Version="13.0.0-preview.1.25560.3" />
<PackageVersion Include="Aspire.Hosting.AppHost" Version="$(AspireAppHostSdkVersion)" />
<PackageVersion Include="Aspire.Hosting.Azure.CognitiveServices" Version="$(AspireAppHostSdkVersion)" />
<PackageVersion Include="Aspire.Microsoft.Azure.Cosmos" Version="$(AspireAppHostSdkVersion)" />
<PackageVersion Include="CommunityToolkit.Aspire.OllamaSharp" Version="13.0.0-beta.440" />
<PackageVersion Include="CommunityToolkit.Aspire.OllamaSharp" Version="13.0.0" />
<!-- Azure.* -->
<PackageVersion Include="Azure.AI.Projects" Version="1.2.0-beta.5" />
<PackageVersion Include="Azure.AI.Projects.OpenAI" Version="1.0.0-beta.5" />
@@ -100,7 +100,7 @@
<!-- MCP -->
<PackageVersion Include="ModelContextProtocol" Version="0.4.0-preview.3" />
<!-- Inference SDKs -->
<PackageVersion Include="AWSSDK.Extensions.Bedrock.MEAI" Version="4.0.4.11" />
<PackageVersion Include="AWSSDK.Extensions.Bedrock.MEAI" Version="4.0.5" />
<PackageVersion Include="Microsoft.ML.OnnxRuntimeGenAI" Version="0.10.0" />
<PackageVersion Include="OllamaSharp" Version="5.4.8" />
<PackageVersion Include="OpenAI" Version="2.8.0" />
@@ -45,7 +45,7 @@ namespace SampleApp
}
// Clone the input messages and turn them into response messages with upper case text.
List<ChatMessage> responseMessages = CloneAndToUpperCase(messages, this.DisplayName).ToList();
List<ChatMessage> responseMessages = CloneAndToUpperCase(messages, this.Name).ToList();
// Notify the thread of the input and output messages.
await typedThread.MessageStore.AddMessagesAsync(messages.Concat(responseMessages), cancellationToken);
@@ -69,7 +69,7 @@ namespace SampleApp
}
// Clone the input messages and turn them into response messages with upper case text.
List<ChatMessage> responseMessages = CloneAndToUpperCase(messages, this.DisplayName).ToList();
List<ChatMessage> responseMessages = CloneAndToUpperCase(messages, this.Name).ToList();
// Notify the thread of the input and output messages.
await typedThread.MessageStore.AddMessagesAsync(messages.Concat(responseMessages), cancellationToken);
@@ -79,7 +79,7 @@ namespace SampleApp
yield return new AgentRunResponseUpdate
{
AgentId = this.Id,
AuthorName = this.DisplayName,
AuthorName = message.AuthorName,
Role = ChatRole.Assistant,
Contents = message.Contents,
ResponseId = Guid.NewGuid().ToString("N"),
@@ -88,7 +88,7 @@ namespace SampleApp
}
}
private static IEnumerable<ChatMessage> CloneAndToUpperCase(IEnumerable<ChatMessage> messages, string agentName) => messages.Select(x =>
private static IEnumerable<ChatMessage> CloneAndToUpperCase(IEnumerable<ChatMessage> messages, string? agentName) => messages.Select(x =>
{
// Clone the message and update its author to be the agent.
var messageClone = x.Clone();
@@ -35,7 +35,7 @@
</ItemGroup>
<ItemGroup>
<PackageReference Include="Azure.AI.AgentServer.AgentFramework" Version="1.0.0-beta.4" />
<PackageReference Include="Azure.AI.AgentServer.AgentFramework" Version="1.0.0-beta.5" />
<PackageReference Include="Azure.AI.OpenAI" Version="2.7.0-beta.2" />
<PackageReference Include="Azure.Identity" Version="1.17.1" />
<PackageReference Include="Microsoft.Agents.AI.OpenAI" Version="1.0.0-preview.251125.1" />
@@ -35,7 +35,7 @@
</ItemGroup>
<ItemGroup>
<PackageReference Include="Azure.AI.AgentServer.AgentFramework" Version="1.0.0-beta.4" />
<PackageReference Include="Azure.AI.AgentServer.AgentFramework" Version="1.0.0-beta.5" />
<PackageReference Include="Azure.AI.OpenAI" Version="2.7.0-beta.2" />
<PackageReference Include="Azure.Identity" Version="1.17.1" />
<PackageReference Include="Microsoft.Agents.AI.Workflows" Version="1.0.0-preview.251125.1" />
@@ -30,7 +30,6 @@ internal sealed class A2AAgent : AIAgent
private readonly string? _id;
private readonly string? _name;
private readonly string? _description;
private readonly string? _displayName;
private readonly ILogger _logger;
/// <summary>
@@ -40,9 +39,8 @@ internal sealed class A2AAgent : AIAgent
/// <param name="id">The unique identifier for the agent.</param>
/// <param name="name">The the name of the agent.</param>
/// <param name="description">The description of the agent.</param>
/// <param name="displayName">The display name of the agent.</param>
/// <param name="loggerFactory">Optional logger factory to use for logging.</param>
public A2AAgent(A2AClient a2aClient, string? id = null, string? name = null, string? description = null, string? displayName = null, ILoggerFactory? loggerFactory = null)
public A2AAgent(A2AClient a2aClient, string? id = null, string? name = null, string? description = null, ILoggerFactory? loggerFactory = null)
{
_ = Throw.IfNull(a2aClient);
@@ -50,7 +48,6 @@ internal sealed class A2AAgent : AIAgent
this._id = id;
this._name = name;
this._description = description;
this._displayName = displayName;
this._logger = (loggerFactory ?? NullLoggerFactory.Instance).CreateLogger<A2AAgent>();
}
@@ -203,9 +200,6 @@ internal sealed class A2AAgent : AIAgent
/// <inheritdoc/>
public override string? Name => this._name ?? base.Name;
/// <inheritdoc/>
public override string DisplayName => this._displayName ?? base.DisplayName;
/// <inheritdoc/>
public override string? Description => this._description ?? base.Description;
@@ -33,9 +33,8 @@ public static class A2AClientExtensions
/// <param name="id">The unique identifier for the agent.</param>
/// <param name="name">The the name of the agent.</param>
/// <param name="description">The description of the agent.</param>
/// <param name="displayName">The display name of the agent.</param>
/// <param name="loggerFactory">Optional logger factory for enabling logging within the agent.</param>
/// <returns>An <see cref="AIAgent"/> instance backed by the A2A agent.</returns>
public static AIAgent GetAIAgent(this A2AClient client, string? id = null, string? name = null, string? description = null, string? displayName = null, ILoggerFactory? loggerFactory = null) =>
new A2AAgent(client, id, name, description, displayName, loggerFactory);
public static AIAgent GetAIAgent(this A2AClient client, string? id = null, string? name = null, string? description = null, ILoggerFactory? loggerFactory = null) =>
new A2AAgent(client, id, name, description, loggerFactory);
}
@@ -60,18 +60,6 @@ public abstract class AIAgent
/// </remarks>
public virtual string? Name { get; }
/// <summary>
/// Gets a display-friendly name for the agent.
/// </summary>
/// <value>
/// The agent's <see cref="Name"/> if available, otherwise the <see cref="Id"/>.
/// </value>
/// <remarks>
/// This property provides a guaranteed non-null string suitable for display in user interfaces,
/// logs, or other contexts where a readable identifier is needed.
/// </remarks>
public virtual string DisplayName => this.Name ?? this.Id;
/// <summary>
/// Gets a description of the agent's purpose, capabilities, or behavior.
/// </summary>
@@ -25,7 +25,7 @@ namespace Microsoft.Agents.AI;
/// Derived classes can override specific methods to add custom behavior while maintaining compatibility with the agent interface.
/// </para>
/// </remarks>
public class DelegatingAIAgent : AIAgent
public abstract class DelegatingAIAgent : AIAgent
{
/// <summary>
/// Initializes a new instance of the <see cref="DelegatingAIAgent"/> class with the specified inner agent.
@@ -231,7 +231,7 @@ internal static class EntitiesApiExtensions
return new EntityInfo(
Id: entityId,
Type: "agent",
Name: agent.DisplayName,
Name: agent.Name ?? agent.Id,
Description: agent.Description,
Framework: "agent_framework",
Tools: tools,
@@ -61,7 +61,7 @@ public static partial class MicrosoftAgentAIHostingOpenAIEndpointRouteBuilderExt
path ??= $"/{agent.Name}/v1/chat/completions";
var group = endpoints.MapGroup(path);
var endpointAgentName = agent.DisplayName;
var endpointAgentName = agent.Name ?? agent.Id;
group.MapPost("/", async ([FromBody] CreateChatCompletion request, CancellationToken cancellationToken)
=> await AIAgentChatCompletionsProcessor.CreateChatCompletionAsync(agent, request, cancellationToken).ConfigureAwait(false))
@@ -76,7 +76,7 @@ public static partial class MicrosoftAgentAIHostingOpenAIEndpointRouteBuilderExt
var handlers = new ResponsesHttpHandler(responsesService);
var group = endpoints.MapGroup(responsesPath);
var endpointAgentName = agent.DisplayName;
var endpointAgentName = agent.Name ?? agent.Id;
// Create response endpoint
group.MapPost("/", handlers.CreateResponseAsync)
@@ -125,14 +125,14 @@ public sealed class HandoffsWorkflowBuilder
{
Throw.ArgumentException(
nameof(to),
$"The provided target agent '{to.DisplayName}' has no description, name, or instructions, and no handoff description has been provided. " +
$"The provided target agent '{to.Name ?? to.Id}' has no description, name, or instructions, and no handoff description has been provided. " +
"At least one of these is required to register a handoff so that the appropriate target agent can be chosen.");
}
}
if (!handoffs.Add(new(to, handoffReason)))
{
Throw.InvalidOperationException($"A handoff from agent '{from.DisplayName}' to agent '{to.DisplayName}' has already been registered.");
Throw.InvalidOperationException($"A handoff from agent '{from.Name ?? from.Id}' to agent '{to.Name ?? to.Id}' has already been registered.");
}
return this;
@@ -20,7 +20,7 @@ internal sealed class AgentRunStreamingExecutor(AIAgent agent, bool includeInput
protected override async ValueTask TakeTurnAsync(List<ChatMessage> messages, IWorkflowContext context, bool? emitEvents, CancellationToken cancellationToken = default)
{
List<ChatMessage>? roleChanged = messages.ChangeAssistantToUserForOtherParticipants(agent.DisplayName);
List<ChatMessage>? roleChanged = messages.ChangeAssistantToUserForOtherParticipants(agent.Name ?? agent.Id);
List<AgentRunResponseUpdate> updates = [];
await foreach (var update in agent.RunStreamingAsync(messages, cancellationToken: cancellationToken).ConfigureAwait(false))
@@ -67,7 +67,7 @@ internal sealed class HandoffAgentExecutor(
List<AgentRunResponseUpdate> updates = [];
List<ChatMessage> allMessages = handoffState.Messages;
List<ChatMessage>? roleChanges = allMessages.ChangeAssistantToUserForOtherParticipants(this._agent.DisplayName);
List<ChatMessage>? roleChanges = allMessages.ChangeAssistantToUserForOtherParticipants(this._agent.Name ?? this._agent.Id);
await foreach (var update in this._agent.RunStreamingAsync(allMessages,
options: this._agentOptions,
@@ -85,7 +85,7 @@ internal sealed class HandoffAgentExecutor(
new AgentRunResponseUpdate
{
AgentId = this._agent.Id,
AuthorName = this._agent.DisplayName,
AuthorName = this._agent.Name ?? this._agent.Id,
Contents = [new FunctionResultContent(fcc.CallId, "Transferred.")],
CreatedAt = DateTimeOffset.UtcNow,
MessageId = Guid.NewGuid().ToString("N"),
@@ -114,7 +114,9 @@ public sealed class OpenTelemetryAgent : DelegatingAIAgent, IDisposable
// Override information set by OpenTelemetryChatClient to make it specific to invoke_agent.
activity.DisplayName = $"{OpenTelemetryConsts.GenAI.InvokeAgent} {this.DisplayName}";
activity.DisplayName = string.IsNullOrWhiteSpace(this.Name)
? $"{OpenTelemetryConsts.GenAI.InvokeAgent} {this.Id}"
: $"{OpenTelemetryConsts.GenAI.InvokeAgent} {this.Name}({this.Id})";
activity.SetTag(OpenTelemetryConsts.GenAI.Operation.Name, OpenTelemetryConsts.GenAI.InvokeAgent);
if (!string.IsNullOrWhiteSpace(this._providerName))
@@ -42,16 +42,14 @@ public sealed class A2AAgentTests : IDisposable
const string TestId = "test-id";
const string TestName = "test-name";
const string TestDescription = "test-description";
const string TestDisplayName = "test-display-name";
// Act
var agent = new A2AAgent(this._a2aClient, TestId, TestName, TestDescription, TestDisplayName);
var agent = new A2AAgent(this._a2aClient, TestId, TestName, TestDescription);
// Assert
Assert.Equal(TestId, agent.Id);
Assert.Equal(TestName, agent.Name);
Assert.Equal(TestDescription, agent.Description);
Assert.Equal(TestDisplayName, agent.DisplayName);
}
[Fact]
@@ -70,7 +68,6 @@ public sealed class A2AAgentTests : IDisposable
Assert.NotEmpty(agent.Id);
Assert.Null(agent.Name);
Assert.Null(agent.Description);
Assert.Equal(agent.Id, agent.DisplayName);
}
[Fact]
@@ -19,10 +19,9 @@ public sealed class A2AClientExtensionsTests
const string TestId = "test-agent-id";
const string TestName = "Test Agent";
const string TestDescription = "This is a test agent description";
const string TestDisplayName = "Test Display Name";
// Act
var agent = a2aClient.GetAIAgent(TestId, TestName, TestDescription, TestDisplayName);
var agent = a2aClient.GetAIAgent(TestId, TestName, TestDescription);
// Assert
Assert.NotNull(agent);
@@ -30,6 +29,5 @@ public sealed class A2AClientExtensionsTests
Assert.Equal(TestId, agent.Id);
Assert.Equal(TestName, agent.Name);
Assert.Equal(TestDescription, agent.Description);
Assert.Equal(TestDisplayName, agent.DisplayName);
}
}
@@ -42,7 +42,6 @@ public class LoggingAgentTests
Assert.Equal("TestAgent", agent.Name);
Assert.Equal("This is a test agent.", agent.Description);
Assert.Equal(innerAgent.Id, agent.Id);
Assert.Equal(innerAgent.DisplayName, agent.DisplayName);
}
[Fact]
@@ -45,7 +45,6 @@ public class OpenTelemetryAgentTests
Assert.Equal("TestAgent", agent.Name);
Assert.Equal("This is a test agent.", agent.Description);
Assert.Equal(innerAgent.Id, agent.Id);
Assert.Equal(innerAgent.DisplayName, agent.DisplayName);
}
[Fact]
@@ -170,7 +169,7 @@ public class OpenTelemetryAgentTests
Assert.Equal("localhost", activity.GetTagItem("server.address"));
Assert.Equal(12345, (int)activity.GetTagItem("server.port")!);
Assert.Equal("invoke_agent TestAgent", activity.DisplayName);
Assert.Equal($"invoke_agent {agent.Name}({agent.Id})", activity.DisplayName);
Assert.Equal("invoke_agent", activity.GetTagItem("gen_ai.operation.name"));
Assert.Equal("TestAgentProviderFromAIAgentMetadata", activity.GetTagItem("gen_ai.provider.name"));
Assert.Equal(innerAgent.Name, activity.GetTagItem("gen_ai.agent.name"));
@@ -431,7 +430,15 @@ public class OpenTelemetryAgentTests
Assert.Equal("localhost", activity.GetTagItem("server.address"));
Assert.Equal(12345, (int)activity.GetTagItem("server.port")!);
Assert.Equal($"invoke_agent {innerAgent.DisplayName}", activity.DisplayName);
if (string.IsNullOrWhiteSpace(innerAgent.Name))
{
Assert.Equal($"invoke_agent {innerAgent.Id}", activity.DisplayName);
}
else
{
Assert.Equal($"invoke_agent {innerAgent.Name}({innerAgent.Id})", activity.DisplayName);
}
Assert.Equal("invoke_agent", activity.GetTagItem("gen_ai.operation.name"));
Assert.Equal("TestAgentProviderFromAIAgentMetadata", activity.GetTagItem("gen_ai.provider.name"));
Assert.Equal(innerAgent.Name, activity.GetTagItem("gen_ai.agent.name"));
@@ -30,7 +30,7 @@ internal sealed class HandoffTestEchoAgent(string id, string name, string prefix
{
return [new(ChatRole.Assistant, [new FunctionCallContent(Guid.NewGuid().ToString("N"), handoff.Name)])
{
AuthorName = this.DisplayName,
AuthorName = this.Name ?? this.Id,
MessageId = Guid.NewGuid().ToString("N"),
CreatedAt = DateTime.UtcNow
}];
@@ -47,7 +47,7 @@ internal class TestEchoAgent(string? id = null, string? name = null, string? pre
select
UpdateThread(new ChatMessage(ChatRole.Assistant, $"{prefix}{message.Text}")
{
AuthorName = this.DisplayName,
AuthorName = this.Name ?? this.Id,
CreatedAt = DateTimeOffset.Now,
MessageId = Guid.NewGuid().ToString("N")
}, thread as InMemoryAgentThread);
+8
View File
@@ -154,6 +154,14 @@ Example:
chat_completion = OpenAIChatClient(env_file_path="openai.env")
```
# Method naming inside connectors
When naming methods inside connectors, we have a loose preference for using the following conventions:
- Use `_prepare_<object>_for_<purpose>` as a prefix for methods that prepare data for sending to the external service.
- Use `_parse_<object>_from_<source>` as a prefix for methods that process data received from the external service.
This is not a strict rule, but a guideline to help maintain consistency across the codebase.
## Tests
All the tests are located in the `tests` folder of each package. There are tests that are marked with a `@skip_if_..._integration_tests_disabled` decorator, these are integration tests that require an external service to be running, like OpenAI or Azure OpenAI.
@@ -237,14 +237,14 @@ class A2AAgent(BaseAgent):
An agent response item.
"""
messages = self._normalize_messages(messages)
a2a_message = self._chat_message_to_a2a_message(messages[-1])
a2a_message = self._prepare_message_for_a2a(messages[-1])
response_stream = self.client.send_message(a2a_message)
async for item in response_stream:
if isinstance(item, Message):
# Process A2A Message
contents = self._a2a_parts_to_contents(item.parts)
contents = self._parse_contents_from_a2a(item.parts)
yield AgentRunResponseUpdate(
contents=contents,
role=Role.ASSISTANT if item.role == A2ARole.agent else Role.USER,
@@ -255,7 +255,7 @@ class A2AAgent(BaseAgent):
task, _update_event = item
if isinstance(task, Task) and task.status.state in TERMINAL_TASK_STATES:
# Convert Task artifacts to ChatMessages and yield as separate updates
task_messages = self._task_to_chat_messages(task)
task_messages = self._parse_messages_from_task(task)
if task_messages:
for message in task_messages:
# Use the artifact's ID from raw_representation as message_id for unique identification
@@ -280,8 +280,8 @@ class A2AAgent(BaseAgent):
msg = f"Only Message and Task responses are supported from A2A agents. Received: {type(item)}"
raise NotImplementedError(msg)
def _chat_message_to_a2a_message(self, message: ChatMessage) -> A2AMessage:
"""Convert a ChatMessage to an A2A Message.
def _prepare_message_for_a2a(self, message: ChatMessage) -> A2AMessage:
"""Prepare a ChatMessage for the A2A protocol.
Transforms Agent Framework ChatMessage objects into A2A protocol Messages by:
- Converting all message contents to appropriate A2A Part types
@@ -361,8 +361,8 @@ class A2AAgent(BaseAgent):
metadata=cast(dict[str, Any], message.additional_properties),
)
def _a2a_parts_to_contents(self, parts: Sequence[A2APart]) -> list[Contents]:
"""Convert A2A Parts to Agent Framework Contents.
def _parse_contents_from_a2a(self, parts: Sequence[A2APart]) -> list[Contents]:
"""Parse A2A Parts into Agent Framework Contents.
Transforms A2A protocol Parts into framework-native Content objects,
handling text, file (URI/bytes), and data parts with metadata preservation.
@@ -410,17 +410,17 @@ class A2AAgent(BaseAgent):
raise ValueError(f"Unknown Part kind: {inner_part.kind}")
return contents
def _task_to_chat_messages(self, task: Task) -> list[ChatMessage]:
"""Convert A2A Task artifacts to ChatMessages with ASSISTANT role."""
def _parse_messages_from_task(self, task: Task) -> list[ChatMessage]:
"""Parse A2A Task artifacts into ChatMessages with ASSISTANT role."""
messages: list[ChatMessage] = []
if task.artifacts is not None:
for artifact in task.artifacts:
messages.append(self._artifact_to_chat_message(artifact))
messages.append(self._parse_message_from_artifact(artifact))
elif task.history is not None and len(task.history) > 0:
# Include the last history item as the agent response
history_item = task.history[-1]
contents = self._a2a_parts_to_contents(history_item.parts)
contents = self._parse_contents_from_a2a(history_item.parts)
messages.append(
ChatMessage(
role=Role.ASSISTANT if history_item.role == A2ARole.agent else Role.USER,
@@ -431,9 +431,9 @@ class A2AAgent(BaseAgent):
return messages
def _artifact_to_chat_message(self, artifact: Artifact) -> ChatMessage:
"""Convert A2A Artifact to ChatMessage using part contents."""
contents = self._a2a_parts_to_contents(artifact.parts)
def _parse_message_from_artifact(self, artifact: Artifact) -> ChatMessage:
"""Parse A2A Artifact into ChatMessage using part contents."""
contents = self._parse_contents_from_a2a(artifact.parts)
return ChatMessage(
role=Role.ASSISTANT,
contents=contents,
+33 -33
View File
@@ -197,18 +197,18 @@ async def test_run_with_unknown_response_type_raises_error(a2a_agent: A2AAgent,
await a2a_agent.run("Test message")
def test_task_to_chat_messages_empty_artifacts(a2a_agent: A2AAgent) -> None:
"""Test _task_to_chat_messages with task containing no artifacts."""
def test_parse_messages_from_task_empty_artifacts(a2a_agent: A2AAgent) -> None:
"""Test _parse_messages_from_task with task containing no artifacts."""
task = MagicMock()
task.artifacts = None
result = a2a_agent._task_to_chat_messages(task)
result = a2a_agent._parse_messages_from_task(task)
assert len(result) == 0
def test_task_to_chat_messages_with_artifacts(a2a_agent: A2AAgent) -> None:
"""Test _task_to_chat_messages with task containing artifacts."""
def test_parse_messages_from_task_with_artifacts(a2a_agent: A2AAgent) -> None:
"""Test _parse_messages_from_task with task containing artifacts."""
task = MagicMock()
# Create mock artifacts
@@ -232,7 +232,7 @@ def test_task_to_chat_messages_with_artifacts(a2a_agent: A2AAgent) -> None:
task.artifacts = [artifact1, artifact2]
result = a2a_agent._task_to_chat_messages(task)
result = a2a_agent._parse_messages_from_task(task)
assert len(result) == 2
assert result[0].text == "Content 1"
@@ -240,8 +240,8 @@ def test_task_to_chat_messages_with_artifacts(a2a_agent: A2AAgent) -> None:
assert all(msg.role == Role.ASSISTANT for msg in result)
def test_artifact_to_chat_message(a2a_agent: A2AAgent) -> None:
"""Test _artifact_to_chat_message conversion."""
def test_parse_message_from_artifact(a2a_agent: A2AAgent) -> None:
"""Test _parse_message_from_artifact conversion."""
artifact = MagicMock()
artifact.artifact_id = "test-artifact"
@@ -253,7 +253,7 @@ def test_artifact_to_chat_message(a2a_agent: A2AAgent) -> None:
artifact.parts = [text_part]
result = a2a_agent._artifact_to_chat_message(artifact)
result = a2a_agent._parse_message_from_artifact(artifact)
assert isinstance(result, ChatMessage)
assert result.role == Role.ASSISTANT
@@ -276,7 +276,7 @@ def test_get_uri_data_invalid_uri() -> None:
_get_uri_data("not-a-valid-data-uri")
def test_a2a_parts_to_contents_conversion(a2a_agent: A2AAgent) -> None:
def test_parse_contents_from_a2a_conversion(a2a_agent: A2AAgent) -> None:
"""Test A2A parts to contents conversion."""
agent = A2AAgent(name="Test Agent", client=MockA2AClient(), _http_client=None)
@@ -285,7 +285,7 @@ def test_a2a_parts_to_contents_conversion(a2a_agent: A2AAgent) -> None:
parts = [Part(root=TextPart(text="First part")), Part(root=TextPart(text="Second part"))]
# Convert to contents
contents = agent._a2a_parts_to_contents(parts)
contents = agent._parse_contents_from_a2a(parts)
# Verify conversion
assert len(contents) == 2
@@ -295,30 +295,30 @@ def test_a2a_parts_to_contents_conversion(a2a_agent: A2AAgent) -> None:
assert contents[1].text == "Second part"
def test_chat_message_to_a2a_message_with_error_content(a2a_agent: A2AAgent) -> None:
"""Test _chat_message_to_a2a_message with ErrorContent."""
def test_prepare_message_for_a2a_with_error_content(a2a_agent: A2AAgent) -> None:
"""Test _prepare_message_for_a2a with ErrorContent."""
# Create ChatMessage with ErrorContent
error_content = ErrorContent(message="Test error message")
message = ChatMessage(role=Role.USER, contents=[error_content])
# Convert to A2A message
a2a_message = a2a_agent._chat_message_to_a2a_message(message)
a2a_message = a2a_agent._prepare_message_for_a2a(message)
# Verify conversion
assert len(a2a_message.parts) == 1
assert a2a_message.parts[0].root.text == "Test error message"
def test_chat_message_to_a2a_message_with_uri_content(a2a_agent: A2AAgent) -> None:
"""Test _chat_message_to_a2a_message with UriContent."""
def test_prepare_message_for_a2a_with_uri_content(a2a_agent: A2AAgent) -> None:
"""Test _prepare_message_for_a2a with UriContent."""
# Create ChatMessage with UriContent
uri_content = UriContent(uri="http://example.com/file.pdf", media_type="application/pdf")
message = ChatMessage(role=Role.USER, contents=[uri_content])
# Convert to A2A message
a2a_message = a2a_agent._chat_message_to_a2a_message(message)
a2a_message = a2a_agent._prepare_message_for_a2a(message)
# Verify conversion
assert len(a2a_message.parts) == 1
@@ -326,15 +326,15 @@ def test_chat_message_to_a2a_message_with_uri_content(a2a_agent: A2AAgent) -> No
assert a2a_message.parts[0].root.file.mime_type == "application/pdf"
def test_chat_message_to_a2a_message_with_data_content(a2a_agent: A2AAgent) -> None:
"""Test _chat_message_to_a2a_message with DataContent."""
def test_prepare_message_for_a2a_with_data_content(a2a_agent: A2AAgent) -> None:
"""Test _prepare_message_for_a2a with DataContent."""
# Create ChatMessage with DataContent (base64 data URI)
data_content = DataContent(uri="data:text/plain;base64,SGVsbG8gV29ybGQ=", media_type="text/plain")
message = ChatMessage(role=Role.USER, contents=[data_content])
# Convert to A2A message
a2a_message = a2a_agent._chat_message_to_a2a_message(message)
a2a_message = a2a_agent._prepare_message_for_a2a(message)
# Verify conversion
assert len(a2a_message.parts) == 1
@@ -342,14 +342,14 @@ def test_chat_message_to_a2a_message_with_data_content(a2a_agent: A2AAgent) -> N
assert a2a_message.parts[0].root.file.mime_type == "text/plain"
def test_chat_message_to_a2a_message_empty_contents_raises_error(a2a_agent: A2AAgent) -> None:
"""Test _chat_message_to_a2a_message with empty contents raises ValueError."""
def test_prepare_message_for_a2a_empty_contents_raises_error(a2a_agent: A2AAgent) -> None:
"""Test _prepare_message_for_a2a with empty contents raises ValueError."""
# Create ChatMessage with no contents
message = ChatMessage(role=Role.USER, contents=[])
# Should raise ValueError for empty contents
with raises(ValueError, match="ChatMessage.contents is empty"):
a2a_agent._chat_message_to_a2a_message(message)
a2a_agent._prepare_message_for_a2a(message)
async def test_run_stream_with_message_response(a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient) -> None:
@@ -405,7 +405,7 @@ async def test_context_manager_no_cleanup_when_no_http_client() -> None:
pass
def test_chat_message_to_a2a_message_with_multiple_contents() -> None:
def test_prepare_message_for_a2a_with_multiple_contents() -> None:
"""Test conversion of ChatMessage with multiple contents."""
agent = A2AAgent(client=MagicMock(), _http_client=None)
@@ -421,7 +421,7 @@ def test_chat_message_to_a2a_message_with_multiple_contents() -> None:
],
)
result = agent._chat_message_to_a2a_message(message)
result = agent._prepare_message_for_a2a(message)
# Should have converted all 4 contents to parts
assert len(result.parts) == 4
@@ -433,7 +433,7 @@ def test_chat_message_to_a2a_message_with_multiple_contents() -> None:
assert result.parts[3].root.kind == "text" # JSON text remains as text (no parsing)
def test_a2a_parts_to_contents_with_data_part() -> None:
def test_parse_contents_from_a2a_with_data_part() -> None:
"""Test conversion of A2A DataPart."""
agent = A2AAgent(client=MagicMock(), _http_client=None)
@@ -441,7 +441,7 @@ def test_a2a_parts_to_contents_with_data_part() -> None:
# Create DataPart
data_part = Part(root=DataPart(data={"key": "value", "number": 42}, metadata={"source": "test"}))
contents = agent._a2a_parts_to_contents([data_part])
contents = agent._parse_contents_from_a2a([data_part])
assert len(contents) == 1
@@ -450,7 +450,7 @@ def test_a2a_parts_to_contents_with_data_part() -> None:
assert contents[0].additional_properties == {"source": "test"}
def test_a2a_parts_to_contents_unknown_part_kind() -> None:
def test_parse_contents_from_a2a_unknown_part_kind() -> None:
"""Test error handling for unknown A2A part kind."""
agent = A2AAgent(client=MagicMock(), _http_client=None)
@@ -459,10 +459,10 @@ def test_a2a_parts_to_contents_unknown_part_kind() -> None:
mock_part.root.kind = "unknown_kind"
with raises(ValueError, match="Unknown Part kind: unknown_kind"):
agent._a2a_parts_to_contents([mock_part])
agent._parse_contents_from_a2a([mock_part])
def test_chat_message_to_a2a_message_with_hosted_file() -> None:
def test_prepare_message_for_a2a_with_hosted_file() -> None:
"""Test conversion of ChatMessage with HostedFileContent to A2A message."""
agent = A2AAgent(client=MagicMock(), _http_client=None)
@@ -473,7 +473,7 @@ def test_chat_message_to_a2a_message_with_hosted_file() -> None:
contents=[HostedFileContent(file_id="hosted://storage/document.pdf")],
)
result = agent._chat_message_to_a2a_message(message) # noqa: SLF001
result = agent._prepare_message_for_a2a(message) # noqa: SLF001
# Verify the conversion
assert len(result.parts) == 1
@@ -488,7 +488,7 @@ def test_chat_message_to_a2a_message_with_hosted_file() -> None:
assert part.root.file.mime_type is None # HostedFileContent doesn't specify media_type
def test_a2a_parts_to_contents_with_hosted_file_uri() -> None:
def test_parse_contents_from_a2a_with_hosted_file_uri() -> None:
"""Test conversion of A2A FilePart with hosted file URI back to UriContent."""
agent = A2AAgent(client=MagicMock(), _http_client=None)
@@ -503,7 +503,7 @@ def test_a2a_parts_to_contents_with_hosted_file_uri() -> None:
)
)
contents = agent._a2a_parts_to_contents([file_part]) # noqa: SLF001
contents = agent._parse_contents_from_a2a([file_part]) # noqa: SLF001
assert len(contents) == 1
@@ -86,7 +86,7 @@ class ExecutionContext:
def run_id(self) -> str:
"""Get or generate run ID."""
if self._run_id is None:
self._run_id = self.input_data.get("run_id") or str(uuid.uuid4())
self._run_id = self.input_data.get("run_id") or self.input_data.get("runId") or str(uuid.uuid4())
# This should never be None after the if block above, but satisfy type checkers
if self._run_id is None: # pragma: no cover
raise RuntimeError("Failed to initialize run_id")
@@ -96,7 +96,7 @@ class ExecutionContext:
def thread_id(self) -> str:
"""Get or generate thread ID."""
if self._thread_id is None:
self._thread_id = self.input_data.get("thread_id") or str(uuid.uuid4())
self._thread_id = self.input_data.get("thread_id") or self.input_data.get("threadId") or str(uuid.uuid4())
# This should never be None after the if block above, but satisfy type checkers
if self._thread_id is None: # pragma: no cover
raise RuntimeError("Failed to initialize thread_id")
@@ -83,3 +83,71 @@ async def test_default_orchestrator_merges_client_tools() -> None:
assert "server_tool" in tool_names
assert "get_weather" in tool_names
assert agent.chat_client.function_invocation_configuration.additional_tools
async def test_default_orchestrator_with_camel_case_ids() -> None:
"""Client tool is able to extract camelCase IDs."""
agent = DummyAgent()
orchestrator = DefaultOrchestrator()
input_data = {
"runId": "test-camelcase-runid",
"threadId": "test-camelcase-threadid",
"messages": [
{
"role": "user",
"content": [{"type": "input_text", "text": "Hello"}],
}
],
"tools": [],
}
context = ExecutionContext(
input_data=input_data,
agent=agent,
config=AgentConfig(),
)
events = []
async for event in orchestrator.run(context):
events.append(event)
# assert the last event has the expected run_id and thread_id
last_event = events[-1]
assert last_event.run_id == "test-camelcase-runid"
assert last_event.thread_id == "test-camelcase-threadid"
async def test_default_orchestrator_with_snake_case_ids() -> None:
"""Client tool is able to extract snake_case IDs."""
agent = DummyAgent()
orchestrator = DefaultOrchestrator()
input_data = {
"run_id": "test-snakecase-runid",
"thread_id": "test-snakecase-threadid",
"messages": [
{
"role": "user",
"content": [{"type": "input_text", "text": "Hello"}],
}
],
"tools": [],
}
context = ExecutionContext(
input_data=input_data,
agent=agent,
config=AgentConfig(),
)
events = []
async for event in orchestrator.run(context):
events.append(event)
# assert the last event has the expected run_id and thread_id
last_event = events[-1]
assert last_event.run_id == "test-snakecase-runid"
assert last_event.thread_id == "test-snakecase-threadid"
@@ -25,7 +25,6 @@ from agent_framework import (
TextContent,
TextReasoningContent,
TextSpanRegion,
ToolProtocol,
UsageContent,
UsageDetails,
get_logger,
@@ -214,9 +213,11 @@ class AnthropicClient(BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> ChatResponse:
# Extract necessary state from messages and options
run_options = self._create_run_options(messages, chat_options, **kwargs)
# prepare
run_options = self._prepare_options(messages, chat_options, **kwargs)
# execute
message = await self.anthropic_client.beta.messages.create(**run_options, stream=False)
# process
return self._process_message(message)
async def _inner_get_streaming_response(
@@ -226,16 +227,17 @@ class AnthropicClient(BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> AsyncIterable[ChatResponseUpdate]:
# Extract necessary state from messages and options
run_options = self._create_run_options(messages, chat_options, **kwargs)
# prepare
run_options = self._prepare_options(messages, chat_options, **kwargs)
# execute and process
async for chunk in await self.anthropic_client.beta.messages.create(**run_options, stream=True):
parsed_chunk = self._process_stream_event(chunk)
if parsed_chunk:
yield parsed_chunk
# region Create Run Options and Helpers
# region Prep methods
def _create_run_options(
def _prepare_options(
self,
messages: MutableSequence[ChatMessage],
chat_options: ChatOptions,
@@ -251,78 +253,91 @@ class AnthropicClient(BaseChatClient):
Returns:
A dictionary of run options for the Anthropic client.
"""
if chat_options.additional_properties and "additional_beta_flags" in chat_options.additional_properties:
betas = chat_options.additional_properties.pop("additional_beta_flags")
else:
betas = []
run_options: dict[str, Any] = {
"model": chat_options.model_id or self.model_id,
"messages": self._convert_messages_to_anthropic_format(messages),
"max_tokens": chat_options.max_tokens or ANTHROPIC_DEFAULT_MAX_TOKENS,
"extra_headers": {"User-Agent": AGENT_FRAMEWORK_USER_AGENT},
"betas": {*BETA_FLAGS, *self.additional_beta_flags, *betas},
}
run_options: dict[str, Any] = chat_options.to_dict(
exclude={
"type",
"instructions", # handled via system message
"tool_choice", # handled separately
"allow_multiple_tool_calls", # handled via tool_choice
"additional_properties", # handled separately
}
)
# Add any additional options from chat_options or kwargs
if chat_options.temperature is not None:
run_options["temperature"] = chat_options.temperature
if chat_options.top_p is not None:
run_options["top_p"] = chat_options.top_p
if chat_options.stop is not None:
run_options["stop_sequences"] = chat_options.stop
# translations between ChatOptions and Anthropic API
translations = {
"model_id": "model",
"stop": "stop_sequences",
}
for old_key, new_key in translations.items():
if old_key in run_options and old_key != new_key:
run_options[new_key] = run_options.pop(old_key)
# model id
if not run_options.get("model"):
if not self.model_id:
raise ValueError("model_id must be a non-empty string")
run_options["model"] = self.model_id
# max_tokens - Anthropic requires this, default if not provided
if not run_options.get("max_tokens"):
run_options["max_tokens"] = ANTHROPIC_DEFAULT_MAX_TOKENS
# messages
run_options["messages"] = self._prepare_messages_for_anthropic(messages)
# system message - first system message is passed as instructions
if messages and isinstance(messages[0], ChatMessage) and messages[0].role == Role.SYSTEM:
# first system message is passed as instructions
run_options["system"] = messages[0].text
if chat_options.tool_choice is not None:
match (
chat_options.tool_choice if isinstance(chat_options.tool_choice, str) else chat_options.tool_choice.mode
):
case "auto":
run_options["tool_choice"] = {"type": "auto"}
if chat_options.allow_multiple_tool_calls is not None:
run_options["tool_choice"][ # type:ignore[reportArgumentType]
"disable_parallel_tool_use"
] = not chat_options.allow_multiple_tool_calls
case "required":
if chat_options.tool_choice.required_function_name:
run_options["tool_choice"] = {
"type": "tool",
"name": chat_options.tool_choice.required_function_name,
}
if chat_options.allow_multiple_tool_calls is not None:
run_options["tool_choice"][ # type:ignore[reportArgumentType]
"disable_parallel_tool_use"
] = not chat_options.allow_multiple_tool_calls
else:
run_options["tool_choice"] = {"type": "any"}
if chat_options.allow_multiple_tool_calls is not None:
run_options["tool_choice"][ # type:ignore[reportArgumentType]
"disable_parallel_tool_use"
] = not chat_options.allow_multiple_tool_calls
case "none":
run_options["tool_choice"] = {"type": "none"}
case _:
logger.debug(f"Ignoring unsupported tool choice mode: {chat_options.tool_choice.mode} for now")
if tools_and_mcp := self._convert_tools_to_anthropic_format(chat_options.tools):
run_options.update(tools_and_mcp)
if chat_options.additional_properties:
run_options.update(chat_options.additional_properties)
# betas
run_options["betas"] = self._prepare_betas(chat_options)
# extra headers
run_options["extra_headers"] = {"User-Agent": AGENT_FRAMEWORK_USER_AGENT}
# tools, mcp servers and tool choice
if tools_config := self._prepare_tools_for_anthropic(chat_options):
run_options.update(tools_config)
# additional properties
additional_options = {
key: value
for key, value in chat_options.additional_properties.items()
if value is not None and key != "additional_beta_flags"
}
if additional_options:
run_options.update(additional_options)
run_options.update(kwargs)
return run_options
def _convert_messages_to_anthropic_format(self, messages: MutableSequence[ChatMessage]) -> list[dict[str, Any]]:
"""Convert a list of ChatMessages to the format expected by the Anthropic client.
def _prepare_betas(self, chat_options: ChatOptions) -> set[str]:
"""Prepare the beta flags for the Anthropic API request.
Args:
chat_options: The chat options that may contain additional beta flags.
Returns:
A set of beta flag strings to include in the request.
"""
return {
*BETA_FLAGS,
*self.additional_beta_flags,
*chat_options.additional_properties.get("additional_beta_flags", []),
}
def _prepare_messages_for_anthropic(self, messages: MutableSequence[ChatMessage]) -> list[dict[str, Any]]:
"""Prepare a list of ChatMessages for the Anthropic client.
This skips the first message if it is a system message,
as Anthropic expects system instructions as a separate parameter.
"""
# first system message is passed as instructions
if messages and isinstance(messages[0], ChatMessage) and messages[0].role == Role.SYSTEM:
return [self._convert_message_to_anthropic_format(msg) for msg in messages[1:]]
return [self._convert_message_to_anthropic_format(msg) for msg in messages]
return [self._prepare_message_for_anthropic(msg) for msg in messages[1:]]
return [self._prepare_message_for_anthropic(msg) for msg in messages]
def _convert_message_to_anthropic_format(self, message: ChatMessage) -> dict[str, Any]:
"""Convert a ChatMessage to the format expected by the Anthropic client.
def _prepare_message_for_anthropic(self, message: ChatMessage) -> dict[str, Any]:
"""Prepare a ChatMessage for the Anthropic client.
Args:
message: The ChatMessage to convert.
@@ -376,58 +391,96 @@ class AnthropicClient(BaseChatClient):
"content": a_content,
}
def _convert_tools_to_anthropic_format(
self, tools: list[ToolProtocol | MutableMapping[str, Any]] | None
) -> dict[str, Any] | None:
if not tools:
return None
tool_list: list[MutableMapping[str, Any]] = []
mcp_server_list: list[MutableMapping[str, Any]] = []
for tool in tools:
match tool:
case MutableMapping():
tool_list.append(tool)
case AIFunction():
tool_list.append({
"type": "custom",
"name": tool.name,
"description": tool.description,
"input_schema": tool.parameters(),
})
case HostedWebSearchTool():
search_tool: dict[str, Any] = {
"type": "web_search_20250305",
"name": "web_search",
}
if tool.additional_properties:
search_tool.update(tool.additional_properties)
tool_list.append(search_tool)
case HostedCodeInterpreterTool():
code_tool: dict[str, Any] = {
"type": "code_execution_20250825",
"name": "code_execution",
}
tool_list.append(code_tool)
case HostedMCPTool():
server_def: dict[str, Any] = {
"type": "url",
"name": tool.name,
"url": str(tool.url),
}
if tool.allowed_tools:
server_def["tool_configuration"] = {"allowed_tools": list(tool.allowed_tools)}
if tool.headers and (auth := tool.headers.get("authorization")):
server_def["authorization_token"] = auth
mcp_server_list.append(server_def)
case _:
logger.debug(f"Ignoring unsupported tool type: {type(tool)} for now")
def _prepare_tools_for_anthropic(self, chat_options: ChatOptions) -> dict[str, Any] | None:
"""Prepare tools and tool choice configuration for the Anthropic API request.
all_tools: dict[str, list[MutableMapping[str, Any]]] = {}
if tool_list:
all_tools["tools"] = tool_list
if mcp_server_list:
all_tools["mcp_servers"] = mcp_server_list
return all_tools
Args:
chat_options: The chat options containing tools and tool choice settings.
Returns:
A dictionary with tools, mcp_servers, and tool_choice configuration, or None if empty.
"""
result: dict[str, Any] = {}
# Process tools
if chat_options.tools:
tool_list: list[MutableMapping[str, Any]] = []
mcp_server_list: list[MutableMapping[str, Any]] = []
for tool in chat_options.tools:
match tool:
case MutableMapping():
tool_list.append(tool)
case AIFunction():
tool_list.append({
"type": "custom",
"name": tool.name,
"description": tool.description,
"input_schema": tool.parameters(),
})
case HostedWebSearchTool():
search_tool: dict[str, Any] = {
"type": "web_search_20250305",
"name": "web_search",
}
if tool.additional_properties:
search_tool.update(tool.additional_properties)
tool_list.append(search_tool)
case HostedCodeInterpreterTool():
code_tool: dict[str, Any] = {
"type": "code_execution_20250825",
"name": "code_execution",
}
tool_list.append(code_tool)
case HostedMCPTool():
server_def: dict[str, Any] = {
"type": "url",
"name": tool.name,
"url": str(tool.url),
}
if tool.allowed_tools:
server_def["tool_configuration"] = {"allowed_tools": list(tool.allowed_tools)}
if tool.headers and (auth := tool.headers.get("authorization")):
server_def["authorization_token"] = auth
mcp_server_list.append(server_def)
case _:
logger.debug(f"Ignoring unsupported tool type: {type(tool)} for now")
if tool_list:
result["tools"] = tool_list
if mcp_server_list:
result["mcp_servers"] = mcp_server_list
# Process tool choice
if chat_options.tool_choice is not None:
tool_choice_mode = (
chat_options.tool_choice if isinstance(chat_options.tool_choice, str) else chat_options.tool_choice.mode
)
match tool_choice_mode:
case "auto":
tool_choice: dict[str, Any] = {"type": "auto"}
if chat_options.allow_multiple_tool_calls is not None:
tool_choice["disable_parallel_tool_use"] = not chat_options.allow_multiple_tool_calls
result["tool_choice"] = tool_choice
case "required":
if (
not isinstance(chat_options.tool_choice, str)
and chat_options.tool_choice.required_function_name
):
tool_choice = {
"type": "tool",
"name": chat_options.tool_choice.required_function_name,
}
else:
tool_choice = {"type": "any"}
if chat_options.allow_multiple_tool_calls is not None:
tool_choice["disable_parallel_tool_use"] = not chat_options.allow_multiple_tool_calls
result["tool_choice"] = tool_choice
case "none":
result["tool_choice"] = {"type": "none"}
case _:
logger.debug(f"Ignoring unsupported tool choice mode: {tool_choice_mode} for now")
return result or None
# region Response Processing Methods
@@ -445,11 +498,11 @@ class AnthropicClient(BaseChatClient):
messages=[
ChatMessage(
role=Role.ASSISTANT,
contents=self._parse_message_contents(message.content),
contents=self._parse_contents_from_anthropic(message.content),
raw_representation=message,
)
],
usage_details=self._parse_message_usage(message.usage),
usage_details=self._parse_usage_from_anthropic(message.usage),
model_id=message.model,
finish_reason=FINISH_REASON_MAP.get(message.stop_reason) if message.stop_reason else None,
raw_response=message,
@@ -467,12 +520,12 @@ class AnthropicClient(BaseChatClient):
match event.type:
case "message_start":
usage_details: list[UsageContent] = []
if event.message.usage and (details := self._parse_message_usage(event.message.usage)):
if event.message.usage and (details := self._parse_usage_from_anthropic(event.message.usage)):
usage_details.append(UsageContent(details=details))
return ChatResponseUpdate(
response_id=event.message.id,
contents=[*self._parse_message_contents(event.message.content), *usage_details],
contents=[*self._parse_contents_from_anthropic(event.message.content), *usage_details],
model_id=event.message.model,
finish_reason=FINISH_REASON_MAP.get(event.message.stop_reason)
if event.message.stop_reason
@@ -480,7 +533,7 @@ class AnthropicClient(BaseChatClient):
raw_response=event,
)
case "message_delta":
usage = self._parse_message_usage(event.usage)
usage = self._parse_usage_from_anthropic(event.usage)
return ChatResponseUpdate(
contents=[UsageContent(details=usage, raw_representation=event.usage)] if usage else [],
raw_response=event,
@@ -488,13 +541,13 @@ class AnthropicClient(BaseChatClient):
case "message_stop":
logger.debug("Received message_stop event; no content to process.")
case "content_block_start":
contents = self._parse_message_contents([event.content_block])
contents = self._parse_contents_from_anthropic([event.content_block])
return ChatResponseUpdate(
contents=contents,
raw_response=event,
)
case "content_block_delta":
contents = self._parse_message_contents([event.delta])
contents = self._parse_contents_from_anthropic([event.delta])
return ChatResponseUpdate(
contents=contents,
raw_response=event,
@@ -505,7 +558,7 @@ class AnthropicClient(BaseChatClient):
logger.debug(f"Ignoring unsupported event type: {event.type}")
return None
def _parse_message_usage(self, usage: BetaUsage | BetaMessageDeltaUsage | None) -> UsageDetails | None:
def _parse_usage_from_anthropic(self, usage: BetaUsage | BetaMessageDeltaUsage | None) -> UsageDetails | None:
"""Parse usage details from the Anthropic message usage."""
if not usage:
return None
@@ -518,7 +571,7 @@ class AnthropicClient(BaseChatClient):
usage_details.additional_counts["anthropic.cache_read_input_tokens"] = usage.cache_read_input_tokens
return usage_details
def _parse_message_contents(
def _parse_contents_from_anthropic(
self, content: Sequence[BetaContentBlock | BetaRawContentBlockDelta | BetaTextBlock]
) -> list[Contents]:
"""Parse contents from the Anthropic message."""
@@ -530,7 +583,7 @@ class AnthropicClient(BaseChatClient):
TextContent(
text=content_block.text,
raw_representation=content_block,
annotations=self._parse_citations(content_block),
annotations=self._parse_citations_from_anthropic(content_block),
)
)
case "tool_use" | "mcp_tool_use" | "server_tool_use":
@@ -549,7 +602,7 @@ class AnthropicClient(BaseChatClient):
FunctionResultContent(
call_id=content_block.tool_use_id,
name=name if name and call_id == content_block.tool_use_id else "mcp_tool",
result=self._parse_message_contents(content_block.content)
result=self._parse_contents_from_anthropic(content_block.content)
if isinstance(content_block.content, list)
else content_block.content,
raw_representation=content_block,
@@ -608,7 +661,7 @@ class AnthropicClient(BaseChatClient):
logger.debug(f"Ignoring unsupported content type: {content_block.type} for now")
return contents
def _parse_citations(
def _parse_citations_from_anthropic(
self, content_block: BetaContentBlock | BetaRawContentBlockDelta | BetaTextBlock
) -> list[Annotations] | None:
content_citations = getattr(content_block, "citations", None)
@@ -151,12 +151,12 @@ def test_anthropic_client_service_url(mock_anthropic_client: MagicMock) -> None:
# Message Conversion Tests
def test_convert_message_to_anthropic_format_text(mock_anthropic_client: MagicMock) -> None:
def test_prepare_message_for_anthropic_text(mock_anthropic_client: MagicMock) -> None:
"""Test converting text message to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
message = ChatMessage(role=Role.USER, text="Hello, world!")
result = chat_client._convert_message_to_anthropic_format(message)
result = chat_client._prepare_message_for_anthropic(message)
assert result["role"] == "user"
assert len(result["content"]) == 1
@@ -164,7 +164,7 @@ def test_convert_message_to_anthropic_format_text(mock_anthropic_client: MagicMo
assert result["content"][0]["text"] == "Hello, world!"
def test_convert_message_to_anthropic_format_function_call(mock_anthropic_client: MagicMock) -> None:
def test_prepare_message_for_anthropic_function_call(mock_anthropic_client: MagicMock) -> None:
"""Test converting function call message to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
message = ChatMessage(
@@ -178,7 +178,7 @@ def test_convert_message_to_anthropic_format_function_call(mock_anthropic_client
],
)
result = chat_client._convert_message_to_anthropic_format(message)
result = chat_client._prepare_message_for_anthropic(message)
assert result["role"] == "assistant"
assert len(result["content"]) == 1
@@ -188,7 +188,7 @@ def test_convert_message_to_anthropic_format_function_call(mock_anthropic_client
assert result["content"][0]["input"] == {"location": "San Francisco"}
def test_convert_message_to_anthropic_format_function_result(mock_anthropic_client: MagicMock) -> None:
def test_prepare_message_for_anthropic_function_result(mock_anthropic_client: MagicMock) -> None:
"""Test converting function result message to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
message = ChatMessage(
@@ -202,7 +202,7 @@ def test_convert_message_to_anthropic_format_function_result(mock_anthropic_clie
],
)
result = chat_client._convert_message_to_anthropic_format(message)
result = chat_client._prepare_message_for_anthropic(message)
assert result["role"] == "user"
assert len(result["content"]) == 1
@@ -214,7 +214,7 @@ def test_convert_message_to_anthropic_format_function_result(mock_anthropic_clie
assert result["content"][0]["is_error"] is False
def test_convert_message_to_anthropic_format_text_reasoning(mock_anthropic_client: MagicMock) -> None:
def test_prepare_message_for_anthropic_text_reasoning(mock_anthropic_client: MagicMock) -> None:
"""Test converting text reasoning message to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
message = ChatMessage(
@@ -222,7 +222,7 @@ def test_convert_message_to_anthropic_format_text_reasoning(mock_anthropic_clien
contents=[TextReasoningContent(text="Let me think about this...")],
)
result = chat_client._convert_message_to_anthropic_format(message)
result = chat_client._prepare_message_for_anthropic(message)
assert result["role"] == "assistant"
assert len(result["content"]) == 1
@@ -230,7 +230,7 @@ def test_convert_message_to_anthropic_format_text_reasoning(mock_anthropic_clien
assert result["content"][0]["thinking"] == "Let me think about this..."
def test_convert_messages_to_anthropic_format_with_system(mock_anthropic_client: MagicMock) -> None:
def test_prepare_messages_for_anthropic_with_system(mock_anthropic_client: MagicMock) -> None:
"""Test converting messages list with system message."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [
@@ -238,7 +238,7 @@ def test_convert_messages_to_anthropic_format_with_system(mock_anthropic_client:
ChatMessage(role=Role.USER, text="Hello!"),
]
result = chat_client._convert_messages_to_anthropic_format(messages)
result = chat_client._prepare_messages_for_anthropic(messages)
# System message should be skipped
assert len(result) == 1
@@ -246,7 +246,7 @@ def test_convert_messages_to_anthropic_format_with_system(mock_anthropic_client:
assert result[0]["content"][0]["text"] == "Hello!"
def test_convert_messages_to_anthropic_format_without_system(mock_anthropic_client: MagicMock) -> None:
def test_prepare_messages_for_anthropic_without_system(mock_anthropic_client: MagicMock) -> None:
"""Test converting messages list without system message."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [
@@ -254,7 +254,7 @@ def test_convert_messages_to_anthropic_format_without_system(mock_anthropic_clie
ChatMessage(role=Role.ASSISTANT, text="Hi there!"),
]
result = chat_client._convert_messages_to_anthropic_format(messages)
result = chat_client._prepare_messages_for_anthropic(messages)
assert len(result) == 2
assert result[0]["role"] == "user"
@@ -264,7 +264,7 @@ def test_convert_messages_to_anthropic_format_without_system(mock_anthropic_clie
# Tool Conversion Tests
def test_convert_tools_to_anthropic_format_ai_function(mock_anthropic_client: MagicMock) -> None:
def test_prepare_tools_for_anthropic_ai_function(mock_anthropic_client: MagicMock) -> None:
"""Test converting AIFunction to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
@@ -273,9 +273,8 @@ def test_convert_tools_to_anthropic_format_ai_function(mock_anthropic_client: Ma
"""Get weather for a location."""
return f"Weather for {location}"
tools = [get_weather]
result = chat_client._convert_tools_to_anthropic_format(tools)
chat_options = ChatOptions(tools=[get_weather])
result = chat_client._prepare_tools_for_anthropic(chat_options)
assert result is not None
assert "tools" in result
@@ -285,12 +284,12 @@ def test_convert_tools_to_anthropic_format_ai_function(mock_anthropic_client: Ma
assert "Get weather for a location" in result["tools"][0]["description"]
def test_convert_tools_to_anthropic_format_web_search(mock_anthropic_client: MagicMock) -> None:
def test_prepare_tools_for_anthropic_web_search(mock_anthropic_client: MagicMock) -> None:
"""Test converting HostedWebSearchTool to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
tools = [HostedWebSearchTool()]
chat_options = ChatOptions(tools=[HostedWebSearchTool()])
result = chat_client._convert_tools_to_anthropic_format(tools)
result = chat_client._prepare_tools_for_anthropic(chat_options)
assert result is not None
assert "tools" in result
@@ -299,12 +298,12 @@ def test_convert_tools_to_anthropic_format_web_search(mock_anthropic_client: Mag
assert result["tools"][0]["name"] == "web_search"
def test_convert_tools_to_anthropic_format_code_interpreter(mock_anthropic_client: MagicMock) -> None:
def test_prepare_tools_for_anthropic_code_interpreter(mock_anthropic_client: MagicMock) -> None:
"""Test converting HostedCodeInterpreterTool to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
tools = [HostedCodeInterpreterTool()]
chat_options = ChatOptions(tools=[HostedCodeInterpreterTool()])
result = chat_client._convert_tools_to_anthropic_format(tools)
result = chat_client._prepare_tools_for_anthropic(chat_options)
assert result is not None
assert "tools" in result
@@ -313,12 +312,12 @@ def test_convert_tools_to_anthropic_format_code_interpreter(mock_anthropic_clien
assert result["tools"][0]["name"] == "code_execution"
def test_convert_tools_to_anthropic_format_mcp_tool(mock_anthropic_client: MagicMock) -> None:
def test_prepare_tools_for_anthropic_mcp_tool(mock_anthropic_client: MagicMock) -> None:
"""Test converting HostedMCPTool to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
tools = [HostedMCPTool(name="test-mcp", url="https://example.com/mcp")]
chat_options = ChatOptions(tools=[HostedMCPTool(name="test-mcp", url="https://example.com/mcp")])
result = chat_client._convert_tools_to_anthropic_format(tools)
result = chat_client._prepare_tools_for_anthropic(chat_options)
assert result is not None
assert "mcp_servers" in result
@@ -328,18 +327,20 @@ def test_convert_tools_to_anthropic_format_mcp_tool(mock_anthropic_client: Magic
assert result["mcp_servers"][0]["url"] == "https://example.com/mcp"
def test_convert_tools_to_anthropic_format_mcp_with_auth(mock_anthropic_client: MagicMock) -> None:
def test_prepare_tools_for_anthropic_mcp_with_auth(mock_anthropic_client: MagicMock) -> None:
"""Test converting HostedMCPTool with authorization headers."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
tools = [
HostedMCPTool(
name="test-mcp",
url="https://example.com/mcp",
headers={"authorization": "Bearer token123"},
)
]
chat_options = ChatOptions(
tools=[
HostedMCPTool(
name="test-mcp",
url="https://example.com/mcp",
headers={"authorization": "Bearer token123"},
)
]
)
result = chat_client._convert_tools_to_anthropic_format(tools)
result = chat_client._prepare_tools_for_anthropic(chat_options)
assert result is not None
assert "mcp_servers" in result
@@ -348,12 +349,12 @@ def test_convert_tools_to_anthropic_format_mcp_with_auth(mock_anthropic_client:
assert result["mcp_servers"][0]["authorization_token"] == "Bearer token123"
def test_convert_tools_to_anthropic_format_dict_tool(mock_anthropic_client: MagicMock) -> None:
def test_prepare_tools_for_anthropic_dict_tool(mock_anthropic_client: MagicMock) -> None:
"""Test converting dict tool to Anthropic format."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
tools = [{"type": "custom", "name": "custom_tool", "description": "A custom tool"}]
chat_options = ChatOptions(tools=[{"type": "custom", "name": "custom_tool", "description": "A custom tool"}])
result = chat_client._convert_tools_to_anthropic_format(tools)
result = chat_client._prepare_tools_for_anthropic(chat_options)
assert result is not None
assert "tools" in result
@@ -361,11 +362,12 @@ def test_convert_tools_to_anthropic_format_dict_tool(mock_anthropic_client: Magi
assert result["tools"][0]["name"] == "custom_tool"
def test_convert_tools_to_anthropic_format_none(mock_anthropic_client: MagicMock) -> None:
def test_prepare_tools_for_anthropic_none(mock_anthropic_client: MagicMock) -> None:
"""Test converting None tools."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
chat_options = ChatOptions()
result = chat_client._convert_tools_to_anthropic_format(None)
result = chat_client._prepare_tools_for_anthropic(chat_options)
assert result is None
@@ -373,14 +375,14 @@ def test_convert_tools_to_anthropic_format_none(mock_anthropic_client: MagicMock
# Run Options Tests
async def test_create_run_options_basic(mock_anthropic_client: MagicMock) -> None:
"""Test _create_run_options with basic ChatOptions."""
async def test_prepare_options_basic(mock_anthropic_client: MagicMock) -> None:
"""Test _prepare_options with basic ChatOptions."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [ChatMessage(role=Role.USER, text="Hello")]
chat_options = ChatOptions(max_tokens=100, temperature=0.7)
run_options = chat_client._create_run_options(messages, chat_options)
run_options = chat_client._prepare_options(messages, chat_options)
assert run_options["model"] == chat_client.model_id
assert run_options["max_tokens"] == 100
@@ -388,8 +390,8 @@ async def test_create_run_options_basic(mock_anthropic_client: MagicMock) -> Non
assert "messages" in run_options
async def test_create_run_options_with_system_message(mock_anthropic_client: MagicMock) -> None:
"""Test _create_run_options with system message."""
async def test_prepare_options_with_system_message(mock_anthropic_client: MagicMock) -> None:
"""Test _prepare_options with system message."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [
@@ -398,52 +400,52 @@ async def test_create_run_options_with_system_message(mock_anthropic_client: Mag
]
chat_options = ChatOptions()
run_options = chat_client._create_run_options(messages, chat_options)
run_options = chat_client._prepare_options(messages, chat_options)
assert run_options["system"] == "You are helpful."
assert len(run_options["messages"]) == 1 # System message not in messages list
async def test_create_run_options_with_tool_choice_auto(mock_anthropic_client: MagicMock) -> None:
"""Test _create_run_options with auto tool choice."""
async def test_prepare_options_with_tool_choice_auto(mock_anthropic_client: MagicMock) -> None:
"""Test _prepare_options with auto tool choice."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [ChatMessage(role=Role.USER, text="Hello")]
chat_options = ChatOptions(tool_choice="auto")
run_options = chat_client._create_run_options(messages, chat_options)
run_options = chat_client._prepare_options(messages, chat_options)
assert run_options["tool_choice"]["type"] == "auto"
async def test_create_run_options_with_tool_choice_required(mock_anthropic_client: MagicMock) -> None:
"""Test _create_run_options with required tool choice."""
async def test_prepare_options_with_tool_choice_required(mock_anthropic_client: MagicMock) -> None:
"""Test _prepare_options with required tool choice."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [ChatMessage(role=Role.USER, text="Hello")]
# For required with specific function, need to pass as dict
chat_options = ChatOptions(tool_choice={"mode": "required", "required_function_name": "get_weather"})
run_options = chat_client._create_run_options(messages, chat_options)
run_options = chat_client._prepare_options(messages, chat_options)
assert run_options["tool_choice"]["type"] == "tool"
assert run_options["tool_choice"]["name"] == "get_weather"
async def test_create_run_options_with_tool_choice_none(mock_anthropic_client: MagicMock) -> None:
"""Test _create_run_options with none tool choice."""
async def test_prepare_options_with_tool_choice_none(mock_anthropic_client: MagicMock) -> None:
"""Test _prepare_options with none tool choice."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [ChatMessage(role=Role.USER, text="Hello")]
chat_options = ChatOptions(tool_choice="none")
run_options = chat_client._create_run_options(messages, chat_options)
run_options = chat_client._prepare_options(messages, chat_options)
assert run_options["tool_choice"]["type"] == "none"
async def test_create_run_options_with_tools(mock_anthropic_client: MagicMock) -> None:
"""Test _create_run_options with tools."""
async def test_prepare_options_with_tools(mock_anthropic_client: MagicMock) -> None:
"""Test _prepare_options with tools."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
@ai_function
@@ -454,32 +456,32 @@ async def test_create_run_options_with_tools(mock_anthropic_client: MagicMock) -
messages = [ChatMessage(role=Role.USER, text="Hello")]
chat_options = ChatOptions(tools=[get_weather])
run_options = chat_client._create_run_options(messages, chat_options)
run_options = chat_client._prepare_options(messages, chat_options)
assert "tools" in run_options
assert len(run_options["tools"]) == 1
async def test_create_run_options_with_stop_sequences(mock_anthropic_client: MagicMock) -> None:
"""Test _create_run_options with stop sequences."""
async def test_prepare_options_with_stop_sequences(mock_anthropic_client: MagicMock) -> None:
"""Test _prepare_options with stop sequences."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [ChatMessage(role=Role.USER, text="Hello")]
chat_options = ChatOptions(stop=["STOP", "END"])
run_options = chat_client._create_run_options(messages, chat_options)
run_options = chat_client._prepare_options(messages, chat_options)
assert run_options["stop_sequences"] == ["STOP", "END"]
async def test_create_run_options_with_top_p(mock_anthropic_client: MagicMock) -> None:
"""Test _create_run_options with top_p."""
async def test_prepare_options_with_top_p(mock_anthropic_client: MagicMock) -> None:
"""Test _prepare_options with top_p."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
messages = [ChatMessage(role=Role.USER, text="Hello")]
chat_options = ChatOptions(top_p=0.9)
run_options = chat_client._create_run_options(messages, chat_options)
run_options = chat_client._prepare_options(messages, chat_options)
assert run_options["top_p"] == 0.9
@@ -540,41 +542,41 @@ def test_process_message_with_tool_use(mock_anthropic_client: MagicMock) -> None
assert response.finish_reason == FinishReason.TOOL_CALLS
def test_parse_message_usage_basic(mock_anthropic_client: MagicMock) -> None:
"""Test _parse_message_usage with basic usage."""
def test_parse_usage_from_anthropic_basic(mock_anthropic_client: MagicMock) -> None:
"""Test _parse_usage_from_anthropic with basic usage."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
usage = BetaUsage(input_tokens=10, output_tokens=5)
result = chat_client._parse_message_usage(usage)
result = chat_client._parse_usage_from_anthropic(usage)
assert result is not None
assert result.input_token_count == 10
assert result.output_token_count == 5
def test_parse_message_usage_none(mock_anthropic_client: MagicMock) -> None:
"""Test _parse_message_usage with None usage."""
def test_parse_usage_from_anthropic_none(mock_anthropic_client: MagicMock) -> None:
"""Test _parse_usage_from_anthropic with None usage."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
result = chat_client._parse_message_usage(None)
result = chat_client._parse_usage_from_anthropic(None)
assert result is None
def test_parse_message_contents_text(mock_anthropic_client: MagicMock) -> None:
"""Test _parse_message_contents with text content."""
def test_parse_contents_from_anthropic_text(mock_anthropic_client: MagicMock) -> None:
"""Test _parse_contents_from_anthropic with text content."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
content = [BetaTextBlock(type="text", text="Hello!")]
result = chat_client._parse_message_contents(content)
result = chat_client._parse_contents_from_anthropic(content)
assert len(result) == 1
assert isinstance(result[0], TextContent)
assert result[0].text == "Hello!"
def test_parse_message_contents_tool_use(mock_anthropic_client: MagicMock) -> None:
"""Test _parse_message_contents with tool use."""
def test_parse_contents_from_anthropic_tool_use(mock_anthropic_client: MagicMock) -> None:
"""Test _parse_contents_from_anthropic with tool use."""
chat_client = create_test_anthropic_client(mock_anthropic_client)
content = [
@@ -585,7 +587,7 @@ def test_parse_message_contents_tool_use(mock_anthropic_client: MagicMock) -> No
input={"location": "SF"},
)
]
result = chat_client._parse_message_contents(content)
result = chat_client._parse_contents_from_anthropic(content)
assert len(result) == 1
assert isinstance(result[0], FunctionCallContent)
@@ -278,22 +278,13 @@ class AzureAIAgentClient(BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> AsyncIterable[ChatResponseUpdate]:
# Extract necessary state from messages and options
run_options, required_action_results = await self._create_run_options(messages, chat_options, **kwargs)
# Get the thread ID
thread_id: str | None = (
chat_options.conversation_id
if chat_options.conversation_id is not None
else run_options.get("conversation_id", self.thread_id)
)
# Determine which agent to use and create if needed
# prepare
run_options, required_action_results = await self._prepare_options(messages, chat_options, **kwargs)
agent_id = await self._get_agent_id_or_create(run_options)
# Process and yield each update from the stream
# execute and process
async for update in self._process_stream(
*(await self._create_agent_stream(thread_id, agent_id, run_options, required_action_results))
*(await self._create_agent_stream(agent_id, run_options, required_action_results))
):
yield update
@@ -342,7 +333,6 @@ class AzureAIAgentClient(BaseChatClient):
async def _create_agent_stream(
self,
thread_id: str | None,
agent_id: str,
run_options: dict[str, Any],
required_action_results: list[FunctionResultContent | FunctionApprovalResponseContent] | None,
@@ -352,14 +342,14 @@ class AzureAIAgentClient(BaseChatClient):
Returns:
tuple: (stream, final_thread_id)
"""
thread_id = run_options.pop("thread_id", None)
# Get any active run for this thread
thread_run = await self._get_active_thread_run(thread_id)
stream: AsyncAgentRunStream[AsyncAgentEventHandler[Any]] | AsyncAgentEventHandler[Any]
handler: AsyncAgentEventHandler[Any] = AsyncAgentEventHandler()
tool_run_id, tool_outputs, tool_approvals = self._convert_required_action_to_tool_output(
required_action_results
)
tool_run_id, tool_outputs, tool_approvals = self._prepare_tool_outputs_for_azure_ai(required_action_results)
if (
thread_run is not None
@@ -421,19 +411,11 @@ class AzureAIAgentClient(BaseChatClient):
# No thread ID was provided, so create a new thread.
thread = await self.agents_client.threads.create(
tool_resources=run_options.get("tool_resources"), metadata=run_options.get("metadata")
tool_resources=run_options.get("tool_resources"),
metadata=run_options.get("metadata"),
messages=run_options.get("additional_messages"),
)
thread_id = thread.id
# workaround for: https://github.com/Azure/azure-sdk-for-python/issues/42805
# this occurs when otel is enabled
# once fixed, in the function above, readd:
# `messages=run_options.pop("additional_messages")`
for msg in run_options.pop("additional_messages", []):
await self.agents_client.messages.create(
thread_id=thread_id, role=msg.role, content=msg.content, metadata=msg.metadata
)
# and remove until here.
return thread_id
return thread.id
def _extract_url_citations(
self, message_delta_chunk: MessageDeltaChunk, azure_search_tool_calls: list[dict[str, Any]]
@@ -611,7 +593,7 @@ class AzureAIAgentClient(BaseChatClient):
"submit_tool_outputs",
"submit_tool_approval",
]:
function_call_contents = self._create_function_call_contents(
function_call_contents = self._parse_function_calls_from_azure_ai(
event_data, response_id
)
if function_call_contents:
@@ -753,8 +735,8 @@ class AzureAIAgentClient(BaseChatClient):
except Exception as ex:
logger.debug(f"Failed to capture Azure AI Search tool call: {ex}")
def _create_function_call_contents(self, event_data: ThreadRun, response_id: str | None) -> list[Contents]:
"""Create function call contents from a tool action event."""
def _parse_function_calls_from_azure_ai(self, event_data: ThreadRun, response_id: str | None) -> list[Contents]:
"""Parse function call contents from an Azure AI tool action event."""
if isinstance(event_data, ThreadRun) and event_data.required_action is not None:
if isinstance(event_data.required_action, SubmitToolOutputsAction):
return [
@@ -815,117 +797,197 @@ class AzureAIAgentClient(BaseChatClient):
chat_options.tool_choice = chat_tool_mode
async def _create_run_options(
async def _prepare_options(
self,
messages: MutableSequence[ChatMessage],
chat_options: ChatOptions | None,
chat_options: ChatOptions,
**kwargs: Any,
) -> tuple[dict[str, Any], list[FunctionResultContent | FunctionApprovalResponseContent] | None]:
run_options: dict[str, Any] = {**kwargs}
agent_definition = await self._load_agent_definition_if_needed()
if chat_options is not None:
run_options["max_completion_tokens"] = chat_options.max_tokens
if chat_options.model_id is not None:
run_options["model"] = chat_options.model_id
else:
run_options["model"] = self.model_id
run_options["top_p"] = chat_options.top_p
run_options["temperature"] = chat_options.temperature
run_options["parallel_tool_calls"] = chat_options.allow_multiple_tool_calls
# Use to_dict with exclusions for properties handled separately
run_options: dict[str, Any] = chat_options.to_dict(
exclude={
"type",
"instructions", # handled via messages
"tools", # handled separately
"tool_choice", # handled separately
"response_format", # handled separately
"additional_properties", # handled separately
"frequency_penalty", # not supported
"presence_penalty", # not supported
"user", # not supported
"stop", # not supported
"logit_bias", # not supported
"seed", # not supported
"store", # not supported
}
)
tool_definitions: list[ToolDefinition | dict[str, Any]] = []
# Translation between ChatOptions and Azure AI Agents API
translations = {
"model_id": "model",
"allow_multiple_tool_calls": "parallel_tool_calls",
"max_tokens": "max_completion_tokens",
}
for old_key, new_key in translations.items():
if old_key in run_options and old_key != new_key:
run_options[new_key] = run_options.pop(old_key)
# Add tools from existing agent
if agent_definition is not None:
# Don't include function tools, since they will be passed through chat_options.tools
agent_tools = [tool for tool in agent_definition.tools if not isinstance(tool, FunctionToolDefinition)]
if agent_tools:
tool_definitions.extend(agent_tools)
if agent_definition.tool_resources:
run_options["tool_resources"] = agent_definition.tool_resources
# model id fallback
if not run_options.get("model"):
run_options["model"] = self.model_id
if chat_options.tool_choice is not None:
if chat_options.tool_choice != "none" and chat_options.tools:
# Add run tools
tool_definitions.extend(await self._prep_tools(chat_options.tools, run_options))
# tools and tool_choice
if tool_definitions := await self._prepare_tool_definitions_and_resources(
chat_options, agent_definition, run_options
):
run_options["tools"] = tool_definitions
# Handle MCP tool resources for approval mode
mcp_tools = [tool for tool in chat_options.tools if isinstance(tool, HostedMCPTool)]
if mcp_tools:
mcp_resources = []
for mcp_tool in mcp_tools:
server_label = mcp_tool.name.replace(" ", "_")
mcp_resource: dict[str, Any] = {"server_label": server_label}
if tool_choice := self._prepare_tool_choice_mode(chat_options):
run_options["tool_choice"] = tool_choice
# Add headers if they exist
if mcp_tool.headers:
mcp_resource["headers"] = mcp_tool.headers
if mcp_tool.approval_mode is not None:
match mcp_tool.approval_mode:
case str():
# Map agent framework approval modes to Azure AI approval modes
approval_mode = (
"always" if mcp_tool.approval_mode == "always_require" else "never"
)
mcp_resource["require_approval"] = approval_mode
case _:
if "always_require_approval" in mcp_tool.approval_mode:
mcp_resource["require_approval"] = {
"always": mcp_tool.approval_mode["always_require_approval"]
}
elif "never_require_approval" in mcp_tool.approval_mode:
mcp_resource["require_approval"] = {
"never": mcp_tool.approval_mode["never_require_approval"]
}
mcp_resources.append(mcp_resource)
# Add MCP resources to tool_resources
if "tool_resources" not in run_options:
run_options["tool_resources"] = {}
run_options["tool_resources"]["mcp"] = mcp_resources
if chat_options.tool_choice == "none":
run_options["tool_choice"] = AgentsToolChoiceOptionMode.NONE
elif chat_options.tool_choice == "auto":
run_options["tool_choice"] = AgentsToolChoiceOptionMode.AUTO
elif (
isinstance(chat_options.tool_choice, ToolMode)
and chat_options.tool_choice == "required"
and chat_options.tool_choice.required_function_name is not None
):
run_options["tool_choice"] = AgentsNamedToolChoice(
type=AgentsNamedToolChoiceType.FUNCTION,
function=FunctionName(name=chat_options.tool_choice.required_function_name),
)
if tool_definitions:
run_options["tools"] = tool_definitions
if chat_options.response_format is not None:
run_options["response_format"] = ResponseFormatJsonSchemaType(
json_schema=ResponseFormatJsonSchema(
name=chat_options.response_format.__name__,
schema=chat_options.response_format.model_json_schema(),
)
# response format
if chat_options.response_format is not None:
run_options["response_format"] = ResponseFormatJsonSchemaType(
json_schema=ResponseFormatJsonSchema(
name=chat_options.response_format.__name__,
schema=chat_options.response_format.model_json_schema(),
)
)
# messages
additional_messages, instructions, required_action_results = self._prepare_messages(messages)
if additional_messages:
run_options["additional_messages"] = additional_messages
# Add instruction from existing agent at the beginning
if (
agent_definition is not None
and agent_definition.instructions
and agent_definition.instructions not in instructions
):
instructions.insert(0, agent_definition.instructions)
if instructions:
run_options["instructions"] = "\n".join(instructions)
# thread_id resolution (conversation_id takes precedence, then kwargs, then instance default)
run_options["thread_id"] = chat_options.conversation_id or kwargs.get("conversation_id") or self.thread_id
return run_options, required_action_results
def _prepare_tool_choice_mode(
self, chat_options: ChatOptions
) -> AgentsToolChoiceOptionMode | AgentsNamedToolChoice | None:
"""Prepare the tool choice mode for Azure AI Agents API."""
if chat_options.tool_choice is None:
return None
if chat_options.tool_choice == "none":
return AgentsToolChoiceOptionMode.NONE
if chat_options.tool_choice == "auto":
return AgentsToolChoiceOptionMode.AUTO
if (
isinstance(chat_options.tool_choice, ToolMode)
and chat_options.tool_choice == "required"
and chat_options.tool_choice.required_function_name is not None
):
return AgentsNamedToolChoice(
type=AgentsNamedToolChoiceType.FUNCTION,
function=FunctionName(name=chat_options.tool_choice.required_function_name),
)
return None
async def _prepare_tool_definitions_and_resources(
self,
chat_options: ChatOptions,
agent_definition: Agent | None,
run_options: dict[str, Any],
) -> list[ToolDefinition | dict[str, Any]]:
"""Prepare tool definitions and resources for the run options."""
tool_definitions: list[ToolDefinition | dict[str, Any]] = []
# Add tools from existing agent (exclude function tools - passed via chat_options.tools)
if agent_definition is not None:
agent_tools = [tool for tool in agent_definition.tools if not isinstance(tool, FunctionToolDefinition)]
if agent_tools:
tool_definitions.extend(agent_tools)
if agent_definition.tool_resources:
run_options["tool_resources"] = agent_definition.tool_resources
# Add run tools if tool_choice allows
if chat_options.tool_choice is not None and chat_options.tool_choice != "none" and chat_options.tools:
tool_definitions.extend(await self._prepare_tools_for_azure_ai(chat_options.tools, run_options))
# Handle MCP tool resources
mcp_resources = self._prepare_mcp_resources(chat_options.tools)
if mcp_resources:
if "tool_resources" not in run_options:
run_options["tool_resources"] = {}
run_options["tool_resources"]["mcp"] = mcp_resources
return tool_definitions
def _prepare_mcp_resources(
self, tools: Sequence["ToolProtocol | MutableMapping[str, Any]"]
) -> list[dict[str, Any]]:
"""Prepare MCP tool resources for approval mode configuration."""
mcp_tools = [tool for tool in tools if isinstance(tool, HostedMCPTool)]
if not mcp_tools:
return []
mcp_resources: list[dict[str, Any]] = []
for mcp_tool in mcp_tools:
server_label = mcp_tool.name.replace(" ", "_")
mcp_resource: dict[str, Any] = {"server_label": server_label}
if mcp_tool.headers:
mcp_resource["headers"] = mcp_tool.headers
if mcp_tool.approval_mode is not None:
match mcp_tool.approval_mode:
case str():
# Map agent framework approval modes to Azure AI approval modes
approval_mode = "always" if mcp_tool.approval_mode == "always_require" else "never"
mcp_resource["require_approval"] = approval_mode
case _:
if "always_require_approval" in mcp_tool.approval_mode:
mcp_resource["require_approval"] = {
"always": mcp_tool.approval_mode["always_require_approval"]
}
elif "never_require_approval" in mcp_tool.approval_mode:
mcp_resource["require_approval"] = {
"never": mcp_tool.approval_mode["never_require_approval"]
}
mcp_resources.append(mcp_resource)
return mcp_resources
def _prepare_messages(
self, messages: MutableSequence[ChatMessage]
) -> tuple[
list[ThreadMessageOptions] | None,
list[str],
list[FunctionResultContent | FunctionApprovalResponseContent] | None,
]:
"""Prepare messages for Azure AI Agents API.
System/developer messages are turned into instructions, since there is no such message roles in Azure AI.
All other messages are added 1:1, treating assistant messages as agent messages
and everything else as user messages.
Returns:
Tuple of (additional_messages, instructions, required_action_results)
"""
instructions: list[str] = []
required_action_results: list[FunctionResultContent | FunctionApprovalResponseContent] | None = None
additional_messages: list[ThreadMessageOptions] | None = None
# System/developer messages are turned into instructions, since there is no such message roles in Azure AI.
# All other messages are added 1:1, treating assistant messages as agent messages
# and everything else as user messages.
for chat_message in messages:
if chat_message.role.value in ["system", "developer"]:
for text_content in [content for content in chat_message.contents if isinstance(content, TextContent)]:
instructions.append(text_content.text)
continue
message_contents: list[MessageInputContentBlock] = []
@@ -942,7 +1004,7 @@ class AzureAIAgentClient(BaseChatClient):
elif isinstance(content.raw_representation, MessageInputContentBlock):
message_contents.append(content.raw_representation)
if len(message_contents) > 0:
if message_contents:
if additional_messages is None:
additional_messages = []
additional_messages.append(
@@ -952,26 +1014,12 @@ class AzureAIAgentClient(BaseChatClient):
)
)
if additional_messages is not None:
run_options["additional_messages"] = additional_messages
return additional_messages, instructions, required_action_results
# Add instruction from existing agent at the beginning
if (
agent_definition is not None
and agent_definition.instructions
and agent_definition.instructions not in instructions
):
instructions.insert(0, agent_definition.instructions)
if len(instructions) > 0:
run_options["instructions"] = "".join(instructions)
return run_options, required_action_results
async def _prep_tools(
async def _prepare_tools_for_azure_ai(
self, tools: Sequence["ToolProtocol | MutableMapping[str, Any]"], run_options: dict[str, Any] | None = None
) -> list[ToolDefinition | dict[str, Any]]:
"""Prepare tool definitions for the run options."""
"""Prepare tool definitions for the Azure AI Agents API."""
tool_definitions: list[ToolDefinition | dict[str, Any]] = []
for tool in tools:
match tool:
@@ -1044,10 +1092,11 @@ class AzureAIAgentClient(BaseChatClient):
raise ServiceInitializationError(f"Unsupported tool type: {type(tool)}")
return tool_definitions
def _convert_required_action_to_tool_output(
def _prepare_tool_outputs_for_azure_ai(
self,
required_action_results: list[FunctionResultContent | FunctionApprovalResponseContent] | None,
) -> tuple[str | None, list[ToolOutput] | None, list[ToolApproval] | None]:
"""Prepare function results and approvals for submission to the Azure AI API."""
run_id: str | None = None
tool_outputs: list[ToolOutput] | None = None
tool_approvals: list[ToolApproval] | None = None
@@ -28,10 +28,6 @@ from azure.ai.projects.models import (
)
from azure.core.credentials_async import AsyncTokenCredential
from azure.core.exceptions import ResourceNotFoundError
from openai.types.responses.parsed_response import (
ParsedResponse,
)
from openai.types.responses.response import Response as OpenAIResponse
from pydantic import BaseModel, ValidationError
from ._shared import AzureAISettings
@@ -41,6 +37,11 @@ if sys.version_info >= (3, 11):
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
logger = get_logger("agent_framework.azure")
@@ -335,6 +336,10 @@ class AzureAIClient(OpenAIBaseResponsesClient):
if "tools" in run_options:
args["tools"] = run_options["tools"]
if "temperature" in run_options:
args["temperature"] = run_options["temperature"]
if "top_p" in run_options:
args["top_p"] = run_options["top_p"]
if "response_format" in run_options:
response_format = run_options["response_format"]
@@ -364,7 +369,38 @@ class AzureAIClient(OpenAIBaseResponsesClient):
if self._should_close_client:
await self.project_client.close()
def _prepare_input(self, messages: MutableSequence[ChatMessage]) -> tuple[list[ChatMessage], str | None]:
@override
async def _prepare_options(
self,
messages: MutableSequence[ChatMessage],
chat_options: ChatOptions,
**kwargs: Any,
) -> dict[str, Any]:
"""Take ChatOptions and create the specific options for Azure AI."""
prepared_messages, instructions = self._prepare_messages_for_azure_ai(messages)
run_options = await super()._prepare_options(prepared_messages, chat_options, **kwargs)
if not self._is_application_endpoint:
# Application-scoped response APIs do not support "agent" property.
agent_reference = await self._get_agent_reference_or_create(run_options, instructions)
run_options["extra_body"] = {"agent": agent_reference}
# Remove properties that are not supported on request level
# but were configured on agent level
exclude = ["model", "tools", "response_format", "temperature", "top_p"]
for property in exclude:
run_options.pop(property, None)
return run_options
@override
def _get_current_conversation_id(self, chat_options: ChatOptions, **kwargs: Any) -> str | None:
"""Get the current conversation ID from chat options or kwargs."""
return chat_options.conversation_id or kwargs.get("conversation_id") or self.conversation_id
def _prepare_messages_for_azure_ai(
self, messages: MutableSequence[ChatMessage]
) -> tuple[list[ChatMessage], str | None]:
"""Prepare input from messages and convert system/developer messages to instructions."""
result: list[ChatMessage] = []
instructions_list: list[str] = []
@@ -383,44 +419,7 @@ class AzureAIClient(OpenAIBaseResponsesClient):
return result, instructions
async def prepare_options(
self,
messages: MutableSequence[ChatMessage],
chat_options: ChatOptions,
**kwargs: Any,
) -> dict[str, Any]:
"""Take ChatOptions and create the specific options for Azure AI."""
prepared_messages, instructions = self._prepare_input(messages)
run_options = await super().prepare_options(prepared_messages, chat_options, **kwargs)
if not self._is_application_endpoint:
# Application-scoped response APIs do not support "agent" property.
agent_reference = await self._get_agent_reference_or_create(run_options, instructions)
run_options["extra_body"] = {"agent": agent_reference}
conversation_id = chat_options.conversation_id or self.conversation_id
# Handle different conversation ID formats
if conversation_id:
if conversation_id.startswith("resp_"):
# For response IDs, set previous_response_id and remove conversation property
run_options.pop("conversation", None)
run_options["previous_response_id"] = conversation_id
elif conversation_id.startswith("conv_"):
# For conversation IDs, set conversation and remove previous_response_id property
run_options.pop("previous_response_id", None)
run_options["conversation"] = conversation_id
# Remove properties that are not supported on request level
# but were configured on agent level
exclude = ["model", "tools", "response_format"]
for property in exclude:
run_options.pop(property, None)
return run_options
async def initialize_client(self) -> None:
async def _initialize_client(self) -> None:
"""Initialize OpenAI client."""
self.client = self.project_client.get_openai_client() # type: ignore
@@ -438,7 +437,8 @@ class AzureAIClient(OpenAIBaseResponsesClient):
if description and not self.agent_description:
self.agent_description = description
def get_mcp_tool(self, tool: HostedMCPTool) -> Any:
@staticmethod
def _prepare_mcp_tool(tool: HostedMCPTool) -> MCPTool: # type: ignore[override]
"""Get MCP tool from HostedMCPTool."""
mcp = MCPTool(server_label=tool.name.replace(" ", "_"), server_url=str(tool.url))
@@ -456,17 +456,3 @@ class AzureAIClient(OpenAIBaseResponsesClient):
mcp["require_approval"] = {"never": {"tool_names": list(never_require_approvals)}}
return mcp
def get_conversation_id(
self, response: OpenAIResponse | ParsedResponse[BaseModel], store: bool | None
) -> str | None:
"""Get the conversation ID from the response if store is True."""
if store is False:
return None
# If conversation ID exists, it means that we operate with conversation
# so we use conversation ID as input and output.
if response.conversation and response.conversation.id:
return response.conversation.id
# If conversation ID doesn't exist, we operate with responses
# so we use response ID as input and output.
return response.id
@@ -367,33 +367,33 @@ async def test_azure_ai_chat_client_get_agent_id_or_create_missing_model(
await chat_client._get_agent_id_or_create() # type: ignore
async def test_azure_ai_chat_client_create_run_options_basic(mock_agents_client: MagicMock) -> None:
"""Test _create_run_options with basic ChatOptions."""
async def test_azure_ai_chat_client_prepare_options_basic(mock_agents_client: MagicMock) -> None:
"""Test _prepare_options with basic ChatOptions."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
messages = [ChatMessage(role=Role.USER, text="Hello")]
chat_options = ChatOptions(max_tokens=100, temperature=0.7)
run_options, tool_results = await chat_client._create_run_options(messages, chat_options) # type: ignore
run_options, tool_results = await chat_client._prepare_options(messages, chat_options) # type: ignore
assert run_options is not None
assert tool_results is None
async def test_azure_ai_chat_client_create_run_options_no_chat_options(mock_agents_client: MagicMock) -> None:
"""Test _create_run_options with no ChatOptions."""
async def test_azure_ai_chat_client_prepare_options_no_chat_options(mock_agents_client: MagicMock) -> None:
"""Test _prepare_options with default ChatOptions."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
messages = [ChatMessage(role=Role.USER, text="Hello")]
run_options, tool_results = await chat_client._create_run_options(messages, None) # type: ignore
run_options, tool_results = await chat_client._prepare_options(messages, ChatOptions()) # type: ignore
assert run_options is not None
assert tool_results is None
async def test_azure_ai_chat_client_create_run_options_with_image_content(mock_agents_client: MagicMock) -> None:
"""Test _create_run_options with image content."""
async def test_azure_ai_chat_client_prepare_options_with_image_content(mock_agents_client: MagicMock) -> None:
"""Test _prepare_options with image content."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
@@ -403,7 +403,7 @@ async def test_azure_ai_chat_client_create_run_options_with_image_content(mock_a
image_content = UriContent(uri="https://example.com/image.jpg", media_type="image/jpeg")
messages = [ChatMessage(role=Role.USER, contents=[image_content])]
run_options, _ = await chat_client._create_run_options(messages, None) # type: ignore
run_options, _ = await chat_client._prepare_options(messages, ChatOptions()) # type: ignore
assert "additional_messages" in run_options
assert len(run_options["additional_messages"]) == 1
@@ -412,11 +412,11 @@ async def test_azure_ai_chat_client_create_run_options_with_image_content(mock_a
assert len(message.content) == 1
def test_azure_ai_chat_client_convert_function_results_to_tool_output_none(mock_agents_client: MagicMock) -> None:
"""Test _convert_required_action_to_tool_output with None input."""
def test_azure_ai_chat_client_prepare_tool_outputs_for_azure_ai_none(mock_agents_client: MagicMock) -> None:
"""Test _prepare_tool_outputs_for_azure_ai with None input."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
run_id, tool_outputs, tool_approvals = chat_client._convert_required_action_to_tool_output(None) # type: ignore
run_id, tool_outputs, tool_approvals = chat_client._prepare_tool_outputs_for_azure_ai(None) # type: ignore
assert run_id is None
assert tool_outputs is None
@@ -484,8 +484,8 @@ def test_azure_ai_chat_client_update_agent_name_and_description_with_none_input(
assert chat_client.agent_description is None
async def test_azure_ai_chat_client_create_run_options_with_messages(mock_agents_client: MagicMock) -> None:
"""Test _create_run_options with different message types."""
async def test_azure_ai_chat_client_prepare_options_with_messages(mock_agents_client: MagicMock) -> None:
"""Test _prepare_options with different message types."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
# Test with system message (becomes instruction)
@@ -494,7 +494,7 @@ async def test_azure_ai_chat_client_create_run_options_with_messages(mock_agents
ChatMessage(role=Role.USER, text="Hello"),
]
run_options, _ = await chat_client._create_run_options(messages, None) # type: ignore
run_options, _ = await chat_client._prepare_options(messages, ChatOptions()) # type: ignore
assert "instructions" in run_options
assert "You are a helpful assistant" in run_options["instructions"]
@@ -565,8 +565,8 @@ async def test_azure_ai_chat_client_prepare_thread_cancels_active_run(mock_agent
mock_agents_client.runs.cancel.assert_called_once_with("test-thread", "run_123")
def test_azure_ai_chat_client_create_function_call_contents_basic(mock_agents_client: MagicMock) -> None:
"""Test _create_function_call_contents with basic function call."""
def test_azure_ai_chat_client_parse_function_calls_from_azure_ai_basic(mock_agents_client: MagicMock) -> None:
"""Test _parse_function_calls_from_azure_ai with basic function call."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
mock_tool_call = MagicMock(spec=RequiredFunctionToolCall)
@@ -580,7 +580,7 @@ def test_azure_ai_chat_client_create_function_call_contents_basic(mock_agents_cl
mock_event_data = MagicMock(spec=ThreadRun)
mock_event_data.required_action = mock_submit_action
result = chat_client._create_function_call_contents(mock_event_data, "response_123") # type: ignore
result = chat_client._parse_function_calls_from_azure_ai(mock_event_data, "response_123") # type: ignore
assert len(result) == 1
assert isinstance(result[0], FunctionCallContent)
@@ -588,22 +588,24 @@ def test_azure_ai_chat_client_create_function_call_contents_basic(mock_agents_cl
assert result[0].call_id == '["response_123", "call_123"]'
def test_azure_ai_chat_client_create_function_call_contents_no_submit_action(mock_agents_client: MagicMock) -> None:
"""Test _create_function_call_contents when required_action is not SubmitToolOutputsAction."""
def test_azure_ai_chat_client_parse_function_calls_from_azure_ai_no_submit_action(
mock_agents_client: MagicMock,
) -> None:
"""Test _parse_function_calls_from_azure_ai when required_action is not SubmitToolOutputsAction."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
mock_event_data = MagicMock(spec=ThreadRun)
mock_event_data.required_action = MagicMock()
result = chat_client._create_function_call_contents(mock_event_data, "response_123") # type: ignore
result = chat_client._parse_function_calls_from_azure_ai(mock_event_data, "response_123") # type: ignore
assert result == []
def test_azure_ai_chat_client_create_function_call_contents_non_function_tool_call(
def test_azure_ai_chat_client_parse_function_calls_from_azure_ai_non_function_tool_call(
mock_agents_client: MagicMock,
) -> None:
"""Test _create_function_call_contents with non-function tool call."""
"""Test _parse_function_calls_from_azure_ai with non-function tool call."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
mock_tool_call = MagicMock()
@@ -614,37 +616,37 @@ def test_azure_ai_chat_client_create_function_call_contents_non_function_tool_ca
mock_event_data = MagicMock(spec=ThreadRun)
mock_event_data.required_action = mock_submit_action
result = chat_client._create_function_call_contents(mock_event_data, "response_123") # type: ignore
result = chat_client._parse_function_calls_from_azure_ai(mock_event_data, "response_123") # type: ignore
assert result == []
async def test_azure_ai_chat_client_create_run_options_with_none_tool_choice(
async def test_azure_ai_chat_client_prepare_options_with_none_tool_choice(
mock_agents_client: MagicMock,
) -> None:
"""Test _create_run_options with tool_choice set to 'none'."""
"""Test _prepare_options with tool_choice set to 'none'."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
chat_options = ChatOptions()
chat_options.tool_choice = "none"
run_options, _ = await chat_client._create_run_options([], chat_options) # type: ignore
run_options, _ = await chat_client._prepare_options([], chat_options) # type: ignore
from azure.ai.agents.models import AgentsToolChoiceOptionMode
assert run_options["tool_choice"] == AgentsToolChoiceOptionMode.NONE
async def test_azure_ai_chat_client_create_run_options_with_auto_tool_choice(
async def test_azure_ai_chat_client_prepare_options_with_auto_tool_choice(
mock_agents_client: MagicMock,
) -> None:
"""Test _create_run_options with tool_choice set to 'auto'."""
"""Test _prepare_options with tool_choice set to 'auto'."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
chat_options = ChatOptions()
chat_options.tool_choice = "auto"
run_options, _ = await chat_client._create_run_options([], chat_options) # type: ignore
run_options, _ = await chat_client._prepare_options([], chat_options) # type: ignore
from azure.ai.agents.models import AgentsToolChoiceOptionMode
@@ -669,10 +671,10 @@ async def test_azure_ai_chat_client_prepare_tool_choice_none_string(
assert chat_options.tool_choice == ToolMode.NONE.mode
async def test_azure_ai_chat_client_create_run_options_tool_choice_required_specific_function(
async def test_azure_ai_chat_client_prepare_options_tool_choice_required_specific_function(
mock_agents_client: MagicMock,
) -> None:
"""Test _create_run_options with ToolMode.REQUIRED specifying a specific function name."""
"""Test _prepare_options with ToolMode.REQUIRED specifying a specific function name."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
required_tool_mode = ToolMode.REQUIRED("specific_function_name")
@@ -682,7 +684,7 @@ async def test_azure_ai_chat_client_create_run_options_tool_choice_required_spec
chat_options = ChatOptions(tools=[dict_tool], tool_choice=required_tool_mode)
messages = [ChatMessage(role=Role.USER, text="Hello")]
run_options, _ = await chat_client._create_run_options(messages, chat_options) # type: ignore
run_options, _ = await chat_client._prepare_options(messages, chat_options) # type: ignore
# Verify tool_choice is set to the specific named function
assert "tool_choice" in run_options
@@ -692,10 +694,10 @@ async def test_azure_ai_chat_client_create_run_options_tool_choice_required_spec
assert tool_choice.function.name == "specific_function_name" # type: ignore
async def test_azure_ai_chat_client_create_run_options_with_response_format(
async def test_azure_ai_chat_client_prepare_options_with_response_format(
mock_agents_client: MagicMock,
) -> None:
"""Test _create_run_options with response_format configured."""
"""Test _prepare_options with response_format configured."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
class TestResponseModel(BaseModel):
@@ -704,7 +706,7 @@ async def test_azure_ai_chat_client_create_run_options_with_response_format(
chat_options = ChatOptions()
chat_options.response_format = TestResponseModel
run_options, _ = await chat_client._create_run_options([], chat_options) # type: ignore
run_options, _ = await chat_client._prepare_options([], chat_options) # type: ignore
assert "response_format" in run_options
response_format = run_options["response_format"]
@@ -720,8 +722,8 @@ def test_azure_ai_chat_client_service_url_method(mock_agents_client: MagicMock)
assert url == "https://test-endpoint.com/"
async def test_azure_ai_chat_client_prep_tools_ai_function(mock_agents_client: MagicMock) -> None:
"""Test _prep_tools with AIFunction tool."""
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_ai_function(mock_agents_client: MagicMock) -> None:
"""Test _prepare_tools_for_azure_ai with AIFunction tool."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
@@ -729,28 +731,28 @@ async def test_azure_ai_chat_client_prep_tools_ai_function(mock_agents_client: M
mock_ai_function = MagicMock(spec=AIFunction)
mock_ai_function.to_json_schema_spec.return_value = {"type": "function", "function": {"name": "test_function"}}
result = await chat_client._prep_tools([mock_ai_function]) # type: ignore
result = await chat_client._prepare_tools_for_azure_ai([mock_ai_function]) # type: ignore
assert len(result) == 1
assert result[0] == {"type": "function", "function": {"name": "test_function"}}
mock_ai_function.to_json_schema_spec.assert_called_once()
async def test_azure_ai_chat_client_prep_tools_code_interpreter(mock_agents_client: MagicMock) -> None:
"""Test _prep_tools with HostedCodeInterpreterTool."""
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_code_interpreter(mock_agents_client: MagicMock) -> None:
"""Test _prepare_tools_for_azure_ai with HostedCodeInterpreterTool."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
code_interpreter_tool = HostedCodeInterpreterTool()
result = await chat_client._prep_tools([code_interpreter_tool]) # type: ignore
result = await chat_client._prepare_tools_for_azure_ai([code_interpreter_tool]) # type: ignore
assert len(result) == 1
assert isinstance(result[0], CodeInterpreterToolDefinition)
async def test_azure_ai_chat_client_prep_tools_mcp_tool(mock_agents_client: MagicMock) -> None:
"""Test _prep_tools with HostedMCPTool."""
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_mcp_tool(mock_agents_client: MagicMock) -> None:
"""Test _prepare_tools_for_azure_ai with HostedMCPTool."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
@@ -762,7 +764,7 @@ async def test_azure_ai_chat_client_prep_tools_mcp_tool(mock_agents_client: Magi
mock_mcp_tool.definitions = [{"type": "mcp", "name": "test_mcp"}]
mock_mcp_tool_class.return_value = mock_mcp_tool
result = await chat_client._prep_tools([mcp_tool]) # type: ignore
result = await chat_client._prepare_tools_for_azure_ai([mcp_tool]) # type: ignore
assert len(result) == 1
assert result[0] == {"type": "mcp", "name": "test_mcp"}
@@ -774,8 +776,8 @@ async def test_azure_ai_chat_client_prep_tools_mcp_tool(mock_agents_client: Magi
assert set(call_args["allowed_tools"]) == {"tool1", "tool2"}
async def test_azure_ai_chat_client_create_run_options_mcp_never_require(mock_agents_client: MagicMock) -> None:
"""Test _create_run_options with HostedMCPTool having never_require approval mode."""
async def test_azure_ai_chat_client_prepare_options_mcp_never_require(mock_agents_client: MagicMock) -> None:
"""Test _prepare_options with HostedMCPTool having never_require approval mode."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
mcp_tool = HostedMCPTool(name="Test MCP Tool", url="https://example.com/mcp", approval_mode="never_require")
@@ -784,12 +786,12 @@ async def test_azure_ai_chat_client_create_run_options_mcp_never_require(mock_ag
chat_options = ChatOptions(tools=[mcp_tool], tool_choice="auto")
with patch("agent_framework_azure_ai._chat_client.McpTool") as mock_mcp_tool_class:
# Mock _prep_tools to avoid actual tool preparation
# Mock _prepare_tools_for_azure_ai to avoid actual tool preparation
mock_mcp_tool_instance = MagicMock()
mock_mcp_tool_instance.definitions = [{"type": "mcp", "name": "test_mcp"}]
mock_mcp_tool_class.return_value = mock_mcp_tool_instance
run_options, _ = await chat_client._create_run_options(messages, chat_options) # type: ignore
run_options, _ = await chat_client._prepare_options(messages, chat_options) # type: ignore
# Verify tool_resources is created with correct MCP approval structure
assert "tool_resources" in run_options, (
@@ -803,8 +805,8 @@ async def test_azure_ai_chat_client_create_run_options_mcp_never_require(mock_ag
assert mcp_resource["require_approval"] == "never"
async def test_azure_ai_chat_client_create_run_options_mcp_with_headers(mock_agents_client: MagicMock) -> None:
"""Test _create_run_options with HostedMCPTool having headers."""
async def test_azure_ai_chat_client_prepare_options_mcp_with_headers(mock_agents_client: MagicMock) -> None:
"""Test _prepare_options with HostedMCPTool having headers."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client)
# Test with headers
@@ -817,12 +819,12 @@ async def test_azure_ai_chat_client_create_run_options_mcp_with_headers(mock_age
chat_options = ChatOptions(tools=[mcp_tool], tool_choice="auto")
with patch("agent_framework_azure_ai._chat_client.McpTool") as mock_mcp_tool_class:
# Mock _prep_tools to avoid actual tool preparation
# Mock _prepare_tools_for_azure_ai to avoid actual tool preparation
mock_mcp_tool_instance = MagicMock()
mock_mcp_tool_instance.definitions = [{"type": "mcp", "name": "test_mcp"}]
mock_mcp_tool_class.return_value = mock_mcp_tool_instance
run_options, _ = await chat_client._create_run_options(messages, chat_options) # type: ignore
run_options, _ = await chat_client._prepare_options(messages, chat_options) # type: ignore
# Verify tool_resources is created with headers
assert "tool_resources" in run_options
@@ -835,8 +837,10 @@ async def test_azure_ai_chat_client_create_run_options_mcp_with_headers(mock_age
assert mcp_resource["headers"] == headers
async def test_azure_ai_chat_client_prep_tools_web_search_bing_grounding(mock_agents_client: MagicMock) -> None:
"""Test _prep_tools with HostedWebSearchTool using Bing Grounding."""
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_web_search_bing_grounding(
mock_agents_client: MagicMock,
) -> None:
"""Test _prepare_tools_for_azure_ai with HostedWebSearchTool using Bing Grounding."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
@@ -856,7 +860,7 @@ async def test_azure_ai_chat_client_prep_tools_web_search_bing_grounding(mock_ag
mock_bing_tool.definitions = [{"type": "bing_grounding"}]
mock_bing_grounding.return_value = mock_bing_tool
result = await chat_client._prep_tools([web_search_tool]) # type: ignore
result = await chat_client._prepare_tools_for_azure_ai([web_search_tool]) # type: ignore
assert len(result) == 1
assert result[0] == {"type": "bing_grounding"}
@@ -868,10 +872,10 @@ async def test_azure_ai_chat_client_prep_tools_web_search_bing_grounding(mock_ag
assert "connection_id" in call_args
async def test_azure_ai_chat_client_prep_tools_web_search_bing_grounding_with_connection_id(
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_web_search_bing_grounding_with_connection_id(
mock_agents_client: MagicMock,
) -> None:
"""Test _prep_tools with HostedWebSearchTool using Bing Grounding with connection_id (no HTTP call)."""
"""Test _prepare_tools_... with HostedWebSearchTool using Bing Grounding with connection_id (no HTTP call)."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
@@ -888,15 +892,17 @@ async def test_azure_ai_chat_client_prep_tools_web_search_bing_grounding_with_co
mock_bing_tool.definitions = [{"type": "bing_grounding"}]
mock_bing_grounding.return_value = mock_bing_tool
result = await chat_client._prep_tools([web_search_tool]) # type: ignore
result = await chat_client._prepare_tools_for_azure_ai([web_search_tool]) # type: ignore
assert len(result) == 1
assert result[0] == {"type": "bing_grounding"}
mock_bing_grounding.assert_called_once_with(connection_id="direct-connection-id", count=3)
async def test_azure_ai_chat_client_prep_tools_web_search_custom_bing(mock_agents_client: MagicMock) -> None:
"""Test _prep_tools with HostedWebSearchTool using Custom Bing Search."""
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_web_search_custom_bing(
mock_agents_client: MagicMock,
) -> None:
"""Test _prepare_tools_for_azure_ai with HostedWebSearchTool using Custom Bing Search."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
@@ -914,16 +920,16 @@ async def test_azure_ai_chat_client_prep_tools_web_search_custom_bing(mock_agent
mock_custom_tool.definitions = [{"type": "bing_custom_search"}]
mock_custom_bing.return_value = mock_custom_tool
result = await chat_client._prep_tools([web_search_tool]) # type: ignore
result = await chat_client._prepare_tools_for_azure_ai([web_search_tool]) # type: ignore
assert len(result) == 1
assert result[0] == {"type": "bing_custom_search"}
async def test_azure_ai_chat_client_prep_tools_file_search_with_vector_stores(
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_file_search_with_vector_stores(
mock_agents_client: MagicMock,
) -> None:
"""Test _prep_tools with HostedFileSearchTool using vector stores."""
"""Test _prepare_tools_for_azure_ai with HostedFileSearchTool using vector stores."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
@@ -938,7 +944,7 @@ async def test_azure_ai_chat_client_prep_tools_file_search_with_vector_stores(
mock_file_search.return_value = mock_file_tool
run_options = {}
result = await chat_client._prep_tools([file_search_tool], run_options) # type: ignore
result = await chat_client._prepare_tools_for_azure_ai([file_search_tool], run_options) # type: ignore
assert len(result) == 1
assert result[0] == {"type": "file_search"}
@@ -973,7 +979,7 @@ async def test_azure_ai_chat_client_create_agent_stream_submit_tool_approvals(
with patch("azure.ai.agents.models.AsyncAgentEventHandler", return_value=mock_handler):
stream, final_thread_id = await chat_client._create_agent_stream( # type: ignore
"test-thread", "test-agent", {}, [approval_response]
"test-agent", {"thread_id": "test-thread"}, [approval_response]
)
# Verify the approvals path was taken
@@ -987,26 +993,26 @@ async def test_azure_ai_chat_client_create_agent_stream_submit_tool_approvals(
assert call_args["tool_approvals"][0].approve is True
async def test_azure_ai_chat_client_prep_tools_dict_tool(mock_agents_client: MagicMock) -> None:
"""Test _prep_tools with dictionary tool definition."""
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_dict_tool(mock_agents_client: MagicMock) -> None:
"""Test _prepare_tools_for_azure_ai with dictionary tool definition."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
dict_tool = {"type": "custom_tool", "config": {"param": "value"}}
result = await chat_client._prep_tools([dict_tool]) # type: ignore
result = await chat_client._prepare_tools_for_azure_ai([dict_tool]) # type: ignore
assert len(result) == 1
assert result[0] == dict_tool
async def test_azure_ai_chat_client_prep_tools_unsupported_tool(mock_agents_client: MagicMock) -> None:
"""Test _prep_tools with unsupported tool type."""
async def test_azure_ai_chat_client_prepare_tools_for_azure_ai_unsupported_tool(mock_agents_client: MagicMock) -> None:
"""Test _prepare_tools_for_azure_ai with unsupported tool type."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
unsupported_tool = "not_a_tool"
with pytest.raises(ServiceInitializationError, match="Unsupported tool type: <class 'str'>"):
await chat_client._prep_tools([unsupported_tool]) # type: ignore
await chat_client._prepare_tools_for_azure_ai([unsupported_tool]) # type: ignore
async def test_azure_ai_chat_client_get_active_thread_run_with_active_run(mock_agents_client: MagicMock) -> None:
@@ -1072,16 +1078,16 @@ async def test_azure_ai_chat_client_service_url(mock_agents_client: MagicMock) -
assert result == "https://test-endpoint.com/"
async def test_azure_ai_chat_client_convert_required_action_to_tool_output_function_result(
async def test_azure_ai_chat_client_prepare_tool_outputs_for_azure_ai_function_result(
mock_agents_client: MagicMock,
) -> None:
"""Test _convert_required_action_to_tool_output with FunctionResultContent."""
"""Test _prepare_tool_outputs_for_azure_ai with FunctionResultContent."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
# Test with simple result
function_result = FunctionResultContent(call_id='["run_123", "call_456"]', result="Simple result")
run_id, tool_outputs, tool_approvals = chat_client._convert_required_action_to_tool_output([function_result]) # type: ignore
run_id, tool_outputs, tool_approvals = chat_client._prepare_tool_outputs_for_azure_ai([function_result]) # type: ignore
assert run_id == "run_123"
assert tool_approvals is None
@@ -1092,7 +1098,7 @@ async def test_azure_ai_chat_client_convert_required_action_to_tool_output_funct
async def test_azure_ai_chat_client_convert_required_action_invalid_call_id(mock_agents_client: MagicMock) -> None:
"""Test _convert_required_action_to_tool_output with invalid call_id format."""
"""Test _prepare_tool_outputs_for_azure_ai with invalid call_id format."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
@@ -1100,19 +1106,19 @@ async def test_azure_ai_chat_client_convert_required_action_invalid_call_id(mock
function_result = FunctionResultContent(call_id="invalid_json", result="result")
with pytest.raises(json.JSONDecodeError):
chat_client._convert_required_action_to_tool_output([function_result]) # type: ignore
chat_client._prepare_tool_outputs_for_azure_ai([function_result]) # type: ignore
async def test_azure_ai_chat_client_convert_required_action_invalid_structure(
mock_agents_client: MagicMock,
) -> None:
"""Test _convert_required_action_to_tool_output with invalid call_id structure."""
"""Test _prepare_tool_outputs_for_azure_ai with invalid call_id structure."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
# Valid JSON but invalid structure (missing second element)
function_result = FunctionResultContent(call_id='["run_123"]', result="result")
run_id, tool_outputs, tool_approvals = chat_client._convert_required_action_to_tool_output([function_result]) # type: ignore
run_id, tool_outputs, tool_approvals = chat_client._prepare_tool_outputs_for_azure_ai([function_result]) # type: ignore
# Should return None values when structure is invalid
assert run_id is None
@@ -1123,7 +1129,7 @@ async def test_azure_ai_chat_client_convert_required_action_invalid_structure(
async def test_azure_ai_chat_client_convert_required_action_serde_model_results(
mock_agents_client: MagicMock,
) -> None:
"""Test _convert_required_action_to_tool_output with BaseModel results."""
"""Test _prepare_tool_outputs_for_azure_ai with BaseModel results."""
class MockResult(SerializationMixin):
def __init__(self, name: str, value: int):
@@ -1136,7 +1142,7 @@ async def test_azure_ai_chat_client_convert_required_action_serde_model_results(
mock_result = MockResult(name="test", value=42)
function_result = FunctionResultContent(call_id='["run_123", "call_456"]', result=mock_result)
run_id, tool_outputs, tool_approvals = chat_client._convert_required_action_to_tool_output([function_result]) # type: ignore
run_id, tool_outputs, tool_approvals = chat_client._prepare_tool_outputs_for_azure_ai([function_result]) # type: ignore
assert run_id == "run_123"
assert tool_approvals is None
@@ -1151,7 +1157,7 @@ async def test_azure_ai_chat_client_convert_required_action_serde_model_results(
async def test_azure_ai_chat_client_convert_required_action_multiple_results(
mock_agents_client: MagicMock,
) -> None:
"""Test _convert_required_action_to_tool_output with multiple results."""
"""Test _prepare_tool_outputs_for_azure_ai with multiple results."""
class MockResult(SerializationMixin):
def __init__(self, data: str):
@@ -1164,7 +1170,7 @@ async def test_azure_ai_chat_client_convert_required_action_multiple_results(
results_list = [mock_basemodel, {"key": "value"}, "string_result"]
function_result = FunctionResultContent(call_id='["run_123", "call_456"]', result=results_list)
run_id, tool_outputs, tool_approvals = chat_client._convert_required_action_to_tool_output([function_result]) # type: ignore
run_id, tool_outputs, tool_approvals = chat_client._prepare_tool_outputs_for_azure_ai([function_result]) # type: ignore
assert run_id == "run_123"
assert tool_outputs is not None
@@ -1184,7 +1190,7 @@ async def test_azure_ai_chat_client_convert_required_action_multiple_results(
async def test_azure_ai_chat_client_convert_required_action_approval_response(
mock_agents_client: MagicMock,
) -> None:
"""Test _convert_required_action_to_tool_output with FunctionApprovalResponseContent."""
"""Test _prepare_tool_outputs_for_azure_ai with FunctionApprovalResponseContent."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
# Test with approval response - need to provide required fields
@@ -1194,7 +1200,7 @@ async def test_azure_ai_chat_client_convert_required_action_approval_response(
approved=True,
)
run_id, tool_outputs, tool_approvals = chat_client._convert_required_action_to_tool_output([approval_response]) # type: ignore
run_id, tool_outputs, tool_approvals = chat_client._prepare_tool_outputs_for_azure_ai([approval_response]) # type: ignore
assert run_id == "run_123"
assert tool_outputs is None
@@ -1204,10 +1210,10 @@ async def test_azure_ai_chat_client_convert_required_action_approval_response(
assert tool_approvals[0].approve is True
async def test_azure_ai_chat_client_create_function_call_contents_approval_request(
async def test_azure_ai_chat_client_parse_function_calls_from_azure_ai_approval_request(
mock_agents_client: MagicMock,
) -> None:
"""Test _create_function_call_contents with approval action."""
"""Test _parse_function_calls_from_azure_ai with approval action."""
chat_client = create_test_azure_ai_chat_client(mock_agents_client, agent_id="test-agent")
# Mock SubmitToolApprovalAction with RequiredMcpToolCall
@@ -1222,7 +1228,7 @@ async def test_azure_ai_chat_client_create_function_call_contents_approval_reque
mock_event_data = MagicMock(spec=ThreadRun)
mock_event_data.required_action = mock_approval_action
result = chat_client._create_function_call_contents(mock_event_data, "response_123") # type: ignore
result = chat_client._parse_function_calls_from_azure_ai(mock_event_data, "response_123") # type: ignore
assert len(result) == 1
assert isinstance(result[0], FunctionApprovalRequestContent)
@@ -1312,7 +1318,7 @@ async def test_azure_ai_chat_client_create_agent_stream_submit_tool_outputs(
with patch("azure.ai.agents.models.AsyncAgentEventHandler", return_value=mock_handler):
stream, final_thread_id = await chat_client._create_agent_stream( # type: ignore
thread_id="test-thread", agent_id="test-agent", run_options={}, required_action_results=[function_result]
agent_id="test-agent", run_options={"thread_id": "test-thread"}, required_action_results=[function_result]
)
# Should call submit_tool_outputs_stream since we have matching run ID
@@ -249,10 +249,10 @@ async def test_azure_ai_client_get_agent_reference_missing_model(
await client._get_agent_reference_or_create({}, None) # type: ignore
async def test_azure_ai_client_prepare_input_with_system_messages(
async def test_azure_ai_client_prepare_messages_for_azure_ai_with_system_messages(
mock_project_client: MagicMock,
) -> None:
"""Test _prepare_input converts system/developer messages to instructions."""
"""Test _prepare_messages_for_azure_ai converts system/developer messages to instructions."""
client = create_test_azure_ai_client(mock_project_client)
messages = [
@@ -261,7 +261,7 @@ async def test_azure_ai_client_prepare_input_with_system_messages(
ChatMessage(role=Role.ASSISTANT, contents=[TextContent(text="System response")]),
]
result_messages, instructions = client._prepare_input(messages) # type: ignore
result_messages, instructions = client._prepare_messages_for_azure_ai(messages) # type: ignore
assert len(result_messages) == 2
assert result_messages[0].role == Role.USER
@@ -269,10 +269,10 @@ async def test_azure_ai_client_prepare_input_with_system_messages(
assert instructions == "You are a helpful assistant."
async def test_azure_ai_client_prepare_input_no_system_messages(
async def test_azure_ai_client_prepare_messages_for_azure_ai_no_system_messages(
mock_project_client: MagicMock,
) -> None:
"""Test _prepare_input with no system/developer messages."""
"""Test _prepare_messages_for_azure_ai with no system/developer messages."""
client = create_test_azure_ai_client(mock_project_client)
messages = [
@@ -280,7 +280,7 @@ async def test_azure_ai_client_prepare_input_no_system_messages(
ChatMessage(role=Role.ASSISTANT, contents=[TextContent(text="Hi there!")]),
]
result_messages, instructions = client._prepare_input(messages) # type: ignore
result_messages, instructions = client._prepare_messages_for_azure_ai(messages) # type: ignore
assert len(result_messages) == 2
assert instructions is None
@@ -294,14 +294,14 @@ async def test_azure_ai_client_prepare_options_basic(mock_project_client: MagicM
chat_options = ChatOptions()
with (
patch.object(client.__class__.__bases__[0], "prepare_options", return_value={"model": "test-model"}),
patch.object(client.__class__.__bases__[0], "_prepare_options", return_value={"model": "test-model"}),
patch.object(
client,
"_get_agent_reference_or_create",
return_value={"name": "test-agent", "version": "1.0", "type": "agent_reference"},
),
):
run_options = await client.prepare_options(messages, chat_options)
run_options = await client._prepare_options(messages, chat_options)
assert "extra_body" in run_options
assert run_options["extra_body"]["agent"]["name"] == "test-agent"
@@ -329,14 +329,14 @@ async def test_azure_ai_client_prepare_options_with_application_endpoint(
chat_options = ChatOptions()
with (
patch.object(client.__class__.__bases__[0], "prepare_options", return_value={"model": "test-model"}),
patch.object(client.__class__.__bases__[0], "_prepare_options", return_value={"model": "test-model"}),
patch.object(
client,
"_get_agent_reference_or_create",
return_value={"name": "test-agent", "version": "1", "type": "agent_reference"},
),
):
run_options = await client.prepare_options(messages, chat_options)
run_options = await client._prepare_options(messages, chat_options)
if expects_agent:
assert "extra_body" in run_options
@@ -369,14 +369,14 @@ async def test_azure_ai_client_prepare_options_with_application_project_client(
chat_options = ChatOptions()
with (
patch.object(client.__class__.__bases__[0], "prepare_options", return_value={"model": "test-model"}),
patch.object(client.__class__.__bases__[0], "_prepare_options", return_value={"model": "test-model"}),
patch.object(
client,
"_get_agent_reference_or_create",
return_value={"name": "test-agent", "version": "1", "type": "agent_reference"},
),
):
run_options = await client.prepare_options(messages, chat_options)
run_options = await client._prepare_options(messages, chat_options)
if expects_agent:
assert "extra_body" in run_options
@@ -386,13 +386,13 @@ async def test_azure_ai_client_prepare_options_with_application_project_client(
async def test_azure_ai_client_initialize_client(mock_project_client: MagicMock) -> None:
"""Test initialize_client method."""
"""Test _initialize_client method."""
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)
await client.initialize_client()
await client._initialize_client()
assert client.client is mock_openai_client
mock_project_client.get_openai_client.assert_called_once()
@@ -477,6 +477,30 @@ async def test_azure_ai_client_agent_creation_with_instructions(
assert call_args[1]["definition"].instructions == "Message instructions. Option instructions. "
async def test_azure_ai_client_agent_creation_with_additional_args(
mock_project_client: MagicMock,
) -> None:
"""Test agent creation with additional arguments."""
client = create_test_azure_ai_client(mock_project_client, agent_name="test-agent")
# Mock agent creation response
mock_agent = MagicMock()
mock_agent.name = "test-agent"
mock_agent.version = "1.0"
mock_project_client.agents.create_version = AsyncMock(return_value=mock_agent)
run_options = {"model": "test-model", "temperature": 0.9, "top_p": 0.8}
messages_instructions = "Message instructions. "
await client._get_agent_reference_or_create(run_options, messages_instructions) # type: ignore
# Verify agent was created with provided arguments
call_args = mock_project_client.agents.create_version.call_args
definition = call_args[1]["definition"]
assert definition.temperature == 0.9
assert definition.top_p == 0.8
async def test_azure_ai_client_agent_creation_with_tools(
mock_project_client: MagicMock,
) -> None:
@@ -703,7 +727,7 @@ async def test_azure_ai_client_prepare_options_excludes_response_format(
with (
patch.object(
client.__class__.__bases__[0],
"prepare_options",
"_prepare_options",
return_value={"model": "test-model", "response_format": ResponseFormatModel},
),
patch.object(
@@ -712,7 +736,7 @@ async def test_azure_ai_client_prepare_options_excludes_response_format(
return_value={"name": "test-agent", "version": "1.0", "type": "agent_reference"},
),
):
run_options = await client.prepare_options(messages, chat_options)
run_options = await client._prepare_options(messages, chat_options)
# response_format should be excluded from final run options
assert "response_format" not in run_options
@@ -721,94 +745,8 @@ async def test_azure_ai_client_prepare_options_excludes_response_format(
assert run_options["extra_body"]["agent"]["name"] == "test-agent"
async def test_azure_ai_client_prepare_options_with_resp_conversation_id(
mock_project_client: MagicMock,
) -> None:
"""Test prepare_options with conversation ID starting with 'resp_'."""
client = create_test_azure_ai_client(mock_project_client, agent_name="test-agent", agent_version="1.0")
messages = [ChatMessage(role=Role.USER, contents=[TextContent(text="Hello")])]
chat_options = ChatOptions(conversation_id="resp_12345")
with (
patch.object(
client.__class__.__bases__[0],
"prepare_options",
return_value={"model": "test-model", "previous_response_id": "old_value", "conversation": "old_conv"},
),
patch.object(
client,
"_get_agent_reference_or_create",
return_value={"name": "test-agent", "version": "1.0", "type": "agent_reference"},
),
):
run_options = await client.prepare_options(messages, chat_options)
# Should set previous_response_id and remove conversation property
assert run_options["previous_response_id"] == "resp_12345"
assert "conversation" not in run_options
async def test_azure_ai_client_prepare_options_with_conv_conversation_id(
mock_project_client: MagicMock,
) -> None:
"""Test prepare_options with conversation ID starting with 'conv_'."""
client = create_test_azure_ai_client(mock_project_client, agent_name="test-agent", agent_version="1.0")
messages = [ChatMessage(role=Role.USER, contents=[TextContent(text="Hello")])]
chat_options = ChatOptions(conversation_id="conv_67890")
with (
patch.object(
client.__class__.__bases__[0],
"prepare_options",
return_value={"model": "test-model", "previous_response_id": "old_value", "conversation": "old_conv"},
),
patch.object(
client,
"_get_agent_reference_or_create",
return_value={"name": "test-agent", "version": "1.0", "type": "agent_reference"},
),
):
run_options = await client.prepare_options(messages, chat_options)
# Should set conversation and remove previous_response_id property
assert run_options["conversation"] == "conv_67890"
assert "previous_response_id" not in run_options
async def test_azure_ai_client_prepare_options_with_client_conversation_id(
mock_project_client: MagicMock,
) -> None:
"""Test prepare_options using client's default conversation ID when chat options don't have one."""
client = create_test_azure_ai_client(
mock_project_client, agent_name="test-agent", agent_version="1.0", conversation_id="resp_client_default"
)
messages = [ChatMessage(role=Role.USER, contents=[TextContent(text="Hello")])]
chat_options = ChatOptions() # No conversation_id specified
with (
patch.object(
client.__class__.__bases__[0],
"prepare_options",
return_value={"model": "test-model", "previous_response_id": "old_value", "conversation": "old_conv"},
),
patch.object(
client,
"_get_agent_reference_or_create",
return_value={"name": "test-agent", "version": "1.0", "type": "agent_reference"},
),
):
run_options = await client.prepare_options(messages, chat_options)
# Should use client's default conversation_id and set previous_response_id
assert run_options["previous_response_id"] == "resp_client_default"
assert "conversation" not in run_options
def test_get_conversation_id_with_store_true_and_conversation_id() -> None:
"""Test get_conversation_id returns conversation ID when store is True and conversation exists."""
"""Test _get_conversation_id returns conversation ID when store is True and conversation exists."""
client = create_test_azure_ai_client(MagicMock())
# Mock OpenAI response with conversation
@@ -818,13 +756,13 @@ def test_get_conversation_id_with_store_true_and_conversation_id() -> None:
mock_conversation.id = "conv_67890"
mock_response.conversation = mock_conversation
result = client.get_conversation_id(mock_response, store=True)
result = client._get_conversation_id(mock_response, store=True)
assert result == "conv_67890"
def test_get_conversation_id_with_store_true_and_no_conversation() -> None:
"""Test get_conversation_id returns response ID when store is True and no conversation exists."""
"""Test _get_conversation_id returns response ID when store is True and no conversation exists."""
client = create_test_azure_ai_client(MagicMock())
# Mock OpenAI response without conversation
@@ -832,13 +770,13 @@ def test_get_conversation_id_with_store_true_and_no_conversation() -> None:
mock_response.id = "resp_12345"
mock_response.conversation = None
result = client.get_conversation_id(mock_response, store=True)
result = client._get_conversation_id(mock_response, store=True)
assert result == "resp_12345"
def test_get_conversation_id_with_store_true_and_empty_conversation_id() -> None:
"""Test get_conversation_id returns response ID when store is True and conversation ID is empty."""
"""Test _get_conversation_id returns response ID when store is True and conversation ID is empty."""
client = create_test_azure_ai_client(MagicMock())
# Mock OpenAI response with conversation but empty ID
@@ -848,13 +786,13 @@ def test_get_conversation_id_with_store_true_and_empty_conversation_id() -> None
mock_conversation.id = ""
mock_response.conversation = mock_conversation
result = client.get_conversation_id(mock_response, store=True)
result = client._get_conversation_id(mock_response, store=True)
assert result == "resp_12345"
def test_get_conversation_id_with_store_false() -> None:
"""Test get_conversation_id returns None when store is False."""
"""Test _get_conversation_id returns None when store is False."""
client = create_test_azure_ai_client(MagicMock())
# Mock OpenAI response with conversation
@@ -864,13 +802,13 @@ def test_get_conversation_id_with_store_false() -> None:
mock_conversation.id = "conv_67890"
mock_response.conversation = mock_conversation
result = client.get_conversation_id(mock_response, store=False)
result = client._get_conversation_id(mock_response, store=False)
assert result is None
def test_get_conversation_id_with_parsed_response_and_store_true() -> None:
"""Test get_conversation_id works with ParsedResponse when store is True."""
"""Test _get_conversation_id works with ParsedResponse when store is True."""
client = create_test_azure_ai_client(MagicMock())
# Mock ParsedResponse with conversation
@@ -880,13 +818,13 @@ def test_get_conversation_id_with_parsed_response_and_store_true() -> None:
mock_conversation.id = "conv_parsed_67890"
mock_response.conversation = mock_conversation
result = client.get_conversation_id(mock_response, store=True)
result = client._get_conversation_id(mock_response, store=True)
assert result == "conv_parsed_67890"
def test_get_conversation_id_with_parsed_response_no_conversation() -> None:
"""Test get_conversation_id returns response ID with ParsedResponse when no conversation exists."""
"""Test _get_conversation_id returns response ID with ParsedResponse when no conversation exists."""
client = create_test_azure_ai_client(MagicMock())
# Mock ParsedResponse without conversation
@@ -894,7 +832,7 @@ def test_get_conversation_id_with_parsed_response_no_conversation() -> None:
mock_response.id = "resp_parsed_12345"
mock_response.conversation = None
result = client.get_conversation_id(mock_response, store=True)
result = client._get_conversation_id(mock_response, store=True)
assert result == "resp_parsed_12345"
@@ -501,7 +501,7 @@ class BaseChatClient(SerializationMixin, ABC):
stop: str | Sequence[str] | None = None,
store: bool | None = None,
temperature: float | None = None,
tool_choice: ToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = None,
tool_choice: ToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
tools: ToolProtocol
| Callable[..., Any]
| MutableMapping[str, Any]
@@ -535,6 +535,7 @@ class BaseChatClient(SerializationMixin, ABC):
store: Whether to store the response.
temperature: The sampling temperature to use.
tool_choice: The tool choice for the request.
Default is `auto`.
tools: The tools to use for the request.
top_p: The nucleus sampling probability to use.
user: The user to associate with the request.
@@ -595,7 +596,7 @@ class BaseChatClient(SerializationMixin, ABC):
stop: str | Sequence[str] | None = None,
store: bool | None = None,
temperature: float | None = None,
tool_choice: ToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = None,
tool_choice: ToolMode | Literal["auto", "required", "none"] | dict[str, Any] | None = "auto",
tools: ToolProtocol
| Callable[..., Any]
| MutableMapping[str, Any]
@@ -629,6 +630,7 @@ class BaseChatClient(SerializationMixin, ABC):
store: Whether to store the response.
temperature: The sampling temperature to use.
tool_choice: The tool choice for the request.
Default is `auto`.
tools: The tools to use for the request.
top_p: The nucleus sampling probability to use.
user: The user to associate with the request.
+19 -19
View File
@@ -63,21 +63,21 @@ __all__ = [
]
def _mcp_prompt_message_to_chat_message(
def _parse_message_from_mcp(
mcp_type: types.PromptMessage | types.SamplingMessage,
) -> ChatMessage:
"""Convert a MCP container type to a Agent Framework type."""
"""Parse an MCP container type into an Agent Framework type."""
return ChatMessage(
role=Role(value=mcp_type.role),
contents=_mcp_type_to_ai_content(mcp_type.content),
contents=_parse_content_from_mcp(mcp_type.content),
raw_representation=mcp_type,
)
def _mcp_call_tool_result_to_ai_contents(
def _parse_contents_from_mcp_tool_result(
mcp_type: types.CallToolResult,
) -> list[Contents]:
"""Convert a MCP container type to a Agent Framework type.
"""Parse an MCP CallToolResult into Agent Framework content types.
This function extracts the complete _meta field from CallToolResult objects
and merges all metadata into the additional_properties field of converted
@@ -111,7 +111,7 @@ def _mcp_call_tool_result_to_ai_contents(
# Convert each content item and merge metadata
result_contents = []
for item in mcp_type.content:
contents = _mcp_type_to_ai_content(item)
contents = _parse_content_from_mcp(item)
if merged_meta_props:
for content in contents:
@@ -124,7 +124,7 @@ def _mcp_call_tool_result_to_ai_contents(
return result_contents
def _mcp_type_to_ai_content(
def _parse_content_from_mcp(
mcp_type: types.ImageContent
| types.TextContent
| types.AudioContent
@@ -142,7 +142,7 @@ def _mcp_type_to_ai_content(
| types.ToolResultContent
],
) -> list[Contents]:
"""Convert a MCP type to a Agent Framework type."""
"""Parse an MCP type into an Agent Framework type."""
mcp_types = mcp_type if isinstance(mcp_type, Sequence) else [mcp_type]
return_types: list[Contents] = []
for mcp_type in mcp_types:
@@ -152,7 +152,7 @@ def _mcp_type_to_ai_content(
case types.ImageContent() | types.AudioContent():
return_types.append(
DataContent(
uri=mcp_type.data,
data=mcp_type.data,
media_type=mcp_type.mimeType,
raw_representation=mcp_type,
)
@@ -178,7 +178,7 @@ def _mcp_type_to_ai_content(
return_types.append(
FunctionResultContent(
call_id=mcp_type.toolUseId,
result=_mcp_type_to_ai_content(mcp_type.content)
result=_parse_content_from_mcp(mcp_type.content)
if mcp_type.content
else mcp_type.structuredContent,
exception=Exception() if mcp_type.isError else None,
@@ -211,10 +211,10 @@ def _mcp_type_to_ai_content(
return return_types
def _ai_content_to_mcp_types(
def _prepare_content_for_mcp(
content: Contents,
) -> types.TextContent | types.ImageContent | types.AudioContent | types.EmbeddedResource | types.ResourceLink | None:
"""Convert a BaseContent type to a MCP type."""
"""Prepare an Agent Framework content type for MCP."""
match content:
case TextContent():
return types.TextContent(type="text", text=content.text)
@@ -253,15 +253,15 @@ def _ai_content_to_mcp_types(
return None
def _chat_message_to_mcp_types(
def _prepare_message_for_mcp(
content: ChatMessage,
) -> list[types.TextContent | types.ImageContent | types.AudioContent | types.EmbeddedResource | types.ResourceLink]:
"""Convert a ChatMessage to a list of MCP types."""
"""Prepare a ChatMessage for MCP format."""
messages: list[
types.TextContent | types.ImageContent | types.AudioContent | types.EmbeddedResource | types.ResourceLink
] = []
for item in content.contents:
mcp_content = _ai_content_to_mcp_types(item)
mcp_content = _prepare_content_for_mcp(item)
if mcp_content:
messages.append(mcp_content)
return messages
@@ -469,7 +469,7 @@ class MCPTool:
logger.debug("Sampling callback called with params: %s", params)
messages: list[ChatMessage] = []
for msg in params.messages:
messages.append(_mcp_prompt_message_to_chat_message(msg))
messages.append(_parse_message_from_mcp(msg))
try:
response = await self.chat_client.get_response(
messages,
@@ -487,7 +487,7 @@ class MCPTool:
code=types.INTERNAL_ERROR,
message="Failed to get chat message content.",
)
mcp_contents = _chat_message_to_mcp_types(response.messages[0])
mcp_contents = _prepare_message_for_mcp(response.messages[0])
# grab the first content that is of type TextContent or ImageContent
mcp_content = next(
(content for content in mcp_contents if isinstance(content, (types.TextContent, types.ImageContent))),
@@ -692,7 +692,7 @@ class MCPTool:
k: v for k, v in kwargs.items() if k not in {"chat_options", "tools", "tool_choice", "thread"}
}
try:
return _mcp_call_tool_result_to_ai_contents(
return _parse_contents_from_mcp_tool_result(
await self.session.call_tool(tool_name, arguments=filtered_kwargs)
)
except McpError as mcp_exc:
@@ -724,7 +724,7 @@ class MCPTool:
)
try:
prompt_result = await self.session.get_prompt(prompt_name, arguments=kwargs)
return [_mcp_prompt_message_to_chat_message(message) for message in prompt_result.messages]
return [_parse_message_from_mcp(message) for message in prompt_result.messages]
except McpError as mcp_exc:
raise ToolExecutionException(mcp_exc.error.message, inner_exception=mcp_exc) from mcp_exc
except Exception as ex:
@@ -886,6 +886,8 @@ def _parse_annotation(annotation: Any) -> Any:
If the second annotation (after the type) is a string, then we convert that to a Pydantic Field description.
The rest are returned as-is, allowing for multiple annotations.
Literal types are returned as-is to preserve their enum-like values.
Args:
annotation: The type annotation to parse.
@@ -894,6 +896,12 @@ def _parse_annotation(annotation: Any) -> Any:
"""
origin = get_origin(annotation)
if origin is not None:
# Literal types should be returned as-is - their args are the allowed values,
# not type annotations to be parsed. For example, Literal["Data", "Security"]
# has args ("Data", "Security") which are the valid string values.
if origin is Literal:
return annotation
args = get_args(annotation)
# For other generics, return the origin type (e.g., list for List[int])
if len(args) > 1 and isinstance(args[1], str):
@@ -1771,11 +1779,6 @@ def _handle_function_calls_response(
response: "ChatResponse | None" = None
fcc_messages: "list[ChatMessage]" = []
# If tools are provided but tool_choice is not set, default to "auto" for function invocation
tools = _extract_tools(kwargs)
if tools and kwargs.get("tool_choice") is None:
kwargs["tool_choice"] = "auto"
for attempt_idx in range(config.max_iterations if config.enabled else 0):
fcc_todo = _collect_approval_responses(prepped_messages)
if fcc_todo:
+38 -4
View File
@@ -925,6 +925,10 @@ class DataContent(BaseContent):
image_data = b"raw image bytes"
data_content = DataContent(data=image_data, media_type="image/png")
# Create from base64-encoded string
base64_string = "iVBORw0KGgoAAAANS..."
data_content = DataContent(data=base64_string, media_type="image/png")
# Create from data URI
data_uri = "data:image/png;base64,iVBORw0KGgoAAAANS..."
data_content = DataContent(uri=data_uri)
@@ -986,11 +990,38 @@ class DataContent(BaseContent):
**kwargs: Any additional keyword arguments.
"""
@overload
def __init__(
self,
*,
data: str,
media_type: str,
annotations: Sequence[Annotations | MutableMapping[str, Any]] | None = None,
additional_properties: dict[str, Any] | None = None,
raw_representation: Any | None = None,
**kwargs: Any,
) -> None:
"""Initializes a DataContent instance with base64-encoded string data.
Important:
This is for binary data that is represented as a data URI, not for online resources.
Use ``UriContent`` for online resources.
Keyword Args:
data: The base64-encoded string data represented by this instance.
The data is used directly to construct a data URI.
media_type: The media type of the data.
annotations: Optional annotations associated with the content.
additional_properties: Optional additional properties associated with the content.
raw_representation: Optional raw representation of the content.
**kwargs: Any additional keyword arguments.
"""
def __init__(
self,
*,
uri: str | None = None,
data: bytes | None = None,
data: bytes | str | None = None,
media_type: str | None = None,
annotations: Sequence[Annotations | MutableMapping[str, Any]] | None = None,
additional_properties: dict[str, Any] | None = None,
@@ -1006,8 +1037,9 @@ class DataContent(BaseContent):
Keyword Args:
uri: The URI of the data represented by this instance.
Should be in the form: "data:{media_type};base64,{base64_data}".
data: The binary data represented by this instance.
The data is transformed into a base64-encoded data URI.
data: The binary data or base64-encoded string represented by this instance.
If bytes, the data is transformed into a base64-encoded data URI.
If str, it is assumed to be already base64-encoded and used directly.
media_type: The media type of the data.
annotations: Optional annotations associated with the content.
additional_properties: Optional additional properties associated with the content.
@@ -1017,7 +1049,9 @@ class DataContent(BaseContent):
if uri is None:
if data is None or media_type is None:
raise ValueError("Either 'data' and 'media_type' or 'uri' must be provided.")
uri = f"data:{media_type};base64,{base64.b64encode(data).decode('utf-8')}"
base64_data: str = base64.b64encode(data).decode("utf-8") if isinstance(data, bytes) else data
uri = f"data:{media_type};base64,{base64_data}"
# Validate URI format and extract media type if not provided
validated_uri = self._validate_uri(uri)
@@ -26,6 +26,7 @@ from agent_framework import (
)
from ..exceptions import AgentExecutionException
from ._agent_executor import AgentExecutor
from ._checkpoint import CheckpointStorage
from ._events import (
AgentRunUpdateEvent,
@@ -141,7 +142,8 @@ class WorkflowAgent(BaseAgent):
checkpoint_storage: Runtime checkpoint storage. When provided with checkpoint_id,
used to load and restore the checkpoint. When provided without checkpoint_id,
enables checkpointing for this run.
**kwargs: Additional keyword arguments.
**kwargs: Additional keyword arguments passed through to underlying workflow
and ai_function tools.
Returns:
The final workflow response as an AgentRunResponse.
@@ -153,7 +155,7 @@ class WorkflowAgent(BaseAgent):
response_id = str(uuid.uuid4())
async for update in self._run_stream_impl(
input_messages, response_id, thread, checkpoint_id, checkpoint_storage
input_messages, response_id, thread, checkpoint_id, checkpoint_storage, **kwargs
):
response_updates.append(update)
@@ -187,7 +189,8 @@ class WorkflowAgent(BaseAgent):
checkpoint_storage: Runtime checkpoint storage. When provided with checkpoint_id,
used to load and restore the checkpoint. When provided without checkpoint_id,
enables checkpointing for this run.
**kwargs: Additional keyword arguments.
**kwargs: Additional keyword arguments passed through to underlying workflow
and ai_function tools.
Yields:
AgentRunResponseUpdate objects representing the workflow execution progress.
@@ -198,7 +201,7 @@ class WorkflowAgent(BaseAgent):
response_id = str(uuid.uuid4())
async for update in self._run_stream_impl(
input_messages, response_id, thread, checkpoint_id, checkpoint_storage
input_messages, response_id, thread, checkpoint_id, checkpoint_storage, **kwargs
):
response_updates.append(update)
yield update
@@ -216,6 +219,7 @@ class WorkflowAgent(BaseAgent):
thread: AgentThread,
checkpoint_id: str | None = None,
checkpoint_storage: CheckpointStorage | None = None,
**kwargs: Any,
) -> AsyncIterable[AgentRunResponseUpdate]:
"""Internal implementation of streaming execution.
@@ -225,6 +229,8 @@ class WorkflowAgent(BaseAgent):
thread: The conversation thread containing message history.
checkpoint_id: ID of checkpoint to restore from.
checkpoint_storage: Runtime checkpoint storage.
**kwargs: Additional keyword arguments passed through to the underlying
workflow and ai_function tools.
Yields:
AgentRunResponseUpdate objects representing the workflow execution progress.
@@ -255,6 +261,7 @@ class WorkflowAgent(BaseAgent):
message=None,
checkpoint_id=checkpoint_id,
checkpoint_storage=checkpoint_storage,
**kwargs,
)
else:
# Execute workflow with streaming (initial run or no function responses)
@@ -268,6 +275,7 @@ class WorkflowAgent(BaseAgent):
event_stream = self.workflow.run_stream(
message=conversation_messages,
checkpoint_storage=checkpoint_storage,
**kwargs,
)
# Process events from the stream
@@ -286,10 +294,20 @@ class WorkflowAgent(BaseAgent):
AgentRunUpdateEvent, RequestInfoEvent, and WorkflowOutputEvent are processed.
Other workflow events are ignored as they are workflow-internal.
For AgentRunUpdateEvent from AgentExecutor instances, only events from executors
with output_response=True are converted to agent updates. This prevents agent
responses from executors that were not explicitly marked to surface their output.
Non-AgentExecutor executors that emit AgentRunUpdateEvent directly are allowed
through since they explicitly chose to emit the event.
"""
match event:
case AgentRunUpdateEvent(data=update):
# Direct pass-through of update in an agent streaming event
case AgentRunUpdateEvent(data=update, executor_id=executor_id):
# For AgentExecutor instances, only pass through if output_response=True.
# Non-AgentExecutor executors that emit AgentRunUpdateEvent are allowed through.
executor = self.workflow.executors.get(executor_id)
if isinstance(executor, AgentExecutor) and not executor.output_response:
return None
if update:
return update
return None
@@ -297,11 +315,17 @@ class WorkflowAgent(BaseAgent):
case WorkflowOutputEvent(data=data, source_executor_id=source_executor_id):
# Convert workflow output to an agent response update.
# Handle different data types appropriately.
# Skip AgentRunResponse from AgentExecutor with output_response=True
# since streaming events already surfaced the content.
if isinstance(data, AgentRunResponse):
executor = self.workflow.executors.get(source_executor_id)
if isinstance(executor, AgentExecutor) and executor.output_response:
return None
if isinstance(data, AgentRunResponseUpdate):
# Already an update, pass through
return data
if isinstance(data, ChatMessage):
# Convert ChatMessage to update
return AgentRunResponseUpdate(
contents=list(data.contents),
role=data.role,
@@ -311,15 +335,9 @@ class WorkflowAgent(BaseAgent):
created_at=datetime.now(tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
raw_representation=data,
)
# Determine contents based on data type
if isinstance(data, BaseContent):
# Already a content type (TextContent, ImageContent, etc.)
contents: list[Contents] = [cast(Contents, data)]
elif isinstance(data, str):
contents = [TextContent(text=data)]
else:
# Fallback: convert to string representation
contents = [TextContent(text=str(data))]
contents = self._extract_contents(data)
if not contents:
return None
return AgentRunResponseUpdate(
contents=contents,
role=Role.ASSISTANT,
@@ -405,6 +423,18 @@ class WorkflowAgent(BaseAgent):
raise AgentExecutionException("Unexpected content type while awaiting request info responses.")
return function_responses
def _extract_contents(self, data: Any) -> list[Contents]:
"""Recursively extract Contents from workflow output data."""
if isinstance(data, ChatMessage):
return list(data.contents)
if isinstance(data, list):
return [c for item in data for c in self._extract_contents(item)]
if isinstance(data, BaseContent):
return [cast(Contents, data)]
if isinstance(data, str):
return [TextContent(text=data)]
return [TextContent(text=str(data))]
class _ResponseState(TypedDict):
"""State for grouping response updates by message_id."""
@@ -99,6 +99,11 @@ class AgentExecutor(Executor):
self._output_response = output_response
self._cache: list[ChatMessage] = []
@property
def output_response(self) -> bool:
"""Whether this executor yields AgentRunResponse as workflow output when complete."""
return self._output_response
@property
def workflow_output_types(self) -> list[type[Any]]:
# Override to declare AgentRunResponse as a possible output type only if enabled.
@@ -7,16 +7,16 @@ import uuid
from pathlib import Path
from typing import Literal
from ._edge import FanInEdgeGroup
from ._edge import FanInEdgeGroup, InternalEdgeGroup
from ._workflow import Workflow
# Import of WorkflowExecutor is performed lazily inside methods to avoid cycles
"""Workflow visualization module using graphviz."""
"""Workflow visualization module using graphviz and Mermaid."""
class WorkflowViz:
"""A class for visualizing workflows using graphviz."""
"""A class for visualizing workflows using graphviz and Mermaid."""
def __init__(self, workflow: Workflow):
"""Initialize the WorkflowViz with a workflow.
@@ -26,9 +26,13 @@ class WorkflowViz:
"""
self._workflow = workflow
def to_digraph(self) -> str:
def to_digraph(self, include_internal_executors: bool = False) -> str:
"""Export the workflow as a DOT format digraph string.
Args:
include_internal_executors (bool): Whether to include internal executors in the visualization.
Default is False.
Returns:
A string representation of the workflow in DOT format.
"""
@@ -39,20 +43,37 @@ class WorkflowViz:
lines.append("")
# Emit the top-level workflow nodes/edges
self._emit_workflow_digraph(self._workflow, lines, indent=" ")
self._emit_workflow_digraph(
self._workflow,
lines,
indent=" ",
include_internal_executors=include_internal_executors,
)
# Emit sub-workflows hosted by WorkflowExecutor as nested clusters
self._emit_sub_workflows_digraph(self._workflow, lines, indent=" ")
self._emit_sub_workflows_digraph(
self._workflow,
lines,
indent=" ",
include_internal_executors=include_internal_executors,
)
lines.append("}")
return "\n".join(lines)
def export(self, format: Literal["svg", "png", "pdf", "dot"] = "svg", filename: str | None = None) -> str:
def export(
self,
format: Literal["svg", "png", "pdf", "dot"] = "svg",
filename: str | None = None,
include_internal_executors: bool = False,
) -> str:
"""Export the workflow visualization to a file or return the file path.
Args:
format: The output format. Supported formats: 'svg', 'png', 'pdf', 'dot'.
filename: Optional filename to save the output. If None, creates a temporary file.
include_internal_executors (bool): Whether to include internal executors in the visualization.
Default is False.
Returns:
The path to the saved file.
@@ -66,7 +87,7 @@ class WorkflowViz:
raise ValueError(f"Unsupported format: {format}. Supported formats: svg, png, pdf, dot")
if format == "dot":
content = self.to_digraph()
content = self.to_digraph(include_internal_executors=include_internal_executors)
if filename:
with open(filename, "w", encoding="utf-8") as f:
f.write(content)
@@ -87,7 +108,7 @@ class WorkflowViz:
) from e
# Create a temporary graphviz Source object
dot_content = self.to_digraph()
dot_content = self.to_digraph(include_internal_executors=include_internal_executors)
source = graphviz.Source(dot_content)
try:
@@ -99,7 +120,7 @@ class WorkflowViz:
# Remove extension if present since graphviz.render() adds it
base_name = str(output_path.with_suffix(""))
source.render(base_name, format=format, cleanup=True)
source.render(base_name, format=format, cleanup=True) # type: ignore
# Return the actual filename with extension
return f"{base_name}.{format}"
@@ -108,7 +129,7 @@ class WorkflowViz:
temp_path = Path(temp_file.name)
base_name = str(temp_path.with_suffix(""))
source.render(base_name, format=format, cleanup=True)
source.render(base_name, format=format, cleanup=True) # type: ignore
return f"{base_name}.{format}"
except graphviz.backend.execute.ExecutableNotFound as e:
raise ImportError(
@@ -118,60 +139,72 @@ class WorkflowViz:
"brew install graphviz on macOS, or download from https://graphviz.org/download/ for other platforms."
) from e
def save_svg(self, filename: str) -> str:
def save_svg(self, filename: str, include_internal_executors: bool = False) -> str:
"""Convenience method to save as SVG.
Args:
filename: The filename to save the SVG file.
include_internal_executors (bool): Whether to include internal executors in the visualization.
Default is False.
Returns:
The path to the saved SVG file.
"""
return self.export(format="svg", filename=filename)
return self.export(format="svg", filename=filename, include_internal_executors=include_internal_executors)
def save_png(self, filename: str) -> str:
def save_png(self, filename: str, include_internal_executors: bool = False) -> str:
"""Convenience method to save as PNG.
Args:
filename: The filename to save the PNG file.
include_internal_executors (bool): Whether to include internal executors in the visualization.
Default is False.
Returns:
The path to the saved PNG file.
"""
return self.export(format="png", filename=filename)
return self.export(format="png", filename=filename, include_internal_executors=include_internal_executors)
def save_pdf(self, filename: str) -> str:
def save_pdf(self, filename: str, include_internal_executors: bool = False) -> str:
"""Convenience method to save as PDF.
Args:
filename: The filename to save the PDF file.
include_internal_executors (bool): Whether to include internal executors in the visualization.
Default is False.
Returns:
The path to the saved PDF file.
"""
return self.export(format="pdf", filename=filename)
return self.export(format="pdf", filename=filename, include_internal_executors=include_internal_executors)
def to_mermaid(self) -> str:
def to_mermaid(self, include_internal_executors: bool = False) -> str:
"""Export the workflow as a Mermaid flowchart string.
Args:
include_internal_executors (bool): Whether to include internal executors in the visualization.
Default is False.
Returns:
A string representation of the workflow in Mermaid flowchart syntax.
"""
def _san(s: str) -> str:
"""Sanitize an ID for Mermaid (alphanumeric and underscore, start with letter)."""
s2 = re.sub(r"[^0-9A-Za-z_]", "_", s)
if not s2 or not s2[0].isalpha():
s2 = f"n_{s2}"
return s2
lines: list[str] = ["flowchart TD"]
# Emit top-level workflow
self._emit_workflow_mermaid(self._workflow, lines, indent=" ")
self._emit_workflow_mermaid(
self._workflow,
lines,
indent=" ",
include_internal_executors=include_internal_executors,
)
# Emit sub-workflows as Mermaid subgraphs
self._emit_sub_workflows_mermaid(self._workflow, lines, indent=" ")
self._emit_sub_workflows_mermaid(
self._workflow,
lines,
indent=" ",
include_internal_executors=include_internal_executors,
)
return "\n".join(lines)
@@ -181,13 +214,13 @@ class WorkflowViz:
sources_sorted = sorted(sources)
return hashlib.sha256((target + "|" + "|".join(sources_sorted)).encode("utf-8")).hexdigest()[:8]
def _compute_fan_in_descriptors(self, wf: Workflow | None = None) -> list[tuple[str, list[str], str]]:
def _compute_fan_in_descriptors(self, workflow: Workflow | None = None) -> list[tuple[str, list[str], str]]:
"""Return list of (node_id, sources, target) for fan-in groups.
node_id is DOT-oriented: fan_in::target::digest
"""
result: list[tuple[str, list[str], str]] = []
workflow = wf or self._workflow
workflow = workflow or self._workflow
for group in workflow.edge_groups:
if isinstance(group, FanInEdgeGroup):
target = group.target_executor_ids[0]
@@ -197,13 +230,19 @@ class WorkflowViz:
result.append((node_id, sorted(sources), target))
return result
def _compute_normal_edges(self, wf: Workflow | None = None) -> list[tuple[str, str, bool]]:
def _compute_normal_edges(
self,
workflow: Workflow | None = None,
include_internal_executors: bool = False,
) -> list[tuple[str, str, bool]]:
"""Return list of (source_id, target_id, is_conditional) for non-fan-in groups."""
edges: list[tuple[str, str, bool]] = []
workflow = wf or self._workflow
workflow = workflow or self._workflow
for group in workflow.edge_groups:
if isinstance(group, FanInEdgeGroup):
continue
if isinstance(group, InternalEdgeGroup) and not include_internal_executors:
continue
for edge in group.edges:
is_cond = getattr(edge, "_condition", None) is not None
edges.append((edge.source_id, edge.target_id, is_cond))
@@ -213,7 +252,14 @@ class WorkflowViz:
# region Internal emitters (DOT)
def _emit_workflow_digraph(self, wf: Workflow, lines: list[str], indent: str, ns: str | None = None) -> None:
def _emit_workflow_digraph(
self,
workflow: Workflow,
lines: list[str],
indent: str,
ns: str | None = None,
include_internal_executors: bool = False,
) -> None:
"""Emit DOT nodes/edges for the given workflow.
If ns (namespace) is provided, node ids are prefixed with f"{ns}/" for uniqueness,
@@ -224,16 +270,16 @@ class WorkflowViz:
return f"{ns}/{x}" if ns else x
# Nodes
start_executor_id = wf.start_executor_id
start_executor_id = workflow.start_executor_id
lines.append(
f'{indent}"{map_id(start_executor_id)}" [fillcolor=lightgreen, label="{start_executor_id}\\n(Start)"];'
)
for executor_id in wf.executors:
for executor_id in workflow.executors:
if executor_id != start_executor_id:
lines.append(f'{indent}"{map_id(executor_id)}" [label="{executor_id}"];')
# Fan-in nodes
fan_in_nodes = self._compute_fan_in_descriptors(wf)
fan_in_nodes = self._compute_fan_in_descriptors(workflow)
if fan_in_nodes:
lines.append("")
for node_id, _, _ in fan_in_nodes:
@@ -246,11 +292,19 @@ class WorkflowViz:
lines.append(f'{indent}"{map_id(node_id)}" -> "{map_id(target)}";')
# Normal edges
for src, tgt, is_cond in self._compute_normal_edges(wf):
for src, tgt, is_cond in self._compute_normal_edges(
workflow, include_internal_executors=include_internal_executors
):
edge_attr = ' [style=dashed, label="conditional"]' if is_cond else ""
lines.append(f'{indent}"{map_id(src)}" -> "{map_id(tgt)}"{edge_attr};')
def _emit_sub_workflows_digraph(self, wf: Workflow, lines: list[str], indent: str) -> None:
def _emit_sub_workflows_digraph(
self,
workflow: Workflow,
lines: list[str],
indent: str,
include_internal_executors: bool = False,
) -> None:
"""Emit DOT subgraphs for any WorkflowExecutor instances found in the workflow."""
# Lazy import to avoid any potential import cycles
try:
@@ -258,7 +312,7 @@ class WorkflowViz:
except ImportError: # pragma: no cover - best-effort; if unavailable, skip subgraphs
return
for exec_id, exec_obj in wf.executors.items():
for exec_id, exec_obj in workflow.executors.items():
if isinstance(exec_obj, WorkflowExecutor) and hasattr(exec_obj, "workflow") and exec_obj.workflow:
subgraph_id = f"cluster_{uuid.uuid5(uuid.NAMESPACE_OID, exec_id).hex[:8]}"
lines.append(f"{indent}subgraph {subgraph_id} {{")
@@ -267,10 +321,21 @@ class WorkflowViz:
# Emit the nested workflow inside this cluster using a namespace
ns = exec_id
self._emit_workflow_digraph(exec_obj.workflow, lines, indent=f"{indent} ", ns=ns)
self._emit_workflow_digraph(
exec_obj.workflow,
lines,
indent=f"{indent} ",
ns=ns,
include_internal_executors=include_internal_executors,
)
# Recurse into deeper nested sub-workflows
self._emit_sub_workflows_digraph(exec_obj.workflow, lines, indent=f"{indent} ")
self._emit_sub_workflows_digraph(
exec_obj.workflow,
lines,
indent=f"{indent} ",
include_internal_executors=include_internal_executors,
)
lines.append(f"{indent}}}")
@@ -278,7 +343,14 @@ class WorkflowViz:
# region Internal emitters (Mermaid)
def _emit_workflow_mermaid(self, wf: Workflow, lines: list[str], indent: str, ns: str | None = None) -> None:
def _emit_workflow_mermaid(
self,
workflow: Workflow,
lines: list[str],
indent: str,
ns: str | None = None,
include_internal_executors: bool = False,
) -> None:
def _san(s: str) -> str:
s2 = re.sub(r"[^0-9A-Za-z_]", "_", s)
if not s2 or not s2[0].isalpha():
@@ -291,15 +363,15 @@ class WorkflowViz:
return _san(x)
# Nodes
start_executor_id = wf.start_executor_id
start_executor_id = workflow.start_executor_id
lines.append(f'{indent}{map_id(start_executor_id)}["{start_executor_id} (Start)"];')
for executor_id in wf.executors:
for executor_id in workflow.executors:
if executor_id == start_executor_id:
continue
lines.append(f'{indent}{map_id(executor_id)}["{executor_id}"];')
# Fan-in nodes
fan_in_nodes_dot = self._compute_fan_in_descriptors(wf)
fan_in_nodes_dot = self._compute_fan_in_descriptors(workflow)
fan_in_nodes: list[tuple[str, list[str], str]] = []
for dot_node_id, sources, target in fan_in_nodes_dot:
digest = dot_node_id.split("::")[-1]
@@ -318,7 +390,9 @@ class WorkflowViz:
lines.append(f"{indent}{fan_node_id} --> {map_id(target)};")
# Normal edges
for src, tgt, is_cond in self._compute_normal_edges(wf):
for src, tgt, is_cond in self._compute_normal_edges(
workflow, include_internal_executors=include_internal_executors
):
s = map_id(src)
t = map_id(tgt)
if is_cond:
@@ -326,7 +400,13 @@ class WorkflowViz:
else:
lines.append(f"{indent}{s} --> {t};")
def _emit_sub_workflows_mermaid(self, wf: Workflow, lines: list[str], indent: str) -> None:
def _emit_sub_workflows_mermaid(
self,
workflow: Workflow,
lines: list[str],
indent: str,
include_internal_executors: bool = False,
) -> None:
try:
from ._workflow_executor import WorkflowExecutor # type: ignore
except ImportError: # pragma: no cover
@@ -338,14 +418,25 @@ class WorkflowViz:
s2 = f"n_{s2}"
return s2
for exec_id, exec_obj in wf.executors.items():
for exec_id, exec_obj in workflow.executors.items():
if isinstance(exec_obj, WorkflowExecutor) and hasattr(exec_obj, "workflow") and exec_obj.workflow:
sg_id = _san(exec_id)
lines.append(f"{indent}subgraph {sg_id}")
# Render nested workflow within this subgraph using namespacing
self._emit_workflow_mermaid(exec_obj.workflow, lines, indent=f"{indent} ", ns=exec_id)
self._emit_workflow_mermaid(
exec_obj.workflow,
lines,
indent=f"{indent} ",
ns=exec_id,
include_internal_executors=include_internal_executors,
)
# Recurse into deeper sub-workflows
self._emit_sub_workflows_mermaid(exec_obj.workflow, lines, indent=f"{indent} ")
self._emit_sub_workflows_mermaid(
exec_obj.workflow,
lines,
indent=f"{indent} ",
include_internal_executors=include_internal_executors,
)
lines.append(f"{indent}end")
# endregion
@@ -11,6 +11,7 @@ if TYPE_CHECKING:
from ._workflow import Workflow
from ._checkpoint_encoding import decode_checkpoint_value, encode_checkpoint_value
from ._const import WORKFLOW_RUN_KWARGS_KEY
from ._events import (
RequestInfoEvent,
WorkflowErrorEvent,
@@ -366,8 +367,11 @@ class WorkflowExecutor(Executor):
logger.debug(f"WorkflowExecutor {self.id} starting sub-workflow {self.workflow.id} execution {execution_id}")
try:
# Run the sub-workflow and collect all events
result = await self.workflow.run(input_data)
# Get kwargs from parent workflow's SharedState to propagate to subworkflow
parent_kwargs: dict[str, Any] = await ctx.get_shared_state(WORKFLOW_RUN_KWARGS_KEY) or {}
# Run the sub-workflow and collect all events, passing parent kwargs
result = await self.workflow.run(input_data, **parent_kwargs)
logger.debug(
f"WorkflowExecutor {self.id} sub-workflow {self.workflow.id} "
@@ -154,7 +154,7 @@ class AzureOpenAIChatClient(AzureOpenAIConfigMixin, OpenAIBaseChatClient):
)
@override
def _parse_text_from_choice(self, choice: Choice | ChunkChoice) -> TextContent | None:
def _parse_text_from_openai(self, choice: Choice | ChunkChoice) -> TextContent | None:
"""Parse the choice into a TextContent object.
Overwritten from OpenAIBaseChatClient to deal with Azure On Your Data function.
@@ -0,0 +1,23 @@
# Copyright (c) Microsoft. All rights reserved.
import importlib
from typing import Any
IMPORT_PATH = "agent_framework_ollama"
PACKAGE_NAME = "agent-framework-ollama"
_IMPORTS = ["__version__", "OllamaChatClient", "OllamaSettings"]
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
@@ -0,0 +1,13 @@
# Copyright (c) Microsoft. All rights reserved.
from agent_framework_ollama import (
OllamaChatClient,
OllamaSettings,
__version__,
)
__all__ = [
"OllamaChatClient",
"OllamaSettings",
"__version__",
]
@@ -164,7 +164,7 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
async def close(self) -> None:
"""Clean up any assistants we created."""
if self._should_delete_assistant and self.assistant_id is not None:
client = await self.ensure_client()
client = await self._ensure_client()
await client.beta.assistants.delete(self.assistant_id)
object.__setattr__(self, "assistant_id", None)
object.__setattr__(self, "_should_delete_assistant", False)
@@ -188,7 +188,7 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> AsyncIterable[ChatResponseUpdate]:
# Extract necessary state from messages and options
# prepare
run_options, tool_results = self._prepare_options(messages, chat_options, **kwargs)
# Get the thread ID
@@ -204,10 +204,10 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
# Determine which assistant to use and create if needed
assistant_id = await self._get_assistant_id_or_create()
# Create the streaming response
# execute
stream, thread_id = await self._create_assistant_stream(thread_id, assistant_id, run_options, tool_results)
# Process and yield each update from the stream
# process
async for update in self._process_stream_events(stream, thread_id):
yield update
@@ -222,7 +222,7 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
if not self.model_id:
raise ServiceInitializationError("Parameter 'model_id' is required for assistant creation.")
client = await self.ensure_client()
client = await self._ensure_client()
created_assistant = await client.beta.assistants.create(
model=self.model_id,
description=self.assistant_description,
@@ -245,11 +245,11 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
Returns:
tuple: (stream, final_thread_id)
"""
client = await self.ensure_client()
client = await self._ensure_client()
# Get any active run for this thread
thread_run = await self._get_active_thread_run(thread_id)
tool_run_id, tool_outputs = self._convert_function_results_to_tool_output(tool_results)
tool_run_id, tool_outputs = self._prepare_tool_outputs_for_assistants(tool_results)
if thread_run is not None and tool_run_id is not None and tool_run_id == thread_run.id and tool_outputs:
# There's an active run and we have tool results to submit, so submit the results.
@@ -270,7 +270,7 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
async def _get_active_thread_run(self, thread_id: str | None) -> Run | None:
"""Get any active run for the given thread."""
client = await self.ensure_client()
client = await self._ensure_client()
if thread_id is None:
return None
@@ -281,7 +281,7 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
async def _prepare_thread(self, thread_id: str | None, thread_run: Run | None, run_options: dict[str, Any]) -> str:
"""Prepare the thread for a new run, creating or cleaning up as needed."""
client = await self.ensure_client()
client = await self._ensure_client()
if thread_id is None:
# No thread ID was provided, so create a new thread.
thread = await client.beta.threads.create( # type: ignore[reportDeprecated]
@@ -330,7 +330,7 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
response_id=response_id,
)
elif response.event == "thread.run.requires_action" and isinstance(response.data, Run):
contents = self._create_function_call_contents(response.data, response_id)
contents = self._parse_function_calls_from_assistants(response.data, response_id)
if contents:
yield ChatResponseUpdate(
role=Role.ASSISTANT,
@@ -371,8 +371,8 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
role=Role.ASSISTANT,
)
def _create_function_call_contents(self, event_data: Run, response_id: str | None) -> list[Contents]:
"""Create function call contents from a tool action event."""
def _parse_function_calls_from_assistants(self, event_data: Run, response_id: str | None) -> list[Contents]:
"""Parse function call contents from an assistants tool action event."""
contents: list[Contents] = []
if event_data.required_action is not None:
@@ -437,7 +437,10 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
if chat_options.response_format is not None:
run_options["response_format"] = {
"type": "json_schema",
"json_schema": chat_options.response_format.model_json_schema(),
"json_schema": {
"name": chat_options.response_format.__name__,
"schema": chat_options.response_format.model_json_schema(),
},
}
instructions: list[str] = []
@@ -487,10 +490,11 @@ class OpenAIAssistantsClient(OpenAIConfigMixin, BaseChatClient):
return run_options, tool_results
def _convert_function_results_to_tool_output(
def _prepare_tool_outputs_for_assistants(
self,
tool_results: list[FunctionResultContent] | None,
) -> tuple[str | None, list[ToolOutput] | None]:
"""Prepare function results for submission to the assistants API."""
run_id: str | None = None
tool_outputs: list[ToolOutput] | None = None
@@ -14,7 +14,7 @@ from openai.types.chat.chat_completion import ChatCompletion, Choice
from openai.types.chat.chat_completion_chunk import ChatCompletionChunk
from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice
from openai.types.chat.chat_completion_message_custom_tool_call import ChatCompletionMessageCustomToolCall
from pydantic import BaseModel, ValidationError
from pydantic import ValidationError
from .._clients import BaseChatClient
from .._logging import get_logger
@@ -69,10 +69,12 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> ChatResponse:
client = await self.ensure_client()
client = await self._ensure_client()
# prepare
options_dict = self._prepare_options(messages, chat_options)
try:
return self._create_chat_response(
# execute and process
return self._parse_response_from_openai(
await client.chat.completions.create(stream=False, **options_dict), chat_options
)
except BadRequestError as ex:
@@ -98,14 +100,16 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> AsyncIterable[ChatResponseUpdate]:
client = await self.ensure_client()
client = await self._ensure_client()
# prepare
options_dict = self._prepare_options(messages, chat_options)
options_dict["stream_options"] = {"include_usage": True}
try:
# execute and process
async for chunk in await client.chat.completions.create(stream=True, **options_dict):
if len(chunk.choices) == 0 and chunk.usage is None:
continue
yield self._create_chat_response_update(chunk)
yield self._parse_response_update_from_openai(chunk)
except BadRequestError as ex:
if ex.code == "content_filter":
raise OpenAIContentFilterException(
@@ -124,7 +128,9 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
# region content creation
def _chat_to_tool_spec(self, tools: Sequence[ToolProtocol | MutableMapping[str, Any]]) -> list[dict[str, Any]]:
def _prepare_tools_for_openai(
self, tools: Sequence[ToolProtocol | MutableMapping[str, Any]]
) -> list[dict[str, Any]]:
chat_tools: list[dict[str, Any]] = []
for tool in tools:
if isinstance(tool, ToolProtocol):
@@ -157,51 +163,65 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
return None
def _prepare_options(self, messages: MutableSequence[ChatMessage], chat_options: ChatOptions) -> dict[str, Any]:
# Preprocess web search tool if it exists
options_dict = chat_options.to_dict(
run_options = chat_options.to_dict(
exclude={
"type",
"instructions", # included as system message
"allow_multiple_tool_calls", # handled separately
"response_format", # handled separately
"additional_properties", # handled separately
}
)
if messages and "messages" not in options_dict:
options_dict["messages"] = self._prepare_chat_history_for_request(messages)
if "messages" not in options_dict:
# messages
if messages and "messages" not in run_options:
run_options["messages"] = self._prepare_messages_for_openai(messages)
if "messages" not in run_options:
raise ServiceInvalidRequestError("Messages are required for chat completions")
# Translation between ChatOptions and Chat Completion API
translations = {
"model_id": "model",
"allow_multiple_tool_calls": "parallel_tool_calls",
"max_tokens": "max_output_tokens",
}
for old_key, new_key in translations.items():
if old_key in run_options and old_key != new_key:
run_options[new_key] = run_options.pop(old_key)
# model id
if not run_options.get("model"):
if not self.model_id:
raise ValueError("model_id must be a non-empty string")
run_options["model"] = self.model_id
# tools
if chat_options.tools is not None:
web_search_options = self._process_web_search_tool(chat_options.tools)
if web_search_options:
options_dict["web_search_options"] = web_search_options
options_dict["tools"] = self._chat_to_tool_spec(chat_options.tools)
if chat_options.allow_multiple_tool_calls is not None:
options_dict["parallel_tool_calls"] = chat_options.allow_multiple_tool_calls
if not options_dict.get("tools", None):
options_dict.pop("tools", None)
options_dict.pop("parallel_tool_calls", None)
options_dict.pop("tool_choice", None)
# Preprocess web search tool if it exists
if web_search_options := self._process_web_search_tool(chat_options.tools):
run_options["web_search_options"] = web_search_options
run_options["tools"] = self._prepare_tools_for_openai(chat_options.tools)
if not run_options.get("tools", None):
run_options.pop("tools", None)
run_options.pop("parallel_tool_calls", None)
run_options.pop("tool_choice", None)
# tool choice when `tool_choice` is a dict with single key `mode`, extract the mode value
if (tool_choice := run_options.get("tool_choice")) and len(tool_choice.keys()) == 1:
run_options["tool_choice"] = tool_choice["mode"]
if "model_id" not in options_dict:
options_dict["model"] = self.model_id
else:
options_dict["model"] = options_dict.pop("model_id")
if (
chat_options.response_format
and isinstance(chat_options.response_format, type)
and issubclass(chat_options.response_format, BaseModel)
):
options_dict["response_format"] = type_to_response_format_param(chat_options.response_format)
if additional_properties := options_dict.pop("additional_properties", None):
for key, value in additional_properties.items():
if value is not None:
options_dict[key] = value
if (tool_choice := options_dict.get("tool_choice")) and len(tool_choice.keys()) == 1:
options_dict["tool_choice"] = tool_choice["mode"]
return options_dict
# response format
if chat_options.response_format:
run_options["response_format"] = type_to_response_format_param(chat_options.response_format)
def _create_chat_response(self, response: ChatCompletion, chat_options: ChatOptions) -> "ChatResponse":
"""Create a chat message content object from a choice."""
# additional properties
additional_options = {
key: value for key, value in chat_options.additional_properties.items() if value is not None
}
if additional_options:
run_options.update(additional_options)
return run_options
def _parse_response_from_openai(self, response: ChatCompletion, chat_options: ChatOptions) -> "ChatResponse":
"""Parse a response from OpenAI into a ChatResponse."""
response_metadata = self._get_metadata_from_chat_response(response)
messages: list[ChatMessage] = []
finish_reason: FinishReason | None = None
@@ -210,15 +230,15 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
if choice.finish_reason:
finish_reason = FinishReason(value=choice.finish_reason)
contents: list[Contents] = []
if text_content := self._parse_text_from_choice(choice):
if text_content := self._parse_text_from_openai(choice):
contents.append(text_content)
if parsed_tool_calls := [tool for tool in self._get_tool_calls_from_chat_choice(choice)]:
if parsed_tool_calls := [tool for tool in self._parse_tool_calls_from_openai(choice)]:
contents.extend(parsed_tool_calls)
messages.append(ChatMessage(role="assistant", contents=contents))
return ChatResponse(
response_id=response.id,
created_at=datetime.fromtimestamp(response.created, tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
usage_details=self._usage_details_from_openai(response.usage) if response.usage else None,
usage_details=self._parse_usage_from_openai(response.usage) if response.usage else None,
messages=messages,
model_id=response.model,
additional_properties=response_metadata,
@@ -226,16 +246,16 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
response_format=chat_options.response_format,
)
def _create_chat_response_update(
def _parse_response_update_from_openai(
self,
chunk: ChatCompletionChunk,
) -> ChatResponseUpdate:
"""Create a streaming chat message content object from a choice."""
"""Parse a streaming response update from OpenAI."""
chunk_metadata = self._get_metadata_from_streaming_chat_response(chunk)
if chunk.usage:
return ChatResponseUpdate(
role=Role.ASSISTANT,
contents=[UsageContent(details=self._usage_details_from_openai(chunk.usage), raw_representation=chunk)],
contents=[UsageContent(details=self._parse_usage_from_openai(chunk.usage), raw_representation=chunk)],
model_id=chunk.model,
additional_properties=chunk_metadata,
response_id=chunk.id,
@@ -245,11 +265,11 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
finish_reason: FinishReason | None = None
for choice in chunk.choices:
chunk_metadata.update(self._get_metadata_from_chat_choice(choice))
contents.extend(self._get_tool_calls_from_chat_choice(choice))
contents.extend(self._parse_tool_calls_from_openai(choice))
if choice.finish_reason:
finish_reason = FinishReason(value=choice.finish_reason)
if text_content := self._parse_text_from_choice(choice):
if text_content := self._parse_text_from_openai(choice):
contents.append(text_content)
return ChatResponseUpdate(
created_at=datetime.fromtimestamp(chunk.created, tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%S.%fZ"),
@@ -263,7 +283,7 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
message_id=chunk.id,
)
def _usage_details_from_openai(self, usage: CompletionUsage) -> UsageDetails:
def _parse_usage_from_openai(self, usage: CompletionUsage) -> UsageDetails:
details = UsageDetails(
input_token_count=usage.prompt_tokens,
output_token_count=usage.completion_tokens,
@@ -285,7 +305,7 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
details["prompt/cached_tokens"] = tokens
return details
def _parse_text_from_choice(self, choice: Choice | ChunkChoice) -> TextContent | None:
def _parse_text_from_openai(self, choice: Choice | ChunkChoice) -> TextContent | None:
"""Parse the choice into a TextContent object."""
message = choice.message if isinstance(choice, Choice) else choice.delta
if message.content:
@@ -312,8 +332,8 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
"logprobs": getattr(choice, "logprobs", None),
}
def _get_tool_calls_from_chat_choice(self, choice: Choice | ChunkChoice) -> list[Contents]:
"""Get tool calls from a chat choice."""
def _parse_tool_calls_from_openai(self, choice: Choice | ChunkChoice) -> list[Contents]:
"""Parse tool calls from an OpenAI response choice."""
resp: list[Contents] = []
content = choice.message if isinstance(choice, Choice) else choice.delta
if content and content.tool_calls:
@@ -331,13 +351,13 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
# When you enable asynchronous content filtering in Azure OpenAI, you may receive empty deltas
return resp
def _prepare_chat_history_for_request(
def _prepare_messages_for_openai(
self,
chat_messages: Sequence[ChatMessage],
role_key: str = "role",
content_key: str = "content",
) -> list[dict[str, Any]]:
"""Prepare the chat history for a request.
"""Prepare the chat history for an OpenAI request.
Allowing customization of the key names for role/author, and optionally overriding the role.
@@ -355,14 +375,14 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
Returns:
prepared_chat_history (Any): The prepared chat history for a request.
"""
list_of_list = [self._openai_chat_message_parser(message) for message in chat_messages]
list_of_list = [self._prepare_message_for_openai(message) for message in chat_messages]
# Flatten the list of lists into a single list
return list(chain.from_iterable(list_of_list))
# region Parsers
def _openai_chat_message_parser(self, message: ChatMessage) -> list[dict[str, Any]]:
"""Parse a chat message into the openai format."""
def _prepare_message_for_openai(self, message: ChatMessage) -> list[dict[str, Any]]:
"""Prepare a chat message for OpenAI."""
all_messages: list[dict[str, Any]] = []
for content in message.contents:
# Skip approval content - it's internal framework state, not for the LLM
@@ -372,13 +392,15 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
args: dict[str, Any] = {
"role": message.role.value if isinstance(message.role, Role) else message.role,
}
if message.author_name and message.role != Role.TOOL:
args["name"] = message.author_name
match content:
case FunctionCallContent():
if all_messages and "tool_calls" in all_messages[-1]:
# If the last message already has tool calls, append to it
all_messages[-1]["tool_calls"].append(self._openai_content_parser(content))
all_messages[-1]["tool_calls"].append(self._prepare_content_for_openai(content))
else:
args["tool_calls"] = [self._openai_content_parser(content)] # type: ignore
args["tool_calls"] = [self._prepare_content_for_openai(content)] # type: ignore
case FunctionResultContent():
args["tool_call_id"] = content.call_id
if content.result is not None:
@@ -387,13 +409,13 @@ class OpenAIBaseChatClient(OpenAIBase, BaseChatClient):
if "content" not in args:
args["content"] = []
# this is a list to allow multi-modal content
args["content"].append(self._openai_content_parser(content)) # type: ignore
args["content"].append(self._prepare_content_for_openai(content)) # type: ignore
if "content" in args or "tool_calls" in args:
all_messages.append(args)
return all_messages
def _openai_content_parser(self, content: Contents) -> dict[str, Any]:
"""Parse contents into the openai format."""
def _prepare_content_for_openai(self, content: Contents) -> dict[str, Any]:
"""Prepare content for OpenAI."""
match content:
case FunctionCallContent():
args = json.dumps(content.arguments) if isinstance(content.arguments, Mapping) else content.arguments
@@ -89,28 +89,16 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> ChatResponse:
client = await self.ensure_client()
run_options = await self.prepare_options(messages, chat_options, **kwargs)
response_format = run_options.pop("response_format", None)
text_config = run_options.pop("text", None)
text_format, text_config = self._prepare_text_config(response_format=response_format, text_config=text_config)
if text_config:
run_options["text"] = text_config
client = await self._ensure_client()
# prepare
run_options = await self._prepare_options(messages, chat_options, **kwargs)
try:
if not text_format:
response = await client.responses.create(
stream=False,
**run_options,
)
chat_options.conversation_id = self.get_conversation_id(response, chat_options.store)
return self._create_response_content(response, chat_options=chat_options)
parsed_response: ParsedResponse[BaseModel] = await client.responses.parse(
text_format=text_format,
stream=False,
**run_options,
)
chat_options.conversation_id = self.get_conversation_id(parsed_response, chat_options.store)
return self._create_response_content(parsed_response, chat_options=chat_options)
# execute and process
if "text_format" in run_options:
response = await client.responses.parse(stream=False, **run_options)
else:
response = await client.responses.create(stream=False, **run_options)
return self._parse_response_from_openai(response, chat_options=chat_options)
except BadRequestError as ex:
if ex.code == "content_filter":
raise OpenAIContentFilterException(
@@ -134,35 +122,23 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> AsyncIterable[ChatResponseUpdate]:
client = await self.ensure_client()
run_options = await self.prepare_options(messages, chat_options, **kwargs)
client = await self._ensure_client()
# prepare
run_options = await self._prepare_options(messages, chat_options, **kwargs)
function_call_ids: dict[int, tuple[str, str]] = {} # output_index: (call_id, name)
response_format = run_options.pop("response_format", None)
text_config = run_options.pop("text", None)
text_format, text_config = self._prepare_text_config(response_format=response_format, text_config=text_config)
if text_config:
run_options["text"] = text_config
try:
if not text_format:
response = await client.responses.create(
stream=True,
**run_options,
)
async for chunk in response:
update = self._create_streaming_response_content(
# execute and process
if "text_format" not in run_options:
async for chunk in await client.responses.create(stream=True, **run_options):
yield self._parse_chunk_from_openai(
chunk, chat_options=chat_options, function_call_ids=function_call_ids
)
yield update
return
async with client.responses.stream(
text_format=text_format,
**run_options,
) as response:
async with client.responses.stream(**run_options) as response:
async for chunk in response:
update = self._create_streaming_response_content(
yield self._parse_chunk_from_openai(
chunk, chat_options=chat_options, function_call_ids=function_call_ids
)
yield update
except BadRequestError as ex:
if ex.code == "content_filter":
raise OpenAIContentFilterException(
@@ -179,33 +155,33 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
inner_exception=ex,
) from ex
def _prepare_text_config(
def _prepare_response_and_text_format(
self,
*,
response_format: Any,
text_config: MutableMapping[str, Any] | None,
) -> tuple[type[BaseModel] | None, dict[str, Any] | None]:
"""Normalize response_format into Responses text configuration and parse target."""
prepared_text = dict(text_config) if isinstance(text_config, MutableMapping) else None
if text_config is not None and not isinstance(text_config, MutableMapping):
raise ServiceInvalidRequestError("text must be a mapping when provided.")
text_config = cast(dict[str, Any], text_config) if isinstance(text_config, MutableMapping) else None
if response_format is None:
return None, prepared_text
return None, text_config
if isinstance(response_format, type) and issubclass(response_format, BaseModel):
if prepared_text and "format" in prepared_text:
if text_config and "format" in text_config:
raise ServiceInvalidRequestError("response_format cannot be combined with explicit text.format.")
return response_format, prepared_text
return response_format, text_config
if isinstance(response_format, Mapping):
format_config = self._convert_response_format(cast("Mapping[str, Any]", response_format))
if prepared_text is None:
prepared_text = {}
elif "format" in prepared_text and prepared_text["format"] != format_config:
if text_config is None:
text_config = {}
elif "format" in text_config and text_config["format"] != format_config:
raise ServiceInvalidRequestError("Conflicting response_format definitions detected.")
prepared_text["format"] = format_config
return None, prepared_text
text_config["format"] = format_config
return None, text_config
raise ServiceInvalidRequestError("response_format must be a Pydantic model or mapping.")
@@ -245,23 +221,33 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
raise ServiceInvalidRequestError("Unsupported response_format provided for Responses client.")
def get_conversation_id(
def _get_conversation_id(
self, response: OpenAIResponse | ParsedResponse[BaseModel], store: bool | None
) -> str | None:
"""Get the conversation ID from the response if store is True."""
return None if store is False else response.id
if store is False:
return None
# If conversation ID exists, it means that we operate with conversation
# so we use conversation ID as input and output.
if response.conversation and response.conversation.id:
return response.conversation.id
# If conversation ID doesn't exist, we operate with responses
# so we use response ID as input and output.
return response.id
# region Prep methods
def _tools_to_response_tools(
self, tools: Sequence[ToolProtocol | MutableMapping[str, Any]]
def _prepare_tools_for_openai(
self, tools: Sequence[ToolProtocol | MutableMapping[str, Any]] | None
) -> list[ToolParam | dict[str, Any]]:
response_tools: list[ToolParam | dict[str, Any]] = []
if not tools:
return response_tools
for tool in tools:
if isinstance(tool, ToolProtocol):
match tool:
case HostedMCPTool():
response_tools.append(self.get_mcp_tool(tool))
response_tools.append(self._prepare_mcp_tool(tool))
case HostedCodeInterpreterTool():
tool_args: CodeInterpreterContainerCodeInterpreterToolAuto = {"type": "auto"}
if tool.inputs:
@@ -363,7 +349,8 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
response_tools.append(tool_dict)
return response_tools
def get_mcp_tool(self, tool: HostedMCPTool) -> Any:
@staticmethod
def _prepare_mcp_tool(tool: HostedMCPTool) -> Mcp:
"""Get MCP tool from HostedMCPTool."""
mcp: Mcp = {
"type": "mcp",
@@ -386,18 +373,13 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
return mcp
async def prepare_options(
async def _prepare_options(
self,
messages: MutableSequence[ChatMessage],
chat_options: ChatOptions,
**kwargs: Any,
) -> dict[str, Any]:
"""Take ChatOptions and create the specific options for Responses API."""
conversation_id = kwargs.pop("conversation_id", None)
if conversation_id:
chat_options.conversation_id = conversation_id
run_options: dict[str, Any] = chat_options.to_dict(
exclude={
"type",
@@ -407,12 +389,24 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
"seed", # not supported
"stop", # not supported
"instructions", # already added as system message
"response_format", # handled separately
"conversation_id", # handled separately
"additional_properties", # handled separately
}
)
# messages
request_input = self._prepare_messages_for_openai(messages)
if not request_input:
raise ServiceInvalidRequestError("Messages are required for chat completions")
run_options["input"] = request_input
if chat_options.response_format:
run_options["response_format"] = chat_options.response_format
# model id
if not run_options.get("model"):
if not self.model_id:
raise ValueError("model_id must be a non-empty string")
run_options["model"] = self.model_id
# translations between ChatOptions and Responses API
translations = {
"model_id": "model",
"allow_multiple_tool_calls": "parallel_tool_calls",
@@ -423,34 +417,53 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
if old_key in run_options and old_key != new_key:
run_options[new_key] = run_options.pop(old_key)
# Handle different conversation ID formats
if conversation_id := self._get_current_conversation_id(chat_options, **kwargs):
if conversation_id.startswith("resp_"):
# For response IDs, set previous_response_id and remove conversation property
run_options["previous_response_id"] = conversation_id
elif conversation_id.startswith("conv_"):
# For conversation IDs, set conversation and remove previous_response_id property
run_options["conversation"] = conversation_id
else:
# If the format is unrecognized, default to previous_response_id
run_options["previous_response_id"] = conversation_id
# tools
if chat_options.tools is None:
run_options.pop("parallel_tool_calls", None)
if tools := self._prepare_tools_for_openai(chat_options.tools):
run_options["tools"] = tools
else:
run_options["tools"] = self._tools_to_response_tools(chat_options.tools)
# model id
if not run_options.get("model"):
if not self.model_id:
raise ValueError("model_id must be a non-empty string")
run_options["model"] = self.model_id
# messages
request_input = self._prepare_chat_messages_for_request(messages)
if not request_input:
raise ServiceInvalidRequestError("Messages are required for chat completions")
run_options["input"] = request_input
# additional provider specific settings
if additional_properties := run_options.pop("additional_properties", None):
for key, value in additional_properties.items():
if value is not None:
run_options[key] = value
run_options.pop("parallel_tool_calls", None)
run_options.pop("tool_choice", None)
# tool choice when `tool_choice` is a dict with single key `mode`, extract the mode value
if (tool_choice := run_options.get("tool_choice")) and len(tool_choice.keys()) == 1:
run_options["tool_choice"] = tool_choice["mode"]
# additional properties
additional_options = {
key: value for key, value in chat_options.additional_properties.items() if value is not None
}
if additional_options:
run_options.update(additional_options)
# response format and text config (after additional_properties so user can pass text via additional_properties)
response_format = chat_options.response_format
text_config = run_options.pop("text", None)
response_format, text_config = self._prepare_response_and_text_format(
response_format=response_format, text_config=text_config
)
if text_config:
run_options["text"] = text_config
if response_format:
run_options["text_format"] = response_format
return run_options
def _prepare_chat_messages_for_request(self, chat_messages: Sequence[ChatMessage]) -> list[dict[str, Any]]:
def _get_current_conversation_id(self, chat_options: ChatOptions, **kwargs: Any) -> str | None:
"""Get the current conversation ID from chat options or kwargs."""
return chat_options.conversation_id or kwargs.get("conversation_id")
def _prepare_messages_for_openai(self, chat_messages: Sequence[ChatMessage]) -> list[dict[str, Any]]:
"""Prepare the chat messages for a request.
Allowing customization of the key names for role/author, and optionally overriding the role.
@@ -476,16 +489,16 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
and "fc_id" in content.additional_properties
):
call_id_to_id[content.call_id] = content.additional_properties["fc_id"]
list_of_list = [self._openai_chat_message_parser(message, call_id_to_id) for message in chat_messages]
list_of_list = [self._prepare_message_for_openai(message, call_id_to_id) for message in chat_messages]
# Flatten the list of lists into a single list
return list(chain.from_iterable(list_of_list))
def _openai_chat_message_parser(
def _prepare_message_for_openai(
self,
message: ChatMessage,
call_id_to_id: dict[str, str],
) -> list[dict[str, Any]]:
"""Parse a chat message into the openai format."""
"""Prepare a chat message for the OpenAI Responses API format."""
all_messages: list[dict[str, Any]] = []
args: dict[str, Any] = {
"role": message.role.value if isinstance(message.role, Role) else message.role,
@@ -497,28 +510,28 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
continue
case FunctionResultContent():
new_args: dict[str, Any] = {}
new_args.update(self._openai_content_parser(message.role, content, call_id_to_id))
new_args.update(self._prepare_content_for_openai(message.role, content, call_id_to_id))
all_messages.append(new_args)
case FunctionCallContent():
function_call = self._openai_content_parser(message.role, content, call_id_to_id)
function_call = self._prepare_content_for_openai(message.role, content, call_id_to_id)
all_messages.append(function_call) # type: ignore
case FunctionApprovalResponseContent() | FunctionApprovalRequestContent():
all_messages.append(self._openai_content_parser(message.role, content, call_id_to_id)) # type: ignore
all_messages.append(self._prepare_content_for_openai(message.role, content, call_id_to_id)) # type: ignore
case _:
if "content" not in args:
args["content"] = []
args["content"].append(self._openai_content_parser(message.role, content, call_id_to_id)) # type: ignore
args["content"].append(self._prepare_content_for_openai(message.role, content, call_id_to_id)) # type: ignore
if "content" in args or "tool_calls" in args:
all_messages.append(args)
return all_messages
def _openai_content_parser(
def _prepare_content_for_openai(
self,
role: Role,
content: Contents,
call_id_to_id: dict[str, str],
) -> dict[str, Any]:
"""Parse contents into the openai format."""
"""Prepare content for the OpenAI Responses API format."""
match content:
case TextContent():
return {
@@ -625,14 +638,13 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
logger.debug("Unsupported content type passed (type: %s)", type(content))
return {}
# region Response creation methods
def _create_response_content(
# region Parse methods
def _parse_response_from_openai(
self,
response: OpenAIResponse | ParsedResponse[BaseModel],
chat_options: ChatOptions,
) -> "ChatResponse":
"""Create a chat message content object from a choice."""
"""Parse an OpenAI Responses API response into a ChatResponse."""
structured_response: BaseModel | None = response.output_parsed if isinstance(response, ParsedResponse) else None # type: ignore[reportUnknownMemberType]
metadata: dict[str, Any] = response.metadata or {}
@@ -826,11 +838,9 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
"raw_representation": response,
}
conversation_id = self.get_conversation_id(response, chat_options.store) # type: ignore[reportArgumentType]
if conversation_id:
if conversation_id := self._get_conversation_id(response, chat_options.store):
args["conversation_id"] = conversation_id
if response.usage and (usage_details := self._usage_details_from_openai(response.usage)):
if response.usage and (usage_details := self._parse_usage_from_openai(response.usage)):
args["usage_details"] = usage_details
if structured_response:
args["value"] = structured_response
@@ -838,13 +848,13 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
args["response_format"] = chat_options.response_format
return ChatResponse(**args)
def _create_streaming_response_content(
def _parse_chunk_from_openai(
self,
event: OpenAIResponseStreamEvent,
chat_options: ChatOptions,
function_call_ids: dict[int, tuple[str, str]],
) -> ChatResponseUpdate:
"""Create a streaming chat message content object from a choice."""
"""Parse an OpenAI Responses API streaming event into a ChatResponseUpdate."""
metadata: dict[str, Any] = {}
contents: list[Contents] = []
conversation_id: str | None = None
@@ -931,10 +941,10 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
contents.append(TextReasoningContent(text=event.text, raw_representation=event))
metadata.update(self._get_metadata_from_response(event))
case "response.completed":
conversation_id = self.get_conversation_id(event.response, chat_options.store)
conversation_id = self._get_conversation_id(event.response, chat_options.store)
model = event.response.model
if event.response.usage:
usage = self._usage_details_from_openai(event.response.usage)
usage = self._parse_usage_from_openai(event.response.usage)
if usage:
contents.append(UsageContent(details=usage, raw_representation=event))
case "response.output_item.added":
@@ -1102,7 +1112,7 @@ class OpenAIBaseResponsesClient(OpenAIBase, BaseChatClient):
raw_representation=event,
)
def _usage_details_from_openai(self, usage: ResponseUsage) -> UsageDetails | None:
def _parse_usage_from_openai(self, usage: ResponseUsage) -> UsageDetails | None:
details = UsageDetails(
input_token_count=usage.input_tokens,
output_token_count=usage.output_tokens,
@@ -160,16 +160,16 @@ class OpenAIBase(SerializationMixin):
for key, value in kwargs.items():
setattr(self, key, value)
async def initialize_client(self) -> None:
async def _initialize_client(self) -> None:
"""Initialize OpenAI client asynchronously.
Override in subclasses to initialize the OpenAI client asynchronously.
"""
pass
async def ensure_client(self) -> AsyncOpenAI:
async def _ensure_client(self) -> AsyncOpenAI:
"""Ensure OpenAI client is initialized."""
await self.initialize_client()
await self._initialize_client()
if self.client is None:
raise ServiceInitializationError("OpenAI client is not initialized")
+1
View File
@@ -52,6 +52,7 @@ all = [
"agent-framework-devui",
"agent-framework-lab",
"agent-framework-mem0",
"agent-framework-ollama",
"agent-framework-purview",
"agent-framework-redis",
]
@@ -193,7 +193,7 @@ async def test_cmc(
mock_create.assert_awaited_once_with(
model=azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"],
stream=False,
messages=azure_chat_client._prepare_chat_history_for_request(chat_history), # type: ignore
messages=azure_chat_client._prepare_messages_for_openai(chat_history), # type: ignore
)
@@ -216,7 +216,7 @@ async def test_cmc_with_logit_bias(
mock_create.assert_awaited_once_with(
model=azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"],
messages=azure_chat_client._prepare_chat_history_for_request(chat_history), # type: ignore
messages=azure_chat_client._prepare_messages_for_openai(chat_history), # type: ignore
stream=False,
logit_bias=token_bias,
)
@@ -241,7 +241,7 @@ async def test_cmc_with_stop(
mock_create.assert_awaited_once_with(
model=azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"],
messages=azure_chat_client._prepare_chat_history_for_request(chat_history), # type: ignore
messages=azure_chat_client._prepare_messages_for_openai(chat_history), # type: ignore
stream=False,
stop=stop,
)
@@ -311,7 +311,7 @@ async def test_azure_on_your_data(
mock_create.assert_awaited_once_with(
model=azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"],
messages=azure_chat_client._prepare_chat_history_for_request(messages_out), # type: ignore
messages=azure_chat_client._prepare_messages_for_openai(messages_out), # type: ignore
stream=False,
extra_body=expected_data_settings,
)
@@ -381,7 +381,7 @@ async def test_azure_on_your_data_string(
mock_create.assert_awaited_once_with(
model=azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"],
messages=azure_chat_client._prepare_chat_history_for_request(messages_out), # type: ignore
messages=azure_chat_client._prepare_messages_for_openai(messages_out), # type: ignore
stream=False,
extra_body=expected_data_settings,
)
@@ -438,7 +438,7 @@ async def test_azure_on_your_data_fail(
mock_create.assert_awaited_once_with(
model=azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"],
messages=azure_chat_client._prepare_chat_history_for_request(messages_out), # type: ignore
messages=azure_chat_client._prepare_messages_for_openai(messages_out), # type: ignore
stream=False,
extra_body=expected_data_settings,
)
@@ -584,7 +584,7 @@ async def test_get_streaming(
mock_create.assert_awaited_once_with(
model=azure_openai_unit_test_env["AZURE_OPENAI_CHAT_DEPLOYMENT_NAME"],
stream=True,
messages=azure_chat_client._prepare_chat_history_for_request(chat_history), # type: ignore
messages=azure_chat_client._prepare_messages_for_openai(chat_history), # type: ignore
# NOTE: The `stream_options={"include_usage": True}` is explicitly enforced in
# `OpenAIChatCompletionBase._inner_get_streaming_response`.
# To ensure consistency, we align the arguments here accordingly.
+36 -31
View File
@@ -24,14 +24,14 @@ from agent_framework import (
)
from agent_framework._mcp import (
MCPTool,
_ai_content_to_mcp_types,
_chat_message_to_mcp_types,
_get_input_model_from_mcp_prompt,
_get_input_model_from_mcp_tool,
_mcp_call_tool_result_to_ai_contents,
_mcp_prompt_message_to_chat_message,
_mcp_type_to_ai_content,
_normalize_mcp_name,
_parse_content_from_mcp,
_parse_contents_from_mcp_tool_result,
_parse_message_from_mcp,
_prepare_content_for_mcp,
_prepare_message_for_mcp,
)
from agent_framework.exceptions import ToolException, ToolExecutionException
@@ -60,7 +60,7 @@ def test_normalize_mcp_name():
def test_mcp_prompt_message_to_ai_content():
"""Test conversion from MCP prompt message to AI content."""
mcp_message = types.PromptMessage(role="user", content=types.TextContent(type="text", text="Hello, world!"))
ai_content = _mcp_prompt_message_to_chat_message(mcp_message)
ai_content = _parse_message_from_mcp(mcp_message)
assert isinstance(ai_content, ChatMessage)
assert ai_content.role.value == "user"
@@ -70,22 +70,26 @@ def test_mcp_prompt_message_to_ai_content():
assert ai_content.raw_representation == mcp_message
def test_mcp_call_tool_result_to_ai_contents():
def test_parse_contents_from_mcp_tool_result():
"""Test conversion from MCP tool result to AI contents."""
mcp_result = types.CallToolResult(
content=[
types.TextContent(type="text", text="Result text"),
types.ImageContent(type="image", data="data:image/png;base64,xyz", mimeType="image/png"),
types.ImageContent(type="image", data="xyz", mimeType="image/png"),
types.ImageContent(type="image", data=b"abc", mimeType="image/webp"),
]
)
ai_contents = _mcp_call_tool_result_to_ai_contents(mcp_result)
ai_contents = _parse_contents_from_mcp_tool_result(mcp_result)
assert len(ai_contents) == 2
assert len(ai_contents) == 3
assert isinstance(ai_contents[0], TextContent)
assert ai_contents[0].text == "Result text"
assert isinstance(ai_contents[1], DataContent)
assert ai_contents[1].uri == "data:image/png;base64,xyz"
assert ai_contents[1].media_type == "image/png"
assert isinstance(ai_contents[2], DataContent)
assert ai_contents[2].uri == "data:image/webp;base64,abc"
assert ai_contents[2].media_type == "image/webp"
def test_mcp_call_tool_result_with_meta_error():
@@ -96,7 +100,7 @@ def test_mcp_call_tool_result_with_meta_error():
_meta={"isError": True, "errorCode": "TOOL_ERROR", "errorMessage": "Tool execution failed"},
)
ai_contents = _mcp_call_tool_result_to_ai_contents(mcp_result)
ai_contents = _parse_contents_from_mcp_tool_result(mcp_result)
assert len(ai_contents) == 1
assert isinstance(ai_contents[0], TextContent)
@@ -127,7 +131,7 @@ def test_mcp_call_tool_result_with_meta_arbitrary_data():
},
)
ai_contents = _mcp_call_tool_result_to_ai_contents(mcp_result)
ai_contents = _parse_contents_from_mcp_tool_result(mcp_result)
assert len(ai_contents) == 1
assert isinstance(ai_contents[0], TextContent)
@@ -149,7 +153,7 @@ def test_mcp_call_tool_result_with_meta_merging_existing_properties():
text_content = types.TextContent(type="text", text="Test content")
mcp_result = types.CallToolResult(content=[text_content], _meta={"newField": "newValue", "isError": False})
ai_contents = _mcp_call_tool_result_to_ai_contents(mcp_result)
ai_contents = _parse_contents_from_mcp_tool_result(mcp_result)
assert len(ai_contents) == 1
content = ai_contents[0]
@@ -165,7 +169,7 @@ def test_mcp_call_tool_result_with_meta_none():
mcp_result = types.CallToolResult(content=[types.TextContent(type="text", text="No meta test")])
# No _meta field set
ai_contents = _mcp_call_tool_result_to_ai_contents(mcp_result)
ai_contents = _parse_contents_from_mcp_tool_result(mcp_result)
assert len(ai_contents) == 1
assert isinstance(ai_contents[0], TextContent)
@@ -183,11 +187,11 @@ def test_mcp_call_tool_result_regression_successful_workflow():
mcp_result = types.CallToolResult(
content=[
types.TextContent(type="text", text="Success message"),
types.ImageContent(type="image", data="data:image/jpeg;base64,abc123", mimeType="image/jpeg"),
types.ImageContent(type="image", data="abc123", mimeType="image/jpeg"),
]
)
ai_contents = _mcp_call_tool_result_to_ai_contents(mcp_result)
ai_contents = _parse_contents_from_mcp_tool_result(mcp_result)
# Verify basic conversion still works correctly
assert len(ai_contents) == 2
@@ -209,7 +213,7 @@ def test_mcp_call_tool_result_regression_successful_workflow():
def test_mcp_content_types_to_ai_content_text():
"""Test conversion of MCP text content to AI content."""
mcp_content = types.TextContent(type="text", text="Sample text")
ai_content = _mcp_type_to_ai_content(mcp_content)[0]
ai_content = _parse_content_from_mcp(mcp_content)[0]
assert isinstance(ai_content, TextContent)
assert ai_content.text == "Sample text"
@@ -218,8 +222,9 @@ def test_mcp_content_types_to_ai_content_text():
def test_mcp_content_types_to_ai_content_image():
"""Test conversion of MCP image content to AI content."""
mcp_content = types.ImageContent(type="image", data="data:image/jpeg;base64,abc", mimeType="image/jpeg")
ai_content = _mcp_type_to_ai_content(mcp_content)[0]
mcp_content = types.ImageContent(type="image", data="abc", mimeType="image/jpeg")
mcp_content = types.ImageContent(type="image", data=b"abc", mimeType="image/jpeg")
ai_content = _parse_content_from_mcp(mcp_content)[0]
assert isinstance(ai_content, DataContent)
assert ai_content.uri == "data:image/jpeg;base64,abc"
@@ -229,8 +234,8 @@ def test_mcp_content_types_to_ai_content_image():
def test_mcp_content_types_to_ai_content_audio():
"""Test conversion of MCP audio content to AI content."""
mcp_content = types.AudioContent(type="audio", data="data:audio/wav;base64,def", mimeType="audio/wav")
ai_content = _mcp_type_to_ai_content(mcp_content)[0]
mcp_content = types.AudioContent(type="audio", data="def", mimeType="audio/wav")
ai_content = _parse_content_from_mcp(mcp_content)[0]
assert isinstance(ai_content, DataContent)
assert ai_content.uri == "data:audio/wav;base64,def"
@@ -246,7 +251,7 @@ def test_mcp_content_types_to_ai_content_resource_link():
name="test_resource",
mimeType="application/json",
)
ai_content = _mcp_type_to_ai_content(mcp_content)[0]
ai_content = _parse_content_from_mcp(mcp_content)[0]
assert isinstance(ai_content, UriContent)
assert ai_content.uri == "https://example.com/resource"
@@ -262,7 +267,7 @@ def test_mcp_content_types_to_ai_content_embedded_resource_text():
text="Embedded text content",
)
mcp_content = types.EmbeddedResource(type="resource", resource=text_resource)
ai_content = _mcp_type_to_ai_content(mcp_content)[0]
ai_content = _parse_content_from_mcp(mcp_content)[0]
assert isinstance(ai_content, TextContent)
assert ai_content.text == "Embedded text content"
@@ -278,7 +283,7 @@ def test_mcp_content_types_to_ai_content_embedded_resource_blob():
blob="data:application/octet-stream;base64,dGVzdCBkYXRh",
)
mcp_content = types.EmbeddedResource(type="resource", resource=blob_resource)
ai_content = _mcp_type_to_ai_content(mcp_content)[0]
ai_content = _parse_content_from_mcp(mcp_content)[0]
assert isinstance(ai_content, DataContent)
assert ai_content.uri == "data:application/octet-stream;base64,dGVzdCBkYXRh"
@@ -289,7 +294,7 @@ def test_mcp_content_types_to_ai_content_embedded_resource_blob():
def test_ai_content_to_mcp_content_types_text():
"""Test conversion of AI text content to MCP content."""
ai_content = TextContent(text="Sample text")
mcp_content = _ai_content_to_mcp_types(ai_content)
mcp_content = _prepare_content_for_mcp(ai_content)
assert isinstance(mcp_content, types.TextContent)
assert mcp_content.type == "text"
@@ -299,7 +304,7 @@ def test_ai_content_to_mcp_content_types_text():
def test_ai_content_to_mcp_content_types_data_image():
"""Test conversion of AI data content to MCP content."""
ai_content = DataContent(uri="data:image/png;base64,xyz", media_type="image/png")
mcp_content = _ai_content_to_mcp_types(ai_content)
mcp_content = _prepare_content_for_mcp(ai_content)
assert isinstance(mcp_content, types.ImageContent)
assert mcp_content.type == "image"
@@ -310,7 +315,7 @@ def test_ai_content_to_mcp_content_types_data_image():
def test_ai_content_to_mcp_content_types_data_audio():
"""Test conversion of AI data content to MCP content."""
ai_content = DataContent(uri="data:audio/mpeg;base64,xyz", media_type="audio/mpeg")
mcp_content = _ai_content_to_mcp_types(ai_content)
mcp_content = _prepare_content_for_mcp(ai_content)
assert isinstance(mcp_content, types.AudioContent)
assert mcp_content.type == "audio"
@@ -324,7 +329,7 @@ def test_ai_content_to_mcp_content_types_data_binary():
uri="data:application/octet-stream;base64,xyz",
media_type="application/octet-stream",
)
mcp_content = _ai_content_to_mcp_types(ai_content)
mcp_content = _prepare_content_for_mcp(ai_content)
assert isinstance(mcp_content, types.EmbeddedResource)
assert mcp_content.type == "resource"
@@ -335,7 +340,7 @@ def test_ai_content_to_mcp_content_types_data_binary():
def test_ai_content_to_mcp_content_types_uri():
"""Test conversion of AI URI content to MCP content."""
ai_content = UriContent(uri="https://example.com/resource", media_type="application/json")
mcp_content = _ai_content_to_mcp_types(ai_content)
mcp_content = _prepare_content_for_mcp(ai_content)
assert isinstance(mcp_content, types.ResourceLink)
assert mcp_content.type == "resource_link"
@@ -343,7 +348,7 @@ def test_ai_content_to_mcp_content_types_uri():
assert mcp_content.mimeType == "application/json"
def test_chat_message_to_mcp_types():
def test_prepare_message_for_mcp():
message = ChatMessage(
role="user",
contents=[
@@ -351,7 +356,7 @@ def test_chat_message_to_mcp_types():
DataContent(uri="data:image/png;base64,xyz", media_type="image/png"),
],
)
mcp_contents = _chat_message_to_mcp_types(message)
mcp_contents = _prepare_message_for_mcp(message)
assert len(mcp_contents) == 2
assert isinstance(mcp_contents[0], types.TextContent)
assert isinstance(mcp_contents[1], types.ImageContent)
+158 -2
View File
@@ -1,5 +1,5 @@
# Copyright (c) Microsoft. All rights reserved.
from typing import Any
from typing import Annotated, Any, Literal
from unittest.mock import Mock
import pytest
@@ -14,7 +14,7 @@ from agent_framework import (
ToolProtocol,
ai_function,
)
from agent_framework._tools import _parse_inputs
from agent_framework._tools import _parse_annotation, _parse_inputs
from agent_framework.exceptions import ToolException
from agent_framework.observability import OtelAttr
@@ -128,6 +128,95 @@ def test_ai_function_decorator_in_class():
assert test_tool(1, 2) == 3
def test_ai_function_with_literal_type_parameter():
"""Test ai_function decorator with Literal type parameter (issue #2891)."""
@ai_function
def search_flows(category: Literal["Data", "Security", "Network"], issue: str) -> str:
"""Search flows by category."""
return f"{category}: {issue}"
assert isinstance(search_flows, AIFunction)
schema = search_flows.parameters()
assert schema == {
"properties": {
"category": {"enum": ["Data", "Security", "Network"], "title": "Category", "type": "string"},
"issue": {"title": "Issue", "type": "string"},
},
"required": ["category", "issue"],
"title": "search_flows_input",
"type": "object",
}
# Verify invocation works
assert search_flows("Data", "test issue") == "Data: test issue"
def test_ai_function_with_literal_type_in_class_method():
"""Test ai_function decorator with Literal type parameter in a class method (issue #2891)."""
class MyTools:
@ai_function
def search_flows(self, category: Literal["Data", "Security", "Network"], issue: str) -> str:
"""Search flows by category."""
return f"{category}: {issue}"
tools = MyTools()
search_tool = tools.search_flows
assert isinstance(search_tool, AIFunction)
schema = search_tool.parameters()
assert schema == {
"properties": {
"category": {"enum": ["Data", "Security", "Network"], "title": "Category", "type": "string"},
"issue": {"title": "Issue", "type": "string"},
},
"required": ["category", "issue"],
"title": "search_flows_input",
"type": "object",
}
# Verify invocation works
assert search_tool("Security", "test issue") == "Security: test issue"
def test_ai_function_with_literal_int_type():
"""Test ai_function decorator with Literal int type parameter."""
@ai_function
def set_priority(priority: Literal[1, 2, 3], task: str) -> str:
"""Set priority for a task."""
return f"Priority {priority}: {task}"
assert isinstance(set_priority, AIFunction)
schema = set_priority.parameters()
assert schema == {
"properties": {
"priority": {"enum": [1, 2, 3], "title": "Priority", "type": "integer"},
"task": {"title": "Task", "type": "string"},
},
"required": ["priority", "task"],
"title": "set_priority_input",
"type": "object",
}
assert set_priority(1, "important task") == "Priority 1: important task"
def test_ai_function_with_literal_and_annotated():
"""Test ai_function decorator with Literal type combined with Annotated for description."""
@ai_function
def categorize(
category: Annotated[Literal["A", "B", "C"], "The category to assign"],
name: str,
) -> str:
"""Categorize an item."""
return f"{category}: {name}"
assert isinstance(categorize, AIFunction)
schema = categorize.parameters()
# Literal type inside Annotated should preserve enum values
assert schema["properties"]["category"]["enum"] == ["A", "B", "C"]
assert categorize("A", "test") == "A: test"
async def test_ai_function_decorator_shared_state():
"""Test that decorated methods maintain shared state across multiple calls and tool usage."""
@@ -1368,3 +1457,70 @@ async def test_ai_function_with_kwargs_injection():
arguments=tool_with_kwargs.input_model(x=10),
)
assert result_default == "x=10, user=unknown"
# region _parse_annotation tests
def test_parse_annotation_with_literal_type():
"""Test that _parse_annotation returns Literal types unchanged (issue #2891)."""
from typing import get_args, get_origin
# Literal with string values
literal_annotation = Literal["Data", "Security", "Network"]
result = _parse_annotation(literal_annotation)
assert result is literal_annotation
assert get_origin(result) is Literal
assert get_args(result) == ("Data", "Security", "Network")
def test_parse_annotation_with_literal_int_type():
"""Test that _parse_annotation returns Literal int types unchanged."""
from typing import get_args, get_origin
literal_annotation = Literal[1, 2, 3]
result = _parse_annotation(literal_annotation)
assert result is literal_annotation
assert get_origin(result) is Literal
assert get_args(result) == (1, 2, 3)
def test_parse_annotation_with_literal_bool_type():
"""Test that _parse_annotation returns Literal bool types unchanged."""
from typing import get_args, get_origin
literal_annotation = Literal[True, False]
result = _parse_annotation(literal_annotation)
assert result is literal_annotation
assert get_origin(result) is Literal
assert get_args(result) == (True, False)
def test_parse_annotation_with_simple_types():
"""Test that _parse_annotation returns simple types unchanged."""
assert _parse_annotation(str) is str
assert _parse_annotation(int) is int
assert _parse_annotation(float) is float
assert _parse_annotation(bool) is bool
def test_parse_annotation_with_annotated_and_literal():
"""Test that Annotated[Literal[...], description] works correctly."""
from typing import get_args, get_origin
# When Literal is inside Annotated, it should still be preserved
annotated_literal = Annotated[Literal["A", "B", "C"], "The category"]
result = _parse_annotation(annotated_literal)
# The Annotated type should be preserved
origin = get_origin(result)
assert origin is Annotated
args = get_args(result)
# First arg is the Literal type
literal_type = args[0]
assert get_origin(literal_type) is Literal
assert get_args(literal_type) == ("A", "B", "C")
# endregion
@@ -463,9 +463,9 @@ async def test_openai_assistants_client_process_stream_events_requires_action(mo
"""Test _process_stream_events with thread.run.requires_action event."""
chat_client = create_test_openai_assistants_client(mock_async_openai)
# Mock the _create_function_call_contents method to return test content
# Mock the _parse_function_calls_from_assistants method to return test content
test_function_content = FunctionCallContent(call_id="call-123", name="test_func", arguments={"arg": "value"})
chat_client._create_function_call_contents = MagicMock(return_value=[test_function_content]) # type: ignore
chat_client._parse_function_calls_from_assistants = MagicMock(return_value=[test_function_content]) # type: ignore
# Create a mock Run object
mock_run = MagicMock(spec=Run)
@@ -498,8 +498,8 @@ async def test_openai_assistants_client_process_stream_events_requires_action(mo
assert update.contents[0] == test_function_content
assert update.raw_representation == mock_run
# Verify _create_function_call_contents was called correctly
chat_client._create_function_call_contents.assert_called_once_with(mock_run, None) # type: ignore
# Verify _parse_function_calls_from_assistants was called correctly
chat_client._parse_function_calls_from_assistants.assert_called_once_with(mock_run, None) # type: ignore
async def test_openai_assistants_client_process_stream_events_run_step_created(mock_async_openai: MagicMock) -> None:
@@ -585,8 +585,8 @@ async def test_openai_assistants_client_process_stream_events_run_completed_with
assert update.raw_representation == mock_run
def test_openai_assistants_client_create_function_call_contents_basic(mock_async_openai: MagicMock) -> None:
"""Test _create_function_call_contents with a simple function call."""
def test_openai_assistants_client_parse_function_calls_from_assistants_basic(mock_async_openai: MagicMock) -> None:
"""Test _parse_function_calls_from_assistants with a simple function call."""
chat_client = create_test_openai_assistants_client(mock_async_openai)
@@ -605,7 +605,7 @@ def test_openai_assistants_client_create_function_call_contents_basic(mock_async
# Call the method
response_id = "response_456"
contents = chat_client._create_function_call_contents(mock_run, response_id) # type: ignore
contents = chat_client._parse_function_calls_from_assistants(mock_run, response_id) # type: ignore
# Test that one function call content was created
assert len(contents) == 1
@@ -825,24 +825,24 @@ def test_openai_assistants_client_prepare_options_with_image_content(mock_async_
assert message["content"][0]["image_url"]["url"] == "https://example.com/image.jpg"
def test_openai_assistants_client_convert_function_results_to_tool_output_empty(mock_async_openai: MagicMock) -> None:
"""Test _convert_function_results_to_tool_output with empty list."""
def test_openai_assistants_client_prepare_tool_outputs_for_assistants_empty(mock_async_openai: MagicMock) -> None:
"""Test _prepare_tool_outputs_for_assistants with empty list."""
chat_client = create_test_openai_assistants_client(mock_async_openai)
run_id, tool_outputs = chat_client._convert_function_results_to_tool_output([]) # type: ignore
run_id, tool_outputs = chat_client._prepare_tool_outputs_for_assistants([]) # type: ignore
assert run_id is None
assert tool_outputs is None
def test_openai_assistants_client_convert_function_results_to_tool_output_valid(mock_async_openai: MagicMock) -> None:
"""Test _convert_function_results_to_tool_output with valid function results."""
def test_openai_assistants_client_prepare_tool_outputs_for_assistants_valid(mock_async_openai: MagicMock) -> None:
"""Test _prepare_tool_outputs_for_assistants with valid function results."""
chat_client = create_test_openai_assistants_client(mock_async_openai)
call_id = json.dumps(["run-123", "call-456"])
function_result = FunctionResultContent(call_id=call_id, result="Function executed successfully")
run_id, tool_outputs = chat_client._convert_function_results_to_tool_output([function_result]) # type: ignore
run_id, tool_outputs = chat_client._prepare_tool_outputs_for_assistants([function_result]) # type: ignore
assert run_id == "run-123"
assert tool_outputs is not None
@@ -851,10 +851,10 @@ def test_openai_assistants_client_convert_function_results_to_tool_output_valid(
assert tool_outputs[0].get("output") == "Function executed successfully"
def test_openai_assistants_client_convert_function_results_to_tool_output_mismatched_run_ids(
def test_openai_assistants_client_prepare_tool_outputs_for_assistants_mismatched_run_ids(
mock_async_openai: MagicMock,
) -> None:
"""Test _convert_function_results_to_tool_output with mismatched run IDs."""
"""Test _prepare_tool_outputs_for_assistants with mismatched run IDs."""
chat_client = create_test_openai_assistants_client(mock_async_openai)
# Create function results with different run IDs
@@ -863,7 +863,7 @@ def test_openai_assistants_client_convert_function_results_to_tool_output_mismat
function_result1 = FunctionResultContent(call_id=call_id1, result="Result 1")
function_result2 = FunctionResultContent(call_id=call_id2, result="Result 2")
run_id, tool_outputs = chat_client._convert_function_results_to_tool_output([function_result1, function_result2]) # type: ignore
run_id, tool_outputs = chat_client._prepare_tool_outputs_for_assistants([function_result1, function_result2]) # type: ignore
# Should only process the first one since run IDs don't match
assert run_id == "run-123"
@@ -182,12 +182,12 @@ def test_unsupported_tool_handling(openai_unit_test_env: dict[str, str]) -> None
unsupported_tool.__class__.__name__ = "UnsupportedAITool"
# This should ignore the unsupported ToolProtocol and return empty list
result = client._chat_to_tool_spec([unsupported_tool]) # type: ignore
result = client._prepare_tools_for_openai([unsupported_tool]) # type: ignore
assert result == []
# Also test with a non-ToolProtocol that should be converted to dict
dict_tool = {"type": "function", "name": "test"}
result = client._chat_to_tool_spec([dict_tool]) # type: ignore
result = client._prepare_tools_for_openai([dict_tool]) # type: ignore
assert result == [dict_tool]
@@ -637,7 +637,7 @@ def test_chat_response_content_order_text_before_tool_calls(openai_unit_test_env
)
client = OpenAIChatClient()
response = client._create_chat_response(mock_response, ChatOptions())
response = client._parse_response_from_openai(mock_response, ChatOptions())
# Verify we have both text and tool call content
assert len(response.messages) == 1
@@ -658,7 +658,7 @@ def test_function_result_falsy_values_handling(openai_unit_test_env: dict[str, s
# Test with empty list (falsy but not None)
message_with_empty_list = ChatMessage(role="tool", contents=[FunctionResultContent(call_id="call-123", result=[])])
openai_messages = client._openai_chat_message_parser(message_with_empty_list)
openai_messages = client._prepare_message_for_openai(message_with_empty_list)
assert len(openai_messages) == 1
assert openai_messages[0]["content"] == "[]" # Empty list should be JSON serialized
@@ -667,14 +667,14 @@ def test_function_result_falsy_values_handling(openai_unit_test_env: dict[str, s
role="tool", contents=[FunctionResultContent(call_id="call-456", result="")]
)
openai_messages = client._openai_chat_message_parser(message_with_empty_string)
openai_messages = client._prepare_message_for_openai(message_with_empty_string)
assert len(openai_messages) == 1
assert openai_messages[0]["content"] == "" # Empty string should be preserved
# Test with False (falsy but not None)
message_with_false = ChatMessage(role="tool", contents=[FunctionResultContent(call_id="call-789", result=False)])
openai_messages = client._openai_chat_message_parser(message_with_false)
openai_messages = client._prepare_message_for_openai(message_with_false)
assert len(openai_messages) == 1
assert openai_messages[0]["content"] == "false" # False should be JSON serialized
@@ -695,7 +695,7 @@ def test_function_result_exception_handling(openai_unit_test_env: dict[str, str]
],
)
openai_messages = client._openai_chat_message_parser(message_with_exception)
openai_messages = client._prepare_message_for_openai(message_with_exception)
assert len(openai_messages) == 1
assert openai_messages[0]["content"] == "Error: Function failed."
assert openai_messages[0]["tool_call_id"] == "call-123"
@@ -708,8 +708,8 @@ def test_prepare_function_call_results_string_passthrough():
assert isinstance(result, str)
def test_openai_content_parser_data_content_image(openai_unit_test_env: dict[str, str]) -> None:
"""Test _openai_content_parser converts DataContent with image media type to OpenAI format."""
def test_prepare_content_for_openai_data_content_image(openai_unit_test_env: dict[str, str]) -> None:
"""Test _prepare_content_for_openai converts DataContent with image media type to OpenAI format."""
client = OpenAIChatClient()
# Test DataContent with image media type
@@ -718,7 +718,7 @@ def test_openai_content_parser_data_content_image(openai_unit_test_env: dict[str
media_type="image/png",
)
result = client._openai_content_parser(image_data_content) # type: ignore
result = client._prepare_content_for_openai(image_data_content) # type: ignore
# Should convert to OpenAI image_url format
assert result["type"] == "image_url"
@@ -727,7 +727,7 @@ def test_openai_content_parser_data_content_image(openai_unit_test_env: dict[str
# Test DataContent with non-image media type should use default model_dump
text_data_content = DataContent(uri="data:text/plain;base64,SGVsbG8gV29ybGQ=", media_type="text/plain")
result = client._openai_content_parser(text_data_content) # type: ignore
result = client._prepare_content_for_openai(text_data_content) # type: ignore
# Should use default model_dump format
assert result["type"] == "data"
@@ -740,7 +740,7 @@ def test_openai_content_parser_data_content_image(openai_unit_test_env: dict[str
media_type="audio/wav",
)
result = client._openai_content_parser(audio_data_content) # type: ignore
result = client._prepare_content_for_openai(audio_data_content) # type: ignore
# Should convert to OpenAI input_audio format
assert result["type"] == "input_audio"
@@ -751,7 +751,7 @@ def test_openai_content_parser_data_content_image(openai_unit_test_env: dict[str
# Test DataContent with MP3 audio
mp3_data_content = DataContent(uri="data:audio/mp3;base64,//uQAAAAWGluZwAAAA8AAAACAAACcQ==", media_type="audio/mp3")
result = client._openai_content_parser(mp3_data_content) # type: ignore
result = client._prepare_content_for_openai(mp3_data_content) # type: ignore
# Should convert to OpenAI input_audio format with mp3
assert result["type"] == "input_audio"
@@ -760,8 +760,8 @@ def test_openai_content_parser_data_content_image(openai_unit_test_env: dict[str
assert result["input_audio"]["format"] == "mp3"
def test_openai_content_parser_document_file_mapping(openai_unit_test_env: dict[str, str]) -> None:
"""Test _openai_content_parser converts document files (PDF, DOCX, etc.) to OpenAI file format."""
def test_prepare_content_for_openai_document_file_mapping(openai_unit_test_env: dict[str, str]) -> None:
"""Test _prepare_content_for_openai converts document files (PDF, DOCX, etc.) to OpenAI file format."""
client = OpenAIChatClient()
# Test PDF without filename - should omit filename in OpenAI payload
@@ -770,7 +770,7 @@ def test_openai_content_parser_document_file_mapping(openai_unit_test_env: dict[
media_type="application/pdf",
)
result = client._openai_content_parser(pdf_data_content) # type: ignore
result = client._prepare_content_for_openai(pdf_data_content) # type: ignore
# Should convert to OpenAI file format without filename
assert result["type"] == "file"
@@ -787,7 +787,7 @@ def test_openai_content_parser_document_file_mapping(openai_unit_test_env: dict[
additional_properties={"filename": "report.pdf"},
)
result = client._openai_content_parser(pdf_with_filename) # type: ignore
result = client._prepare_content_for_openai(pdf_with_filename) # type: ignore
# Should use custom filename
assert result["type"] == "file"
@@ -820,7 +820,7 @@ def test_openai_content_parser_document_file_mapping(openai_unit_test_env: dict[
media_type=case["media_type"],
)
result = client._openai_content_parser(doc_content) # type: ignore
result = client._prepare_content_for_openai(doc_content) # type: ignore
# All application/* types should now be mapped to file format
assert result["type"] == "file"
@@ -834,7 +834,7 @@ def test_openai_content_parser_document_file_mapping(openai_unit_test_env: dict[
additional_properties={"filename": case["filename"]},
)
result = client._openai_content_parser(doc_with_filename) # type: ignore
result = client._prepare_content_for_openai(doc_with_filename) # type: ignore
# Should now use file format with filename
assert result["type"] == "file"
@@ -848,7 +848,7 @@ def test_openai_content_parser_document_file_mapping(openai_unit_test_env: dict[
additional_properties={},
)
result = client._openai_content_parser(pdf_empty_props) # type: ignore
result = client._prepare_content_for_openai(pdf_empty_props) # type: ignore
assert result["type"] == "file"
assert "filename" not in result["file"]
@@ -860,7 +860,7 @@ def test_openai_content_parser_document_file_mapping(openai_unit_test_env: dict[
additional_properties={"filename": None},
)
result = client._openai_content_parser(pdf_none_filename) # type: ignore
result = client._prepare_content_for_openai(pdf_none_filename) # type: ignore
assert result["type"] == "file"
assert "filename" not in result["file"] # None filename should be omitted
@@ -76,7 +76,7 @@ async def test_cmc(
mock_create.assert_awaited_once_with(
model=openai_unit_test_env["OPENAI_CHAT_MODEL_ID"],
stream=False,
messages=openai_chat_completion._prepare_chat_history_for_request(chat_history), # type: ignore
messages=openai_chat_completion._prepare_messages_for_openai(chat_history), # type: ignore
)
@@ -97,7 +97,7 @@ async def test_cmc_chat_options(
mock_create.assert_awaited_once_with(
model=openai_unit_test_env["OPENAI_CHAT_MODEL_ID"],
stream=False,
messages=openai_chat_completion._prepare_chat_history_for_request(chat_history), # type: ignore
messages=openai_chat_completion._prepare_messages_for_openai(chat_history), # type: ignore
)
@@ -120,7 +120,7 @@ async def test_cmc_no_fcc_in_response(
mock_create.assert_awaited_once_with(
model=openai_unit_test_env["OPENAI_CHAT_MODEL_ID"],
stream=False,
messages=openai_chat_completion._prepare_chat_history_for_request(orig_chat_history), # type: ignore
messages=openai_chat_completion._prepare_messages_for_openai(orig_chat_history), # type: ignore
)
@@ -167,7 +167,7 @@ async def test_scmc_chat_options(
model=openai_unit_test_env["OPENAI_CHAT_MODEL_ID"],
stream=True,
stream_options={"include_usage": True},
messages=openai_chat_completion._prepare_chat_history_for_request(chat_history), # type: ignore
messages=openai_chat_completion._prepare_messages_for_openai(chat_history), # type: ignore
)
@@ -203,7 +203,7 @@ async def test_cmc_additional_properties(
mock_create.assert_awaited_once_with(
model=openai_unit_test_env["OPENAI_CHAT_MODEL_ID"],
stream=False,
messages=openai_chat_completion._prepare_chat_history_for_request(chat_history), # type: ignore
messages=openai_chat_completion._prepare_messages_for_openai(chat_history), # type: ignore
reasoning_effort="low",
)
@@ -246,7 +246,7 @@ async def test_get_streaming(
model=openai_unit_test_env["OPENAI_CHAT_MODEL_ID"],
stream=True,
stream_options={"include_usage": True},
messages=openai_chat_completion._prepare_chat_history_for_request(orig_chat_history), # type: ignore
messages=openai_chat_completion._prepare_messages_for_openai(orig_chat_history), # type: ignore
)
@@ -285,7 +285,7 @@ async def test_get_streaming_singular(
model=openai_unit_test_env["OPENAI_CHAT_MODEL_ID"],
stream=True,
stream_options={"include_usage": True},
messages=openai_chat_completion._prepare_chat_history_for_request(orig_chat_history), # type: ignore
messages=openai_chat_completion._prepare_messages_for_openai(orig_chat_history), # type: ignore
)
@@ -349,7 +349,7 @@ async def test_get_streaming_no_fcc_in_response(
model=openai_unit_test_env["OPENAI_CHAT_MODEL_ID"],
stream=True,
stream_options={"include_usage": True},
messages=openai_chat_completion._prepare_chat_history_for_request(orig_chat_history), # type: ignore
messages=openai_chat_completion._prepare_messages_for_openai(orig_chat_history), # type: ignore
)
@@ -399,7 +399,7 @@ def test_chat_response_created_at_uses_utc(openai_unit_test_env: dict[str, str])
)
client = OpenAIChatClient()
response = client._create_chat_response(mock_response, ChatOptions())
response = client._parse_response_from_openai(mock_response, ChatOptions())
# Verify that created_at is correctly formatted as UTC
assert response.created_at is not None
@@ -431,7 +431,7 @@ def test_chat_response_update_created_at_uses_utc(openai_unit_test_env: dict[str
)
client = OpenAIChatClient()
response_update = client._create_chat_response_update(mock_chunk)
response_update = client._parse_response_update_from_openai(mock_chunk)
# Verify that created_at is correctly formatted as UTC
assert response_update.created_at is not None
@@ -368,6 +368,7 @@ async def test_response_format_parse_path() -> None:
mock_parsed_response.output_parsed = None
mock_parsed_response.usage = None
mock_parsed_response.finish_reason = None
mock_parsed_response.conversation = None # No conversation object
with patch.object(client.client.responses, "parse", return_value=mock_parsed_response):
response = await client.get_response(
@@ -454,7 +455,7 @@ async def test_get_streaming_response_with_all_parameters() -> None:
def test_response_content_creation_with_annotations() -> None:
"""Test _create_response_content with different annotation types."""
"""Test _parse_response_from_openai with different annotation types."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Create a mock response with annotated text content
@@ -485,7 +486,7 @@ def test_response_content_creation_with_annotations() -> None:
mock_response.output = [mock_message_item]
with patch.object(client, "_get_metadata_from_response", return_value={}):
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
assert len(response.messages[0].contents) >= 1
assert isinstance(response.messages[0].contents[0], TextContent)
@@ -494,7 +495,7 @@ def test_response_content_creation_with_annotations() -> None:
def test_response_content_creation_with_refusal() -> None:
"""Test _create_response_content with refusal content."""
"""Test _parse_response_from_openai with refusal content."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Create a mock response with refusal content
@@ -516,7 +517,7 @@ def test_response_content_creation_with_refusal() -> None:
mock_response.output = [mock_message_item]
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
assert len(response.messages[0].contents) == 1
assert isinstance(response.messages[0].contents[0], TextContent)
@@ -524,7 +525,7 @@ def test_response_content_creation_with_refusal() -> None:
def test_response_content_creation_with_reasoning() -> None:
"""Test _create_response_content with reasoning content."""
"""Test _parse_response_from_openai with reasoning content."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Create a mock response with reasoning content
@@ -546,7 +547,7 @@ def test_response_content_creation_with_reasoning() -> None:
mock_response.output = [mock_reasoning_item]
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
assert len(response.messages[0].contents) == 2
assert isinstance(response.messages[0].contents[0], TextReasoningContent)
@@ -554,7 +555,7 @@ def test_response_content_creation_with_reasoning() -> None:
def test_response_content_creation_with_code_interpreter() -> None:
"""Test _create_response_content with code interpreter outputs."""
"""Test _parse_response_from_openai with code interpreter outputs."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -582,7 +583,7 @@ def test_response_content_creation_with_code_interpreter() -> None:
mock_response.output = [mock_code_interpreter_item]
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
assert len(response.messages[0].contents) == 2
assert isinstance(response.messages[0].contents[0], TextContent)
@@ -593,7 +594,7 @@ def test_response_content_creation_with_code_interpreter() -> None:
def test_response_content_creation_with_function_call() -> None:
"""Test _create_response_content with function call content."""
"""Test _parse_response_from_openai with function call content."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Create a mock response with function call
@@ -614,7 +615,7 @@ def test_response_content_creation_with_function_call() -> None:
mock_response.output = [mock_function_call_item]
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
assert len(response.messages[0].contents) == 1
assert isinstance(response.messages[0].contents[0], FunctionCallContent)
@@ -624,7 +625,7 @@ def test_response_content_creation_with_function_call() -> None:
assert function_call.arguments == '{"location": "Seattle"}'
def test_tools_to_response_tools_with_hosted_mcp() -> None:
def test_prepare_tools_for_openai_with_hosted_mcp() -> None:
"""Test that HostedMCPTool is converted to the correct response tool dict."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -638,7 +639,7 @@ def test_tools_to_response_tools_with_hosted_mcp() -> None:
additional_properties={"custom": "value"},
)
resp_tools = client._tools_to_response_tools([tool])
resp_tools = client._prepare_tools_for_openai([tool])
assert isinstance(resp_tools, list)
assert len(resp_tools) == 1
mcp = resp_tools[0]
@@ -654,7 +655,7 @@ def test_tools_to_response_tools_with_hosted_mcp() -> None:
assert "require_approval" in mcp
def test_create_response_content_with_mcp_approval_request() -> None:
def test_parse_response_from_openai_with_mcp_approval_request() -> None:
"""Test that a non-streaming mcp_approval_request is parsed into FunctionApprovalRequestContent."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -675,7 +676,7 @@ def test_create_response_content_with_mcp_approval_request() -> None:
mock_response.output = [mock_item]
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
assert isinstance(response.messages[0].contents[0], FunctionApprovalRequestContent)
req = response.messages[0].contents[0]
@@ -716,7 +717,7 @@ def test_responses_client_created_at_uses_utc(openai_unit_test_env: dict[str, st
mock_response.output = [mock_message_item]
with patch.object(client, "_get_metadata_from_response", return_value={}):
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
# Verify that created_at is correctly formatted as UTC
assert response.created_at is not None
@@ -730,7 +731,7 @@ def test_responses_client_created_at_uses_utc(openai_unit_test_env: dict[str, st
)
def test_tools_to_response_tools_with_raw_image_generation() -> None:
def test_prepare_tools_for_openai_with_raw_image_generation() -> None:
"""Test that raw image_generation tool dict is handled correctly with parameter mapping."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -744,7 +745,7 @@ def test_tools_to_response_tools_with_raw_image_generation() -> None:
"background": "transparent",
}
resp_tools = client._tools_to_response_tools([tool])
resp_tools = client._prepare_tools_for_openai([tool])
assert isinstance(resp_tools, list)
assert len(resp_tools) == 1
@@ -759,7 +760,7 @@ def test_tools_to_response_tools_with_raw_image_generation() -> None:
assert image_tool["output_compression"] == 75
def test_tools_to_response_tools_with_raw_image_generation_openai_responses_params() -> None:
def test_prepare_tools_for_openai_with_raw_image_generation_openai_responses_params() -> None:
"""Test raw image_generation tool with OpenAI-specific parameters."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -773,7 +774,7 @@ def test_tools_to_response_tools_with_raw_image_generation_openai_responses_para
"partial_images": 2, # Should be integer 0-3
}
resp_tools = client._tools_to_response_tools([tool])
resp_tools = client._prepare_tools_for_openai([tool])
assert isinstance(resp_tools, list)
assert len(resp_tools) == 1
@@ -791,14 +792,14 @@ def test_tools_to_response_tools_with_raw_image_generation_openai_responses_para
assert tool_dict["partial_images"] == 2
def test_tools_to_response_tools_with_raw_image_generation_minimal() -> None:
def test_prepare_tools_for_openai_with_raw_image_generation_minimal() -> None:
"""Test raw image_generation tool with minimal configuration."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Test with minimal parameters (just type)
tool = {"type": "image_generation"}
resp_tools = client._tools_to_response_tools([tool])
resp_tools = client._prepare_tools_for_openai([tool])
assert isinstance(resp_tools, list)
assert len(resp_tools) == 1
@@ -809,7 +810,7 @@ def test_tools_to_response_tools_with_raw_image_generation_minimal() -> None:
assert len(image_tool) == 1
def test_create_streaming_response_content_with_mcp_approval_request() -> None:
def test_parse_chunk_from_openai_with_mcp_approval_request() -> None:
"""Test that a streaming mcp_approval_request event is parsed into FunctionApprovalRequestContent."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
chat_options = ChatOptions()
@@ -825,7 +826,7 @@ def test_create_streaming_response_content_with_mcp_approval_request() -> None:
mock_item.server_label = "My_MCP"
mock_event.item = mock_item
update = client._create_streaming_response_content(mock_event, chat_options, function_call_ids)
update = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
assert any(isinstance(c, FunctionApprovalRequestContent) for c in update.contents)
fa = next(c for c in update.contents if isinstance(c, FunctionApprovalRequestContent))
assert fa.id == "approval-stream-1"
@@ -901,7 +902,7 @@ async def test_end_to_end_mcp_approval_flow(span_exporter) -> None:
def test_usage_details_basic() -> None:
"""Test _usage_details_from_openai without cached or reasoning tokens."""
"""Test _parse_usage_from_openai without cached or reasoning tokens."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
mock_usage = MagicMock()
@@ -911,7 +912,7 @@ def test_usage_details_basic() -> None:
mock_usage.input_tokens_details = None
mock_usage.output_tokens_details = None
details = client._usage_details_from_openai(mock_usage) # type: ignore
details = client._parse_usage_from_openai(mock_usage) # type: ignore
assert details is not None
assert details.input_token_count == 100
assert details.output_token_count == 50
@@ -919,7 +920,7 @@ def test_usage_details_basic() -> None:
def test_usage_details_with_cached_tokens() -> None:
"""Test _usage_details_from_openai with cached input tokens."""
"""Test _parse_usage_from_openai with cached input tokens."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
mock_usage = MagicMock()
@@ -930,14 +931,14 @@ def test_usage_details_with_cached_tokens() -> None:
mock_usage.input_tokens_details.cached_tokens = 25
mock_usage.output_tokens_details = None
details = client._usage_details_from_openai(mock_usage) # type: ignore
details = client._parse_usage_from_openai(mock_usage) # type: ignore
assert details is not None
assert details.input_token_count == 200
assert details.additional_counts["openai.cached_input_tokens"] == 25
def test_usage_details_with_reasoning_tokens() -> None:
"""Test _usage_details_from_openai with reasoning tokens."""
"""Test _parse_usage_from_openai with reasoning tokens."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
mock_usage = MagicMock()
@@ -948,7 +949,7 @@ def test_usage_details_with_reasoning_tokens() -> None:
mock_usage.output_tokens_details = MagicMock()
mock_usage.output_tokens_details.reasoning_tokens = 30
details = client._usage_details_from_openai(mock_usage) # type: ignore
details = client._parse_usage_from_openai(mock_usage) # type: ignore
assert details is not None
assert details.output_token_count == 80
assert details.additional_counts["openai.reasoning_tokens"] == 30
@@ -975,7 +976,7 @@ def test_get_metadata_from_response() -> None:
def test_streaming_response_basic_structure() -> None:
"""Test that _create_streaming_response_content returns proper structure."""
"""Test that _parse_chunk_from_openai returns proper structure."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
chat_options = ChatOptions(store=True)
function_call_ids: dict[int, tuple[str, str]] = {}
@@ -983,7 +984,7 @@ def test_streaming_response_basic_structure() -> None:
# Test with a basic mock event to ensure the method returns proper structure
mock_event = MagicMock()
response = client._create_streaming_response_content(mock_event, chat_options, function_call_ids) # type: ignore
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids) # type: ignore
# Should get a valid ChatResponseUpdate structure
assert isinstance(response, ChatResponseUpdate)
@@ -1008,7 +1009,7 @@ def test_streaming_annotation_added_with_file_path() -> None:
"index": 42,
}
response = client._create_streaming_response_content(mock_event, chat_options, function_call_ids)
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
assert len(response.contents) == 1
content = response.contents[0]
@@ -1035,7 +1036,7 @@ def test_streaming_annotation_added_with_file_citation() -> None:
"index": 15,
}
response = client._create_streaming_response_content(mock_event, chat_options, function_call_ids)
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
assert len(response.contents) == 1
content = response.contents[0]
@@ -1064,7 +1065,7 @@ def test_streaming_annotation_added_with_container_file_citation() -> None:
"end_index": 50,
}
response = client._create_streaming_response_content(mock_event, chat_options, function_call_ids)
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
assert len(response.contents) == 1
content = response.contents[0]
@@ -1091,7 +1092,7 @@ def test_streaming_annotation_added_with_unknown_type() -> None:
"url": "https://example.com",
}
response = client._create_streaming_response_content(mock_event, chat_options, function_call_ids)
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
# url_citation should not produce HostedFileContent
assert len(response.contents) == 0
@@ -1137,8 +1138,8 @@ def test_get_streaming_response_with_response_format() -> None:
asyncio.run(run_streaming())
def test_openai_content_parser_image_content() -> None:
"""Test _openai_content_parser with image content variations."""
def test_prepare_content_for_openai_image_content() -> None:
"""Test _prepare_content_for_openai with image content variations."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Test image content with detail parameter and file_id
@@ -1147,7 +1148,7 @@ def test_openai_content_parser_image_content() -> None:
media_type="image/jpeg",
additional_properties={"detail": "high", "file_id": "file_123"},
)
result = client._openai_content_parser(Role.USER, image_content_with_detail, {}) # type: ignore
result = client._prepare_content_for_openai(Role.USER, image_content_with_detail, {}) # type: ignore
assert result["type"] == "input_image"
assert result["image_url"] == "https://example.com/image.jpg"
assert result["detail"] == "high"
@@ -1155,47 +1156,47 @@ def test_openai_content_parser_image_content() -> None:
# Test image content without additional properties (defaults)
image_content_basic = UriContent(uri="https://example.com/basic.png", media_type="image/png")
result = client._openai_content_parser(Role.USER, image_content_basic, {}) # type: ignore
result = client._prepare_content_for_openai(Role.USER, image_content_basic, {}) # type: ignore
assert result["type"] == "input_image"
assert result["detail"] == "auto"
assert result["file_id"] is None
def test_openai_content_parser_audio_content() -> None:
"""Test _openai_content_parser with audio content variations."""
def test_prepare_content_for_openai_audio_content() -> None:
"""Test _prepare_content_for_openai with audio content variations."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Test WAV audio content
wav_content = UriContent(uri="data:audio/wav;base64,abc123", media_type="audio/wav")
result = client._openai_content_parser(Role.USER, wav_content, {}) # type: ignore
result = client._prepare_content_for_openai(Role.USER, wav_content, {}) # type: ignore
assert result["type"] == "input_audio"
assert result["input_audio"]["data"] == "data:audio/wav;base64,abc123"
assert result["input_audio"]["format"] == "wav"
# Test MP3 audio content
mp3_content = UriContent(uri="data:audio/mp3;base64,def456", media_type="audio/mp3")
result = client._openai_content_parser(Role.USER, mp3_content, {}) # type: ignore
result = client._prepare_content_for_openai(Role.USER, mp3_content, {}) # type: ignore
assert result["type"] == "input_audio"
assert result["input_audio"]["format"] == "mp3"
def test_openai_content_parser_unsupported_content() -> None:
"""Test _openai_content_parser with unsupported content types."""
def test_prepare_content_for_openai_unsupported_content() -> None:
"""Test _prepare_content_for_openai with unsupported content types."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Test unsupported audio format
unsupported_audio = UriContent(uri="data:audio/ogg;base64,ghi789", media_type="audio/ogg")
result = client._openai_content_parser(Role.USER, unsupported_audio, {}) # type: ignore
result = client._prepare_content_for_openai(Role.USER, unsupported_audio, {}) # type: ignore
assert result == {}
# Test non-media content
text_uri_content = UriContent(uri="https://example.com/document.txt", media_type="text/plain")
result = client._openai_content_parser(Role.USER, text_uri_content, {}) # type: ignore
result = client._prepare_content_for_openai(Role.USER, text_uri_content, {}) # type: ignore
assert result == {}
def test_create_streaming_response_content_code_interpreter() -> None:
"""Test _create_streaming_response_content with code_interpreter_call."""
def test_parse_chunk_from_openai_code_interpreter() -> None:
"""Test _parse_chunk_from_openai with code_interpreter_call."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
chat_options = ChatOptions()
function_call_ids: dict[int, tuple[str, str]] = {}
@@ -1211,15 +1212,15 @@ def test_create_streaming_response_content_code_interpreter() -> None:
mock_item_image.code = None
mock_event_image.item = mock_item_image
result = client._create_streaming_response_content(mock_event_image, chat_options, function_call_ids) # type: ignore
result = client._parse_chunk_from_openai(mock_event_image, chat_options, function_call_ids) # type: ignore
assert len(result.contents) == 1
assert isinstance(result.contents[0], UriContent)
assert result.contents[0].uri == "https://example.com/plot.png"
assert result.contents[0].media_type == "image"
def test_create_streaming_response_content_reasoning() -> None:
"""Test _create_streaming_response_content with reasoning content."""
def test_parse_chunk_from_openai_reasoning() -> None:
"""Test _parse_chunk_from_openai with reasoning content."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
chat_options = ChatOptions()
function_call_ids: dict[int, tuple[str, str]] = {}
@@ -1234,7 +1235,7 @@ def test_create_streaming_response_content_reasoning() -> None:
mock_item_reasoning.summary = ["Problem analysis summary"]
mock_event_reasoning.item = mock_item_reasoning
result = client._create_streaming_response_content(mock_event_reasoning, chat_options, function_call_ids) # type: ignore
result = client._parse_chunk_from_openai(mock_event_reasoning, chat_options, function_call_ids) # type: ignore
assert len(result.contents) == 1
assert isinstance(result.contents[0], TextReasoningContent)
assert result.contents[0].text == "Analyzing the problem step by step..."
@@ -1242,8 +1243,8 @@ def test_create_streaming_response_content_reasoning() -> None:
assert result.contents[0].additional_properties["summary"] == "Problem analysis summary"
def test_openai_content_parser_text_reasoning_comprehensive() -> None:
"""Test _openai_content_parser with TextReasoningContent all additional properties."""
def test_prepare_content_for_openai_text_reasoning_comprehensive() -> None:
"""Test _prepare_content_for_openai with TextReasoningContent all additional properties."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
# Test TextReasoningContent with all additional properties
@@ -1255,7 +1256,7 @@ def test_openai_content_parser_text_reasoning_comprehensive() -> None:
"encrypted_content": "secure_data_456",
},
)
result = client._openai_content_parser(Role.ASSISTANT, comprehensive_reasoning, {}) # type: ignore
result = client._prepare_content_for_openai(Role.ASSISTANT, comprehensive_reasoning, {}) # type: ignore
assert result["type"] == "reasoning"
assert result["summary"]["text"] == "Comprehensive reasoning summary"
assert result["status"] == "in_progress"
@@ -1280,7 +1281,7 @@ def test_streaming_reasoning_text_delta_event() -> None:
)
with patch.object(client, "_get_metadata_from_response", return_value={}) as mock_metadata:
response = client._create_streaming_response_content(event, chat_options, function_call_ids) # type: ignore
response = client._parse_chunk_from_openai(event, chat_options, function_call_ids) # type: ignore
assert len(response.contents) == 1
assert isinstance(response.contents[0], TextReasoningContent)
@@ -1305,7 +1306,7 @@ def test_streaming_reasoning_text_done_event() -> None:
)
with patch.object(client, "_get_metadata_from_response", return_value={"test": "data"}) as mock_metadata:
response = client._create_streaming_response_content(event, chat_options, function_call_ids) # type: ignore
response = client._parse_chunk_from_openai(event, chat_options, function_call_ids) # type: ignore
assert len(response.contents) == 1
assert isinstance(response.contents[0], TextReasoningContent)
@@ -1331,7 +1332,7 @@ def test_streaming_reasoning_summary_text_delta_event() -> None:
)
with patch.object(client, "_get_metadata_from_response", return_value={}) as mock_metadata:
response = client._create_streaming_response_content(event, chat_options, function_call_ids) # type: ignore
response = client._parse_chunk_from_openai(event, chat_options, function_call_ids) # type: ignore
assert len(response.contents) == 1
assert isinstance(response.contents[0], TextReasoningContent)
@@ -1356,7 +1357,7 @@ def test_streaming_reasoning_summary_text_done_event() -> None:
)
with patch.object(client, "_get_metadata_from_response", return_value={"custom": "meta"}) as mock_metadata:
response = client._create_streaming_response_content(event, chat_options, function_call_ids) # type: ignore
response = client._parse_chunk_from_openai(event, chat_options, function_call_ids) # type: ignore
assert len(response.contents) == 1
assert isinstance(response.contents[0], TextReasoningContent)
@@ -1392,8 +1393,8 @@ def test_streaming_reasoning_events_preserve_metadata() -> None:
)
with patch.object(client, "_get_metadata_from_response", return_value={"test": "metadata"}):
text_response = client._create_streaming_response_content(text_event, chat_options, function_call_ids) # type: ignore
reasoning_response = client._create_streaming_response_content(reasoning_event, chat_options, function_call_ids) # type: ignore
text_response = client._parse_chunk_from_openai(text_event, chat_options, function_call_ids) # type: ignore
reasoning_response = client._parse_chunk_from_openai(reasoning_event, chat_options, function_call_ids) # type: ignore
# Both should preserve metadata
assert text_response.additional_properties == {"test": "metadata"}
@@ -1404,7 +1405,7 @@ def test_streaming_reasoning_events_preserve_metadata() -> None:
assert isinstance(reasoning_response.contents[0], TextReasoningContent)
def test_create_response_content_image_generation_raw_base64():
def test_parse_response_from_openai_image_generation_raw_base64():
"""Test image generation response parsing with raw base64 string."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -1428,7 +1429,7 @@ def test_create_response_content_image_generation_raw_base64():
mock_response.output = [mock_item]
with patch.object(client, "_get_metadata_from_response", return_value={}):
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
# Verify the response contains DataContent with proper URI and media_type
assert len(response.messages[0].contents) == 1
@@ -1438,7 +1439,7 @@ def test_create_response_content_image_generation_raw_base64():
assert content.media_type == "image/png"
def test_create_response_content_image_generation_existing_data_uri():
def test_parse_response_from_openai_image_generation_existing_data_uri():
"""Test image generation response parsing with existing data URI."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -1461,7 +1462,7 @@ def test_create_response_content_image_generation_existing_data_uri():
mock_response.output = [mock_item]
with patch.object(client, "_get_metadata_from_response", return_value={}):
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
# Verify the response contains DataContent with proper media_type parsed from URI
assert len(response.messages[0].contents) == 1
@@ -1471,7 +1472,7 @@ def test_create_response_content_image_generation_existing_data_uri():
assert content.media_type == "image/webp"
def test_create_response_content_image_generation_format_detection():
def test_parse_response_from_openai_image_generation_format_detection():
"""Test different image format detection from base64 data."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -1493,7 +1494,7 @@ def test_create_response_content_image_generation_format_detection():
mock_response_jpeg.output = [mock_item_jpeg]
with patch.object(client, "_get_metadata_from_response", return_value={}):
response_jpeg = client._create_response_content(mock_response_jpeg, chat_options=ChatOptions()) # type: ignore
response_jpeg = client._parse_response_from_openai(mock_response_jpeg, chat_options=ChatOptions()) # type: ignore
content_jpeg = response_jpeg.messages[0].contents[0]
assert isinstance(content_jpeg, DataContent)
assert content_jpeg.media_type == "image/jpeg"
@@ -1517,14 +1518,14 @@ def test_create_response_content_image_generation_format_detection():
mock_response_webp.output = [mock_item_webp]
with patch.object(client, "_get_metadata_from_response", return_value={}):
response_webp = client._create_response_content(mock_response_webp, chat_options=ChatOptions()) # type: ignore
response_webp = client._parse_response_from_openai(mock_response_webp, chat_options=ChatOptions()) # type: ignore
content_webp = response_webp.messages[0].contents[0]
assert isinstance(content_webp, DataContent)
assert content_webp.media_type == "image/webp"
assert "data:image/webp;base64," in content_webp.uri
def test_create_response_content_image_generation_fallback():
def test_parse_response_from_openai_image_generation_fallback():
"""Test image generation with invalid base64 falls back to PNG."""
client = OpenAIResponsesClient(model_id="test-model", api_key="test-key")
@@ -1547,7 +1548,7 @@ def test_create_response_content_image_generation_fallback():
mock_response.output = [mock_item]
with patch.object(client, "_get_metadata_from_response", return_value={}):
response = client._create_response_content(mock_response, chat_options=ChatOptions()) # type: ignore
response = client._parse_response_from_openai(mock_response, chat_options=ChatOptions()) # type: ignore
# Verify it falls back to PNG format for unrecognized binary data
assert len(response.messages[0].contents) == 1
@@ -1563,21 +1564,21 @@ async def test_prepare_options_store_parameter_handling() -> None:
test_conversation_id = "test-conversation-123"
chat_options = ChatOptions(store=True, conversation_id=test_conversation_id)
options = await client.prepare_options(messages, chat_options)
options = await client._prepare_options(messages, chat_options) # type: ignore
assert options["store"] is True
assert options["previous_response_id"] == test_conversation_id
chat_options = ChatOptions(store=False, conversation_id="")
options = await client.prepare_options(messages, chat_options)
options = await client._prepare_options(messages, chat_options) # type: ignore
assert options["store"] is False
chat_options = ChatOptions(store=None, conversation_id=None)
options = await client.prepare_options(messages, chat_options)
options = await client._prepare_options(messages, chat_options) # type: ignore
assert "store" not in options
assert "previous_response_id" not in options
chat_options = ChatOptions()
options = await client.prepare_options(messages, chat_options)
options = await client._prepare_options(messages, chat_options) # type: ignore
assert "store" not in options
assert "previous_response_id" not in options
@@ -1,11 +1,13 @@
# Copyright (c) Microsoft. All rights reserved.
import uuid
from collections.abc import AsyncIterable
from typing import Any
import pytest
from agent_framework import (
AgentProtocol,
AgentRunResponse,
AgentRunResponseUpdate,
AgentRunUpdateEvent,
@@ -422,6 +424,48 @@ class TestWorkflowAgent:
assert isinstance(updates[2].raw_representation, CustomData)
assert updates[2].raw_representation.value == 42
async def test_workflow_as_agent_yield_output_with_list_of_chat_messages(self) -> None:
"""Test that yield_output with list[ChatMessage] extracts contents from all messages.
Note: TextContent items are coalesced by _finalize_response, so multiple text contents
become a single merged TextContent in the final response.
"""
@executor
async def list_yielding_executor(messages: list[ChatMessage], ctx: WorkflowContext) -> None:
# Yield a list of ChatMessages (as SequentialBuilder does)
msg_list = [
ChatMessage(role=Role.USER, contents=[TextContent(text="first message")]),
ChatMessage(role=Role.ASSISTANT, contents=[TextContent(text="second message")]),
ChatMessage(
role=Role.ASSISTANT,
contents=[TextContent(text="third"), TextContent(text="fourth")],
),
]
await ctx.yield_output(msg_list)
workflow = WorkflowBuilder().set_start_executor(list_yielding_executor).build()
agent = workflow.as_agent("list-msg-agent")
# Verify streaming returns the update with all 4 contents before coalescing
updates: list[AgentRunResponseUpdate] = []
async for update in agent.run_stream("test"):
updates.append(update)
assert len(updates) == 1
assert len(updates[0].contents) == 4
texts = [c.text for c in updates[0].contents if isinstance(c, TextContent)]
assert texts == ["first message", "second message", "third", "fourth"]
# Verify run() coalesces text contents (expected behavior)
result = await agent.run("test")
assert isinstance(result, AgentRunResponse)
assert len(result.messages) == 1
# TextContent items are coalesced into one
assert len(result.messages[0].contents) == 1
assert result.messages[0].text == "first messagesecond messagethirdfourth"
async def test_thread_conversation_history_included_in_workflow_run(self) -> None:
"""Test that conversation history from thread is included when running WorkflowAgent.
@@ -521,6 +565,142 @@ class TestWorkflowAgent:
checkpoints = await checkpoint_storage.list_checkpoints(workflow.id)
assert len(checkpoints) > 0, "Checkpoints should have been created when checkpoint_storage is provided"
async def test_agent_executor_output_response_false_filters_streaming_events(self):
"""Test that AgentExecutor with output_response=False does not surface streaming events."""
class MockAgent(AgentProtocol):
"""Mock agent for testing."""
def __init__(self, name: str, response_text: str) -> None:
self._name = name
self._response_text = response_text
self._description: str | None = None
@property
def name(self) -> str | None:
return self._name
@property
def description(self) -> str | None:
return self._description
def get_new_thread(self) -> AgentThread:
return AgentThread()
async def run(self, messages: Any, *, thread: AgentThread | None = None, **kwargs: Any) -> AgentRunResponse:
return AgentRunResponse(
messages=[ChatMessage(role=Role.ASSISTANT, text=self._response_text)],
text=self._response_text,
)
async def run_stream(
self, messages: Any, *, thread: AgentThread | None = None, **kwargs: Any
) -> AsyncIterable[AgentRunResponseUpdate]:
for word in self._response_text.split():
yield AgentRunResponseUpdate(
contents=[TextContent(text=word + " ")],
role=Role.ASSISTANT,
author_name=self._name,
)
@executor
async def start_executor(messages: list[ChatMessage], ctx: WorkflowContext) -> None:
from agent_framework import AgentExecutorRequest
await ctx.yield_output("Start output")
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
# Build workflow: start -> agent1 (no output) -> agent2 (output_response=True)
workflow = (
WorkflowBuilder()
.register_executor(lambda: start_executor, "start")
.register_agent(lambda: MockAgent("agent1", "Agent1 output - should NOT appear"), "agent1")
.register_agent(
lambda: MockAgent("agent2", "Agent2 output - SHOULD appear"), "agent2", output_response=True
)
.set_start_executor("start")
.add_edge("start", "agent1")
.add_edge("agent1", "agent2")
.build()
)
agent = WorkflowAgent(workflow=workflow, name="Test Agent")
result = await agent.run("Test input")
# Collect all message texts
texts = [msg.text for msg in result.messages if msg.text]
# Start output should appear (from yield_output)
assert any("Start output" in t for t in texts), "Start output should appear"
# Agent1 output should NOT appear (output_response=False)
assert not any("Agent1" in t for t in texts), "Agent1 output should NOT appear"
# Agent2 output should appear (output_response=True)
assert any("Agent2" in t for t in texts), "Agent2 output should appear"
async def test_agent_executor_output_response_no_duplicate_from_workflow_output_event(self):
"""Test that AgentExecutor with output_response=True does not duplicate content."""
class MockAgent(AgentProtocol):
"""Mock agent for testing."""
def __init__(self, name: str, response_text: str) -> None:
self._name = name
self._response_text = response_text
self._description: str | None = None
@property
def name(self) -> str | None:
return self._name
@property
def description(self) -> str | None:
return self._description
def get_new_thread(self) -> AgentThread:
return AgentThread()
async def run(self, messages: Any, *, thread: AgentThread | None = None, **kwargs: Any) -> AgentRunResponse:
return AgentRunResponse(
messages=[ChatMessage(role=Role.ASSISTANT, text=self._response_text)],
text=self._response_text,
)
async def run_stream(
self, messages: Any, *, thread: AgentThread | None = None, **kwargs: Any
) -> AsyncIterable[AgentRunResponseUpdate]:
yield AgentRunResponseUpdate(
contents=[TextContent(text=self._response_text)],
role=Role.ASSISTANT,
author_name=self._name,
)
@executor
async def start_executor(messages: list[ChatMessage], ctx: WorkflowContext) -> None:
from agent_framework import AgentExecutorRequest
await ctx.send_message(AgentExecutorRequest(messages=messages, should_respond=True))
# Build workflow with single agent that has output_response=True
workflow = (
WorkflowBuilder()
.register_executor(lambda: start_executor, "start")
.register_agent(lambda: MockAgent("agent", "Unique response text"), "agent", output_response=True)
.set_start_executor("start")
.add_edge("start", "agent")
.build()
)
agent = WorkflowAgent(workflow=workflow, name="Test Agent")
result = await agent.run("Test input")
# Count occurrences of the unique response text
unique_text_count = sum(1 for msg in result.messages if msg.text and "Unique response text" in msg.text)
# Should appear exactly once (not duplicated from both streaming and WorkflowOutputEvent)
assert unique_text_count == 1, f"Response should appear exactly once, but appeared {unique_text_count} times"
class TestWorkflowAgentMergeUpdates:
"""Test cases specifically for the WorkflowAgent.merge_updates static method."""
@@ -490,3 +490,266 @@ async def test_magentic_kwargs_stored_in_shared_state() -> None:
# endregion
# region WorkflowAgent (as_agent) kwargs Tests
async def test_workflow_as_agent_run_propagates_kwargs_to_underlying_agent() -> None:
"""Test that kwargs passed to workflow_agent.run() flow through to the underlying agents."""
agent = _KwargsCapturingAgent(name="inner_agent")
workflow = SequentialBuilder().participants([agent]).build()
workflow_agent = workflow.as_agent(name="TestWorkflowAgent")
custom_data = {"endpoint": "https://api.example.com", "version": "v1"}
user_token = {"user_name": "alice", "access_level": "admin"}
_ = await workflow_agent.run(
"test message",
custom_data=custom_data,
user_token=user_token,
)
# Verify inner agent received kwargs
assert len(agent.captured_kwargs) >= 1, "Inner agent should have been invoked at least once"
received = agent.captured_kwargs[0]
assert "custom_data" in received, "Inner agent should receive custom_data kwarg"
assert "user_token" in received, "Inner agent should receive user_token kwarg"
assert received["custom_data"] == custom_data
assert received["user_token"] == user_token
async def test_workflow_as_agent_run_stream_propagates_kwargs_to_underlying_agent() -> None:
"""Test that kwargs passed to workflow_agent.run_stream() flow through to the underlying agents."""
agent = _KwargsCapturingAgent(name="inner_agent")
workflow = SequentialBuilder().participants([agent]).build()
workflow_agent = workflow.as_agent(name="TestWorkflowAgent")
custom_data = {"session_id": "xyz123"}
api_token = "secret-token"
async for _ in workflow_agent.run_stream(
"test message",
custom_data=custom_data,
api_token=api_token,
):
pass
# Verify inner agent received kwargs
assert len(agent.captured_kwargs) >= 1, "Inner agent should have been invoked at least once"
received = agent.captured_kwargs[0]
assert "custom_data" in received, "Inner agent should receive custom_data kwarg"
assert "api_token" in received, "Inner agent should receive api_token kwarg"
assert received["custom_data"] == custom_data
assert received["api_token"] == api_token
async def test_workflow_as_agent_propagates_kwargs_to_multiple_agents() -> None:
"""Test that kwargs flow to all agents when using workflow.as_agent()."""
agent1 = _KwargsCapturingAgent(name="agent1")
agent2 = _KwargsCapturingAgent(name="agent2")
workflow = SequentialBuilder().participants([agent1, agent2]).build()
workflow_agent = workflow.as_agent(name="MultiAgentWorkflow")
custom_data = {"batch_id": "batch-001"}
_ = await workflow_agent.run("test message", custom_data=custom_data)
# Both agents should have received kwargs
assert len(agent1.captured_kwargs) >= 1, "First agent should be invoked"
assert len(agent2.captured_kwargs) >= 1, "Second agent should be invoked"
assert agent1.captured_kwargs[0].get("custom_data") == custom_data
assert agent2.captured_kwargs[0].get("custom_data") == custom_data
async def test_workflow_as_agent_kwargs_with_none_values() -> None:
"""Test that kwargs with None values are passed through correctly via as_agent()."""
agent = _KwargsCapturingAgent(name="none_test_agent")
workflow = SequentialBuilder().participants([agent]).build()
workflow_agent = workflow.as_agent(name="NoneTestWorkflow")
_ = await workflow_agent.run("test", optional_param=None, other_param="value")
assert len(agent.captured_kwargs) >= 1
received = agent.captured_kwargs[0]
assert "optional_param" in received
assert received["optional_param"] is None
assert received["other_param"] == "value"
async def test_workflow_as_agent_kwargs_with_complex_nested_data() -> None:
"""Test that complex nested data structures flow through correctly via as_agent()."""
agent = _KwargsCapturingAgent(name="nested_agent")
workflow = SequentialBuilder().participants([agent]).build()
workflow_agent = workflow.as_agent(name="NestedDataWorkflow")
complex_data = {
"level1": {
"level2": {
"level3": ["a", "b", "c"],
"number": 42,
},
"list": [1, 2, {"nested": True}],
},
}
_ = await workflow_agent.run("test", complex_data=complex_data)
assert len(agent.captured_kwargs) >= 1
received = agent.captured_kwargs[0]
assert received.get("complex_data") == complex_data
# endregion
# region SubWorkflow (WorkflowExecutor) Tests
async def test_subworkflow_kwargs_propagation() -> None:
"""Test that kwargs are propagated to subworkflows.
Verifies kwargs passed to parent workflow.run_stream() flow through to agents
in subworkflows wrapped by WorkflowExecutor.
"""
from agent_framework._workflows._workflow_executor import WorkflowExecutor
# Create an agent inside the subworkflow that captures kwargs
inner_agent = _KwargsCapturingAgent(name="inner_agent")
# Build the inner (sub) workflow with the agent
inner_workflow = SequentialBuilder().participants([inner_agent]).build()
# Wrap the inner workflow in a WorkflowExecutor so it can be used as a subworkflow
subworkflow_executor = WorkflowExecutor(workflow=inner_workflow, id="subworkflow_executor")
# Build the outer (parent) workflow containing the subworkflow
outer_workflow = SequentialBuilder().participants([subworkflow_executor]).build()
# Define kwargs that should propagate to subworkflow
custom_data = {"api_key": "secret123", "endpoint": "https://api.example.com"}
user_token = {"user_name": "alice", "access_level": "admin"}
# Run the outer workflow with kwargs
async for event in outer_workflow.run_stream(
"test message for subworkflow",
custom_data=custom_data,
user_token=user_token,
):
if isinstance(event, WorkflowStatusEvent) and event.state == WorkflowRunState.IDLE:
break
# Verify that the inner agent was called
assert len(inner_agent.captured_kwargs) >= 1, "Inner agent in subworkflow should have been invoked"
received_kwargs = inner_agent.captured_kwargs[0]
# Verify kwargs were propagated from parent workflow to subworkflow agent
assert "custom_data" in received_kwargs, (
f"Subworkflow agent should receive 'custom_data' kwarg. Received keys: {list(received_kwargs.keys())}"
)
assert "user_token" in received_kwargs, (
f"Subworkflow agent should receive 'user_token' kwarg. Received keys: {list(received_kwargs.keys())}"
)
assert received_kwargs.get("custom_data") == custom_data, (
f"Expected custom_data={custom_data}, got {received_kwargs.get('custom_data')}"
)
assert received_kwargs.get("user_token") == user_token, (
f"Expected user_token={user_token}, got {received_kwargs.get('user_token')}"
)
async def test_subworkflow_kwargs_accessible_via_shared_state() -> None:
"""Test that kwargs are accessible via SharedState within subworkflow.
Verifies that WORKFLOW_RUN_KWARGS_KEY is populated in the subworkflow's SharedState
with kwargs from the parent workflow.
"""
from agent_framework import Executor, WorkflowContext, handler
from agent_framework._workflows._workflow_executor import WorkflowExecutor
captured_kwargs_from_state: list[dict[str, Any]] = []
class _SharedStateReader(Executor):
"""Executor that reads kwargs from SharedState for verification."""
@handler
async def read_kwargs(self, msgs: list[ChatMessage], ctx: WorkflowContext[list[ChatMessage]]) -> None:
kwargs_from_state = await ctx.get_shared_state(WORKFLOW_RUN_KWARGS_KEY)
captured_kwargs_from_state.append(kwargs_from_state or {})
await ctx.send_message(msgs)
# Build inner workflow with SharedState reader
state_reader = _SharedStateReader(id="state_reader")
inner_workflow = SequentialBuilder().participants([state_reader]).build()
# Wrap as subworkflow
subworkflow_executor = WorkflowExecutor(workflow=inner_workflow, id="subworkflow")
# Build outer workflow
outer_workflow = SequentialBuilder().participants([subworkflow_executor]).build()
# Run with kwargs
async for event in outer_workflow.run_stream(
"test",
my_custom_kwarg="should_be_propagated",
another_kwarg=42,
):
if isinstance(event, WorkflowStatusEvent) and event.state == WorkflowRunState.IDLE:
break
# Verify the state reader was invoked
assert len(captured_kwargs_from_state) >= 1, "SharedState reader should have been invoked"
kwargs_in_subworkflow = captured_kwargs_from_state[0]
assert kwargs_in_subworkflow.get("my_custom_kwarg") == "should_be_propagated", (
f"Expected 'my_custom_kwarg' in subworkflow SharedState, got: {kwargs_in_subworkflow}"
)
assert kwargs_in_subworkflow.get("another_kwarg") == 42, (
f"Expected 'another_kwarg'=42 in subworkflow SharedState, got: {kwargs_in_subworkflow}"
)
async def test_nested_subworkflow_kwargs_propagation() -> None:
"""Test kwargs propagation through multiple levels of nested subworkflows.
Verifies kwargs flow through 3 levels:
- Outer workflow
- Middle subworkflow (WorkflowExecutor)
- Inner subworkflow (WorkflowExecutor) with agent
"""
from agent_framework._workflows._workflow_executor import WorkflowExecutor
# Innermost agent
inner_agent = _KwargsCapturingAgent(name="deeply_nested_agent")
# Build inner workflow
inner_workflow = SequentialBuilder().participants([inner_agent]).build()
inner_executor = WorkflowExecutor(workflow=inner_workflow, id="inner_executor")
# Build middle workflow containing inner
middle_workflow = SequentialBuilder().participants([inner_executor]).build()
middle_executor = WorkflowExecutor(workflow=middle_workflow, id="middle_executor")
# Build outer workflow containing middle
outer_workflow = SequentialBuilder().participants([middle_executor]).build()
# Run with kwargs
async for event in outer_workflow.run_stream(
"deeply nested test",
deep_kwarg="should_reach_inner",
):
if isinstance(event, WorkflowStatusEvent) and event.state == WorkflowRunState.IDLE:
break
# Verify inner agent was called
assert len(inner_agent.captured_kwargs) >= 1, "Deeply nested agent should be invoked"
received = inner_agent.captured_kwargs[0]
assert received.get("deep_kwarg") == "should_reach_inner", (
f"Deeply nested agent should receive 'deep_kwarg'. Got: {received}"
)
# endregion
@@ -248,7 +248,7 @@ class AgentFrameworkExecutor:
# Get thread from conversation parameter (OpenAI standard!)
thread = None
conversation_id = request.get_conversation_id()
conversation_id = request._get_conversation_id()
if conversation_id:
thread = self.conversation_store.get_thread(conversation_id)
if thread:
@@ -324,7 +324,7 @@ class AgentFrameworkExecutor:
entity_id = request.get_entity_id() or "unknown"
# Get or create session conversation for checkpoint storage
conversation_id = request.get_conversation_id()
conversation_id = request._get_conversation_id()
if not conversation_id:
# Create default session if not provided
import time
@@ -324,7 +324,7 @@ class AgentFrameworkRequest(BaseModel):
return self.metadata.get("entity_id")
return None
def get_conversation_id(self) -> str | None:
def _get_conversation_id(self) -> str | None:
"""Extract conversation_id from conversation parameter.
Supports both string and object forms:
@@ -117,9 +117,11 @@ class OllamaChatClient(BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> ChatResponse:
# prepare
options_dict = self._prepare_options(messages, chat_options)
try:
# execute
response: OllamaChatResponse = await self.client.chat( # type: ignore[misc]
stream=False,
**options_dict,
@@ -128,7 +130,8 @@ class OllamaChatClient(BaseChatClient):
except Exception as ex:
raise ServiceResponseException(f"Ollama chat request failed : {ex}", ex) from ex
return self._ollama_response_to_agent_framework_response(response)
# process
return self._parse_response_from_ollama(response)
async def _inner_get_streaming_response(
self,
@@ -137,9 +140,11 @@ class OllamaChatClient(BaseChatClient):
chat_options: ChatOptions,
**kwargs: Any,
) -> AsyncIterable[ChatResponseUpdate]:
# prepare
options_dict = self._prepare_options(messages, chat_options)
try:
# execute
response_object: AsyncIterable[OllamaChatResponse] = await self.client.chat( # type: ignore[misc]
stream=True,
**options_dict,
@@ -148,49 +153,61 @@ class OllamaChatClient(BaseChatClient):
except Exception as ex:
raise ServiceResponseException(f"Ollama streaming chat request failed : {ex}", ex) from ex
# process
async for part in response_object:
yield self._ollama_streaming_response_to_agent_framework_response(part)
yield self._parse_streaming_response_from_ollama(part)
def _prepare_options(self, messages: MutableSequence[ChatMessage], chat_options: ChatOptions) -> dict[str, Any]:
# Preprocess web search tool if it exists
options_dict = chat_options.to_dict(exclude={"instructions", "type"})
# Promote additional_properties to the top level of options_dict
additional_props = options_dict.pop("additional_properties", {})
options_dict.update(additional_props)
# Prepare Messages from Agent Framework format to Ollama format
if messages and "messages" not in options_dict:
options_dict["messages"] = self._prepare_chat_history_for_request(messages)
if "messages" not in options_dict:
raise ServiceInvalidRequestError("Messages are required for chat completions")
# Prepare Tools from Agent Framework format to Json Schema format
if chat_options.tools:
options_dict["tools"] = self._chat_to_tool_spec(chat_options.tools)
# Currently Ollama only supports auto tool choice
# tool choice - Currently Ollama only supports auto tool choice
if chat_options.tool_choice == "required":
raise ServiceInvalidRequestError("Ollama does not support required tool choice.")
# Always auto: remove tool_choice since Ollama does not expose configuration to force or disable tools.
if "tool_choice" in options_dict:
del options_dict["tool_choice"]
# Rename model_id to model for Ollama API, if no model is provided use the one from client initialization
if "model_id" in options_dict:
options_dict["model"] = options_dict.pop("model_id")
run_options = chat_options.to_dict(
exclude={
"type",
"instructions",
"tool_choice", # Ollama does not support tool_choice configuration
"additional_properties", # handled separately
}
)
if "model_id" not in options_dict:
options_dict["model"] = self.model_id
# messages
if messages and "messages" not in run_options:
run_options["messages"] = self._prepare_messages_for_ollama(messages)
if "messages" not in run_options:
raise ServiceInvalidRequestError("Messages are required for chat completions")
return options_dict
# translations between ChatOptions and Ollama API
translations = {"model_id": "model"}
for old_key, new_key in translations.items():
if old_key in run_options and old_key != new_key:
run_options[new_key] = run_options.pop(old_key)
def _prepare_chat_history_for_request(self, messages: MutableSequence[ChatMessage]) -> list[OllamaMessage]:
ollama_messages = [self._agent_framework_message_to_ollama_message(msg) for msg in messages]
# model id
if not run_options.get("model"):
if not self.model_id:
raise ValueError("model_id must be a non-empty string")
run_options["model"] = self.model_id
# tools
if chat_options.tools and (tools := self._prepare_tools_for_ollama(chat_options.tools)):
run_options["tools"] = tools
# additional properties
additional_options = {
key: value for key, value in chat_options.additional_properties.items() if value is not None
}
if additional_options:
run_options.update(additional_options)
return run_options
def _prepare_messages_for_ollama(self, messages: MutableSequence[ChatMessage]) -> list[OllamaMessage]:
ollama_messages = [self._prepare_message_for_ollama(msg) for msg in messages]
# Flatten the list of lists into a single list
return list(chain.from_iterable(ollama_messages))
def _agent_framework_message_to_ollama_message(self, message: ChatMessage) -> list[OllamaMessage]:
def _prepare_message_for_ollama(self, message: ChatMessage) -> list[OllamaMessage]:
message_converters: dict[str, Callable[[ChatMessage], list[OllamaMessage]]] = {
Role.SYSTEM.value: self._format_system_message,
Role.USER.value: self._format_user_message,
@@ -250,21 +267,19 @@ class OllamaChatClient(BaseChatClient):
if isinstance(item, FunctionResultContent)
]
def _ollama_response_to_agent_framework_content(self, response: OllamaChatResponse) -> list[Contents]:
def _parse_contents_from_ollama(self, response: OllamaChatResponse) -> list[Contents]:
contents: list[Contents] = []
if response.message.thinking:
contents.append(TextReasoningContent(text=response.message.thinking))
if response.message.content:
contents.append(TextContent(text=response.message.content))
if response.message.tool_calls:
tool_calls = self._parse_ollama_tool_calls(response.message.tool_calls)
tool_calls = self._parse_tool_calls_from_ollama(response.message.tool_calls)
contents.extend(tool_calls)
return contents
def _ollama_streaming_response_to_agent_framework_response(
self, response: OllamaChatResponse
) -> ChatResponseUpdate:
contents = self._ollama_response_to_agent_framework_content(response)
def _parse_streaming_response_from_ollama(self, response: OllamaChatResponse) -> ChatResponseUpdate:
contents = self._parse_contents_from_ollama(response)
return ChatResponseUpdate(
contents=contents,
role=Role.ASSISTANT,
@@ -272,8 +287,8 @@ class OllamaChatClient(BaseChatClient):
created_at=response.created_at,
)
def _ollama_response_to_agent_framework_response(self, response: OllamaChatResponse) -> ChatResponse:
contents = self._ollama_response_to_agent_framework_content(response)
def _parse_response_from_ollama(self, response: OllamaChatResponse) -> ChatResponse:
contents = self._parse_contents_from_ollama(response)
return ChatResponse(
messages=[ChatMessage(role=Role.ASSISTANT, contents=contents)],
@@ -285,7 +300,7 @@ class OllamaChatClient(BaseChatClient):
),
)
def _parse_ollama_tool_calls(self, tool_calls: Sequence[OllamaMessage.ToolCall]) -> list[Contents]:
def _parse_tool_calls_from_ollama(self, tool_calls: Sequence[OllamaMessage.ToolCall]) -> list[Contents]:
resp: list[Contents] = []
for tool in tool_calls:
fcc = FunctionCallContent(
@@ -297,7 +312,7 @@ class OllamaChatClient(BaseChatClient):
resp.append(fcc)
return resp
def _chat_to_tool_spec(self, tools: list[ToolProtocol | MutableMapping[str, Any]]) -> list[dict[str, Any]]:
def _prepare_tools_for_ollama(self, tools: list[ToolProtocol | MutableMapping[str, Any]]) -> list[dict[str, Any]]:
chat_tools: list[dict[str, Any]] = []
for tool in tools:
if isinstance(tool, ToolProtocol):
@@ -1,38 +0,0 @@
# Ollama Examples
This folder contains examples demonstrating how to use Ollama models with the Agent Framework.
## Prerequisites
1. **Install Ollama**: Download and install Ollama from [ollama.com](https://ollama.com/)
2. **Start Ollama**: Ensure Ollama is running on your local machine
3. **Pull a model**: Run `ollama pull mistral` (or any other model you prefer)
- For function calling examples, use models that support tool calling like `mistral` or `qwen2.5`
- For reasoning examples, use models that support reasoning like `qwen2.5:8b`
- For Multimodality you can use models like `gemma3:4b`
> **Note**: Not all models support all features. Function calling and reasoning capabilities depend on the specific model you're using.
## Examples
| File | Description |
|------|-------------|
| [`ollama_agent_basic.py`](ollama_agent_basic.py) | Demonstrates basic Ollama agent usage with the native Ollama Chat Client. Shows both streaming and non-streaming responses with tool calling capabilities. |
| [`ollama_agent_reasoning.py`](ollama_agent_reasoning.py) | Demonstrates Ollama agent with reasoning capabilities using the native Ollama Chat Client. Shows how to enable thinking/reasoning mode. |
| [`ollama_chat_client.py`](ollama_chat_client.py) | Ollama Chat Client with native Ollama Chat Client |
| [`ollama_chat_multimodal.py`](ollama_chat_multimodal.py) | Ollama Chat with multimodal native Ollama Chat Client |
## Configuration
The examples use environment variables for configuration. Set the appropriate variables based on which example you're running:
### For Native Ollama Examples (`ollama_agent_basic.py`, `ollama_agent_reasoning.py`)
Set the following environment variables:
- `OLLAMA_HOST`: The base URL for your Ollama server (optional, defaults to `http://localhost:11434`)
- Example: `export OLLAMA_HOST="http://localhost:11434"`
- `OLLAMA_CHAT_MODEL_ID`: The model name to use
- Example: `export OLLAMA_CHAT_MODEL_ID="qwen2.5:8b"`
- Must be a model you have pulled with Ollama
+1 -1
View File
@@ -4,7 +4,7 @@ description = "Ollama integration for Microsoft Agent Framework."
authors = [{ name = "Microsoft", email = "af-support@microsoft.com"}]
readme = "README.md"
requires-python = ">=3.10"
version = "0.1.0b1"
version = "1.0.0b251216"
license-files = ["LICENSE"]
urls.homepage = "https://learn.microsoft.com/en-us/agent-framework/"
urls.source = "https://github.com/microsoft/agent-framework/tree/main/python"
@@ -9,6 +9,7 @@ from uuid import uuid4
import redis.asyncio as redis
from agent_framework import ChatMessage
from agent_framework._serialization import SerializationMixin
from redis.credentials import CredentialProvider
class RedisStoreState(SerializationMixin):
@@ -55,6 +56,11 @@ class RedisChatMessageStore:
def __init__(
self,
redis_url: str | None = None,
credential_provider: CredentialProvider | None = None,
host: str | None = None,
port: int = 6380,
ssl: bool = True,
username: str | None = None,
thread_id: str | None = None,
key_prefix: str = "chat_messages",
max_messages: int | None = None,
@@ -63,12 +69,19 @@ class RedisChatMessageStore:
"""Initialize the Redis chat message store.
Creates a Redis-backed chat message store for a specific conversation thread.
The store will automatically create a Redis connection and manage message
persistence using Redis List operations.
Supports both traditional URL-based authentication and Azure Managed Redis
with credential provider.
Args:
redis_url: Redis connection URL (e.g., "redis://localhost:6379").
Required for establishing Redis connection.
Used for traditional authentication. Mutually exclusive with credential_provider.
credential_provider: Redis credential provider (redis.credentials.CredentialProvider) for
Azure AD authentication. Requires host parameter. Mutually exclusive with redis_url.
host: Redis host name (e.g., "myredis.redis.cache.windows.net").
Required when using credential_provider.
port: Redis port number. Defaults to 6380 (Azure Redis SSL port).
ssl: Enable SSL/TLS connection. Defaults to True.
username: Redis username. Defaults to None.
thread_id: Unique identifier for this conversation thread.
If not provided, a UUID will be auto-generated.
This becomes part of the Redis key: {key_prefix}:{thread_id}
@@ -82,23 +95,58 @@ class RedisChatMessageStore:
Useful for resuming conversations or seeding with context.
Raises:
ValueError: If redis_url is None (Redis connection is required).
redis.ConnectionError: If unable to connect to Redis server.
ValueError: If neither redis_url nor credential_provider is provided.
ValueError: If both redis_url and credential_provider are provided.
ValueError: If credential_provider is used without host parameter.
Examples:
Traditional connection:
store = RedisChatMessageStore(
redis_url="redis://localhost:6379",
thread_id="conversation_123"
)
Azure Managed Redis with credential provider:
from redis.credentials import CredentialProvider
from azure.identity.aio import DefaultAzureCredential
store = RedisChatMessageStore(
credential_provider=CredentialProvider(DefaultAzureCredential()),
host="myredis.redis.cache.windows.net",
thread_id="conversation_123"
)
"""
# Validate required parameters
if redis_url is None:
raise ValueError("redis_url is required for Redis connection")
# Validate connection parameters
if redis_url is None and credential_provider is None:
raise ValueError("Either redis_url or credential_provider must be provided")
if redis_url is not None and credential_provider is not None:
raise ValueError("redis_url and credential_provider are mutually exclusive")
if credential_provider is not None and host is None:
raise ValueError("host is required when using credential_provider")
# Store configuration
self.redis_url = redis_url
self.thread_id = thread_id or f"thread_{uuid4()}"
self.key_prefix = key_prefix
self.max_messages = max_messages
# Initialize Redis client with connection pooling and async support
self._redis_client = redis.from_url(redis_url, decode_responses=True) # type: ignore[no-untyped-call]
# Initialize Redis client based on authentication method
if credential_provider is not None and host is not None:
# Azure AD authentication with credential provider
self.redis_url = None # Not using URL-based auth
self._redis_client = redis.Redis(
host=host,
port=port,
ssl=ssl,
username=username,
credential_provider=credential_provider,
decode_responses=True,
)
else:
# Traditional URL-based authentication
self.redis_url = redis_url
self._redis_client = redis.from_url(redis_url, decode_responses=True) # type: ignore[no-untyped-call]
# Handle initial messages (will be moved to Redis on first access)
self._initial_messages = list(messages) if messages else []
@@ -93,11 +93,118 @@ class TestRedisChatMessageStore:
assert store.max_messages == 100
def test_init_with_redis_url_required(self):
"""Test that redis_url is required for initialization."""
with pytest.raises(ValueError, match="redis_url is required for Redis connection"):
# Should raise an exception since redis_url is required
"""Test that either redis_url or credential_provider is required."""
with pytest.raises(ValueError, match="Either redis_url or credential_provider must be provided"):
RedisChatMessageStore(thread_id="test123")
def test_init_with_credential_provider(self):
"""Test initialization with credential_provider."""
mock_credential_provider = MagicMock()
with patch("agent_framework_redis._chat_message_store.redis.Redis") as mock_redis_class:
mock_redis_instance = MagicMock()
mock_redis_class.return_value = mock_redis_instance
store = RedisChatMessageStore(
credential_provider=mock_credential_provider,
host="myredis.redis.cache.windows.net",
thread_id="test123",
)
# Verify Redis.Redis was called with correct parameters
mock_redis_class.assert_called_once_with(
host="myredis.redis.cache.windows.net",
port=6380,
ssl=True,
username=None,
credential_provider=mock_credential_provider,
decode_responses=True,
)
# Verify store instance is properly initialized
assert store.thread_id == "test123"
assert store.redis_url is None # Should be None for credential provider auth
assert store.key_prefix == "chat_messages"
assert store.max_messages is None
def test_init_with_credential_provider_custom_port(self):
"""Test initialization with credential_provider and custom port."""
mock_credential_provider = MagicMock()
with patch("agent_framework_redis._chat_message_store.redis.Redis") as mock_redis_class:
mock_redis_instance = MagicMock()
mock_redis_class.return_value = mock_redis_instance
store = RedisChatMessageStore(
credential_provider=mock_credential_provider,
host="myredis.redis.cache.windows.net",
port=6379,
ssl=False,
username="admin",
thread_id="test123",
)
# Verify custom parameters were passed
mock_redis_class.assert_called_once_with(
host="myredis.redis.cache.windows.net",
port=6379,
ssl=False,
username="admin",
credential_provider=mock_credential_provider,
decode_responses=True,
)
# Verify store instance is properly initialized
assert store.thread_id == "test123"
assert store.redis_url is None # Should be None for credential provider auth
assert store.key_prefix == "chat_messages"
def test_init_credential_provider_requires_host(self):
"""Test that credential_provider requires host parameter."""
mock_credential_provider = MagicMock()
with pytest.raises(ValueError, match="host is required when using credential_provider"):
RedisChatMessageStore(
credential_provider=mock_credential_provider,
thread_id="test123",
)
def test_init_mutually_exclusive_params(self):
"""Test that redis_url and credential_provider are mutually exclusive."""
mock_credential_provider = MagicMock()
with pytest.raises(ValueError, match="redis_url and credential_provider are mutually exclusive"):
RedisChatMessageStore(
redis_url="redis://localhost:6379",
credential_provider=mock_credential_provider,
host="myredis.redis.cache.windows.net",
thread_id="test123",
)
async def test_serialize_with_credential_provider(self):
"""Test that serialization works correctly with credential provider authentication."""
mock_credential_provider = MagicMock()
with patch("agent_framework_redis._chat_message_store.redis.Redis") as mock_redis_class:
mock_redis_instance = MagicMock()
mock_redis_class.return_value = mock_redis_instance
store = RedisChatMessageStore(
credential_provider=mock_credential_provider,
host="myredis.redis.cache.windows.net",
thread_id="test123",
key_prefix="custom_prefix",
max_messages=100,
)
# Serialize the store state
state = await store.serialize()
# Verify serialization includes correct values
assert state["thread_id"] == "test123"
assert state["redis_url"] is None # Should be None for credential provider auth
assert state["key_prefix"] == "custom_prefix"
assert state["max_messages"] == 100
assert state["type"] == "redis_store_state"
def test_init_with_initial_messages(self, sample_messages):
"""Test initialization with initial messages."""
with patch("agent_framework_redis._chat_message_store.redis.from_url"):
+7 -5
View File
@@ -99,11 +99,15 @@ This directory contains samples demonstrating the capabilities of Microsoft Agen
### Ollama
The recommended way to use Ollama is via the native `OllamaChatClient` from the `agent-framework-ollama` package.
| File | Description |
|------|-------------|
| [`getting_started/agents/ollama/ollama_with_openai_chat_client.py`](./getting_started/agents/ollama/ollama_with_openai_chat_client.py) | Ollama with OpenAI Chat Client Example |
| [`packages/ollama/getting_started/ollama_agent_basic.py`](../packages/ollama/getting_started/ollama_agent_basic.py) | (Experimental) Ollama Agent with native Ollama Chat Client |
| [`packages/ollama/getting_started/ollama_agent_reasoning.py`](../packages/ollama/getting_started/ollama_agent_reasoning.py) | (Experimental) Ollama Reasoning Agent with native Ollama Chat Client |
| [`getting_started/agents/ollama/ollama_agent_basic.py`](./getting_started/agents/ollama/ollama_agent_basic.py) | Basic Ollama Agent with native Ollama Chat Client |
| [`getting_started/agents/ollama/ollama_agent_reasoning.py`](./getting_started/agents/ollama/ollama_agent_reasoning.py) | Ollama Agent with reasoning capabilities |
| [`getting_started/agents/ollama/ollama_chat_client.py`](./getting_started/agents/ollama/ollama_chat_client.py) | Direct usage of Ollama Chat Client |
| [`getting_started/agents/ollama/ollama_chat_multimodal.py`](./getting_started/agents/ollama/ollama_chat_multimodal.py) | Ollama Chat Client with multimodal (image) input |
| [`getting_started/agents/ollama/ollama_with_openai_chat_client.py`](./getting_started/agents/ollama/ollama_with_openai_chat_client.py) | Alternative: Ollama via OpenAI Chat Client |
### OpenAI
@@ -149,7 +153,6 @@ This directory contains samples demonstrating the capabilities of Microsoft Agen
| [`getting_started/chat_client/openai_assistants_client.py`](./getting_started/chat_client/openai_assistants_client.py) | OpenAI Assistants Client Direct Usage Example |
| [`getting_started/chat_client/openai_chat_client.py`](./getting_started/chat_client/openai_chat_client.py) | OpenAI Chat Client Direct Usage Example |
| [`getting_started/chat_client/openai_responses_client.py`](./getting_started/chat_client/openai_responses_client.py) | OpenAI Responses Client Direct Usage Example |
| [`packages/ollama/getting_started/ollama_chat_client.py`](../packages/ollama/getting_started/ollama_chat_client.py) | (Experimental) Ollama Chat Client with native Ollama Chat Client |
## Context Providers
@@ -225,7 +228,6 @@ This directory contains samples demonstrating the capabilities of Microsoft Agen
| [`getting_started/multimodal_input/azure_chat_multimodal.py`](./getting_started/multimodal_input/azure_chat_multimodal.py) | Azure OpenAI Chat with multimodal (image) input example |
| [`getting_started/multimodal_input/azure_responses_multimodal.py`](./getting_started/multimodal_input/azure_responses_multimodal.py) | Azure OpenAI Responses with multimodal (image) input example |
| [`getting_started/multimodal_input/openai_chat_multimodal.py`](./getting_started/multimodal_input/openai_chat_multimodal.py) | OpenAI Chat with multimodal (image) input example |
| [`packages/ollama/getting_started/ollama_chat_multimodal.py`](../packages/ollama/getting_started/ollama_chat_multimodal.py) | (Experimental) Ollama Chat with multimodal native Ollama Chat Client |
## Azure Functions
@@ -9,6 +9,7 @@ This folder contains examples demonstrating different ways to create and use age
| [`azure_ai_basic.py`](azure_ai_basic.py) | The simplest way to create an agent using `ChatAgent` with `AzureAIAgentClient`. It automatically handles all configuration using environment variables. |
| [`azure_ai_with_bing_custom_search.py`](azure_ai_with_bing_custom_search.py) | Shows how to use Bing Custom Search with Azure AI agents to find real-time information from the web using custom search configurations. Demonstrates how to set up and use HostedWebSearchTool with custom search instances. |
| [`azure_ai_with_bing_grounding.py`](azure_ai_with_bing_grounding.py) | Shows how to use Bing Grounding search with Azure AI agents to find real-time information from the web. Demonstrates web search capabilities with proper source citations and comprehensive error handling. |
| [`azure_ai_with_bing_grounding_citations.py`](azure_ai_with_bing_grounding_citations.py) | Demonstrates how to extract and display citations from Bing Grounding search responses. Shows how to collect citation annotations (title, URL, snippet) during streaming responses, enabling users to verify sources and access referenced content. |
| [`azure_ai_with_code_interpreter_file_generation.py`](azure_ai_with_code_interpreter_file_generation.py) | Shows how to retrieve file IDs from code interpreter generated files using both streaming and non-streaming approaches. |
| [`azure_ai_with_code_interpreter.py`](azure_ai_with_code_interpreter.py) | Shows how to use the HostedCodeInterpreterTool with Azure AI agents to write and execute Python code. Includes helper methods for accessing code interpreter data from response chunks. |
| [`azure_ai_with_existing_agent.py`](azure_ai_with_existing_agent.py) | Shows how to work with a pre-existing agent by providing the agent ID to the Azure AI chat client. This example also demonstrates proper cleanup of manually created agents. |
@@ -0,0 +1,86 @@
# Copyright (c) Microsoft. All rights reserved.
import asyncio
from agent_framework import ChatAgent, CitationAnnotation, HostedWebSearchTool
from agent_framework.azure import AzureAIAgentClient
from azure.identity.aio import AzureCliCredential
"""
This sample demonstrates how to create an Azure AI agent that uses Bing Grounding
search to find real-time information from the web with comprehensive citation support.
It shows how to extract and display citations (title, URL, and snippet) from Bing
Grounding responses, enabling users to verify sources and explore referenced content.
Prerequisites:
1. A connected Grounding with Bing Search resource in your Azure AI project
2. Set BING_CONNECTION_ID environment variable
Example: BING_CONNECTION_ID="your-bing-connection-id"
To set up Bing Grounding:
1. Go to Azure AI Foundry portal (https://ai.azure.com)
2. Navigate to your project's "Connected resources" section
3. Add a new connection for "Grounding with Bing Search"
4. Copy the connection ID and set the BING_CONNECTION_ID environment variable
"""
async def main() -> None:
"""Main function demonstrating Azure AI agent with Bing Grounding search."""
# 1. Create Bing Grounding search tool using HostedWebSearchTool
# The connection ID will be automatically picked up from environment variable
bing_search_tool = HostedWebSearchTool(
name="Bing Grounding Search",
description="Search the web for current information using Bing",
)
# 2. Use AzureAIAgentClient as async context manager for automatic cleanup
async with (
AzureAIAgentClient(credential=AzureCliCredential()) as client,
ChatAgent(
chat_client=client,
name="BingSearchAgent",
instructions=(
"You are a helpful assistant that can search the web for current information. "
"Use the Bing search tool to find up-to-date information and provide accurate, "
"well-sourced answers. Always cite your sources when possible."
),
tools=bing_search_tool,
) as agent,
):
# 3. Demonstrate agent capabilities with web search
print("=== Azure AI Agent with Bing Grounding Search ===\n")
user_input = "What is the most popular programming language?"
print(f"User: {user_input}")
print("Agent: ", end="", flush=True)
# Stream the response and collect citations
citations: list[CitationAnnotation] = []
async for chunk in agent.run_stream(user_input):
if chunk.text:
print(chunk.text, end="", flush=True)
# Collect citations from Bing Grounding responses
for content in getattr(chunk, "contents", []):
annotations = getattr(content, "annotations", [])
if annotations:
citations.extend(annotations)
print()
# Display collected citations
if citations:
print("\n\nCitations:")
for i, citation in enumerate(citations, 1):
print(f"[{i}] {citation.title}: {citation.url}")
if citation.snippet:
print(f" Snippet: {citation.snippet}")
else:
print("\nNo citations found in the response.")
print()
if __name__ == "__main__":
asyncio.run(main())
@@ -8,20 +8,41 @@ This folder contains examples demonstrating how to use Ollama models with the Ag
2. **Start Ollama**: Ensure Ollama is running on your local machine
3. **Pull a model**: Run `ollama pull mistral` (or any other model you prefer)
- For function calling examples, use models that support tool calling like `mistral` or `qwen2.5`
- For reasoning examples, use models that support reasoning like `qwen2.5:8b`
- For reasoning examples, use models that support reasoning like `qwen3:8b`
- For multimodal examples, use models like `gemma3:4b`
> **Note**: Not all models support all features. Function calling and reasoning capabilities depend on the specific model you're using.
> **Note**: Not all models support all features. Function calling, reasoning, and multimodal capabilities depend on the specific model you're using.
## Recommended Approach
The recommended way to use Ollama with Agent Framework is via the native `OllamaChatClient` from the `agent-framework-ollama` package. This provides full support for Ollama-specific features like reasoning mode.
Alternatively, you can use the `OpenAIChatClient` configured to point to your local Ollama server, which may be useful if you're already familiar with the OpenAI client interface.
## Examples
| File | Description |
|------|-------------|
| [`ollama_with_openai_chat_client.py`](ollama_with_openai_chat_client.py) | Demonstrates how to configure OpenAI Chat Client to use local Ollama models. Shows both streaming and non-streaming responses with tool calling capabilities. |
| [`ollama_agent_basic.py`](ollama_agent_basic.py) | Basic Ollama agent with tool calling using native Ollama Chat Client. Shows both streaming and non-streaming responses. |
| [`ollama_agent_reasoning.py`](ollama_agent_reasoning.py) | Ollama agent with reasoning capabilities using native Ollama Chat Client. Shows how to enable thinking/reasoning mode. |
| [`ollama_chat_client.py`](ollama_chat_client.py) | Direct usage of the native Ollama Chat Client with tool calling. |
| [`ollama_chat_multimodal.py`](ollama_chat_multimodal.py) | Ollama Chat Client with multimodal (image) input capabilities. |
| [`ollama_with_openai_chat_client.py`](ollama_with_openai_chat_client.py) | Alternative approach using OpenAI Chat Client configured to use local Ollama models. |
## Configuration
The examples use environment variables for configuration. Set the appropriate variables based on which example you're running:
### For Native Ollama Examples
Set the following environment variables:
- `OLLAMA_HOST`: The base URL for your Ollama server (optional, defaults to `http://localhost:11434`)
- Example: `export OLLAMA_HOST="http://localhost:11434"`
- `OLLAMA_CHAT_MODEL_ID`: The model name to use
- Example: `export OLLAMA_CHAT_MODEL_ID="qwen2.5:8b"`
- Must be a model you have pulled with Ollama
### For OpenAI Client with Ollama (`ollama_with_openai_chat_client.py`)
@@ -3,7 +3,7 @@
import asyncio
from datetime import datetime
from agent_framework_ollama import OllamaChatClient
from agent_framework.ollama import OllamaChatClient
"""
Ollama Agent Basic Example
@@ -3,8 +3,7 @@
import asyncio
from agent_framework import TextReasoningContent
from agent_framework_ollama import OllamaChatClient
from agent_framework.ollama import OllamaChatClient
"""
Ollama Agent Reasoning Example
@@ -3,12 +3,19 @@
import asyncio
from datetime import datetime
from agent_framework_ollama import OllamaChatClient
from agent_framework.ollama import OllamaChatClient
# Ensure to install Ollama and have a model running locally before running the sample
# Not all Models support function calling, to test function calling try llama3.2
# Set the model to use via the OLLAMA_CHAT_MODEL_ID environment variable or modify the code below.
# https://ollama.com/
"""
Ollama Chat Client Example
This sample demonstrates using the native Ollama Chat Client directly.
Ensure to install Ollama and have a model running locally before running the sample.
Not all Models support function calling, to test function calling try llama3.2
Set the model to use via the OLLAMA_CHAT_MODEL_ID environment variable or modify the code below.
https://ollama.com/
"""
def get_time():
@@ -3,8 +3,7 @@
import asyncio
from agent_framework import ChatMessage, DataContent, Role, TextContent
from agent_framework_ollama import OllamaChatClient
from agent_framework.ollama import OllamaChatClient
"""
Ollama Agent Multimodal Example
@@ -8,8 +8,11 @@ This folder contains an example demonstrating how to use the Redis context provi
| File | Description |
|------|-------------|
| [`azure_redis_conversation.py`](azure_redis_conversation.py) | Demonstrates conversation persistence with RedisChatMessageStore and Azure Redis with Azure AD (Entra ID) authentication using credential provider. |
| [`redis_basics.py`](redis_basics.py) | Shows standalone provider usage and agent integration. Demonstrates writing messages to Redis, retrieving context via full‑text or hybrid vector search, and persisting preferences across threads. Also includes a simple tool example whose outputs are remembered. |
| [`redis_threads.py`](redis_threads.py) | Demonstrates thread scoping. Includes: (1) global thread scope with a fixed `thread_id` shared across operations; (2) per‑operation thread scope where `scope_to_per_operation_thread_id=True` binds memory to a single thread for the provider’s lifetime; and (3) multiple agents with isolated memory via different `agent_id` values. |
| [`redis_conversation.py`](redis_conversation.py) | Simple example showing conversation persistence with RedisChatMessageStore using traditional connection string authentication. |
| [`redis_threads.py`](redis_threads.py) | Demonstrates thread scoping. Includes: (1) global thread scope with a fixed `thread_id` shared across operations; (2) per‑operation thread scope where `scope_to_per_operation_thread_id=True` binds memory to a single thread for the provider's lifetime; and (3) multiple agents with isolated memory via different `agent_id` values. |
## Prerequisites
@@ -0,0 +1,123 @@
# Copyright (c) Microsoft. All rights reserved.
"""Azure Managed Redis Chat Message Store with Azure AD Authentication
This example demonstrates how to use Azure Managed Redis with Azure AD authentication
to persist conversational details using RedisChatMessageStore.
Requirements:
- Azure Managed Redis instance with Azure AD authentication enabled
- Azure credentials configured (az login or managed identity)
- agent-framework-redis: pip install agent-framework-redis
- azure-identity: pip install azure-identity
Environment Variables:
- AZURE_REDIS_HOST: Your Azure Managed Redis host (e.g., myredis.redis.cache.windows.net)
- OPENAI_API_KEY: Your OpenAI API key
- OPENAI_CHAT_MODEL_ID: OpenAI model (e.g., gpt-4o-mini)
- AZURE_USER_OBJECT_ID: Your Azure AD User Object ID for authentication
"""
import asyncio
import os
from agent_framework.openai import OpenAIChatClient
from agent_framework.redis import RedisChatMessageStore
from azure.identity.aio import AzureCliCredential
from redis.credentials import CredentialProvider
class AzureCredentialProvider(CredentialProvider):
"""Credential provider for Azure AD authentication with Redis Enterprise."""
def __init__(self, azure_credential: AzureCliCredential, user_object_id: str):
self.azure_credential = azure_credential
self.user_object_id = user_object_id
async def get_credentials_async(self) -> tuple[str] | tuple[str, str]:
"""Get Azure AD token for Redis authentication.
Returns (username, token) where username is the Azure user's Object ID.
"""
token = await self.azure_credential.get_token("https://redis.azure.com/.default")
return (self.user_object_id, token.token)
async def main() -> None:
redis_host = os.environ.get("AZURE_REDIS_HOST")
if not redis_host:
print("ERROR: Set AZURE_REDIS_HOST environment variable")
return
# For Azure Redis with Entra ID, username must be your Object ID
user_object_id = os.environ.get("AZURE_USER_OBJECT_ID")
if not user_object_id:
print("ERROR: Set AZURE_USER_OBJECT_ID environment variable")
print("Get your Object ID from the Azure Portal")
return
# Create Azure CLI credential provider (uses 'az login' credentials)
azure_credential = AzureCliCredential()
credential_provider = AzureCredentialProvider(azure_credential, user_object_id)
thread_id = "azure_test_thread"
# Factory for creating Azure Redis chat message store
chat_message_store_factory = lambda: RedisChatMessageStore(
credential_provider=credential_provider,
host=redis_host,
port=10000,
ssl=True,
thread_id=thread_id,
key_prefix="chat_messages",
max_messages=100,
)
# Create chat client
client = OpenAIChatClient()
# Create agent with Azure Redis store
agent = client.create_agent(
name="AzureRedisAssistant",
instructions="You are a helpful assistant.",
chat_message_store_factory=chat_message_store_factory,
)
# Conversation
query = "Remember that I enjoy gumbo"
result = await agent.run(query)
print("User: ", query)
print("Agent: ", result)
# Ask the agent to recall the stored preference; it should retrieve from memory
query = "What do I enjoy?"
result = await agent.run(query)
print("User: ", query)
print("Agent: ", result)
query = "What did I say to you just now?"
result = await agent.run(query)
print("User: ", query)
print("Agent: ", result)
query = "Remember that I have a meeting at 3pm tomorrow"
result = await agent.run(query)
print("User: ", query)
print("Agent: ", result)
query = "Tulips are red"
result = await agent.run(query)
print("User: ", query)
print("Agent: ", result)
query = "What was the first thing I said to you this conversation?"
result = await agent.run(query)
print("User: ", query)
print("Agent: ", result)
# Cleanup
await azure_credential.close()
if __name__ == "__main__":
asyncio.run(main())
@@ -91,7 +91,7 @@ async def main() -> None:
print("User: ", query)
print("Agent: ", result)
query = "Remember that anyone who does not clean shrimp will be eaten by a shark"
query = "Remember that I have a meeting at 3pm tomorro"
result = await agent.run(query)
print("User: ", query)
print("Agent: ", result)
@@ -45,6 +45,7 @@ Once comfortable with these, explore the rest of the samples below.
| Workflow as Agent (Reflection Pattern) | [agents/workflow_as_agent_reflection_pattern.py](./agents/workflow_as_agent_reflection_pattern.py) | Wrap a workflow so it can behave like an agent (reflection pattern) |
| Workflow as Agent + HITL | [agents/workflow_as_agent_human_in_the_loop.py](./agents/workflow_as_agent_human_in_the_loop.py) | Extend workflow-as-agent with human-in-the-loop capability |
| Workflow as Agent with Thread | [agents/workflow_as_agent_with_thread.py](./agents/workflow_as_agent_with_thread.py) | Use AgentThread to maintain conversation history across workflow-as-agent invocations |
| Workflow as Agent kwargs | [agents/workflow_as_agent_kwargs.py](./agents/workflow_as_agent_kwargs.py) | Pass custom context (data, user tokens) via kwargs through workflow.as_agent() to @ai_function tools |
| Handoff Workflow as Agent | [agents/handoff_workflow_as_agent.py](./agents/handoff_workflow_as_agent.py) | Use a HandoffBuilder workflow as an agent with HITL via FunctionCallContent/FunctionResultContent |
### checkpoint
@@ -64,6 +65,7 @@ Once comfortable with these, explore the rest of the samples below.
| Sub-Workflow (Basics) | [composition/sub_workflow_basics.py](./composition/sub_workflow_basics.py) | Wrap a workflow as an executor and orchestrate sub-workflows |
| Sub-Workflow: Request Interception | [composition/sub_workflow_request_interception.py](./composition/sub_workflow_request_interception.py) | Intercept and forward sub-workflow requests using @handler for SubWorkflowRequestMessage |
| Sub-Workflow: Parallel Requests | [composition/sub_workflow_parallel_requests.py](./composition/sub_workflow_parallel_requests.py) | Multiple specialized interceptors handling different request types from same sub-workflow |
| Sub-Workflow: kwargs Propagation | [composition/sub_workflow_kwargs.py](./composition/sub_workflow_kwargs.py) | Pass custom context (user tokens, config) from parent workflow through to sub-workflow agents |
### control-flow
@@ -75,6 +77,7 @@ Once comfortable with these, explore the rest of the samples below.
| Switch-Case Edge Group | [control-flow/switch_case_edge_group.py](./control-flow/switch_case_edge_group.py) | Switch-case branching using classifier outputs |
| Multi-Selection Edge Group | [control-flow/multi_selection_edge_group.py](./control-flow/multi_selection_edge_group.py) | Select one or many targets dynamically (subset fan-out) |
| Simple Loop | [control-flow/simple_loop.py](./control-flow/simple_loop.py) | Feedback loop where an agent judges ABOVE/BELOW/MATCHED |
| Workflow Cancellation | [control-flow/workflow_cancellation.py](./control-flow/workflow_cancellation.py) | Cancel a running workflow using asyncio tasks |
### human-in-the-loop
@@ -0,0 +1,140 @@
# Copyright (c) Microsoft. All rights reserved.
import asyncio
import json
from typing import Annotated, Any
from agent_framework import SequentialBuilder, ai_function
from agent_framework.openai import OpenAIChatClient
from pydantic import Field
"""
Sample: Workflow as Agent with kwargs Propagation to @ai_function Tools
This sample demonstrates how to flow custom context (skill data, user tokens, etc.)
through a workflow exposed via .as_agent() to @ai_function tools using the **kwargs pattern.
Key Concepts:
- Build a workflow using SequentialBuilder (or any builder pattern)
- Expose the workflow as a reusable agent via workflow.as_agent()
- Pass custom context as kwargs when invoking workflow_agent.run() or run_stream()
- kwargs are stored in SharedState and propagated to all agent invocations
- @ai_function tools receive kwargs via **kwargs parameter
When to use workflow.as_agent():
- To treat an entire workflow orchestration as a single agent
- To compose workflows into higher-level orchestrations
- To maintain a consistent agent interface for callers
Prerequisites:
- OpenAI environment variables configured
"""
# Define tools that accept custom context via **kwargs
@ai_function
def get_user_data(
query: Annotated[str, Field(description="What user data to retrieve")],
**kwargs: Any,
) -> str:
"""Retrieve user-specific data based on the authenticated context."""
user_token = kwargs.get("user_token", {})
user_name = user_token.get("user_name", "anonymous")
access_level = user_token.get("access_level", "none")
print(f"\n[get_user_data] Received kwargs keys: {list(kwargs.keys())}")
print(f"[get_user_data] User: {user_name}")
print(f"[get_user_data] Access level: {access_level}")
return f"Retrieved data for user {user_name} with {access_level} access: {query}"
@ai_function
def call_api(
endpoint_name: Annotated[str, Field(description="Name of the API endpoint to call")],
**kwargs: Any,
) -> str:
"""Call an API using the configured endpoints from custom_data."""
custom_data = kwargs.get("custom_data", {})
api_config = custom_data.get("api_config", {})
base_url = api_config.get("base_url", "unknown")
endpoints = api_config.get("endpoints", {})
print(f"\n[call_api] Received kwargs keys: {list(kwargs.keys())}")
print(f"[call_api] Base URL: {base_url}")
print(f"[call_api] Available endpoints: {list(endpoints.keys())}")
if endpoint_name in endpoints:
return f"Called {base_url}{endpoints[endpoint_name]} successfully"
return f"Endpoint '{endpoint_name}' not found in configuration"
async def main() -> None:
print("=" * 70)
print("Workflow as Agent kwargs Flow Demo")
print("=" * 70)
# Create chat client
chat_client = OpenAIChatClient()
# Create agent with tools that use kwargs
agent = chat_client.create_agent(
name="assistant",
instructions=(
"You are a helpful assistant. Use the available tools to help users. "
"When asked about user data, use get_user_data. "
"When asked to call an API, use call_api."
),
tools=[get_user_data, call_api],
)
# Build a sequential workflow
workflow = SequentialBuilder().participants([agent]).build()
# Expose the workflow as an agent using .as_agent()
workflow_agent = workflow.as_agent(name="WorkflowAgent")
# Define custom context that will flow to ai_functions via kwargs
custom_data = {
"api_config": {
"base_url": "https://api.example.com",
"endpoints": {
"users": "/v1/users",
"orders": "/v1/orders",
"products": "/v1/products",
},
},
}
user_token = {
"user_name": "bob@contoso.com",
"access_level": "admin",
}
print("\nCustom Data being passed:")
print(json.dumps(custom_data, indent=2))
print(f"\nUser: {user_token['user_name']}")
print("\n" + "-" * 70)
print("Workflow Agent Execution (watch for [tool_name] logs showing kwargs received):")
print("-" * 70)
# Run workflow agent with kwargs - these will flow through to ai_functions
# Note: kwargs are passed to workflow_agent.run_stream() just like workflow.run_stream()
print("\n===== Streaming Response =====")
async for update in workflow_agent.run_stream(
"Please get my user data and then call the users API endpoint.",
custom_data=custom_data,
user_token=user_token,
):
if update.text:
print(update.text, end="", flush=True)
print()
print("\n" + "=" * 70)
print("Sample Complete")
print("=" * 70)
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,143 @@
# Copyright (c) Microsoft. All rights reserved.
import asyncio
import json
from typing import Annotated, Any
from agent_framework import (
ChatMessage,
SequentialBuilder,
WorkflowExecutor,
WorkflowOutputEvent,
ai_function,
)
from agent_framework.openai import OpenAIChatClient
"""
Sample: Sub-Workflow kwargs Propagation
This sample demonstrates how custom context (kwargs) flows from a parent workflow
through to agents in sub-workflows. When you pass kwargs to the parent workflow's
run_stream() or run(), they automatically propagate to nested sub-workflows.
Key Concepts:
- kwargs passed to parent workflow.run_stream() propagate to sub-workflows
- Sub-workflow agents receive the same kwargs as the parent workflow
- Works with nested WorkflowExecutor compositions at any depth
- Useful for passing authentication tokens, configuration, or request context
Prerequisites:
- OpenAI environment variables configured
"""
# Define tools that access custom context via **kwargs
@ai_function
def get_authenticated_data(
resource: Annotated[str, "The resource to fetch"],
**kwargs: Any,
) -> str:
"""Fetch data using the authenticated user context from kwargs."""
user_token = kwargs.get("user_token", {})
user_name = user_token.get("user_name", "anonymous")
access_level = user_token.get("access_level", "none")
print(f"\n[get_authenticated_data] kwargs keys: {list(kwargs.keys())}")
print(f"[get_authenticated_data] User: {user_name}, Access: {access_level}")
return f"Fetched '{resource}' for user {user_name} ({access_level} access)"
@ai_function
def call_configured_service(
service_name: Annotated[str, "Name of the service to call"],
**kwargs: Any,
) -> str:
"""Call a service using configuration from kwargs."""
config = kwargs.get("service_config", {})
services = config.get("services", {})
print(f"\n[call_configured_service] kwargs keys: {list(kwargs.keys())}")
print(f"[call_configured_service] Available services: {list(services.keys())}")
if service_name in services:
endpoint = services[service_name]
return f"Called service '{service_name}' at {endpoint}"
return f"Service '{service_name}' not found in configuration"
async def main() -> None:
print("=" * 70)
print("Sub-Workflow kwargs Propagation Demo")
print("=" * 70)
# Create chat client
chat_client = OpenAIChatClient()
# Create an agent with tools that use kwargs
inner_agent = chat_client.create_agent(
name="data_agent",
instructions=(
"You are a data access agent. Use the available tools to help users. "
"When asked to fetch data, use get_authenticated_data. "
"When asked to call a service, use call_configured_service."
),
tools=[get_authenticated_data, call_configured_service],
)
# Build the inner (sub) workflow with the agent
inner_workflow = SequentialBuilder().participants([inner_agent]).build()
# Wrap the inner workflow in a WorkflowExecutor to use it as a sub-workflow
subworkflow_executor = WorkflowExecutor(
workflow=inner_workflow,
id="data_subworkflow",
)
# Build the outer (parent) workflow containing the sub-workflow
outer_workflow = SequentialBuilder().participants([subworkflow_executor]).build()
# Define custom context that will flow through to the sub-workflow's agent
user_token = {
"user_name": "alice@contoso.com",
"access_level": "admin",
"session_id": "sess_12345",
}
service_config = {
"services": {
"users": "https://api.example.com/v1/users",
"orders": "https://api.example.com/v1/orders",
"inventory": "https://api.example.com/v1/inventory",
},
"timeout": 30,
}
print("\nContext being passed to parent workflow:")
print(f" user_token: {json.dumps(user_token, indent=4)}")
print(f" service_config: {json.dumps(service_config, indent=4)}")
print("\n" + "-" * 70)
print("Workflow Execution (kwargs flow: parent -> sub-workflow -> agent -> tool):")
print("-" * 70)
# Run the OUTER workflow with kwargs
# These kwargs will automatically propagate to the inner sub-workflow
async for event in outer_workflow.run_stream(
"Please fetch my profile data and then call the users service.",
user_token=user_token,
service_config=service_config,
):
if isinstance(event, WorkflowOutputEvent):
output_data = event.data
if isinstance(output_data, list):
for item in output_data: # type: ignore
if isinstance(item, ChatMessage) and item.text:
print(f"\n[Final Answer]: {item.text}")
print("\n" + "=" * 70)
print("Sample Complete - kwargs successfully flowed through sub-workflow!")
print("=" * 70)
if __name__ == "__main__":
asyncio.run(main())
@@ -0,0 +1,103 @@
# Copyright (c) Microsoft. All rights reserved.
import asyncio
from agent_framework import WorkflowBuilder, WorkflowContext, executor
from typing_extensions import Never
"""
Sample: Workflow Cancellation
A three-step workflow where each step takes 2 seconds. We cancel it after 3 seconds
to demonstrate mid-execution cancellation using asyncio tasks.
Purpose:
Show how to cancel a running workflow by wrapping it in an asyncio.Task. This pattern
works with both workflow.run() and workflow.run_stream(). Useful for implementing
timeouts, graceful shutdown, or A2A executors that need cancellation support.
Prerequisites:
- No external services required.
"""
@executor(id="step1")
async def step1(text: str, ctx: WorkflowContext[str]) -> None:
"""First step - simulates 2 seconds of work."""
print("[Step1] Starting...")
await asyncio.sleep(2)
print("[Step1] Done")
await ctx.send_message(text.upper())
@executor(id="step2")
async def step2(text: str, ctx: WorkflowContext[str]) -> None:
"""Second step - simulates 2 seconds of work."""
print("[Step2] Starting...")
await asyncio.sleep(2)
print("[Step2] Done")
await ctx.send_message(text + "!")
@executor(id="step3")
async def step3(text: str, ctx: WorkflowContext[Never, str]) -> None:
"""Final step - simulates 2 seconds of work."""
print("[Step3] Starting...")
await asyncio.sleep(2)
print("[Step3] Done")
await ctx.yield_output(f"Result: {text}")
def build_workflow():
"""Build a simple 3-step sequential workflow (~6 seconds total)."""
return (
WorkflowBuilder()
.register_executor(lambda: step1, name="step1")
.register_executor(lambda: step2, name="step2")
.register_executor(lambda: step3, name="step3")
.add_edge("step1", "step2")
.add_edge("step2", "step3")
.set_start_executor("step1")
.build()
)
async def run_with_cancellation() -> None:
"""Cancel the workflow after 3 seconds (mid-execution during Step2)."""
print("=== Run with cancellation ===\n")
workflow = build_workflow()
# Wrap workflow.run() in a task to enable cancellation
task = asyncio.create_task(workflow.run("hello world"))
# Wait 3 seconds (Step1 completes, Step2 is mid-execution), then cancel
await asyncio.sleep(3)
print("\n--- Cancelling workflow ---\n")
task.cancel()
try:
await task
except asyncio.CancelledError:
print("Workflow was cancelled")
async def run_to_completion() -> None:
"""Let the workflow run to completion and get the result."""
print("=== Run to completion ===\n")
workflow = build_workflow()
# Run without cancellation - await the result directly
result = await workflow.run("hello world")
print(f"\nWorkflow completed with output: {result.get_outputs()}")
async def main() -> None:
"""Demonstrate both cancellation and completion scenarios."""
await run_with_cancellation()
print("\n")
await run_to_completion()
if __name__ == "__main__":
asyncio.run(main())
@@ -6,14 +6,12 @@ from dataclasses import dataclass
from agent_framework import (
AgentExecutorRequest,
AgentExecutorResponse,
AgentRunEvent,
ChatAgent,
ChatMessage,
Executor,
Role,
WorkflowBuilder,
WorkflowContext,
WorkflowOutputEvent,
WorkflowViz,
handler,
)
@@ -124,7 +122,7 @@ def create_legal_agent() -> ChatAgent:
async def main() -> None:
"""Build and run the concurrent workflow with visualization."""
# 1) Build a simple fan-out/fan-in workflow
# Build a simple fan-out/fan-in workflow
workflow = (
WorkflowBuilder()
.register_agent(create_researcher_agent, name="researcher")
@@ -138,31 +136,22 @@ async def main() -> None:
.build()
)
# 1.5) Generate workflow visualization
# Generate workflow visualization
print("Generating workflow visualization...")
viz = WorkflowViz(workflow)
# Print out the mermaid string.
print("Mermaid string: \n=======")
print(viz.to_mermaid())
print("=======")
# Print out the DiGraph string.
# Print out the DiGraph string with internal executors.
print("DiGraph string: \n=======")
print(viz.to_digraph())
print(viz.to_digraph(include_internal_executors=True))
print("=======")
# Export the DiGraph visualization as SVG.
svg_file = viz.export(format="svg")
print(f"SVG file saved to: {svg_file}")
# 2) Run with a single prompt
async for event in workflow.run_stream("We are launching a new budget-friendly electric bike for urban commuters."):
if isinstance(event, AgentRunEvent):
# Show which agent ran and what step completed.
print(event)
elif isinstance(event, WorkflowOutputEvent):
print("===== Final Aggregated Output =====")
print(event.data)
if __name__ == "__main__":
asyncio.run(main())
+49 -47
View File
@@ -53,7 +53,7 @@ overrides = [
[[package]]
name = "a2a-sdk"
version = "0.3.21"
version = "0.3.22"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "google-api-core", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
@@ -62,9 +62,9 @@ dependencies = [
{ name = "protobuf", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "pydantic", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/13/d3/bff391f9be2e56bdccb8f1fb7058a88e03582010890f1c3622174244972d/a2a_sdk-0.3.21.tar.gz", hash = "sha256:627aa841187540f5975625c7f0383754795d55fe3e774073f2741c5bdd7c9972", size = 231305, upload-time = "2025-12-12T17:05:31.755Z" }
sdist = { url = "https://files.pythonhosted.org/packages/92/a3/76f2d94a32a1b0dc760432d893a09ec5ed31de5ad51b1ef0f9d199ceb260/a2a_sdk-0.3.22.tar.gz", hash = "sha256:77a5694bfc4f26679c11b70c7f1062522206d430b34bc1215cfbb1eba67b7e7d", size = 231535, upload-time = "2025-12-16T18:39:21.19Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/33/02/26f02c84d7037dbd5c662140434f6e5085e3fe13d9c1e85984852af06930/a2a_sdk-0.3.21-py3-none-any.whl", hash = "sha256:5d817dd424a4396587416241de1342a97d6a8bb1c0b31a47d8dc6546b389ab68", size = 144138, upload-time = "2025-12-12T17:05:29.8Z" },
{ url = "https://files.pythonhosted.org/packages/64/e8/f4e39fd1cf0b3c4537b974637143f3ebfe1158dad7232d9eef15666a81ba/a2a_sdk-0.3.22-py3-none-any.whl", hash = "sha256:b98701135bb90b0ff85d35f31533b6b7a299bf810658c1c65f3814a6c15ea385", size = 144347, upload-time = "2025-12-16T18:39:19.218Z" },
]
[[package]]
@@ -335,6 +335,7 @@ all = [
{ name = "agent-framework-devui", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "agent-framework-lab", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "agent-framework-mem0", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "agent-framework-ollama", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "agent-framework-purview", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "agent-framework-redis", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
@@ -353,6 +354,7 @@ requires-dist = [
{ name = "agent-framework-devui", marker = "extra == 'all'", editable = "packages/devui" },
{ name = "agent-framework-lab", marker = "extra == 'all'", editable = "packages/lab" },
{ name = "agent-framework-mem0", marker = "extra == 'all'", editable = "packages/mem0" },
{ name = "agent-framework-ollama", marker = "extra == 'all'", editable = "packages/ollama" },
{ name = "agent-framework-purview", marker = "extra == 'all'", editable = "packages/purview" },
{ name = "agent-framework-redis", marker = "extra == 'all'", editable = "packages/redis" },
{ name = "azure-identity", specifier = ">=1,<2" },
@@ -535,7 +537,7 @@ requires-dist = [
[[package]]
name = "agent-framework-ollama"
version = "0.1.0b1"
version = "1.0.0b251216"
source = { editable = "packages/ollama" }
dependencies = [
{ name = "agent-framework-core", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
@@ -1341,7 +1343,7 @@ name = "clr-loader"
version = "0.2.9"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "cffi", marker = "(python_full_version < '3.14' and sys_platform == 'darwin') or (python_full_version < '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and sys_platform == 'win32')" },
{ name = "cffi", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/54/c2/da52aaf19424e3f0abec003d08dd1ccae52c88a3b41e31151a03bed18488/clr_loader-0.2.9.tar.gz", hash = "sha256:6af3d582c3de55ce9e9e676d2b3dbf6bc680c4ea8f76c58786739a5bdcf6b52d", size = 84829, upload-time = "2025-12-05T16:57:12.466Z" }
wheels = [
@@ -1820,7 +1822,7 @@ name = "exceptiongroup"
version = "1.3.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions", marker = "(python_full_version < '3.11' and sys_platform == 'darwin') or (python_full_version < '3.11' and sys_platform == 'linux') or (python_full_version < '3.11' and sys_platform == 'win32')" },
{ name = "typing-extensions", marker = "(python_full_version < '3.13' and sys_platform == 'darwin') or (python_full_version < '3.13' and sys_platform == 'linux') or (python_full_version < '3.13' and sys_platform == 'win32')" },
]
sdist = { url = "https://files.pythonhosted.org/packages/50/79/66800aadf48771f6b62f7eb014e352e5d06856655206165d775e675a02c9/exceptiongroup-1.3.1.tar.gz", hash = "sha256:8b412432c6055b0b7d14c310000ae93352ed6754f70fa8f7c34141f91c4e3219", size = 30371, upload-time = "2025-11-21T23:01:54.787Z" }
wheels = [
@@ -2940,7 +2942,7 @@ wheels = [
[[package]]
name = "langfuse"
version = "3.10.6"
version = "3.10.7"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "backoff", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
@@ -2954,9 +2956,9 @@ dependencies = [
{ name = "requests", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "wrapt", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/e6/70/4ff19dd1085bb4d5007f008a696c8cf989a0ad76eabc512a5cd19ee4a0b7/langfuse-3.10.6.tar.gz", hash = "sha256:fced9ca0416ba7499afa45fbedf831afc0ec824cb283719b9cf429bf5713f205", size = 223656, upload-time = "2025-12-12T13:29:24.048Z" }
sdist = { url = "https://files.pythonhosted.org/packages/44/62/f46319500aff363bedf5dbbcb3afa0fdd5788c6faf901eee8fce27f9643c/langfuse-3.10.7.tar.gz", hash = "sha256:64eaec6923e6c61baa62b18516f5f37c011d55caa409b2214c1819fe01cd1056", size = 223808, upload-time = "2025-12-16T15:36:55.959Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/ce/f0/fac7d56ce1136afbbebaddd1dc119fb1b94b5a7489944d0b4c2dcee99ed7/langfuse-3.10.6-py3-none-any.whl", hash = "sha256:36ca490cd64e372b1b94c28063b3fea39b1a8446cabd20172b524d01011a34e1", size = 399347, upload-time = "2025-12-12T13:29:22.462Z" },
{ url = "https://files.pythonhosted.org/packages/ed/10/67fd830dfd4ab33f66a60d6b61395601b478dec78e15065d4d6fb4d74610/langfuse-3.10.7-py3-none-any.whl", hash = "sha256:206dabd786ca64c403b5552488515ff08d4b2d55cebf00255ed1ae3c59794d17", size = 399345, upload-time = "2025-12-16T15:36:54.716Z" },
]
[[package]]
@@ -3661,11 +3663,11 @@ wheels = [
[[package]]
name = "narwhals"
version = "2.13.0"
version = "2.14.0"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/89/ea/f82ef99ced4d03c33bb314c9b84a08a0a86c448aaa11ffd6256b99538aa5/narwhals-2.13.0.tar.gz", hash = "sha256:ee94c97f4cf7cfeebbeca8d274784df8b3d7fd3f955ce418af998d405576fdd9", size = 594555, upload-time = "2025-12-01T13:54:05.329Z" }
sdist = { url = "https://files.pythonhosted.org/packages/4a/84/897fe7b6406d436ef312e57e5a1a13b4a5e7e36d1844e8d934ce8880e3d3/narwhals-2.14.0.tar.gz", hash = "sha256:98be155c3599db4d5c211e565c3190c398c87e7bf5b3cdb157dece67641946e0", size = 600648, upload-time = "2025-12-16T11:29:13.458Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/87/0d/1861d1599571974b15b025e12b142d8e6b42ad66c8a07a89cb0fc21f1e03/narwhals-2.13.0-py3-none-any.whl", hash = "sha256:9b795523c179ca78204e3be53726da374168f906e38de2ff174c2363baaaf481", size = 426407, upload-time = "2025-12-01T13:54:03.861Z" },
{ url = "https://files.pythonhosted.org/packages/79/3e/b8ecc67e178919671695f64374a7ba916cf0adbf86efedc6054f38b5b8ae/narwhals-2.14.0-py3-none-any.whl", hash = "sha256:b56796c9a00179bd757d15282c540024e1d5c910b19b8c9944d836566c030acf", size = 430788, upload-time = "2025-12-16T11:29:11.699Z" },
]
[[package]]
@@ -3863,7 +3865,7 @@ wheels = [
[[package]]
name = "openai"
version = "2.12.0"
version = "2.13.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "anyio", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
@@ -3875,9 +3877,9 @@ dependencies = [
{ name = "tqdm", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "typing-extensions", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/86/f9/fb8abeb4cdba6f24daf3d7781f42ceb1be1ff579eb20705899e617dd95f1/openai-2.12.0.tar.gz", hash = "sha256:cc6dcbcb8bccf05976d983f6516c5c1f447b71c747720f1530b61e8f858bcbc9", size = 626183, upload-time = "2025-12-15T16:17:15.097Z" }
sdist = { url = "https://files.pythonhosted.org/packages/0f/39/8e347e9fda125324d253084bb1b82407e5e3c7777a03dc398f79b2d95626/openai-2.13.0.tar.gz", hash = "sha256:9ff633b07a19469ec476b1e2b5b26c5ef700886524a7a72f65e6f0b5203142d5", size = 626583, upload-time = "2025-12-16T18:19:44.387Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/c3/a1/f055214448cb4b176e89459d889af9615fe7d927634fb5a2cecfb7674bc5/openai-2.12.0-py3-none-any.whl", hash = "sha256:7177998ce49ba3f90bcce8b5769a6666d90b1f328f0518d913aaec701271485a", size = 1066590, upload-time = "2025-12-15T16:17:13.301Z" },
{ url = "https://files.pythonhosted.org/packages/bb/d5/eb52edff49d3d5ea116e225538c118699ddeb7c29fa17ec28af14bc10033/openai-2.13.0-py3-none-any.whl", hash = "sha256:746521065fed68df2f9c2d85613bb50844343ea81f60009b60e6a600c9352c79", size = 1066837, upload-time = "2025-12-16T18:19:43.124Z" },
]
[[package]]
@@ -4453,7 +4455,7 @@ wheels = [
[[package]]
name = "posthog"
version = "7.0.1"
version = "7.4.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "backoff", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
@@ -4463,9 +4465,9 @@ dependencies = [
{ name = "six", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "typing-extensions", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/a2/d4/b9afe855a8a7a1bf4459c28ae4c300b40338122dc850acabefcf2c3df24d/posthog-7.0.1.tar.gz", hash = "sha256:21150562c2630a599c1d7eac94bc5c64eb6f6acbf3ff52ccf1e57345706db05a", size = 126985, upload-time = "2025-11-15T12:44:22.465Z" }
sdist = { url = "https://files.pythonhosted.org/packages/14/e5/5262d1604a3eb19b23d4e896bce87b4603fd39ec366a96b27e19e3299aef/posthog-7.4.0.tar.gz", hash = "sha256:1fb97b11960e24fcf0b80f0a6450b2311478e5a3ee6ea3c6f9284ff89060a876", size = 143780, upload-time = "2025-12-16T23:42:05.829Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/05/0c/8b6b20b0be71725e6e8a32dcd460cdbf62fe6df9bc656a650150dc98fedd/posthog-7.0.1-py3-none-any.whl", hash = "sha256:efe212d8d88a9ba80a20c588eab4baf4b1a5e90e40b551160a5603bb21e96904", size = 145234, upload-time = "2025-11-15T12:44:21.247Z" },
{ url = "https://files.pythonhosted.org/packages/9f/8b/13066693d7a6f94fb5da3407417bbbc3f6aa8487051294d0ef766c1567fa/posthog-7.4.0-py3-none-any.whl", hash = "sha256:f9d4e32c1c0f2110256b1aae7046ed90af312c1dbb1eecc6a9cb427733b22970", size = 166079, upload-time = "2025-12-16T23:42:04.33Z" },
]
[[package]]
@@ -4473,8 +4475,8 @@ name = "powerfx"
version = "0.0.33"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "cffi", marker = "(python_full_version < '3.14' and sys_platform == 'darwin') or (python_full_version < '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and sys_platform == 'win32')" },
{ name = "pythonnet", marker = "(python_full_version < '3.14' and sys_platform == 'darwin') or (python_full_version < '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and sys_platform == 'win32')" },
{ name = "cffi", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "pythonnet", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/5e/41/8f95f72f4f3b7ea54357c449bf5bd94813b6321dec31db9ffcbf578e2fa3/powerfx-0.0.33.tar.gz", hash = "sha256:85e8330bef8a7a207c3e010aa232df0ae38825e94d590c73daf3a3f44115cb09", size = 3236647, upload-time = "2025-11-20T19:31:09.414Z" }
wheels = [
@@ -4483,7 +4485,7 @@ wheels = [
[[package]]
name = "pre-commit"
version = "4.5.0"
version = "4.5.1"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "cfgv", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
@@ -4492,9 +4494,9 @@ dependencies = [
{ name = "pyyaml", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
{ name = "virtualenv", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/f4/9b/6a4ffb4ed980519da959e1cf3122fc6cb41211daa58dbae1c73c0e519a37/pre_commit-4.5.0.tar.gz", hash = "sha256:dc5a065e932b19fc1d4c653c6939068fe54325af8e741e74e88db4d28a4dd66b", size = 198428, upload-time = "2025-11-22T21:02:42.304Z" }
sdist = { url = "https://files.pythonhosted.org/packages/40/f1/6d86a29246dfd2e9b6237f0b5823717f60cad94d47ddc26afa916d21f525/pre_commit-4.5.1.tar.gz", hash = "sha256:eb545fcff725875197837263e977ea257a402056661f09dae08e4b149b030a61", size = 198232, upload-time = "2025-12-16T21:14:33.552Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/5d/c4/b2d28e9d2edf4f1713eb3c29307f1a63f3d67cf09bdda29715a36a68921a/pre_commit-4.5.0-py2.py3-none-any.whl", hash = "sha256:25e2ce09595174d9c97860a95609f9f852c0614ba602de3561e267547f2335e1", size = 226429, upload-time = "2025-11-22T21:02:40.836Z" },
{ url = "https://files.pythonhosted.org/packages/5d/19/fd3ef348460c80af7bb4669ea7926651d1f95c23ff2df18b9d24bab4f3fa/pre_commit-4.5.1-py2.py3-none-any.whl", hash = "sha256:3b3afd891e97337708c1674210f8eba659b52a38ea5f822ff142d10786221f77", size = 226437, upload-time = "2025-12-16T21:14:32.409Z" },
]
[[package]]
@@ -4613,14 +4615,14 @@ wheels = [
[[package]]
name = "proto-plus"
version = "1.26.1"
version = "1.27.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "protobuf", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/f4/ac/87285f15f7cce6d4a008f33f1757fb5a13611ea8914eb58c3d0d26243468/proto_plus-1.26.1.tar.gz", hash = "sha256:21a515a4c4c0088a773899e23c7bbade3d18f9c66c73edd4c7ee3816bc96a012", size = 56142, upload-time = "2025-03-10T15:54:38.843Z" }
sdist = { url = "https://files.pythonhosted.org/packages/01/89/9cbe2f4bba860e149108b683bc2efec21f14d5f7ed6e25562ad86acbc373/proto_plus-1.27.0.tar.gz", hash = "sha256:873af56dd0d7e91836aee871e5799e1c6f1bda86ac9a983e0bb9f0c266a568c4", size = 56158, upload-time = "2025-12-16T13:46:25.729Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/4e/6d/280c4c2ce28b1593a19ad5239c8b826871fc6ec275c21afc8e1820108039/proto_plus-1.26.1-py3-none-any.whl", hash = "sha256:13285478c2dcf2abb829db158e1047e2f1e8d63a077d94263c2b88b043c75a66", size = 50163, upload-time = "2025-03-10T15:54:37.335Z" },
{ url = "https://files.pythonhosted.org/packages/cd/24/3b7a0818484df9c28172857af32c2397b6d8fcd99d9468bd4684f98ebf0a/proto_plus-1.27.0-py3-none-any.whl", hash = "sha256:1baa7f81cf0f8acb8bc1f6d085008ba4171eaf669629d1b6d1673b21ed1c0a82", size = 50205, upload-time = "2025-12-16T13:46:24.76Z" },
]
[[package]]
@@ -5143,7 +5145,7 @@ name = "pythonnet"
version = "3.0.5"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "clr-loader", marker = "(python_full_version < '3.14' and sys_platform == 'darwin') or (python_full_version < '3.14' and sys_platform == 'linux') or (python_full_version < '3.14' and sys_platform == 'win32')" },
{ name = "clr-loader", marker = "sys_platform == 'darwin' or sys_platform == 'linux' or sys_platform == 'win32'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/9a/d6/1afd75edd932306ae9bd2c2d961d603dc2b52fcec51b04afea464f1f6646/pythonnet-3.0.5.tar.gz", hash = "sha256:48e43ca463941b3608b32b4e236db92d8d40db4c58a75ace902985f76dac21cf", size = 239212, upload-time = "2024-12-13T08:30:44.393Z" }
wheels = [
@@ -6551,28 +6553,28 @@ wheels = [
[[package]]
name = "uv"
version = "0.9.17"
version = "0.9.18"
source = { registry = "https://pypi.org/simple" }
sdist = { url = "https://files.pythonhosted.org/packages/52/1a/cb0c37ae8513b253bcbc13d42392feb7d95ea696eb398b37535a28df9040/uv-0.9.17.tar.gz", hash = "sha256:6d93ab9012673e82039cfa7f9f66f69b388bc3f910f9e8a2ebee211353f620aa", size = 3815957, upload-time = "2025-12-09T23:01:21.756Z" }
sdist = { url = "https://files.pythonhosted.org/packages/e3/03/1afff9e6362dc9d3a9e03743da0a4b4c7a0809f859c79eb52bbae31ea582/uv-0.9.18.tar.gz", hash = "sha256:17b5502f7689c4dc1fdeee9d8437a9a6664dcaa8476e70046b5f4753559533f5", size = 3824466, upload-time = "2025-12-16T15:45:11.81Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/2b/e2/b6e2d473bdc37f4d86307151b53c0776e9925de7376ce297e92eab2e8894/uv-0.9.17-py3-none-linux_armv6l.whl", hash = "sha256:c708e6560ae5bc3cda1ba93f0094148ce773b6764240ced433acf88879e57a67", size = 21254511, upload-time = "2025-12-09T23:00:36.604Z" },
{ url = "https://files.pythonhosted.org/packages/d5/40/75f1529a8bf33cc5c885048e64a014c3096db5ac7826c71e20f2b731b588/uv-0.9.17-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:233b3d90f104c59d602abf434898057876b87f64df67a37129877d6dab6e5e10", size = 20384366, upload-time = "2025-12-09T23:01:17.293Z" },
{ url = "https://files.pythonhosted.org/packages/de/30/b3a343893681a569cbb74f8747a1c24e5f18ca9e07de0430aceaf9389ef4/uv-0.9.17-py3-none-macosx_11_0_arm64.whl", hash = "sha256:4b8e5513d48a267bfa180ca7fefaf6f27b1267e191573b3dba059981143e88ef", size = 18924624, upload-time = "2025-12-09T23:01:10.291Z" },
{ url = "https://files.pythonhosted.org/packages/21/56/9daf8bbe4a9a36eb0b9257cf5e1e20f9433d0ce996778ccf1929cbe071a4/uv-0.9.17-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.musllinux_1_1_aarch64.whl", hash = "sha256:8f283488bbcf19754910cc1ae7349c567918d6367c596e5a75d4751e0080eee0", size = 20671687, upload-time = "2025-12-09T23:00:51.927Z" },
{ url = "https://files.pythonhosted.org/packages/9f/c8/4050ff7dc692770092042fcef57223b8852662544f5981a7f6cac8fc488d/uv-0.9.17-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:9cf8052ba669dc17bdba75dae655094d820f4044990ea95c01ec9688c182f1da", size = 20861866, upload-time = "2025-12-09T23:01:12.555Z" },
{ url = "https://files.pythonhosted.org/packages/84/d4/208e62b7db7a65cb3390a11604c59937e387d07ed9f8b63b54edb55e2292/uv-0.9.17-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:06749461b11175a884be193120044e7f632a55e2624d9203398808907d346aad", size = 21858420, upload-time = "2025-12-09T23:01:00.009Z" },
{ url = "https://files.pythonhosted.org/packages/86/2c/91288cd5a04db37dfc1e0dad26ead84787db5832d9836b4cc8e0fa7f3c53/uv-0.9.17-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:35eb1a519688209160e48e1bb8032d36d285948a13b4dd21afe7ec36dc2a9787", size = 23471658, upload-time = "2025-12-09T23:00:49.503Z" },
{ url = "https://files.pythonhosted.org/packages/44/ba/493eba650ffad1df9e04fd8eabfc2d0aebc23e8f378acaaee9d95ca43518/uv-0.9.17-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:2bfb60a533e82690ab17dfe619ff7f294d053415645800d38d13062170230714", size = 23062950, upload-time = "2025-12-09T23:00:39.055Z" },
{ url = "https://files.pythonhosted.org/packages/9a/9e/f7f679503c06843ba59451e3193f35fb7c782ff0afc697020d4718a7de46/uv-0.9.17-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:cd0f3e380ff148aff3d769e95a9743cb29c7f040d7ef2896cafe8063279a6bc1", size = 22080299, upload-time = "2025-12-09T23:00:44.026Z" },
{ url = "https://files.pythonhosted.org/packages/32/2e/76ba33c7d9efe9f17480db1b94d3393025062005e346bb8b3660554526da/uv-0.9.17-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:cd2c3d25fbd8f91b30d0fac69a13b8e2c2cd8e606d7e6e924c1423e4ff84e616", size = 22087554, upload-time = "2025-12-09T23:00:41.715Z" },
{ url = "https://files.pythonhosted.org/packages/14/db/ef4aae4a6c49076db2acd2a7b0278ddf3dbf785d5172b3165018b96ba2fb/uv-0.9.17-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:330e7085857e4205c5196a417aca81cfbfa936a97dd2a0871f6560a88424ebf2", size = 20823225, upload-time = "2025-12-09T23:00:57.041Z" },
{ url = "https://files.pythonhosted.org/packages/11/73/e0f816cacd802a1cb25e71de9d60e57fa1f6c659eb5599cef708668618cc/uv-0.9.17-py3-none-manylinux_2_31_riscv64.whl", hash = "sha256:45880faa9f6cf91e3cda4e5f947da6a1004238fdc0ed4ebc18783a12ce197312", size = 22004893, upload-time = "2025-12-09T23:01:15.011Z" },
{ url = "https://files.pythonhosted.org/packages/15/6b/700f6256ee191136eb06e40d16970a4fc687efdccf5e67c553a258063019/uv-0.9.17-py3-none-musllinux_1_1_armv7l.whl", hash = "sha256:8e775a1b94c6f248e22f0ce2f86ed37c24e10ae31fb98b7e1b9f9a3189d25991", size = 20853850, upload-time = "2025-12-09T23:01:02.694Z" },
{ url = "https://files.pythonhosted.org/packages/bc/6a/13f02e2ed6510223c40f74804586b09e5151d9319f93aab1e49d91db13bb/uv-0.9.17-py3-none-musllinux_1_1_i686.whl", hash = "sha256:8650c894401ec96488a6fd84a5b4675e09be102f5525c902a12ba1c8ef8ff230", size = 21322623, upload-time = "2025-12-09T23:00:46.806Z" },
{ url = "https://files.pythonhosted.org/packages/d0/18/2d19780cebfbec877ea645463410c17859f8070f79c1a34568b153d78e1d/uv-0.9.17-py3-none-musllinux_1_1_x86_64.whl", hash = "sha256:673066b72d8b6c86be0dae6d5f73926bcee8e4810f1690d7b8ce5429d919cde3", size = 22290123, upload-time = "2025-12-09T23:00:54.394Z" },
{ url = "https://files.pythonhosted.org/packages/77/69/ab79bde3f7b6d2ac89f839ea40411a9cf3e67abede2278806305b6ba797e/uv-0.9.17-py3-none-win32.whl", hash = "sha256:7407d45afeae12399de048f7c8c2256546899c94bd7892dbddfae6766616f5a3", size = 20070709, upload-time = "2025-12-09T23:01:05.105Z" },
{ url = "https://files.pythonhosted.org/packages/08/a0/ab5b1850197bf407d095361b214352e40805441791fed35b891621cb1562/uv-0.9.17-py3-none-win_amd64.whl", hash = "sha256:22fcc26755abebdf366becc529b2872a831ce8bb14b36b6a80d443a1d7f84d3b", size = 22122852, upload-time = "2025-12-09T23:01:07.783Z" },
{ url = "https://files.pythonhosted.org/packages/37/ef/813cfedda3c8e49d8b59a41c14fcc652174facfd7a1caf9fee162b40ccbd/uv-0.9.17-py3-none-win_arm64.whl", hash = "sha256:6761076b27a763d0ede2f5e72455d2a46968ff334badf8312bb35988c5254831", size = 20435751, upload-time = "2025-12-09T23:01:19.732Z" },
{ url = "https://files.pythonhosted.org/packages/26/9c/92fad10fcee8ea170b66442d95fd2af308fe9a107909ded4b3cc384fdc69/uv-0.9.18-py3-none-linux_armv6l.whl", hash = "sha256:e9e4915bb280c1f79b9a1c16021e79f61ed7c6382856ceaa99d53258cb0b4951", size = 21345538, upload-time = "2025-12-16T15:45:13.992Z" },
{ url = "https://files.pythonhosted.org/packages/81/b1/b0e5808e05acb54aa118c625d9f7b117df614703b0cbb89d419d03d117f3/uv-0.9.18-py3-none-macosx_10_12_x86_64.whl", hash = "sha256:d91abfd2649987996e3778729140c305ef0f6ff5909f55aac35c3c372544a24f", size = 20439572, upload-time = "2025-12-16T15:45:26.397Z" },
{ url = "https://files.pythonhosted.org/packages/b7/0b/9487d83adf5b7fd1e20ced33f78adf84cb18239c3d7e91f224cedba46c08/uv-0.9.18-py3-none-macosx_11_0_arm64.whl", hash = "sha256:cf33f4146fd97e94cdebe6afc5122208eea8c55b65ca4127f5a5643c9717c8b8", size = 18952907, upload-time = "2025-12-16T15:44:48.399Z" },
{ url = "https://files.pythonhosted.org/packages/58/92/c8f7ae8900eff8e4ce1f7826d2e1e2ad5a95a5f141abdb539865aff79930/uv-0.9.18-py3-none-manylinux_2_17_aarch64.manylinux2014_aarch64.musllinux_1_1_aarch64.whl", hash = "sha256:edf965e9a5c55f74020ac82285eb0dfe7fac4f325ad0a7afc816290269ecfec1", size = 20772495, upload-time = "2025-12-16T15:45:29.614Z" },
{ url = "https://files.pythonhosted.org/packages/5a/28/9831500317c1dd6cde5099e3eb3b22b88ac75e47df7b502f6aef4df5750e/uv-0.9.18-py3-none-manylinux_2_17_armv7l.manylinux2014_armv7l.whl", hash = "sha256:ae10a941bd7ca1ee69edbe3998c34dce0a9fc2d2406d98198343daf7d2078493", size = 20949623, upload-time = "2025-12-16T15:44:57.482Z" },
{ url = "https://files.pythonhosted.org/packages/0c/ff/1fe1ffa69c8910e54dd11f01fb0765d4fd537ceaeb0c05fa584b6b635b82/uv-0.9.18-py3-none-manylinux_2_17_i686.manylinux2014_i686.whl", hash = "sha256:a1669a95b588f613b13dd10e08ced6d5bcd79169bba29a2240eee87532648790", size = 21920580, upload-time = "2025-12-16T15:44:39.009Z" },
{ url = "https://files.pythonhosted.org/packages/d6/ee/eed3ec7679ee80e16316cfc95ed28ef6851700bcc66edacfc583cbd2cc47/uv-0.9.18-py3-none-manylinux_2_17_ppc64.manylinux2014_ppc64.whl", hash = "sha256:11e1e406590d3159138288203a41ff8a8904600b8628a57462f04ff87d62c477", size = 23491234, upload-time = "2025-12-16T15:45:32.59Z" },
{ url = "https://files.pythonhosted.org/packages/78/58/64b15df743c79ad03ea7fbcbd27b146ba16a116c57f557425dd4e44d6684/uv-0.9.18-py3-none-manylinux_2_17_ppc64le.manylinux2014_ppc64le.whl", hash = "sha256:1e82078d3c622cb4c60da87f156168ffa78b9911136db7ffeb8e5b0a040bf30e", size = 23095438, upload-time = "2025-12-16T15:45:17.916Z" },
{ url = "https://files.pythonhosted.org/packages/43/6d/3d3dae71796961603c3871699e10d6b9de2e65a3c327b58d4750610a5f93/uv-0.9.18-py3-none-manylinux_2_17_s390x.manylinux2014_s390x.whl", hash = "sha256:704abaf6e76b4d293fc1f24bef2c289021f1df0de9ed351f476cbbf67a7edae0", size = 22140992, upload-time = "2025-12-16T15:44:45.527Z" },
{ url = "https://files.pythonhosted.org/packages/31/91/1042d0966a30e937df500daed63e1f61018714406ce4023c8a6e6d2dcf7c/uv-0.9.18-py3-none-manylinux_2_17_x86_64.manylinux2014_x86_64.whl", hash = "sha256:3332188fd8d96a68e5001409a52156dced910bf1bc41ec3066534cffcd46eb68", size = 22229626, upload-time = "2025-12-16T15:45:20.712Z" },
{ url = "https://files.pythonhosted.org/packages/5a/1f/0a4a979bb2bf6e1292cc57882955bf1d7757cad40b1862d524c59c2a2ad8/uv-0.9.18-py3-none-manylinux_2_28_aarch64.whl", hash = "sha256:b7295e6d505f1fd61c54b1219e3b18e11907396333a9fa61cefe489c08fc7995", size = 20896524, upload-time = "2025-12-16T15:45:06.799Z" },
{ url = "https://files.pythonhosted.org/packages/a5/3c/24f92e56af00cac7d9bed2888d99a580f8093c8745395ccf6213bfccf20b/uv-0.9.18-py3-none-manylinux_2_31_riscv64.whl", hash = "sha256:62ea0e518dd4ab76e6f06c0f43a25898a6342a3ecf996c12f27f08eb801ef7f1", size = 22077340, upload-time = "2025-12-16T15:44:51.271Z" },
{ url = "https://files.pythonhosted.org/packages/9c/3e/73163116f748800e676bf30cee838448e74ac4cc2f716c750e1705bc3fe4/uv-0.9.18-py3-none-musllinux_1_1_armv7l.whl", hash = "sha256:8bd073e30030211ba01206caa57b4d63714e1adee2c76a1678987dd52f72d44d", size = 20932956, upload-time = "2025-12-16T15:45:00.3Z" },
{ url = "https://files.pythonhosted.org/packages/59/1b/a26990b51a17de1ffe41fbf2e30de3a98f0e0bce40cc60829fb9d9ed1a8a/uv-0.9.18-py3-none-musllinux_1_1_i686.whl", hash = "sha256:f248e013d10e1fc7a41f94310628b4a8130886b6d683c7c85c42b5b36d1bcd02", size = 21357247, upload-time = "2025-12-16T15:45:23.575Z" },
{ url = "https://files.pythonhosted.org/packages/5f/20/b6ba14fdd671e9237b22060d7422aba4a34503e3e42d914dbf925eff19aa/uv-0.9.18-py3-none-musllinux_1_1_x86_64.whl", hash = "sha256:17bedf2b0791e87d889e1c7f125bd5de77e4b7579aec372fa06ba832e07c957e", size = 22443585, upload-time = "2025-12-16T15:44:42.213Z" },
{ url = "https://files.pythonhosted.org/packages/5e/da/1b3dd596964f90a122cfe94dcf5b6b89cf5670eb84434b8c23864382576f/uv-0.9.18-py3-none-win32.whl", hash = "sha256:de6f0bb3e9c18e484545bd1549ec3c956968a141a393d42e2efb25281cb62787", size = 20091088, upload-time = "2025-12-16T15:45:03.225Z" },
{ url = "https://files.pythonhosted.org/packages/11/0b/50e13ebc1eedb36d88524b7740f78351be33213073e3faf81ac8925d0c6e/uv-0.9.18-py3-none-win_amd64.whl", hash = "sha256:c82b0e2e36b33e2146fba5f0ae6906b9679b3b5fe6a712e5d624e45e441e58e9", size = 22181193, upload-time = "2025-12-16T15:44:54.394Z" },
{ url = "https://files.pythonhosted.org/packages/8c/d4/0bf338d863a3d9e5545e268d77a8e6afdd75d26bffc939603042f2e739f9/uv-0.9.18-py3-none-win_arm64.whl", hash = "sha256:4c4ce0ed080440bbda2377488575d426867f94f5922323af6d4728a1cd4d091d", size = 20564933, upload-time = "2025-12-16T15:45:09.819Z" },
]
[[package]]