Skip to content
Open
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
77 changes: 74 additions & 3 deletions src/A2A.AspNetCore/A2AHttpProcessor.cs
Original file line number Diff line number Diff line change
Expand Up @@ -133,6 +133,7 @@ internal static Task<IResult> SendMessageRestAsync(
IA2ARequestHandler requestHandler, ILogger logger, SendMessageRequest request, CancellationToken cancellationToken)
=> WithExceptionHandlingAsync(logger, "REST.SendMessage", async ct =>
{
ValidateSendMessageRequest(request);
var result = await requestHandler.SendMessageAsync(request, ct).ConfigureAwait(false);
return new A2AResponseResult(result);
}, cancellationToken: cancellationToken);
Expand All @@ -142,10 +143,28 @@ internal static IResult SendMessageStreamRest(
IA2ARequestHandler requestHandler, ILogger logger, SendMessageRequest request, CancellationToken cancellationToken)
=> WithExceptionHandling(logger, "REST.SendMessageStream", () =>
{
ValidateSendMessageRequest(request);
var events = requestHandler.SendStreamingMessageAsync(request, cancellationToken);
return new A2AEventStreamResult(events);
});

/// <summary>
/// Validates a REST-bound <see cref="SendMessageRequest"/>, mirroring the JSON-RPC
/// binding validation in <see cref="A2AJsonRpcProcessor"/>. The REST endpoints
/// use <c>[FromBody]</c> model binding, which does not enforce a non-empty message parts
/// list on its own.
/// </summary>
/// <param name="request">The send message request to validate.</param>
/// <exception cref="A2AException">Thrown with <see cref="A2AErrorCode.InvalidParams"/> when the message parts list is empty.</exception>
private static void ValidateSendMessageRequest(SendMessageRequest request)
{
ArgumentNullException.ThrowIfNull(request);
if (request.Message.Parts.Count == 0)
{
throw new A2AException("Message parts cannot be empty", A2AErrorCode.InvalidParams);
}
}

// REST handler: Subscribe to task
internal static IResult SubscribeToTaskRest(
IA2ARequestHandler requestHandler, ILogger logger, string id, CancellationToken cancellationToken)
Expand Down Expand Up @@ -300,13 +319,17 @@ await httpContext.Response.BodyWriter.WriteAsync(
{
// Client disconnected — expected
}
catch (Exception)
catch (Exception ex)
{
// Stream error — response already started, best-effort error event
// Stream error — response already started, best-effort error event.
// Use the same structured error shape (code/message/data) as the JSON-RPC
// SSE stream so clients get consistent errors across transports.
// A2AException error codes are preserved; unexpected errors fall back to -32603.
try
{
var errorJson = BuildErrorJson(ex);
await httpContext.Response.BodyWriter.WriteAsync(
Encoding.UTF8.GetBytes("data: {\"error\":\"An internal error occurred during streaming.\"}\n\n"), httpContext.RequestAborted);
Encoding.UTF8.GetBytes($"data: {errorJson}\n\n"), httpContext.RequestAborted);
await httpContext.Response.BodyWriter.FlushAsync(httpContext.RequestAborted);
}
catch
Expand All @@ -315,4 +338,52 @@ await httpContext.Response.BodyWriter.WriteAsync(
}
}
}

/// <summary>
/// Builds a structured SSE error event payload: <c>{"error":{"code":..,"message":..,"data":..}}</c>.
/// Mirrors the JSON-RPC error object (<see cref="JsonRpcError"/>) so REST and JSON-RPC
/// transports return consistent errors.
/// </summary>
/// <param name="exception">The exception to render as an SSE error event.</param>
/// <returns>The JSON payload for the SSE error event.</returns>
private static string BuildErrorJson(Exception exception)
{
var errorCode = exception is A2AException a2aException
? a2aException.ErrorCode
: A2AErrorCode.InternalError;
var message = exception is A2AException a2aEx
? a2aEx.Message
: "An internal error occurred during streaming.";

using var buffer = new MemoryStream();
using (var writer = new Utf8JsonWriter(buffer))
{
writer.WriteStartObject();
writer.WritePropertyName("error");
writer.WriteStartObject();
writer.WriteNumber("code", (int)errorCode);
writer.WriteString("message", message);

if (A2AErrorCodeMapping.IsA2ASpecificError(errorCode))
{
var reason = A2AErrorCodeMapping.GetReasonString(errorCode);
if (reason is not null)
{
writer.WritePropertyName("data");
writer.WriteStartArray();
writer.WriteStartObject();
writer.WriteString("@type", "type.googleapis.com/google.rpc.ErrorInfo");
writer.WriteString("reason", reason);
writer.WriteString("domain", "a2a-protocol.org");
writer.WriteEndObject();
writer.WriteEndArray();
}
}

writer.WriteEndObject();
writer.WriteEndObject();
}

return Encoding.UTF8.GetString(buffer.ToArray());
}
}
152 changes: 152 additions & 0 deletions tests/A2A.AspNetCore.UnitTests/A2AEventStreamResultTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,152 @@
using Microsoft.AspNetCore.Http;
using System.Runtime.CompilerServices;
using System.Text;
using System.Text.Json;

namespace A2A.AspNetCore.Tests;

public class A2AEventStreamResultTests
{
[Fact]
public async Task ExecuteAsync_A2AException_PreservesErrorCodeAndMessage()
{
// Arrange — A2A-specific error must keep code, message, and structured data
var events = ThrowingAsyncEnumerable(new A2AException("Task not found", A2AErrorCode.TaskNotFound));
var result = new A2AEventStreamResult(events);
var httpContext = CreateHttpContext();

// Act
await result.ExecuteAsync(httpContext);

// Assert
var body = GetResponseBody(httpContext);
using var doc = JsonDocument.Parse(ExtractErrorDataLine(body));
var error = doc.RootElement.GetProperty("error");

Assert.Equal((int)A2AErrorCode.TaskNotFound, error.GetProperty("code").GetInt32());
Assert.Equal("Task not found", error.GetProperty("message").GetString());

// A2A-specific codes carry google.rpc ErrorInfo data (same as JSON-RPC transport)
var data = error.GetProperty("data");
Assert.Equal("TASK_NOT_FOUND", data[0].GetProperty("reason").GetString());
Assert.Equal("a2a-protocol.org", data[0].GetProperty("domain").GetString());
}

[Fact]
public async Task ExecuteAsync_GenericException_ReturnsInternalError_WithoutLeakingMessage()
{
// Arrange
var events = ThrowingAsyncEnumerable(new InvalidOperationException("sensitive internal details"));
var result = new A2AEventStreamResult(events);
var httpContext = CreateHttpContext();

// Act
await result.ExecuteAsync(httpContext);

// Assert — falls back to -32603 with a generic message, never leaks internals
var body = GetResponseBody(httpContext);
using var doc = JsonDocument.Parse(ExtractErrorDataLine(body));
var error = doc.RootElement.GetProperty("error");

Assert.Equal((int)A2AErrorCode.InternalError, error.GetProperty("code").GetInt32());
Assert.Equal("An internal error occurred during streaming.", error.GetProperty("message").GetString());
Assert.DoesNotContain("sensitive internal details", body);
}

[Fact]
public async Task ExecuteAsync_OperationCanceledException_WritesNoErrorEvent()
{
// Arrange
var events = ThrowingAsyncEnumerable(new OperationCanceledException());
var result = new A2AEventStreamResult(events);
var httpContext = CreateHttpContext();

// Act
await result.ExecuteAsync(httpContext);

// Assert — body contains no error SSE data line
var body = GetResponseBody(httpContext);
Assert.DoesNotContain("\"error\"", body);
}

[Fact]
public async Task ExecuteAsync_SetsCorrectResponseHeaders()
{
// Arrange
var events = ThrowingAsyncEnumerable(new A2AException("test", A2AErrorCode.InternalError));
var result = new A2AEventStreamResult(events);
var httpContext = CreateHttpContext();

// Act
await result.ExecuteAsync(httpContext);

// Assert
Assert.Equal("text/event-stream", httpContext.Response.ContentType);
Assert.Equal("no-cache,no-store", httpContext.Response.Headers.CacheControl.ToString());
}

[Fact]
public async Task ExecuteAsync_DataEvents_AreValidSseFrames()
{
// Arrange
var result = new A2AEventStreamResult(DataAsyncEnumerable());
var httpContext = CreateHttpContext();

// Act
await result.ExecuteAsync(httpContext);

// Assert — regular stream events are still bare "data: {StreamResponse}" frames
var body = GetResponseBody(httpContext);
Assert.Contains("data: {", body);
Assert.DoesNotContain("\"error\"", body);
}

// --- Helpers ---

private static DefaultHttpContext CreateHttpContext()
{
var context = new DefaultHttpContext();
context.Response.Body = new MemoryStream();
return context;
}

private static string GetResponseBody(DefaultHttpContext context)
{
context.Response.Body.Position = 0;
using var reader = new StreamReader(context.Response.Body, Encoding.UTF8);
return reader.ReadToEnd();
}

private static string ExtractErrorDataLine(string body)
{
var lines = body.Split('\n');
var dataLine = lines.FirstOrDefault(l => l.StartsWith("data: ", StringComparison.Ordinal) && l.Contains("\"error\""))
?? throw new InvalidOperationException($"No SSE data line with error found in response body:\n{body}");
return dataLine["data: ".Length..];
}

private static async IAsyncEnumerable<StreamResponse> ThrowingAsyncEnumerable(
Exception exception, [EnumeratorCancellation] CancellationToken cancellationToken = default)
{
await Task.CompletedTask; // force async state machine
throw exception;
#pragma warning disable CS0162 // Unreachable code — required to satisfy IAsyncEnumerable<T>
yield break;
#pragma warning restore CS0162
}

private static async IAsyncEnumerable<StreamResponse> DataAsyncEnumerable(
[EnumeratorCancellation] CancellationToken cancellationToken = default)
{
await Task.CompletedTask;
yield return new StreamResponse
{
Task = new AgentTask
{
Id = "t1",
ContextId = "c1",
Status = new TaskStatus { State = TaskState.Working },
},
};
}
}
47 changes: 47 additions & 0 deletions tests/A2A.AspNetCore.UnitTests/A2AHttpProcessorTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -183,4 +183,51 @@ public async Task GetTask_WithUnknownA2AErrorCode_ShouldReturn500InternalServerE
Assert.NotNull(result);
Assert.Equal(StatusCodes.Status500InternalServerError, ((IStatusCodeHttpResult)result).StatusCode);
}

[Fact]
public async Task SendMessageRest_WithEmptyParts_ReturnsInvalidParams()
{
// Arrange — empty Message parts are accepted by [FromBody] model binding,
// so the REST processor must validate them like the JSON-RPC binding does.
var (requestHandler, _) = CreateServer();
var logger = NullLogger.Instance;
var sendRequest = new SendMessageRequest
{
Message = new Message { Role = Role.User, Parts = [] },
};

// Act
var result = await A2AHttpProcessor.SendMessageRestAsync(requestHandler, logger, sendRequest, CancellationToken.None);

// Assert
Assert.NotNull(result);
Assert.Equal(StatusCodes.Status400BadRequest, ((IStatusCodeHttpResult)result).StatusCode);

// Execute and verify the error message is present
var httpContext = new DefaultHttpContext();
httpContext.Response.Body = new MemoryStream();
await result.ExecuteAsync(httpContext);
httpContext.Response.Body.Position = 0;
var body = await new StreamReader(httpContext.Response.Body).ReadToEndAsync();
Assert.Contains("Message parts cannot be empty", body);
}

[Fact]
public async Task SendMessageStreamRest_WithEmptyParts_ReturnsInvalidParams()
{
// Arrange
var (requestHandler, _) = CreateServer();
var logger = NullLogger.Instance;
var sendRequest = new SendMessageRequest
{
Message = new Message { Role = Role.User, Parts = [] },
};

// Act — validation happens synchronously before streaming starts
var result = A2AHttpProcessor.SendMessageStreamRest(requestHandler, logger, sendRequest, CancellationToken.None);

// Assert
Assert.NotNull(result);
Assert.Equal(StatusCodes.Status400BadRequest, ((IStatusCodeHttpResult)result).StatusCode);
}
}