mirror of
https://github.com/microsoft/agent-framework.git
synced 2026-06-16 21:04:09 +08:00
Compare commits
16
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
fee6f1b862 | ||
|
|
f853370621 | ||
|
|
c14beedb3a | ||
|
|
2c9f543100 | ||
|
|
43d98974d3 | ||
|
|
60da0ffb48 | ||
|
|
a2044829b1 | ||
|
|
435c66e9c9 | ||
|
|
52d50be9e0 | ||
|
|
d20f9b5f97 | ||
|
|
87a8fa2a9d | ||
|
|
8f7fd9525d | ||
|
|
69697065ab | ||
|
|
fe4cd3cddc | ||
|
|
611230cc8e | ||
|
|
f112150cfb |
@@ -171,7 +171,7 @@ jobs:
|
||||
-m integration
|
||||
-n logical --dist worksteal
|
||||
--timeout=120 --session-timeout=900 --timeout_method thread
|
||||
--retries 2 --retry-delay 5
|
||||
--retries 2 --retry-delay 30
|
||||
- name: Stop local MCP server
|
||||
if: always()
|
||||
shell: bash
|
||||
|
||||
@@ -287,7 +287,7 @@ jobs:
|
||||
-m integration
|
||||
-n logical --dist worksteal
|
||||
--timeout=120 --session-timeout=900 --timeout_method thread
|
||||
--retries 2 --retry-delay 5
|
||||
--retries 2 --retry-delay 30
|
||||
--junitxml=pytest.xml
|
||||
working-directory: ./python
|
||||
- name: Stop local MCP server
|
||||
|
||||
@@ -4,8 +4,9 @@
|
||||
<!-- https://learn.microsoft.com/en-us/nuget/consume-packages/Central-Package-Management -->
|
||||
<Sdk Name="Microsoft.Build.CentralPackageVersions" Version="2.1.3" />
|
||||
<!-- Only run 'dotnet format' on dev machines, Release builds. Skip on GitHub Actions -->
|
||||
<!-- as this runs in its own Actions job. -->
|
||||
<Target Name="DotnetFormatOnBuild" BeforeTargets="Build" Condition=" '$(Configuration)' == 'Release' AND '$(GITHUB_ACTIONS)' == '' ">
|
||||
<!-- as this runs in its own Actions job. Only run for net10.0 target frameworks since the dotnet format command -->
|
||||
<!-- already formats all target frameworks in project. Otherwise it will run format x times x where x is the number of target frameworks -->
|
||||
<Target Name="DotnetFormatOnBuild" BeforeTargets="Build" Condition=" '$(Configuration)' == 'Release' AND '$(GITHUB_ACTIONS)' == '' AND '$(TargetFramework)' == 'net10.0' ">
|
||||
<Message Text="Running dotnet format" Importance="high" />
|
||||
<Exec Command="dotnet format --no-restore -v diag $(ProjectFileName)" />
|
||||
</Target>
|
||||
|
||||
@@ -11,8 +11,8 @@
|
||||
</PropertyGroup>
|
||||
<ItemGroup>
|
||||
<!-- Aspire.* -->
|
||||
<PackageVersion Include="Anthropic" Version="12.11.0" />
|
||||
<PackageVersion Include="Anthropic.Foundry" Version="0.4.2" />
|
||||
<PackageVersion Include="Anthropic" Version="12.13.0" />
|
||||
<PackageVersion Include="Anthropic.Foundry" Version="0.5.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)" />
|
||||
@@ -32,19 +32,19 @@
|
||||
<!-- Newtonsoft.Json -->
|
||||
<PackageVersion Include="Newtonsoft.Json" Version="13.0.4" />
|
||||
<!-- System.* -->
|
||||
<PackageVersion Include="Microsoft.Bcl.AsyncInterfaces" Version="10.0.4" />
|
||||
<PackageVersion Include="Microsoft.Bcl.AsyncInterfaces" Version="10.0.6" />
|
||||
<PackageVersion Include="Microsoft.Bcl.HashCode" Version="6.0.0" />
|
||||
<PackageVersion Include="Microsoft.Bcl.Memory" Version="10.0.4" />
|
||||
<PackageVersion Include="System.ClientModel" Version="1.10.0" />
|
||||
<PackageVersion Include="System.CodeDom" Version="10.0.0" />
|
||||
<PackageVersion Include="System.Collections.Immutable" Version="10.0.1" />
|
||||
<PackageVersion Include="System.CommandLine" Version="2.0.0-rc.2.25502.107" />
|
||||
<PackageVersion Include="System.Diagnostics.DiagnosticSource" Version="10.0.4" />
|
||||
<PackageVersion Include="System.Diagnostics.DiagnosticSource" Version="10.0.6" />
|
||||
<PackageVersion Include="System.Linq.AsyncEnumerable" Version="10.0.4" />
|
||||
<PackageVersion Include="System.Net.Http.Json" Version="10.0.0" />
|
||||
<PackageVersion Include="System.Net.ServerSentEvents" Version="10.0.4" />
|
||||
<PackageVersion Include="System.Text.Json" Version="10.0.4" />
|
||||
<PackageVersion Include="System.Threading.Channels" Version="10.0.4" />
|
||||
<PackageVersion Include="System.Text.Json" Version="10.0.6" />
|
||||
<PackageVersion Include="System.Threading.Channels" Version="10.0.6" />
|
||||
<PackageVersion Include="System.Threading.Tasks.Extensions" Version="4.6.3" />
|
||||
<PackageVersion Include="System.Net.Security" Version="4.3.2" />
|
||||
<!-- OpenTelemetry -->
|
||||
@@ -63,37 +63,28 @@
|
||||
<PackageVersion Include="Microsoft.AspNetCore.OpenApi" Version="10.0.0" />
|
||||
<PackageVersion Include="Swashbuckle.AspNetCore.SwaggerUI" Version="10.0.0" />
|
||||
<!-- Microsoft.Extensions.* -->
|
||||
<PackageVersion Include="Microsoft.Extensions.AI" Version="10.4.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.AI.Abstractions" Version="10.4.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.AI.Evaluation" Version="10.4.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.AI.Evaluation.Quality" Version="10.4.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.AI.Evaluation.Safety" Version="10.3.0-preview.1.26109.11" />
|
||||
<PackageVersion Include="Microsoft.Extensions.AI.OpenAI" Version="10.4.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.AI" Version="10.5.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.AI.Abstractions" Version="10.5.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.AI.OpenAI" Version="10.5.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Caching.Memory" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Compliance.Abstractions" Version="10.4.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Compliance.Abstractions" Version="10.5.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Configuration" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Configuration.Binder" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Configuration.EnvironmentVariables" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Configuration.Json" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Configuration.UserSecrets" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.DependencyInjection" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.DependencyInjection.Abstractions" Version="10.0.4" />
|
||||
<PackageVersion Include="Microsoft.Extensions.DependencyInjection.Abstractions" Version="10.0.6" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Hosting" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Http.Resilience" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Logging" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Logging.Abstractions" Version="10.0.4" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Logging.Abstractions" Version="10.0.6" />
|
||||
<PackageVersion Include="Microsoft.Extensions.Logging.Console" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.ServiceDiscovery" Version="10.0.0" />
|
||||
<PackageVersion Include="Microsoft.Extensions.VectorData.Abstractions" Version="9.7.0" />
|
||||
<!-- Vector Stores -->
|
||||
<PackageVersion Include="Microsoft.SemanticKernel.Connectors.InMemory" Version="1.67.0-preview" />
|
||||
<PackageVersion Include="Microsoft.SemanticKernel.Connectors.Qdrant" Version="1.67.0-preview" />
|
||||
<!-- Semantic Kernel -->
|
||||
<PackageVersion Include="Microsoft.SemanticKernel" Version="1.67.0" />
|
||||
<PackageVersion Include="Microsoft.SemanticKernel.Agents.Core" Version="1.67.0" />
|
||||
<PackageVersion Include="Microsoft.SemanticKernel.Agents.OpenAI" Version="1.67.0-preview" />
|
||||
<PackageVersion Include="Microsoft.SemanticKernel.Agents.AzureAI" Version="1.67.0-preview" />
|
||||
<PackageVersion Include="Microsoft.SemanticKernel.Plugins.OpenApi" Version="1.67.0" />
|
||||
<!-- Agent SDKs -->
|
||||
<PackageVersion Include="GitHub.Copilot.SDK" Version="0.1.29" />
|
||||
<PackageVersion Include="Microsoft.Agents.CopilotStudio.Client" Version="1.3.171-beta" />
|
||||
@@ -107,11 +98,10 @@
|
||||
<!-- MCP -->
|
||||
<PackageVersion Include="ModelContextProtocol" Version="1.1.0" />
|
||||
<!-- Inference SDKs -->
|
||||
<PackageVersion Include="AWSSDK.Extensions.Bedrock.MEAI" Version="4.0.5.1" />
|
||||
<PackageVersion Include="Microsoft.ML.OnnxRuntimeGenAI" Version="0.10.0" />
|
||||
<PackageVersion Include="Microsoft.ML.Tokenizers" Version="2.0.0" />
|
||||
<PackageVersion Include="OllamaSharp" Version="5.4.8" />
|
||||
<PackageVersion Include="OpenAI" Version="2.9.1" />
|
||||
<PackageVersion Include="OpenAI" Version="2.10.0" />
|
||||
<!-- Identity -->
|
||||
<PackageVersion Include="Microsoft.Identity.Client.Extensions.Msal" Version="4.83.1" />
|
||||
<!-- Workflows -->
|
||||
@@ -126,7 +116,6 @@
|
||||
<PackageVersion Include="Microsoft.DurableTask.Worker.AzureManaged" Version="1.18.0" />
|
||||
<!-- Azure Functions -->
|
||||
<PackageVersion Include="Microsoft.Azure.Functions.Worker" Version="2.50.0" />
|
||||
<PackageVersion Include="Microsoft.Azure.Functions.Worker.ApplicationInsights" Version="2.50.0" />
|
||||
<PackageVersion Include="Microsoft.Azure.Functions.Worker.Extensions.DurableTask" Version="1.12.1" />
|
||||
<PackageVersion Include="Microsoft.Azure.Functions.Worker.Extensions.DurableTask.AzureManaged" Version="1.0.1" />
|
||||
<PackageVersion Include="Microsoft.Azure.Functions.Worker.Extensions.Http" Version="3.3.0" />
|
||||
|
||||
@@ -243,6 +243,9 @@
|
||||
<Folder Name="/Samples/03-workflows/HumanInTheLoop/">
|
||||
<Project Path="samples/03-workflows/HumanInTheLoop/HumanInTheLoopBasic/HumanInTheLoopBasic.csproj" />
|
||||
</Folder>
|
||||
<Folder Name="/Samples/03-workflows/Orchestration/">
|
||||
<Project Path="samples/03-workflows/Orchestration/Handoff/Handoff.csproj" />
|
||||
</Folder>
|
||||
<Folder Name="/Samples/03-workflows/Observability/">
|
||||
<Project Path="samples/03-workflows/Observability/ApplicationInsights/ApplicationInsights.csproj" />
|
||||
<Project Path="samples/03-workflows/Observability/AspireDashboard/AspireDashboard.csproj" />
|
||||
@@ -288,7 +291,7 @@
|
||||
<File Path="samples/04-hosting/A2A/README.md" />
|
||||
<Project Path="samples/04-hosting/A2A/A2AAgent_AsFunctionTools/A2AAgent_AsFunctionTools.csproj" />
|
||||
<Project Path="samples/04-hosting/A2A/A2AAgent_PollingForTaskCompletion/A2AAgent_PollingForTaskCompletion.csproj" />
|
||||
</Folder>
|
||||
</Folder>
|
||||
<Folder Name="/Samples/05-end-to-end/">
|
||||
<Project Path="samples/05-end-to-end/AgentWithPurview/AgentWithPurview.csproj" />
|
||||
<Project Path="samples/05-end-to-end/M365Agent/M365Agent.csproj" />
|
||||
|
||||
@@ -0,0 +1,72 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
using Microsoft.Agents.AI;
|
||||
using Microsoft.Agents.AI.Workflows;
|
||||
using Microsoft.Extensions.AI;
|
||||
|
||||
/// <summary>
|
||||
/// The registry of agents used in the workflow.
|
||||
/// </summary>
|
||||
/// <param name="chatClient">The <see cref="IChatClient"/> to use as the agent backend.</param>
|
||||
internal sealed class AgentRegistry(IChatClient chatClient)
|
||||
{
|
||||
internal const string IntakeAgentName = "Assistant";
|
||||
public AIAgent IntakeAgent { get; } = chatClient.AsAIAgent(
|
||||
instructions:
|
||||
"""
|
||||
You receive a user request and are responsible for routing to the correct initial expert agent.
|
||||
""",
|
||||
IntakeAgentName
|
||||
);
|
||||
|
||||
internal const string LiquidityAnalysisAgentName = "Liquidity Analysis";
|
||||
public AIAgent LiquidityAnalysisAgent { get; } = chatClient.AsAIAgent(
|
||||
instructions:
|
||||
"""
|
||||
You are responsible for Liquidity Analysis.
|
||||
""",
|
||||
LiquidityAnalysisAgentName
|
||||
);
|
||||
|
||||
internal const string TaxAnalysisAgentName = "Tax Analysis";
|
||||
public AIAgent TaxAnalysisAgent { get; } = chatClient.AsAIAgent(
|
||||
instructions:
|
||||
"""
|
||||
You are responsible for Tax Analysis.
|
||||
""",
|
||||
TaxAnalysisAgentName
|
||||
);
|
||||
|
||||
internal const string ForeignExchangeAgentName = "Foreign Exchange Analysis";
|
||||
public AIAgent ForeignExchangeAgent { get; } = chatClient.AsAIAgent(
|
||||
instructions:
|
||||
"""
|
||||
You are responsible for Foreign Exchange Analysis.
|
||||
""",
|
||||
ForeignExchangeAgentName
|
||||
);
|
||||
|
||||
internal const string EquityAgentName = "Equity Analysis";
|
||||
public AIAgent EquityAgent { get; } = chatClient.AsAIAgent(
|
||||
instructions:
|
||||
"""
|
||||
You are responsible for Equity Analysis.
|
||||
""",
|
||||
EquityAgentName
|
||||
);
|
||||
|
||||
public IEnumerable<AIAgent> Experts => [this.LiquidityAnalysisAgent, this.TaxAnalysisAgent, this.ForeignExchangeAgent, this.EquityAgent];
|
||||
|
||||
public HashSet<AIAgent> All
|
||||
{
|
||||
get
|
||||
{
|
||||
if (field == null)
|
||||
{
|
||||
field = [this.IntakeAgent, .. this.Experts];
|
||||
}
|
||||
|
||||
return field;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,29 @@
|
||||
<Project Sdk="Microsoft.NET.Sdk">
|
||||
|
||||
<PropertyGroup>
|
||||
<OutputType>Exe</OutputType>
|
||||
<TargetFrameworks>net10.0</TargetFrameworks>
|
||||
|
||||
<Nullable>enable</Nullable>
|
||||
<ImplicitUsings>enable</ImplicitUsings>
|
||||
|
||||
<NoWarn>MAAIW001</NoWarn>
|
||||
</PropertyGroup>
|
||||
|
||||
<ItemGroup>
|
||||
<PackageReference Include="Azure.AI.Projects" />
|
||||
<PackageReference Include="Azure.Identity" />
|
||||
<PackageReference Include="Microsoft.Extensions.AI.OpenAI" />
|
||||
</ItemGroup>
|
||||
|
||||
<ItemGroup>
|
||||
<ProjectReference Include="..\..\..\..\src\Microsoft.Agents.AI.Workflows\Microsoft.Agents.AI.Workflows.csproj" />
|
||||
<ProjectReference Include="..\..\..\..\src\Microsoft.Agents.AI\Microsoft.Agents.AI.csproj" />
|
||||
<!-- Include Workflows source generator when using [MessageHandler] attribute -->
|
||||
<ProjectReference Include="$(RepoRoot)/dotnet/src/Microsoft.Agents.AI.Workflows.Generators/Microsoft.Agents.AI.Workflows.Generators.csproj"
|
||||
OutputItemType="Analyzer"
|
||||
ReferenceOutputAssembly="false"
|
||||
GlobalPropertiesToRemove="TargetFramework" />
|
||||
</ItemGroup>
|
||||
|
||||
</Project>
|
||||
@@ -0,0 +1,125 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
using Azure.AI.Projects;
|
||||
using Azure.Identity;
|
||||
using Microsoft.Agents.AI;
|
||||
using Microsoft.Agents.AI.Workflows;
|
||||
using Microsoft.Extensions.AI;
|
||||
|
||||
string endpoint = Environment.GetEnvironmentVariable("AZURE_AI_PROJECT_ENDPOINT")
|
||||
?? throw new InvalidOperationException("AZURE_AI_PROJECT_ENDPOINT is not set.");
|
||||
string deploymentName = Environment.GetEnvironmentVariable("AZURE_AI_MODEL_DEPLOYMENT_NAME") ?? "gpt-5.4-mini";
|
||||
|
||||
// WARNING: DefaultAzureCredential is convenient for development but requires careful consideration in production.
|
||||
// In production, consider using a specific credential (e.g., ManagedIdentityCredential) to avoid
|
||||
// latency issues, unintended credential probing, and potential security risks from fallback mechanisms.
|
||||
AIProjectClient projectClient = new(new Uri(endpoint), new DefaultAzureCredential());
|
||||
|
||||
IChatClient chatClient = projectClient.ProjectOpenAIClient
|
||||
.GetChatClient(deploymentName)
|
||||
.AsIChatClient();
|
||||
|
||||
Workflow workflow = CreateWorkflow(chatClient);
|
||||
|
||||
await RunWorkflowAsync(workflow).ConfigureAwait(false);
|
||||
|
||||
static Workflow CreateWorkflow(IChatClient chatClient)
|
||||
{
|
||||
AgentRegistry agents = new(chatClient);
|
||||
|
||||
HandoffWorkflowBuilder handoffBuilder = AgentWorkflowBuilder.CreateHandoffBuilderWith(agents.IntakeAgent);
|
||||
|
||||
// Add a handoff to each of the experts from every agent in the registry (experts + Intake)
|
||||
foreach (AIAgent expert in agents.Experts)
|
||||
{
|
||||
handoffBuilder.WithHandoffs(agents.All.Except([expert]), expert);
|
||||
}
|
||||
|
||||
// Let agents request more user information and return to the asking agent (rather than going back to the intake agent)
|
||||
handoffBuilder.EnableReturnToPrevious();
|
||||
|
||||
return handoffBuilder.Build();
|
||||
}
|
||||
|
||||
static async Task RunWorkflowAsync(Workflow workflow)
|
||||
{
|
||||
using CancellationTokenSource cts = CreateConsoleCancelKeySource();
|
||||
await using StreamingRun run = await InProcessExecution.OpenStreamingAsync(workflow, cancellationToken: cts.Token)
|
||||
.ConfigureAwait(false);
|
||||
|
||||
bool hadError = false;
|
||||
do
|
||||
{
|
||||
Console.Write("> ");
|
||||
string userInput = Console.ReadLine() ?? string.Empty;
|
||||
|
||||
if (userInput.Equals("exit", StringComparison.OrdinalIgnoreCase))
|
||||
{
|
||||
break;
|
||||
}
|
||||
|
||||
await run.TrySendMessageAsync(userInput);
|
||||
string? speakingAgent = null;
|
||||
await foreach (WorkflowEvent evt in run.WatchStreamAsync(cts.Token))
|
||||
{
|
||||
switch (evt)
|
||||
{
|
||||
case AgentResponseUpdateEvent update:
|
||||
{
|
||||
if (speakingAgent == null || speakingAgent != update.Update.AuthorName)
|
||||
{
|
||||
speakingAgent = update.Update.AuthorName;
|
||||
Console.Write($"\n{speakingAgent}: ");
|
||||
}
|
||||
|
||||
Console.Write(update.Update.Text);
|
||||
break;
|
||||
}
|
||||
|
||||
case WorkflowErrorEvent workflowError:
|
||||
{
|
||||
Console.ForegroundColor = ConsoleColor.Red;
|
||||
|
||||
if (workflowError.Exception != null)
|
||||
{
|
||||
Console.WriteLine($"\nWorkflow error: {workflowError.Exception}");
|
||||
}
|
||||
else
|
||||
{
|
||||
Console.WriteLine("\nUnknown workflow error occurred.");
|
||||
}
|
||||
|
||||
Console.ResetColor();
|
||||
|
||||
hadError = true;
|
||||
break;
|
||||
}
|
||||
|
||||
case WorkflowWarningEvent workflowWarning when workflowWarning.Data is string message:
|
||||
{
|
||||
Console.ForegroundColor = ConsoleColor.Yellow;
|
||||
Console.WriteLine(message);
|
||||
Console.ResetColor();
|
||||
break;
|
||||
}
|
||||
}
|
||||
}
|
||||
} while (!hadError);
|
||||
}
|
||||
|
||||
static CancellationTokenSource CreateConsoleCancelKeySource()
|
||||
{
|
||||
CancellationTokenSource cts = new();
|
||||
|
||||
// Normally, support a way to detach events, but in this case this is a termination signal, so cleanup will happen
|
||||
// as part of application shutdown.
|
||||
Console.CancelKeyPress += (s, args) =>
|
||||
{
|
||||
cts.Cancel();
|
||||
|
||||
// We handle cleanup + termination ourselves
|
||||
args.Cancel = true;
|
||||
};
|
||||
|
||||
return cts;
|
||||
}
|
||||
@@ -56,3 +56,9 @@ Once completed, please proceed to the other samples listed below.
|
||||
| [Edge Conditions](./ConditionalEdges/01_EdgeCondition) | Introduces conditional edges for dynamic routing based on executor outputs |
|
||||
| [Switch-Case Routing](./ConditionalEdges/02_SwitchCase) | Extends conditional edges with switch-case routing for multiple paths |
|
||||
| [Multi-Selection Routing](./ConditionalEdges/03_MultiSelection) | Demonstrates multi-selection routing where one executor can trigger multiple downstream executors |
|
||||
|
||||
### Orchestration Patterns
|
||||
|
||||
| Sample | Concepts |
|
||||
|--------|----------|
|
||||
| [Handoff Orchestration](./Orchestration/Handoff) | Introduces the Handoff Orchestration pattern |
|
||||
|
||||
@@ -1,7 +1,7 @@
|
||||
<Project Sdk="Microsoft.NET.Sdk">
|
||||
|
||||
<PropertyGroup>
|
||||
<IsReleaseCandidate>true</IsReleaseCandidate>
|
||||
<IsReleaseCandidate>false</IsReleaseCandidate>
|
||||
<ImplicitUsings>enable</ImplicitUsings>
|
||||
<InjectSharedThrow>true</InjectSharedThrow>
|
||||
</PropertyGroup>
|
||||
|
||||
@@ -419,6 +419,12 @@ internal sealed class InProcessRunnerContext : IRunnerContext
|
||||
.Select(id => this.EnsureExecutorAsync(id, tracer: null).AsTask())
|
||||
.ToArray();
|
||||
|
||||
// Discard queued external deliveries from the superseded timeline so a runtime
|
||||
// restore cannot apply stale responses after importing the checkpoint state.
|
||||
while (this._queuedExternalDeliveries.TryDequeue(out _))
|
||||
{
|
||||
}
|
||||
|
||||
this._nextStep = new StepContext();
|
||||
this._nextStep.ImportMessages(importedState.QueuedMessages);
|
||||
|
||||
|
||||
+8
@@ -483,6 +483,14 @@ public sealed class AnthropicBetaServiceExtensionsTests
|
||||
|
||||
public IBetaMessageService Messages => new Mock<IBetaMessageService>().Object;
|
||||
|
||||
public global::Anthropic.Services.Beta.IAgentService Agents => throw new NotImplementedException();
|
||||
|
||||
public global::Anthropic.Services.Beta.IEnvironmentService Environments => throw new NotImplementedException();
|
||||
|
||||
public global::Anthropic.Services.Beta.ISessionService Sessions => throw new NotImplementedException();
|
||||
|
||||
public global::Anthropic.Services.Beta.IVaultService Vaults => throw new NotImplementedException();
|
||||
|
||||
public IBetaService WithOptions(Func<ClientOptions, ClientOptions> modifier)
|
||||
{
|
||||
throw new NotImplementedException();
|
||||
|
||||
@@ -279,6 +279,48 @@ public class CheckpointResumeTests
|
||||
"the workflow should be able to continue after the runtime restore replay");
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Verifies that restoring a live run clears any queued external responses from the
|
||||
/// superseded timeline before importing checkpoint state.
|
||||
/// </summary>
|
||||
[Fact]
|
||||
internal async Task Checkpoint_Restore_ClearsQueuedExternalResponsesBeforeImportAsync()
|
||||
{
|
||||
Workflow workflow = CreateSimpleRequestWorkflow();
|
||||
CheckpointManager checkpointManager = CheckpointManager.CreateInMemory();
|
||||
InProcessExecutionEnvironment env = ExecutionEnvironment.InProcess_Lockstep.ToWorkflowExecutionEnvironment();
|
||||
|
||||
await using StreamingRun run = await env.WithCheckpointing(checkpointManager)
|
||||
.RunStreamingAsync(workflow, "Hello");
|
||||
|
||||
(ExternalRequest pendingRequest, CheckpointInfo checkpoint) = await CapturePendingRequestAndCheckpointAsync(run);
|
||||
|
||||
await run.SendResponseAsync(pendingRequest.CreateResponse("World"));
|
||||
await run.RestoreCheckpointAsync(checkpoint);
|
||||
|
||||
List<WorkflowEvent> restoredEvents = await ReadToHaltAsync(run);
|
||||
ExternalRequest replayedRequest = restoredEvents.OfType<RequestInfoEvent>()
|
||||
.Select(evt => evt.Request)
|
||||
.Should()
|
||||
.ContainSingle("the restored run should still be waiting for the checkpointed request")
|
||||
.Subject;
|
||||
|
||||
restoredEvents.OfType<WorkflowErrorEvent>().Should().BeEmpty(
|
||||
"a queued response from the superseded timeline should not be processed after restore");
|
||||
RunStatus statusAfterRestore = await run.GetStatusAsync();
|
||||
statusAfterRestore.Should().Be(RunStatus.PendingRequests,
|
||||
"the restored run should remain pending until a post-restore response is sent");
|
||||
|
||||
await run.SendResponseAsync(replayedRequest.CreateResponse("Again"));
|
||||
|
||||
List<WorkflowEvent> completionEvents = await ReadToHaltAsync(run);
|
||||
completionEvents.OfType<WorkflowErrorEvent>().Should().BeEmpty(
|
||||
"the restored request should complete cleanly once a new response is provided");
|
||||
RunStatus finalStatus = await run.GetStatusAsync();
|
||||
finalStatus.Should().Be(RunStatus.Idle,
|
||||
"the workflow should finish once the replayed request receives a fresh response");
|
||||
}
|
||||
|
||||
/// <summary>
|
||||
/// Verifies that a resumed parent workflow re-emits pending requests that originated in a subworkflow.
|
||||
/// </summary>
|
||||
|
||||
@@ -1,8 +1,14 @@
|
||||
// Copyright (c) Microsoft. All rights reserved.
|
||||
|
||||
using System;
|
||||
using System.Collections.Generic;
|
||||
using System.Linq;
|
||||
using System.Runtime.CompilerServices;
|
||||
using System.Threading;
|
||||
using System.Threading.Tasks;
|
||||
using FluentAssertions;
|
||||
using Microsoft.Agents.AI.Workflows.Specialized;
|
||||
using Microsoft.Extensions.AI;
|
||||
|
||||
namespace Microsoft.Agents.AI.Workflows.UnitTests;
|
||||
|
||||
@@ -68,4 +74,98 @@ public class HandoffAgentExecutorTests : AIAgentHostingExecutorTestsBase
|
||||
AgentResponseEvent[] updates = testContext.Events.OfType<AgentResponseEvent>().ToArray();
|
||||
CheckResponseEventsAgainstTestMessages(updates, expectingResponse: executorSetting, agent.GetDescriptiveId());
|
||||
}
|
||||
|
||||
[Fact]
|
||||
public async Task Test_HandoffAgentExecutor_PreservesExistingInstructionsAndToolsAsync()
|
||||
{
|
||||
// Arrange
|
||||
const string BaseInstructions = "BaseInstructions";
|
||||
const string HandoffInstructions = "HandoffInstructions";
|
||||
|
||||
AITool someTool = AIFunctionFactory.CreateDeclaration("BaseTool", null, AIFunctionFactory.Create(() => { }).JsonSchema);
|
||||
|
||||
OptionValidatingChatClient chatClient = new(BaseInstructions, HandoffInstructions, someTool);
|
||||
AIAgent handoffAgent = chatClient.AsAIAgent(BaseInstructions, tools: [someTool]);
|
||||
AIAgent targetAgent = new TestEchoAgent();
|
||||
|
||||
HandoffAgentExecutorOptions options = new(HandoffInstructions, false, null, HandoffToolCallFilteringBehavior.None);
|
||||
HandoffTarget handoff = new(targetAgent);
|
||||
HandoffAgentExecutor executor = new(handoffAgent, [handoff], options);
|
||||
|
||||
TestWorkflowContext testContext = new(executor.Id);
|
||||
HandoffState state = new(new(false), null, [], null);
|
||||
|
||||
// Act / Assert
|
||||
Func<Task> runStreamingAsync = async () => await executor.HandleAsync(state, testContext);
|
||||
await runStreamingAsync.Should().NotThrowAsync();
|
||||
}
|
||||
|
||||
private sealed class OptionValidatingChatClient(string baseInstructions, string handoffInstructions, AITool baseTool) : IChatClient
|
||||
{
|
||||
public void Dispose()
|
||||
{
|
||||
}
|
||||
|
||||
private void CheckOptions(ChatOptions? options)
|
||||
{
|
||||
options.Should().NotBeNull();
|
||||
|
||||
options.Instructions.Should().NotBeNullOrEmpty("Handoff orchestration should preserve and augment instructions.")
|
||||
.And.Contain(baseInstructions, because: "Handoff orchestration should preserve existing instructions.")
|
||||
.And.Contain(handoffInstructions, because: "Handoff orchestration should inject handoff instructions.");
|
||||
|
||||
options.Tools.Should().NotBeNullOrEmpty("Handoff orchestration should preserve and augment tools.")
|
||||
.And.Contain(tool => tool.Name == baseTool.Name, "Handoff orchestration should preserve existing tools.")
|
||||
.And.Contain(tool => tool.Name.StartsWith(HandoffWorkflowBuilder.FunctionPrefix, StringComparison.Ordinal),
|
||||
because: "Handoff orchestration should inject handoff tools.");
|
||||
}
|
||||
|
||||
private List<ChatMessage> ResponseMessages =>
|
||||
[
|
||||
new ChatMessage(ChatRole.Assistant, "Ok")
|
||||
{
|
||||
MessageId = Guid.NewGuid().ToString(),
|
||||
AuthorName = nameof(OptionValidatingChatClient)
|
||||
}
|
||||
];
|
||||
|
||||
public Task<ChatResponse> GetResponseAsync(IEnumerable<ChatMessage> messages, ChatOptions? options = null, CancellationToken cancellationToken = default)
|
||||
{
|
||||
this.CheckOptions(options);
|
||||
|
||||
ChatResponse response = new(this.ResponseMessages)
|
||||
{
|
||||
ResponseId = Guid.NewGuid().ToString("N"),
|
||||
CreatedAt = DateTimeOffset.Now
|
||||
};
|
||||
|
||||
return Task.FromResult(response);
|
||||
}
|
||||
|
||||
public object? GetService(Type serviceType, object? serviceKey = null)
|
||||
{
|
||||
if (serviceType == typeof(OptionValidatingChatClient))
|
||||
{
|
||||
return this;
|
||||
}
|
||||
|
||||
return null;
|
||||
}
|
||||
|
||||
public async IAsyncEnumerable<ChatResponseUpdate> GetStreamingResponseAsync(IEnumerable<ChatMessage> messages, ChatOptions? options = null, [EnumeratorCancellation] CancellationToken cancellationToken = default)
|
||||
{
|
||||
this.CheckOptions(options);
|
||||
|
||||
string responseId = Guid.NewGuid().ToString("N");
|
||||
foreach (ChatMessage message in this.ResponseMessages)
|
||||
{
|
||||
yield return new(message.Role, message.Contents)
|
||||
{
|
||||
ResponseId = responseId,
|
||||
MessageId = message.MessageId,
|
||||
CreatedAt = DateTimeOffset.Now
|
||||
};
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -374,6 +374,7 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
contents=contents,
|
||||
role="assistant" if item.role == A2ARole.agent else "user",
|
||||
response_id=str(getattr(item, "message_id", uuid.uuid4())),
|
||||
additional_properties={"a2a_metadata": item.metadata} if item.metadata else None,
|
||||
raw_representation=item,
|
||||
)
|
||||
all_updates.append(update)
|
||||
@@ -452,13 +453,24 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
role=message.role,
|
||||
response_id=task.id,
|
||||
message_id=getattr(message.raw_representation, "artifact_id", None),
|
||||
additional_properties={"a2a_metadata": merged}
|
||||
if (merged := {**message.additional_properties, **(task.metadata or {})})
|
||||
else None,
|
||||
raw_representation=task,
|
||||
)
|
||||
for message in task_messages
|
||||
]
|
||||
if task.artifacts is not None:
|
||||
return []
|
||||
return [AgentResponseUpdate(contents=[], role="assistant", response_id=task.id, raw_representation=task)]
|
||||
return [
|
||||
AgentResponseUpdate(
|
||||
contents=[],
|
||||
role="assistant",
|
||||
response_id=task.id,
|
||||
additional_properties={"a2a_metadata": task.metadata} if task.metadata else None,
|
||||
raw_representation=task,
|
||||
)
|
||||
]
|
||||
|
||||
if background and status.state in IN_PROGRESS_TASK_STATES:
|
||||
token = self._build_continuation_token(task)
|
||||
@@ -468,6 +480,7 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
role="assistant",
|
||||
response_id=task.id,
|
||||
continuation_token=token,
|
||||
additional_properties={"a2a_metadata": task.metadata} if task.metadata else None,
|
||||
raw_representation=task,
|
||||
)
|
||||
]
|
||||
@@ -488,6 +501,7 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
contents=contents,
|
||||
role="assistant" if status.message.role == A2ARole.agent else "user",
|
||||
response_id=task.id,
|
||||
additional_properties={"a2a_metadata": task.metadata} if task.metadata else None,
|
||||
raw_representation=task,
|
||||
)
|
||||
]
|
||||
@@ -502,12 +516,17 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
contents = self._parse_contents_from_a2a(update_event.artifact.parts)
|
||||
if not contents:
|
||||
return []
|
||||
merged_metadata = {
|
||||
**(update_event.artifact.metadata or {}),
|
||||
**(update_event.metadata or {}),
|
||||
} or None
|
||||
return [
|
||||
AgentResponseUpdate(
|
||||
contents=contents,
|
||||
role="assistant",
|
||||
response_id=update_event.task_id,
|
||||
message_id=update_event.artifact.artifact_id,
|
||||
additional_properties={"a2a_metadata": merged_metadata} if merged_metadata else None,
|
||||
raw_representation=update_event,
|
||||
)
|
||||
]
|
||||
@@ -523,11 +542,16 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
if not contents:
|
||||
return []
|
||||
|
||||
merged_metadata = {
|
||||
**(message.metadata or {}),
|
||||
**(update_event.metadata or {}),
|
||||
} or None
|
||||
return [
|
||||
AgentResponseUpdate(
|
||||
contents=contents,
|
||||
role="assistant" if message.role == A2ARole.agent else "user",
|
||||
response_id=update_event.task_id,
|
||||
additional_properties={"a2a_metadata": merged_metadata} if merged_metadata else None,
|
||||
raw_representation=update_event,
|
||||
)
|
||||
]
|
||||
@@ -642,9 +666,7 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
case _:
|
||||
raise ValueError(f"Unknown content type: {content.type}")
|
||||
|
||||
# Exclude framework-internal keys (e.g. attribution) from wire metadata
|
||||
internal_keys = {"_attribution", "context_id"}
|
||||
metadata = {k: v for k, v in message.additional_properties.items() if k not in internal_keys} or None
|
||||
metadata = message.additional_properties.get("a2a_metadata")
|
||||
|
||||
return A2AMessage(
|
||||
role=A2ARole("user"),
|
||||
@@ -718,6 +740,7 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
Message(
|
||||
role="assistant" if history_item.role == A2ARole.agent else "user",
|
||||
contents=contents,
|
||||
additional_properties=history_item.metadata,
|
||||
raw_representation=history_item,
|
||||
)
|
||||
)
|
||||
@@ -730,5 +753,6 @@ class A2AAgent(AgentTelemetryLayer, BaseAgent):
|
||||
return Message(
|
||||
role="assistant",
|
||||
contents=contents,
|
||||
additional_properties=artifact.metadata,
|
||||
raw_representation=artifact,
|
||||
)
|
||||
|
||||
@@ -530,7 +530,7 @@ def test_prepare_message_for_a2a_forwards_context_id() -> None:
|
||||
message = Message(
|
||||
role="user",
|
||||
contents=[Content.from_text(text="Continue the task")],
|
||||
additional_properties={"context_id": "ctx-123", "trace_id": "trace-456"},
|
||||
additional_properties={"context_id": "ctx-123", "a2a_metadata": {"trace_id": "trace-456"}},
|
||||
)
|
||||
|
||||
result = agent._prepare_message_for_a2a(message)
|
||||
@@ -1385,3 +1385,210 @@ async def test_streaming_terminal_task_only_emits_unstreamed_artifacts(
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
# region Metadata propagation tests
|
||||
|
||||
|
||||
async def test_message_metadata_propagated(a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient) -> None:
|
||||
"""A2AMessage.metadata should appear on response.additional_properties."""
|
||||
msg = A2AMessage(
|
||||
message_id="msg-meta",
|
||||
role=A2ARole.agent,
|
||||
parts=[Part(root=TextPart(text="hi"))],
|
||||
metadata={"source": "server", "trace_id": "abc"},
|
||||
)
|
||||
mock_a2a_client.responses.append(msg)
|
||||
|
||||
response = await a2a_agent.run("hello")
|
||||
assert response.additional_properties["a2a_metadata"]["source"] == "server"
|
||||
assert response.additional_properties["a2a_metadata"]["trace_id"] == "abc"
|
||||
|
||||
|
||||
async def test_artifact_metadata_propagated(a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient) -> None:
|
||||
"""Artifact.metadata should appear on response.additional_properties."""
|
||||
task = Task(
|
||||
id="task-art-meta",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.completed),
|
||||
artifacts=[
|
||||
Artifact(
|
||||
artifact_id="a1",
|
||||
parts=[Part(root=TextPart(text="result"))],
|
||||
metadata={"artifact_key": "artifact_value"},
|
||||
),
|
||||
],
|
||||
)
|
||||
mock_a2a_client.responses.append((task, None))
|
||||
|
||||
response = await a2a_agent.run("go")
|
||||
assert response.additional_properties["a2a_metadata"]["artifact_key"] == "artifact_value"
|
||||
|
||||
|
||||
async def test_task_metadata_propagated_to_response(a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient) -> None:
|
||||
"""Task.metadata should appear on response.additional_properties for terminal tasks."""
|
||||
task = Task(
|
||||
id="task-meta",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.completed),
|
||||
artifacts=[
|
||||
Artifact(artifact_id="a1", parts=[Part(root=TextPart(text="done"))]),
|
||||
],
|
||||
metadata={"task_key": "task_value"},
|
||||
)
|
||||
mock_a2a_client.responses.append((task, None))
|
||||
|
||||
response = await a2a_agent.run("go")
|
||||
assert response.additional_properties["a2a_metadata"]["task_key"] == "task_value"
|
||||
|
||||
|
||||
async def test_task_artifact_update_event_metadata_merged(a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient) -> None:
|
||||
"""TaskArtifactUpdateEvent and Artifact metadata should both appear on the streaming update."""
|
||||
artifact_event = TaskArtifactUpdateEvent(
|
||||
task_id="task-ae",
|
||||
context_id="ctx",
|
||||
artifact=Artifact(
|
||||
artifact_id="a1",
|
||||
parts=[Part(root=TextPart(text="chunk"))],
|
||||
metadata={"from_artifact": True},
|
||||
),
|
||||
metadata={"from_event": True},
|
||||
)
|
||||
working_task = Task(
|
||||
id="task-ae",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.working),
|
||||
)
|
||||
terminal_task = Task(
|
||||
id="task-ae",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.completed),
|
||||
artifacts=[
|
||||
Artifact(artifact_id="a1", parts=[Part(root=TextPart(text="chunk"))]),
|
||||
],
|
||||
)
|
||||
terminal_event = TaskStatusUpdateEvent(
|
||||
task_id="task-ae",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.completed),
|
||||
final=True,
|
||||
)
|
||||
mock_a2a_client.responses.extend([
|
||||
(working_task, artifact_event),
|
||||
(terminal_task, terminal_event),
|
||||
])
|
||||
|
||||
stream = a2a_agent.run("hello", stream=True)
|
||||
updates: list[AgentResponseUpdate] = []
|
||||
async for update in stream:
|
||||
updates.append(update)
|
||||
|
||||
artifact_update = updates[0]
|
||||
assert artifact_update.additional_properties["a2a_metadata"]["from_artifact"] is True
|
||||
assert artifact_update.additional_properties["a2a_metadata"]["from_event"] is True
|
||||
|
||||
|
||||
async def test_task_status_update_event_metadata_merged(a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient) -> None:
|
||||
"""TaskStatusUpdateEvent and its message metadata should both appear on the streaming update."""
|
||||
status_event = TaskStatusUpdateEvent(
|
||||
task_id="task-se",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(
|
||||
state=TaskState.working,
|
||||
message=A2AMessage(
|
||||
message_id="m1",
|
||||
role=A2ARole.agent,
|
||||
parts=[Part(root=TextPart(text="working..."))],
|
||||
metadata={"msg_key": "msg_val"},
|
||||
),
|
||||
),
|
||||
final=False,
|
||||
metadata={"event_key": "event_val"},
|
||||
)
|
||||
working_task = Task(
|
||||
id="task-se",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.working),
|
||||
)
|
||||
terminal_task = Task(
|
||||
id="task-se",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.completed),
|
||||
artifacts=[
|
||||
Artifact(artifact_id="a1", parts=[Part(root=TextPart(text="done"))]),
|
||||
],
|
||||
)
|
||||
terminal_event = TaskStatusUpdateEvent(
|
||||
task_id="task-se",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.completed),
|
||||
final=True,
|
||||
)
|
||||
mock_a2a_client.responses.extend([
|
||||
(working_task, status_event),
|
||||
(terminal_task, terminal_event),
|
||||
])
|
||||
|
||||
stream = a2a_agent.run("hello", stream=True)
|
||||
updates: list[AgentResponseUpdate] = []
|
||||
async for update in stream:
|
||||
updates.append(update)
|
||||
|
||||
status_update = updates[0]
|
||||
assert status_update.additional_properties["a2a_metadata"]["msg_key"] == "msg_val"
|
||||
assert status_update.additional_properties["a2a_metadata"]["event_key"] == "event_val"
|
||||
|
||||
|
||||
async def test_history_message_metadata_propagated(a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient) -> None:
|
||||
"""Metadata on a history Message should appear on response.additional_properties."""
|
||||
task = Task(
|
||||
id="task-hist",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.completed),
|
||||
history=[
|
||||
A2AMessage(
|
||||
message_id="h1",
|
||||
role=A2ARole.agent,
|
||||
parts=[Part(root=TextPart(text="reply"))],
|
||||
metadata={"history_key": "history_value"},
|
||||
),
|
||||
],
|
||||
)
|
||||
mock_a2a_client.responses.append((task, None))
|
||||
|
||||
response = await a2a_agent.run("go")
|
||||
assert response.additional_properties["a2a_metadata"]["history_key"] == "history_value"
|
||||
|
||||
|
||||
async def test_continuation_token_update_carries_task_metadata(
|
||||
a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient
|
||||
) -> None:
|
||||
"""In-progress tasks with background=True should propagate task metadata."""
|
||||
task = Task(
|
||||
id="task-cont",
|
||||
context_id="ctx",
|
||||
status=TaskStatus(state=TaskState.working),
|
||||
metadata={"bg_key": "bg_value"},
|
||||
)
|
||||
mock_a2a_client.responses.append((task, None))
|
||||
|
||||
response = await a2a_agent.run("go", background=True)
|
||||
assert response.continuation_token is not None
|
||||
assert response.additional_properties["a2a_metadata"]["bg_key"] == "bg_value"
|
||||
|
||||
|
||||
async def test_none_metadata_leaves_additional_properties_empty(
|
||||
a2a_agent: A2AAgent, mock_a2a_client: MockA2AClient
|
||||
) -> None:
|
||||
"""When A2A types have no metadata, additional_properties should remain empty/default."""
|
||||
msg = A2AMessage(
|
||||
message_id="msg-none",
|
||||
role=A2ARole.agent,
|
||||
parts=[Part(root=TextPart(text="no meta"))],
|
||||
)
|
||||
mock_a2a_client.responses.append(msg)
|
||||
|
||||
response = await a2a_agent.run("hello")
|
||||
assert not response.additional_properties
|
||||
|
||||
|
||||
# endregion
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
# Copyright (c) Microsoft. All rights reserved.
|
||||
import os
|
||||
import re
|
||||
from pathlib import Path
|
||||
from typing import Annotated, Any
|
||||
from unittest.mock import MagicMock, patch
|
||||
@@ -1503,6 +1504,8 @@ async def test_anthropic_client_integration_function_calling() -> None:
|
||||
@skip_if_anthropic_integration_tests_disabled
|
||||
async def test_anthropic_client_integration_hosted_tools() -> None:
|
||||
"""Integration test for hosted tools."""
|
||||
import anthropic
|
||||
|
||||
client = AnthropicClient()
|
||||
|
||||
messages = [Message(role="user", contents=["What tools do you have available?"])]
|
||||
@@ -1515,10 +1518,18 @@ async def test_anthropic_client_integration_hosted_tools() -> None:
|
||||
),
|
||||
]
|
||||
|
||||
response = await client.get_response(
|
||||
messages=messages,
|
||||
options={"tools": tools, "max_tokens": 100},
|
||||
)
|
||||
try:
|
||||
response = await client.get_response(
|
||||
messages=messages,
|
||||
options={"tools": tools, "max_tokens": 100},
|
||||
)
|
||||
except (
|
||||
anthropic.BadRequestError,
|
||||
anthropic.InternalServerError,
|
||||
anthropic.APIConnectionError,
|
||||
anthropic.APITimeoutError,
|
||||
) as e:
|
||||
pytest.skip(f"Upstream MCP server unavailable: {e}")
|
||||
|
||||
assert response is not None
|
||||
assert response.text is not None
|
||||
@@ -1607,7 +1618,8 @@ async def test_anthropic_client_integration_images() -> None:
|
||||
|
||||
assert response is not None
|
||||
assert response.messages[0].text is not None
|
||||
assert "house" in response.messages[0].text.lower()
|
||||
text = response.messages[0].text.lower()
|
||||
assert re.search(r"\b(house|home|building|cottage|mansion|villa)\b", text)
|
||||
|
||||
|
||||
# Response Format Tests
|
||||
|
||||
@@ -59,6 +59,62 @@ class AgentExecutorResponse:
|
||||
agent_response: AgentResponse
|
||||
full_conversation: list[Message]
|
||||
|
||||
def with_text(self, text: str) -> "AgentExecutorResponse":
|
||||
"""Create a new AgentExecutorResponse with replaced text, preserving the conversation history.
|
||||
|
||||
Use this in custom executors that transform agent output text (e.g. upper-casing, summarising)
|
||||
when you need downstream AgentExecutors to still have access to the full prior conversation.
|
||||
|
||||
Without this helper, sending a plain ``str`` from a custom executor breaks the context chain:
|
||||
the downstream ``AgentExecutor.from_str`` handler only adds that one string to its cache and
|
||||
loses all prior messages. By using ``with_text`` the response type stays
|
||||
``AgentExecutorResponse``, so ``AgentExecutor.from_response`` is invoked instead and the full
|
||||
conversation is preserved.
|
||||
|
||||
Args:
|
||||
text: The replacement assistant message text.
|
||||
|
||||
Returns:
|
||||
A new ``AgentExecutorResponse`` whose ``agent_response`` contains a single assistant
|
||||
message with ``text``, and whose ``full_conversation`` is the prior conversation
|
||||
(everything before the original agent turn) followed by the new assistant message.
|
||||
|
||||
Example:
|
||||
.. code-block:: python
|
||||
|
||||
from agent_framework import AgentExecutorResponse, WorkflowContext, executor
|
||||
|
||||
|
||||
@executor(
|
||||
id="upper_case_executor",
|
||||
input=AgentExecutorResponse,
|
||||
output=AgentExecutorResponse,
|
||||
workflow_output=str,
|
||||
)
|
||||
async def upper_case(
|
||||
response: AgentExecutorResponse,
|
||||
ctx: WorkflowContext[AgentExecutorResponse, str],
|
||||
) -> None:
|
||||
upper_text = response.agent_response.text.upper()
|
||||
await ctx.send_message(response.with_text(upper_text))
|
||||
await ctx.yield_output(upper_text)
|
||||
"""
|
||||
new_message = Message("assistant", [text])
|
||||
new_agent_response = AgentResponse(messages=[new_message])
|
||||
|
||||
# Strip off the original agent turn and replace with the new text.
|
||||
n_agent_messages = len(self.agent_response.messages)
|
||||
prior_messages = (
|
||||
self.full_conversation[:-n_agent_messages] if n_agent_messages else list(self.full_conversation)
|
||||
)
|
||||
new_full_conversation = [*prior_messages, new_message]
|
||||
|
||||
return AgentExecutorResponse(
|
||||
executor_id=self.executor_id,
|
||||
agent_response=new_agent_response,
|
||||
full_conversation=new_full_conversation,
|
||||
)
|
||||
|
||||
|
||||
class AgentExecutor(Executor):
|
||||
"""built-in executor that wraps an agent for handling messages.
|
||||
@@ -183,7 +239,25 @@ class AgentExecutor(Executor):
|
||||
"""Accept a raw user prompt string and run the agent.
|
||||
|
||||
The new string input will be added to the cache which is used as the conversation context for the agent run.
|
||||
|
||||
Warning:
|
||||
If the upstream executor received an ``AgentExecutorResponse`` but emits a plain
|
||||
``str``, this handler will be invoked instead of ``from_response``. This resets
|
||||
the conversation context because only the new string is added to the cache and
|
||||
all prior messages from the upstream agent are lost.
|
||||
|
||||
To preserve the full conversation when transforming agent output in a custom
|
||||
executor, use ``AgentExecutorResponse.with_text(...)`` so that the message type
|
||||
stays ``AgentExecutorResponse`` and ``from_response`` is called instead.
|
||||
"""
|
||||
if not self._cache and ctx.source_executor_ids != ["Workflow"]:
|
||||
logger.warning(
|
||||
"AgentExecutor '%s': from_str handler invoked with an empty cache. "
|
||||
"If you are chaining from an AgentExecutor, the upstream custom executor may be "
|
||||
"emitting a plain str instead of using AgentExecutorResponse.with_text(...), "
|
||||
"which causes the full conversation context to be lost.",
|
||||
self.id,
|
||||
)
|
||||
self._cache.extend(normalize_messages_input(text))
|
||||
await self._run_agent_and_emit(ctx)
|
||||
|
||||
|
||||
@@ -244,10 +244,10 @@ class FileCheckpointStorage:
|
||||
is serialized using pickle and embedded as base64-encoded strings within the JSON. This allows
|
||||
for human-readable checkpoint files while preserving the ability to store complex Python objects.
|
||||
|
||||
By default, checkpoint deserialization is restricted to a built-in set of safe
|
||||
Python types (primitives, datetime, uuid, ...) and all ``agent_framework``
|
||||
internal types. To allow additional application-specific types, pass them via
|
||||
the ``allowed_checkpoint_types`` parameter using ``"module:qualname"`` format.
|
||||
By default, checkpoint deserialization is restricted to a built-in set of safe Python types
|
||||
(primitives, datetime, uuid, ...), all ``agent_framework`` internal types, and OpenAI SDK types
|
||||
(``openai.types``). To allow additional application-specific types, pass them via the
|
||||
``allowed_checkpoint_types`` parameter using ``"module:qualname"`` format.
|
||||
|
||||
Example::
|
||||
|
||||
|
||||
@@ -10,9 +10,9 @@ This hybrid approach provides:
|
||||
When ``allowed_types`` is supplied to :func:`decode_checkpoint_value`, a
|
||||
``RestrictedUnpickler`` is used that limits which classes may be instantiated
|
||||
during deserialization. The default built-in safe set covers common Python
|
||||
value types (primitives, datetime, uuid, ...) and all ``agent_framework``
|
||||
internal types. Callers can extend the set by passing additional
|
||||
``"module:qualname"`` strings.
|
||||
value types (primitives, datetime, uuid, ...), all ``agent_framework`` internal
|
||||
types, and all ``openai.types`` types. Callers can extend the set by passing
|
||||
additional ``"module:qualname"`` strings.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
@@ -37,6 +37,9 @@ _JSON_NATIVE_TYPES = (str, int, float, bool, type(None))
|
||||
# Module prefix for framework-internal types that are always allowed
|
||||
_FRAMEWORK_MODULE_PREFIX = "agent_framework."
|
||||
|
||||
# Module prefix for OpenAI SDK types that are always allowed
|
||||
_OPENAI_MODULE_PREFIX = "openai.types."
|
||||
|
||||
# Built-in types considered safe for checkpoint deserialization.
|
||||
# Each entry is a ``module:qualname`` string matching the format produced by
|
||||
# :func:`_type_to_key`. These are the classes for which pickle's
|
||||
@@ -84,8 +87,9 @@ class _RestrictedUnpickler(pickle.Unpickler): # noqa: S301
|
||||
"""Unpickler that restricts which classes may be instantiated.
|
||||
|
||||
Only classes whose ``module:qualname`` key appears in the combined allow
|
||||
set (built-in safe types + framework types + caller-specified extras) are
|
||||
permitted. All other classes raise :class:`pickle.UnpicklingError`.
|
||||
set (built-in safe types + framework types + OpenAI SDK types +
|
||||
caller-specified extras) are permitted. All other classes raise
|
||||
:class:`pickle.UnpicklingError`.
|
||||
"""
|
||||
|
||||
def __init__(self, data: bytes, allowed_types: frozenset[str]) -> None:
|
||||
@@ -99,6 +103,7 @@ class _RestrictedUnpickler(pickle.Unpickler): # noqa: S301
|
||||
type_key in _BUILTIN_ALLOWED_TYPE_KEYS
|
||||
or type_key in self._allowed_types
|
||||
or module.startswith(_FRAMEWORK_MODULE_PREFIX)
|
||||
or module.startswith(_OPENAI_MODULE_PREFIX)
|
||||
):
|
||||
return super().find_class(module, name) # type: ignore[no-any-return] # nosec
|
||||
|
||||
|
||||
@@ -268,6 +268,19 @@ def executor(
|
||||
forward references. When provided, takes precedence over introspection from the
|
||||
``WorkflowContext`` second generic parameter (W_OutT).
|
||||
|
||||
Warning:
|
||||
When placing a custom ``@executor`` **between** two ``AgentExecutor`` nodes, be
|
||||
careful about the output type. If the custom executor receives an
|
||||
``AgentExecutorResponse`` but emits a plain ``str``, the downstream
|
||||
``AgentExecutor.from_str`` handler is invoked instead of ``from_response``.
|
||||
This resets the conversation context because only the new string is added to
|
||||
the cache and all prior messages from the upstream agent are lost.
|
||||
|
||||
To preserve the full conversation, use
|
||||
``AgentExecutorResponse.with_text(new_text)`` to create a new response that
|
||||
keeps the prior history, and set ``output=AgentExecutorResponse`` on the
|
||||
decorator.
|
||||
|
||||
Returns:
|
||||
A FunctionExecutor instance that can be wired into a Workflow.
|
||||
|
||||
|
||||
@@ -11,11 +11,11 @@ import logging
|
||||
import types
|
||||
import uuid
|
||||
from collections.abc import AsyncIterable, Awaitable, Callable, Mapping, Sequence
|
||||
from typing import Any, Literal, overload
|
||||
from typing import TYPE_CHECKING, Any, Literal, overload
|
||||
|
||||
from .._sessions import ContextProvider
|
||||
from .._types import ResponseStream
|
||||
from ..observability import OtelAttr, capture_exception, create_workflow_span
|
||||
from ._agent import WorkflowAgent
|
||||
from ._checkpoint import CheckpointStorage
|
||||
from ._const import DEFAULT_MAX_ITERATIONS, GLOBAL_KWARGS_KEY, WORKFLOW_RUN_KWARGS_KEY
|
||||
from ._edge import (
|
||||
@@ -35,6 +35,9 @@ from ._runner_context import RunnerContext
|
||||
from ._state import State
|
||||
from ._typing_utils import is_instance_of, try_coerce_to_type
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ._agent import WorkflowAgent
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -910,7 +913,14 @@ class Workflow(DictConvertible):
|
||||
|
||||
return list(output_types)
|
||||
|
||||
def as_agent(self, name: str | None = None) -> WorkflowAgent:
|
||||
def as_agent(
|
||||
self,
|
||||
name: str | None = None,
|
||||
*,
|
||||
description: str | None = None,
|
||||
context_providers: Sequence[ContextProvider] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> WorkflowAgent:
|
||||
"""Create a WorkflowAgent that wraps this workflow.
|
||||
|
||||
The returned agent converts standard agent inputs (strings, Message, or lists of these)
|
||||
@@ -924,7 +934,10 @@ class Workflow(DictConvertible):
|
||||
initialization will fail with a ValueError.
|
||||
|
||||
Args:
|
||||
name: Optional name for the agent. If None, a default name will be generated.
|
||||
name: Optional name for the agent. Defaults to workflow name.
|
||||
description: Optional description of the agent. Defaults to workflow description.
|
||||
context_providers: Optional sequence of context providers for the agent.
|
||||
**kwargs: Additional keyword arguments passed to BaseAgent.
|
||||
|
||||
Returns:
|
||||
A WorkflowAgent instance that wraps this workflow.
|
||||
@@ -935,4 +948,10 @@ class Workflow(DictConvertible):
|
||||
# Import here to avoid circular imports
|
||||
from ._agent import WorkflowAgent
|
||||
|
||||
return WorkflowAgent(workflow=self, name=name)
|
||||
return WorkflowAgent(
|
||||
workflow=self,
|
||||
name=name if name is not None else self.name,
|
||||
description=description if description is not None else self.description,
|
||||
context_providers=context_providers,
|
||||
**kwargs,
|
||||
)
|
||||
|
||||
@@ -216,3 +216,50 @@ def test_restricted_unpickler_raises_pickle_error():
|
||||
unpickler = _RestrictedUnpickler(pickled, frozenset())
|
||||
with pytest.raises(pickle.UnpicklingError, match="deserialization blocked"):
|
||||
unpickler.load()
|
||||
|
||||
|
||||
def test_restricted_decode_allows_openai_types():
|
||||
"""OpenAI SDK types are always allowed during restricted deserialization."""
|
||||
from openai.types.chat.chat_completion import ChatCompletion, Choice
|
||||
from openai.types.chat.chat_completion_message import ChatCompletionMessage
|
||||
from openai.types.completion_usage import CompletionUsage
|
||||
|
||||
completion = ChatCompletion(
|
||||
id="chatcmpl-test",
|
||||
choices=[
|
||||
Choice(
|
||||
finish_reason="stop",
|
||||
index=0,
|
||||
message=ChatCompletionMessage(role="assistant", content="hello"),
|
||||
)
|
||||
],
|
||||
created=1700000000,
|
||||
model="gpt-4",
|
||||
object="chat.completion",
|
||||
usage=CompletionUsage(completion_tokens=1, prompt_tokens=1, total_tokens=2),
|
||||
)
|
||||
encoded = encode_checkpoint_value(completion)
|
||||
decoded = decode_checkpoint_value(encoded, allowed_types=frozenset())
|
||||
|
||||
assert isinstance(decoded, ChatCompletion)
|
||||
assert decoded.id == "chatcmpl-test"
|
||||
assert decoded.choices[0].message.content == "hello"
|
||||
|
||||
|
||||
def test_restricted_decode_allows_openai_response_types():
|
||||
"""OpenAI Responses API types are always allowed during restricted deserialization."""
|
||||
from openai.types.responses.response_usage import InputTokensDetails, OutputTokensDetails, ResponseUsage
|
||||
|
||||
usage = ResponseUsage(
|
||||
input_tokens=10,
|
||||
output_tokens=20,
|
||||
total_tokens=30,
|
||||
input_tokens_details=InputTokensDetails(cached_tokens=0),
|
||||
output_tokens_details=OutputTokensDetails(reasoning_tokens=0),
|
||||
)
|
||||
encoded = encode_checkpoint_value(usage)
|
||||
decoded = decode_checkpoint_value(encoded, allowed_types=frozenset())
|
||||
|
||||
assert isinstance(decoded, ResponseUsage)
|
||||
assert decoded.input_tokens == 10
|
||||
assert decoded.output_tokens == 20
|
||||
|
||||
@@ -23,6 +23,7 @@ from agent_framework import (
|
||||
WorkflowBuilder,
|
||||
WorkflowContext,
|
||||
WorkflowRunState,
|
||||
executor,
|
||||
handler,
|
||||
)
|
||||
from agent_framework.orchestrations import SequentialBuilder
|
||||
@@ -478,3 +479,90 @@ async def test_from_response_preserves_service_session_id() -> None:
|
||||
assert result.get_outputs() is not None
|
||||
|
||||
assert spy_agent._captured_service_session_id == "resp_PREVIOUS_RUN" # pyright: ignore[reportPrivateUsage]
|
||||
|
||||
|
||||
@executor(
|
||||
id="upper_case_executor",
|
||||
input=AgentExecutorResponse,
|
||||
output=AgentExecutorResponse,
|
||||
workflow_output=str,
|
||||
)
|
||||
async def _upper_case_executor(
|
||||
response: AgentExecutorResponse,
|
||||
ctx: WorkflowContext[AgentExecutorResponse, str],
|
||||
) -> None:
|
||||
upper_text = response.agent_response.text.upper()
|
||||
await ctx.send_message(response.with_text(upper_text))
|
||||
await ctx.yield_output(upper_text)
|
||||
|
||||
|
||||
async def test_with_text_preserves_full_conversation_through_custom_executor() -> None:
|
||||
"""Custom executor using with_text must preserve the full conversation chain."""
|
||||
# Mirrors the reproduction from issue #5246:
|
||||
# agent1 ("User likes sky red") -> agent2 ("User likes sky blue") -> upper_case -> agent3 ("User likes sky green")
|
||||
agent1 = AgentExecutor(
|
||||
_SimpleAgent(id="agent1", name="ContextAgent1", reply_text="User likes sky red"), id="agent1"
|
||||
)
|
||||
agent2 = AgentExecutor(
|
||||
_SimpleAgent(id="agent2", name="ContextAgent2", reply_text="User likes sky blue"), id="agent2"
|
||||
)
|
||||
agent3 = AgentExecutor(
|
||||
_SimpleAgent(id="agent3", name="ContextAgent3", reply_text="User likes sky green"), id="agent3"
|
||||
)
|
||||
capturer = _CaptureFullConversation(id="capture")
|
||||
|
||||
wf = (
|
||||
WorkflowBuilder(start_executor=agent1, output_executors=[capturer])
|
||||
.add_chain([agent1, agent2, _upper_case_executor, agent3, capturer])
|
||||
.build()
|
||||
)
|
||||
|
||||
result = await wf.run("")
|
||||
payload = next(o for o in result.get_outputs() if isinstance(o, dict))
|
||||
|
||||
# The final agent must see the full conversation: user, agent1, UPPER(agent2), agent3
|
||||
assert payload["roles"] == ["user", "assistant", "assistant", "assistant"]
|
||||
assert payload["texts"][1] == "User likes sky red"
|
||||
assert payload["texts"][2] == "USER LIKES SKY BLUE"
|
||||
assert payload["texts"][3] == "User likes sky green"
|
||||
|
||||
|
||||
async def test_with_text_does_not_mutate_original() -> None:
|
||||
"""with_text returns a new instance; the original must be unmodified."""
|
||||
original = AgentExecutorResponse(
|
||||
executor_id="test_exec",
|
||||
agent_response=AgentResponse(messages=[Message("assistant", ["original reply"])]),
|
||||
full_conversation=[Message("user", ["prompt"]), Message("assistant", ["original reply"])],
|
||||
)
|
||||
|
||||
new = original.with_text("transformed reply")
|
||||
|
||||
assert new is not original
|
||||
assert new.agent_response.text == "transformed reply"
|
||||
assert new.full_conversation[-1].text == "transformed reply"
|
||||
assert new.full_conversation[-1].role == "assistant"
|
||||
# Original unchanged
|
||||
assert original.agent_response.text == "original reply"
|
||||
assert original.full_conversation[-1].text == "original reply"
|
||||
|
||||
|
||||
async def test_with_text_strips_multi_message_agent_turn() -> None:
|
||||
"""When the agent turn has multiple messages (tool calls), with_text strips all of them."""
|
||||
tool_call = Message("assistant", ["<tool_call>"])
|
||||
tool_result = Message("tool", ["<result>"])
|
||||
final_reply = Message("assistant", ["actual answer"])
|
||||
user_msg = Message("user", ["question"])
|
||||
|
||||
original = AgentExecutorResponse(
|
||||
executor_id="exec",
|
||||
agent_response=AgentResponse(messages=[tool_call, tool_result, final_reply]),
|
||||
full_conversation=[user_msg, tool_call, tool_result, final_reply],
|
||||
)
|
||||
|
||||
new = original.with_text("summarised answer")
|
||||
|
||||
# Only the pre-agent-turn messages should remain, plus the replacement
|
||||
assert len(new.full_conversation) == 2
|
||||
assert new.full_conversation[0].text == "question"
|
||||
assert new.full_conversation[1].text == "summarised answer"
|
||||
assert new.agent_response.text == "summarised answer"
|
||||
|
||||
@@ -313,6 +313,37 @@ class TestWorkflowAgent:
|
||||
assert isinstance(agent_no_name, WorkflowAgent)
|
||||
assert agent_no_name.workflow is workflow
|
||||
|
||||
def test_workflow_as_agent_with_description_and_context_providers(self) -> None:
|
||||
"""Test that Workflow.as_agent() forwards description and context_providers."""
|
||||
executor = SimpleExecutor(id="executor1", response_text="Response")
|
||||
workflow = WorkflowBuilder(start_executor=executor).build()
|
||||
|
||||
history_provider = InMemoryHistoryProvider()
|
||||
agent = workflow.as_agent(
|
||||
name="MyAgent",
|
||||
description="A test agent",
|
||||
context_providers=[history_provider],
|
||||
)
|
||||
|
||||
assert isinstance(agent, WorkflowAgent)
|
||||
assert agent.name == "MyAgent"
|
||||
assert agent.description == "A test agent"
|
||||
assert history_provider in agent.context_providers
|
||||
|
||||
def test_workflow_as_agent_defaults_name_and_description_from_workflow(self) -> None:
|
||||
"""Test that as_agent() defaults name and description to the workflow's own values."""
|
||||
executor = SimpleExecutor(id="executor1", response_text="Response")
|
||||
workflow = WorkflowBuilder(
|
||||
start_executor=executor,
|
||||
name="my-workflow",
|
||||
description="Workflow description",
|
||||
).build()
|
||||
|
||||
agent = workflow.as_agent()
|
||||
|
||||
assert agent.name == "my-workflow"
|
||||
assert agent.description == "Workflow description"
|
||||
|
||||
def test_workflow_as_agent_cannot_handle_agent_inputs(self) -> None:
|
||||
"""Test that Workflow.as_agent() raises an error if the start executor cannot handle agent inputs."""
|
||||
|
||||
|
||||
@@ -2474,6 +2474,29 @@ class RawOpenAIChatClient( # type: ignore[misc]
|
||||
raw_representation=event,
|
||||
)
|
||||
)
|
||||
elif ann_type == "url_citation":
|
||||
ann_url = _get_ann_value("url")
|
||||
if ann_url:
|
||||
ann_start = _get_ann_value("start_index")
|
||||
ann_end = _get_ann_value("end_index")
|
||||
annotation_obj = Annotation(
|
||||
type="citation",
|
||||
title=_get_ann_value("title") or "",
|
||||
url=str(ann_url),
|
||||
additional_properties={"annotation_index": event.annotation_index},
|
||||
raw_representation=annotation,
|
||||
)
|
||||
if ann_start is not None and ann_end is not None:
|
||||
annotation_obj["annotated_regions"] = [
|
||||
TextSpanRegion(
|
||||
type="text_span",
|
||||
start_index=ann_start,
|
||||
end_index=ann_end,
|
||||
)
|
||||
]
|
||||
contents.append(
|
||||
Content.from_text(text="", annotations=[annotation_obj], raw_representation=event)
|
||||
)
|
||||
else:
|
||||
logger.debug("Unparsed annotation type in streaming: %s", ann_type)
|
||||
case "response.output_item.done":
|
||||
|
||||
@@ -2570,8 +2570,65 @@ def test_streaming_annotation_added_with_container_file_citation() -> None:
|
||||
assert content.additional_properties.get("end_index") == 50
|
||||
|
||||
|
||||
def test_streaming_annotation_added_with_unknown_type() -> None:
|
||||
"""Test streaming annotation added event with unknown type is ignored."""
|
||||
def test_streaming_annotation_added_with_url_citation() -> None:
|
||||
"""Test streaming annotation added event with url_citation type produces citation annotation."""
|
||||
client = OpenAIChatClient(model="test-model", api_key="test-key")
|
||||
chat_options = ChatOptions()
|
||||
function_call_ids: dict[int, tuple[str, str]] = {}
|
||||
|
||||
mock_event = MagicMock()
|
||||
mock_event.type = "response.output_text.annotation.added"
|
||||
mock_event.annotation_index = 0
|
||||
mock_event.annotation = {
|
||||
"type": "url_citation",
|
||||
"url": "https://example.sharepoint.com/sites/my-site/doc.pdf",
|
||||
"title": "doc.pdf",
|
||||
"start_index": 100,
|
||||
"end_index": 112,
|
||||
}
|
||||
|
||||
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
|
||||
|
||||
assert len(response.contents) == 1
|
||||
content = response.contents[0]
|
||||
assert content.type == "text"
|
||||
assert content.annotations is not None
|
||||
assert len(content.annotations) == 1
|
||||
annotation = content.annotations[0]
|
||||
assert annotation["type"] == "citation"
|
||||
assert annotation["title"] == "doc.pdf"
|
||||
assert annotation["url"] == "https://example.sharepoint.com/sites/my-site/doc.pdf"
|
||||
assert annotation["additional_properties"]["annotation_index"] == 0
|
||||
assert annotation["raw_representation"] == mock_event.annotation
|
||||
assert annotation["annotated_regions"] is not None
|
||||
assert len(annotation["annotated_regions"]) == 1
|
||||
region = annotation["annotated_regions"][0]
|
||||
assert region["type"] == "text_span"
|
||||
assert region["start_index"] == 100
|
||||
assert region["end_index"] == 112
|
||||
|
||||
|
||||
def test_streaming_annotation_added_with_url_citation_no_url() -> None:
|
||||
"""Test streaming annotation added event with url_citation but missing url is ignored."""
|
||||
client = OpenAIChatClient(model="test-model", api_key="test-key")
|
||||
chat_options = ChatOptions()
|
||||
function_call_ids: dict[int, tuple[str, str]] = {}
|
||||
|
||||
mock_event = MagicMock()
|
||||
mock_event.type = "response.output_text.annotation.added"
|
||||
mock_event.annotation_index = 0
|
||||
mock_event.annotation = {
|
||||
"type": "url_citation",
|
||||
"title": "doc.pdf",
|
||||
}
|
||||
|
||||
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
|
||||
|
||||
assert len(response.contents) == 0
|
||||
|
||||
|
||||
def test_streaming_annotation_added_with_url_citation_no_indices() -> None:
|
||||
"""Test streaming annotation with url_citation that has url but no start_index/end_index."""
|
||||
client = OpenAIChatClient(model="test-model", api_key="test-key")
|
||||
chat_options = ChatOptions()
|
||||
function_call_ids: dict[int, tuple[str, str]] = {}
|
||||
@@ -2582,11 +2639,36 @@ def test_streaming_annotation_added_with_unknown_type() -> None:
|
||||
mock_event.annotation = {
|
||||
"type": "url_citation",
|
||||
"url": "https://example.com",
|
||||
"title": "Example",
|
||||
}
|
||||
|
||||
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
|
||||
|
||||
assert len(response.contents) == 1
|
||||
annotation = response.contents[0].annotations[0]
|
||||
assert annotation["type"] == "citation"
|
||||
assert annotation["title"] == "Example"
|
||||
assert annotation["url"] == "https://example.com"
|
||||
assert annotation["additional_properties"]["annotation_index"] == 0
|
||||
assert "annotated_regions" not in annotation
|
||||
|
||||
|
||||
def test_streaming_annotation_added_with_unknown_type() -> None:
|
||||
"""Test streaming annotation added event with unknown type is ignored."""
|
||||
client = OpenAIChatClient(model="test-model", api_key="test-key")
|
||||
chat_options = ChatOptions()
|
||||
function_call_ids: dict[int, tuple[str, str]] = {}
|
||||
|
||||
mock_event = MagicMock()
|
||||
mock_event.type = "response.output_text.annotation.added"
|
||||
mock_event.annotation_index = 0
|
||||
mock_event.annotation = {
|
||||
"type": "some_future_annotation_type",
|
||||
"data": "test",
|
||||
}
|
||||
|
||||
response = client._parse_chunk_from_openai(mock_event, chat_options, function_call_ids)
|
||||
|
||||
# url_citation should not produce HostedFileContent
|
||||
assert len(response.contents) == 0
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user