Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions dotnet/eng/MSBuild/Shared.props
Original file line number Diff line number Diff line change
Expand Up @@ -38,4 +38,7 @@
<ItemGroup Condition="'$(InjectSharedOriginPinning)' == 'true'">
<Compile Include="$(MSBuildThisFileDirectory)\..\..\src\Shared\OriginPinning\*.cs" LinkBase="Shared\OriginPinning" />
</ItemGroup>
<ItemGroup Condition="'$(InjectSharedHttpHeaderValidation)' == 'true'">
<Compile Include="$(MSBuildThisFileDirectory)\..\..\src\Shared\Http\*.cs" LinkBase="Shared\Http" />
</ItemGroup>
</Project>
Original file line number Diff line number Diff line change
Expand Up @@ -21,20 +21,20 @@ public static void Validate(string name, string value)
}

// Reject transport delimiters before using the name in exception text or a transport API.
if (ContainsProhibitedCharacter(name))
{
throw new ArgumentException("Header name must not contain NUL, carriage-return, or line-feed characters.", nameof(name));
}
HttpHeaderValidation.ValidateNoProhibitedCharacters(
name,
nameof(name),
"Header name must not contain NUL, carriage-return, or line-feed characters.");

if (value.Length == 0)
{
throw new ArgumentException("Header value must not be empty.", nameof(value));
}

if (ContainsProhibitedCharacter(value))
{
throw new ArgumentException("Header value must not contain NUL, carriage-return, or line-feed characters.", nameof(value));
}
HttpHeaderValidation.ValidateNoProhibitedCharacters(
value,
nameof(value),
"Header value must not contain NUL, carriage-return, or line-feed characters.");

if (!name.StartsWith(ClientHeaderPrefix, StringComparison.OrdinalIgnoreCase))
{
Expand All @@ -43,7 +43,4 @@ public static void Validate(string name, string value)
nameof(name));
}
}

private static bool ContainsProhibitedCharacter(string value) =>
value.IndexOf('\0') >= 0 || value.IndexOf('\r') >= 0 || value.IndexOf('\n') >= 0;
}
Original file line number Diff line number Diff line change
Expand Up @@ -108,6 +108,10 @@ internal set
{
_ = Throw.IfNull(session);
_ = Throw.IfNullOrWhitespace(value);
HttpHeaderValidation.ValidateNoProhibitedCharacters(
value,
"userIdentity",
"User identity must not contain NUL, carriage-return, or line-feed characters.");
session.StateBag.SetValue(FoundryHostedAgentUserIdentityKey, value);
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
related per-agent endpoint surface). Flip back to IsReleased=true once Azure.AI.Projects
ships a stable 3.0.0. -->
<InjectSharedFeatureUsageUserAgent>true</InjectSharedFeatureUsageUserAgent>
<InjectSharedHttpHeaderValidation>true</InjectSharedHttpHeaderValidation>
<InjectSharedThrow>true</InjectSharedThrow>
<NoWarn>$(NoWarn);AAIP001;AAIP002;OPENAI001;SCME0001</NoWarn>
</PropertyGroup>
Expand Down
13 changes: 13 additions & 0 deletions dotnet/src/Microsoft.Agents.AI.Foundry/UserIdentityPolicy.cs
Original file line number Diff line number Diff line change
Expand Up @@ -3,6 +3,7 @@
using System.ClientModel.Primitives;
using System.Collections.Generic;
using System.Threading.Tasks;
using Microsoft.Shared.Diagnostics;

namespace Microsoft.Agents.AI.Foundry;

Expand Down Expand Up @@ -33,6 +34,18 @@ public override ValueTask ProcessAsync(PipelineMessage message, IReadOnlyList<Pi
private static void Stamp(PipelineMessage message)
{
var identity = UserIdentityScope.Current;
if (identity is null)
{
return;
}

// Session state can be restored without using the binding API, so validate again at the
// final transport boundary before the value reaches the header collection.
HttpHeaderValidation.ValidateNoProhibitedCharacters(
Comment thread
rogerbarreto marked this conversation as resolved.
identity,
"userIdentity",
"User identity must not contain NUL, carriage-return, or line-feed characters.");
Comment thread
rogerbarreto marked this conversation as resolved.

if (string.IsNullOrWhiteSpace(identity))
{
return;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@
using System.Text;
using System.Threading;
using System.Threading.Tasks;
using Microsoft.Shared.Diagnostics;

namespace Microsoft.Agents.AI.Workflows.Declarative;

Expand Down Expand Up @@ -264,25 +265,22 @@ private static HttpRequestMessage BuildHttpRequestMessage(HttpRequestInfo reques

private static void ValidateHeader(string name, string value)
{
if (ContainsHttpHeaderDelimiter(name))
{
throw new ArgumentException("HTTP header name contains invalid characters.", nameof(name));
}
HttpHeaderValidation.ValidateNoProhibitedCharacters(
name,
nameof(name),
"HTTP header name contains invalid characters.");

ValidateHeaderValue(name, value);
}

private static void ValidateHeaderValue(string name, string value)
{
if (ContainsHttpHeaderDelimiter(value))
{
throw new ArgumentException($"HTTP header '{name}' contains invalid characters.", nameof(value));
}
HttpHeaderValidation.ValidateNoProhibitedCharacters(
value,
nameof(value),
$"HTTP header '{name}' contains invalid characters.");
}

private static bool ContainsHttpHeaderDelimiter(string value) =>
value.IndexOfAny(['\r', '\n', '\0']) >= 0;

private static HttpClient CreateOwnedHttpClient()
{
HttpClientHandler handler = new()
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@
</PropertyGroup>

<PropertyGroup>
<InjectSharedHttpHeaderValidation>true</InjectSharedHttpHeaderValidation>
<InjectSharedThrow>true</InjectSharedThrow>
<InjectSharedDiagnosticIds>true</InjectSharedDiagnosticIds>
<InjectIsExternalInitOnLegacy>true</InjectIsExternalInitOnLegacy>
Expand Down
17 changes: 17 additions & 0 deletions dotnet/src/Shared/Http/HttpHeaderValidation.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
// Copyright (c) Microsoft. All rights reserved.

using System;

namespace Microsoft.Shared.Diagnostics;

/// <summary>Validates values before they reach HTTP header transport APIs.</summary>
internal static class HttpHeaderValidation
{
public static void ValidateNoProhibitedCharacters(string value, string parameterName, string message)
{
if (value.IndexOfAny(['\r', '\n', '\0']) >= 0)
{
throw new ArgumentException(message, parameterName);
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -101,6 +101,169 @@ await Assert.ThrowsAsync<ArgumentException>(
() => agent.CreateFoundryHostedAgentSessionAsync(userIdentity: " "));
}

[Theory]
[InlineData("\0")]
[InlineData("\r")]
[InlineData("\n")]
public async Task CreateFoundryHostedAgentSessionAsync_ProhibitedUserIdentityCharacter_ThrowsAsync(string prohibitedCharacter)
{
// Arrange
FoundryAgent agent = CreateFoundryAgent();

// Act
ArgumentException exception = await Assert.ThrowsAsync<ArgumentException>(
() => agent.CreateFoundryHostedAgentSessionAsync(userIdentity: $"before{prohibitedCharacter}after"));

// Assert
Assert.Equal("userIdentity", exception.ParamName);
}

[Theory]
[InlineData("\0")]
[InlineData("\r")]
[InlineData("\n")]
[InlineData("before\0after")]
[InlineData("before\rafter")]
[InlineData("before\nafter")]
[InlineData(" \r ")]
[InlineData(" \n ")]
public void UserIdentityPolicy_Process_ProhibitedUserIdentityCharacter_Throws(string invalidIdentity)
{
// Arrange
using var handler = new RecordingHandler(MinimalResponseJson());
#pragma warning disable CA5399
using var http = new HttpClient(handler, disposeHandler: false);
#pragma warning restore CA5399
ClientPipeline pipeline = ClientPipeline.Create(
new ClientPipelineOptions { Transport = new HttpClientPipelineTransport(http) },
perCallPolicies: default,
perTryPolicies: default,
beforeTransportPolicies: default);
PipelineMessage message = CreatePipelineMessage(pipeline);
UserIdentityScope.Current = invalidIdentity;

try
{
// Act
ArgumentException exception = Assert.Throws<ArgumentException>(
() => UserIdentityPolicy.Instance.Process(
message,
[UserIdentityPolicy.Instance, TerminalPolicy.Instance],
0));

// Assert
Assert.Equal("userIdentity", exception.ParamName);
}
finally
{
UserIdentityScope.Current = null;
}
}

[Theory]
[InlineData("\0")]
[InlineData("\r")]
[InlineData("\n")]
[InlineData("before\0after")]
[InlineData("before\rafter")]
[InlineData("before\nafter")]
[InlineData(" \r ")]
[InlineData(" \n ")]
public async Task UserIdentityPolicy_ProcessAsync_ProhibitedUserIdentityCharacter_ThrowsAsync(string invalidIdentity)
{
// Arrange
using var handler = new RecordingHandler(MinimalResponseJson());
#pragma warning disable CA5399
using var http = new HttpClient(handler, disposeHandler: false);
#pragma warning restore CA5399
ClientPipeline pipeline = ClientPipeline.Create(
new ClientPipelineOptions { Transport = new HttpClientPipelineTransport(http) },
perCallPolicies: default,
perTryPolicies: default,
beforeTransportPolicies: default);
PipelineMessage message = CreatePipelineMessage(pipeline);
UserIdentityScope.Current = invalidIdentity;

try
{
// Act
ArgumentException exception = await Assert.ThrowsAsync<ArgumentException>(
async () => await UserIdentityPolicy.Instance.ProcessAsync(
message,
[UserIdentityPolicy.Instance, TerminalPolicy.Instance],
0));

// Assert
Assert.Equal("userIdentity", exception.ParamName);
}
finally
{
UserIdentityScope.Current = null;
}
}

[Theory]
[InlineData("\0")]
[InlineData("\r")]
[InlineData("\n")]
[InlineData("before\0after")]
[InlineData("before\rafter")]
[InlineData("before\nafter")]
[InlineData(" \r ")]
[InlineData(" \n ")]
public async Task RunAsync_RestoredProhibitedUserIdentity_DoesNotReachTransportAsync(string invalidIdentity)
{
// Arrange
using var handler = new RecordingHandler(MinimalResponseJson());
(FoundryHostedRequestAgent agent, AgentSession session, HttpClient http) = await CreateEndToEndAgentAsync(handler);
using (http)
{
SetRawUserIdentity(session, invalidIdentity);

// Act
ArgumentException exception = await Assert.ThrowsAsync<ArgumentException>(
() => agent.RunAsync("hi", session));

// Assert
Assert.Equal("userIdentity", exception.ParamName);
Assert.Empty(handler.Requests);
}
}

[Theory]
[InlineData("\0")]
[InlineData("\r")]
[InlineData("\n")]
[InlineData("before\0after")]
[InlineData("before\rafter")]
[InlineData("before\nafter")]
[InlineData(" \r ")]
[InlineData(" \n ")]
public async Task RunStreamingAsync_RestoredProhibitedUserIdentity_DoesNotReachTransportAsync(string invalidIdentity)
{
// Arrange
using var handler = new RecordingHandler(MinimalResponseJson());
(FoundryHostedRequestAgent agent, AgentSession session, HttpClient http) = await CreateEndToEndAgentAsync(handler);
using (http)
{
SetRawUserIdentity(session, invalidIdentity);

// Act
ArgumentException exception = await Assert.ThrowsAsync<ArgumentException>(
async () =>
{
await foreach (AgentResponseUpdate _ in agent.RunStreamingAsync("hi", session))
{
// Drain the stream so transport validation executes.
}
});

// Assert
Assert.Equal("userIdentity", exception.ParamName);
Assert.Empty(handler.Requests);
}
}

[Fact]
public async Task Conflict_SessionAndOptionsHostedIdsDiffer_ThrowsAsync()
{
Expand Down Expand Up @@ -382,6 +545,42 @@ private static FoundryAgent CreateFoundryAgent() =>
model: "gpt-4o-mini",
instructions: "Test");

private static PipelineMessage CreatePipelineMessage(ClientPipeline pipeline)
{
PipelineMessage message = pipeline.CreateMessage();
message.Request.Method = "POST";
message.Request.Uri = new Uri("https://example.test/");
return message;
}

private static async Task<(FoundryHostedRequestAgent Agent, AgentSession Session, HttpClient HttpClient)> CreateEndToEndAgentAsync(
RecordingHandler handler)
{
#pragma warning disable CA5399
var http = new HttpClient(handler, disposeHandler: false);
#pragma warning restore CA5399
var openAIClient = new OpenAIClient(
new ApiKeyCredential("fake"),
new OpenAIClientOptions { Transport = new HttpClientPipelineTransport(http) });
IChatClient chatClient = openAIClient.GetResponsesClient().AsIChatClient();

#pragma warning disable MEAI001
OpenAIRequestPolicies policies = chatClient.GetService<OpenAIRequestPolicies>()!;
OpenAIRequestPoliciesReflection.AddPolicyIfMissing(policies, UserIdentityPolicy.Instance);
#pragma warning restore MEAI001

var chatAgent = new ChatClientAgent(chatClient);
AgentSession session = await chatAgent.CreateSessionAsync();
return (new FoundryHostedRequestAgent(chatAgent), session, http);
}

private static void SetRawUserIdentity(AgentSession session, string userIdentity)
{
// Session deserialization writes state directly, so the final transport boundary must
// reject invalid restored values even when the public session factory was not used.
session.StateBag.SetValue("Microsoft.Agents.AI.Foundry.UserIdentity", userIdentity);
}

private static string MinimalResponseJson() => """
{
"id":"resp_1","object":"response","created_at":1700000000,"status":"completed",
Expand All @@ -391,6 +590,18 @@ private static string MinimalResponseJson() => """

private sealed class TestSession : AgentSession;

private sealed class TerminalPolicy : PipelinePolicy
{
public static TerminalPolicy Instance { get; } = new();

public override void Process(PipelineMessage message, IReadOnlyList<PipelinePolicy> pipeline, int currentIndex)
{
}

public override ValueTask ProcessAsync(PipelineMessage message, IReadOnlyList<PipelinePolicy> pipeline, int currentIndex) =>
default;
}

private sealed class ProbeAgent : AIAgent
{
private readonly Action<AgentRunOptions?>? _onRun;
Expand Down
Loading