diff --git a/clients/dotnet/README.md b/clients/dotnet/README.md
index b5d3839b8..47ab388ba 100644
--- a/clients/dotnet/README.md
+++ b/clients/dotnet/README.md
@@ -50,9 +50,62 @@ var state = new SessionState { /* ... */ };
Reducers.ApplyToSession(state, action); // mutates `state` in place
```
+For TCP accounting, `Reducers.TcpReducer(state, action)` returns a new state
+and throws `InvalidOperationException` on invalid actions.
+
See [`examples/`](examples/) for runnable `ConnectWs` and `ReducersDemo`
console apps.
+## TCP streams
+
+After `InitializeAsync` negotiates `TcpConnections`, the concrete `AhpClient`
+provides an integrated adapter:
+
+```csharp
+await using var tcp = await client.OpenTcpConnectionAsync(sessionUri,
+ new TcpConnectionSubscription
+ {
+ Type = "tcpConnection", Host = "localhost", Port = 3000,
+ Encoding = TcpDataEncoding.Base64,
+ ReceiveWindowBytes = 65536, MaximumChunkSize = 16384,
+ }, cancellationToken);
+await tcp.WriteAsync(requestBytes, cancellationToken);
+await tcp.EndAsync(cancellationToken); // input EOF; output remains readable
+while (await tcp.ReadAsync(cancellationToken) is { } bytes)
+ await ConsumeAsync(bytes);
+```
+
+The SDK owns buffering, flow control, and replay. One reader and one writer may
+run concurrently. `ReadAsync` releases receive credit; `DrainAsync` waits for
+destination consumption. `CloseAsync` stops writes but retains crossing output,
+so keep reading during graceful close. Disposal aborts without draining.
+
+Transport loss suspends the same handles. On a fresh `AhpClient`, call
+`ReconnectTcpConnectionsAsync(reconnectParams, handles, token)` with the original
+`ClientId`; apply its returned replay to ordinary subscriptions. For deliberate
+transport replacement, use `ShutdownAsync(preserveTcpConnections: true)` instead
+of normal shutdown, which terminates streams.
+
+With `MultiHostClient`, use `HostClientHandle.OpenTcpConnectionAsync(session, create)`
+for automatic reconnect. Host removal/shutdown terminates retained streams.
+Missing resources and snapshot fallback fail streams rather than creating new
+sockets. TCP streams do not belong in ordinary state mirrors.
+
+Transport policy, connection limits, and native socket bridges remain
+application-owned. See the [TCP channel contract](../../docs/specification/tcp-channel.md).
+
+## Strict events
+
+For custom loss-sensitive consumers, attach
+`client.CreateEventStream(failOnOverflow: true)` before sending requests and
+retain it for `Events.ReadAllAsync`. Overflow throws `SubscriptionLagException`;
+decode loss throws `AhpTransportException` of kind `"protocol"`. Both terminate
+the receiver rather than skipping events. Capacity uses
+`ClientConfig.SubscriptionBufferCapacity`; ordinary receivers are unchanged.
+These raw receivers are global. The owned TCP adapter instead registers a strict
+child-scoped receiver during creation reply processing and reattaches it per
+child on reconnect. Unrelated traffic cannot exhaust a TCP stream's event buffer.
+
## Dependency injection
Register the services with `AddAgentHostProtocol` (in the
diff --git a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Actions.generated.cs b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Actions.generated.cs
index 751701393..3f0339c2d 100644
--- a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Actions.generated.cs
+++ b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Actions.generated.cs
@@ -231,6 +231,26 @@ public ActionType(string value)
public static readonly ActionType AutomationRunCancelRequested = new ActionType("automationRun/cancelRequested");
+ public static readonly ActionType TcpInput = new ActionType("tcp/input");
+
+ public static readonly ActionType TcpData = new ActionType("tcp/data");
+
+ public static readonly ActionType TcpInputConsumed = new ActionType("tcp/inputConsumed");
+
+ public static readonly ActionType TcpDataConsumed = new ActionType("tcp/dataConsumed");
+
+ public static readonly ActionType TcpInputEof = new ActionType("tcp/inputEof");
+
+ public static readonly ActionType TcpDataEof = new ActionType("tcp/dataEof");
+
+ public static readonly ActionType TcpClientClose = new ActionType("tcp/clientClose");
+
+ public static readonly ActionType TcpHostClose = new ActionType("tcp/hostClose");
+
+ public static readonly ActionType TcpClientReset = new ActionType("tcp/clientReset");
+
+ public static readonly ActionType TcpHostReset = new ActionType("tcp/hostReset");
+
///
public bool Equals(ActionType other) => string.Equals(Value, other.Value, StringComparison.Ordinal);
@@ -2400,6 +2420,92 @@ public sealed record ResourceWatchChangedAction
public JsonElement Changes { get; init; }
}
+/// Client bytes. Never apply optimistically to the authoritative reducer.
+/// Write to the destination only when accepted input.receivedBytes advances.
+public sealed record TcpInputAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpInput;
+
+ /// Absolute decoded-byte offset.
+ public long Offset { get; init; }
+
+ /// Nonempty canonical padded RFC 4648 base64; no whitespace.
+ public required string Data { get; init; }
+}
+
+/// Host bytes. Deliver once, only when output.receivedBytes advances.
+public sealed record TcpDataAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpData;
+
+ /// Absolute decoded-byte offset.
+ public long Offset { get; init; }
+
+ /// Nonempty canonical padded RFC 4648 base64; no whitespace.
+ public required string Data { get; init; }
+}
+
+/// Cumulative input bytes released from the host's bounded write buffer.
+/// Not an acknowledgment that the destination application processed the bytes.
+public sealed record TcpInputConsumedAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpInputConsumed;
+
+ public long ConsumedBytes { get; init; }
+}
+
+/// Cumulative output bytes released by the client's bounded stream consumer.
+public sealed record TcpDataConsumedAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpDataConsumed;
+
+ public long ConsumedBytes { get; init; }
+}
+
+/// Half-close client input after all preceding input bytes have been written.
+public sealed record TcpInputEofAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpInputEof;
+
+ public long FinalOffset { get; init; }
+}
+
+/// Half-close host output after all preceding output bytes have been delivered.
+public sealed record TcpDataEofAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpDataEof;
+
+ public long FinalOffset { get; init; }
+}
+
+/// Client's final close. Respond with hostClose if not already sent.
+public sealed record TcpClientCloseAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpClientClose;
+}
+
+/// Host's final close. Respond with clientClose if not already sent.
+public sealed record TcpHostCloseAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpHostClose;
+}
+
+/// Abort both directions and discard buffered payload.
+public sealed record TcpClientResetAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpClientReset;
+
+ public TcpResetReason Reason { get; init; }
+}
+
+/// Abort both directions and discard buffered payload.
+public sealed record TcpHostResetAction
+{
+ public ActionType Type { get; init; } = ActionType.TcpHostReset;
+
+ public TcpResetReason Reason { get; init; }
+}
+
/// Upsert an {@link Annotation} in the annotations channel — adds a new
/// annotation, or replaces an existing one identified by
/// {@link Annotation.id}.
@@ -2836,6 +2942,16 @@ public StateActionConverter()
["terminal/commandExecuted"] = typeof(TerminalCommandExecutedAction),
["terminal/commandFinished"] = typeof(TerminalCommandFinishedAction),
["resourceWatch/changed"] = typeof(ResourceWatchChangedAction),
+ ["tcp/input"] = typeof(TcpInputAction),
+ ["tcp/data"] = typeof(TcpDataAction),
+ ["tcp/inputConsumed"] = typeof(TcpInputConsumedAction),
+ ["tcp/dataConsumed"] = typeof(TcpDataConsumedAction),
+ ["tcp/inputEof"] = typeof(TcpInputEofAction),
+ ["tcp/dataEof"] = typeof(TcpDataEofAction),
+ ["tcp/clientClose"] = typeof(TcpClientCloseAction),
+ ["tcp/hostClose"] = typeof(TcpHostCloseAction),
+ ["tcp/clientReset"] = typeof(TcpClientResetAction),
+ ["tcp/hostReset"] = typeof(TcpHostResetAction),
["annotations/set"] = typeof(AnnotationsSetAction),
["annotations/removed"] = typeof(AnnotationsRemovedAction),
["annotations/entrySet"] = typeof(AnnotationsEntrySetAction),
diff --git a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Commands.generated.cs b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Commands.generated.cs
index 9d0877928..f1973356e 100644
--- a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Commands.generated.cs
+++ b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Commands.generated.cs
@@ -394,6 +394,10 @@ public sealed record InitializeResult
/// host does not expose an automation catalogue or automation commands.
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
public AutomationCapabilities? Automations { get; init; }
+
+ /// Enables atomic creation of session-scoped, replay-only TCP channels.
+ [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
+ public TcpConnectionsCapability? TcpConnections { get; init; }
}
/// Identifies a protocol implementation — the software (and build) on one end
@@ -571,6 +575,12 @@ public sealed record ReconnectSnapshotResult
/// Fresh snapshots for each subscription
public required List Snapshots { get; init; }
+
+ /// Subscriptions that cannot be restored. Hosts supporting TCP MUST list all
+ /// requested TCP channels here and dispose their sockets on snapshot fallback.
+ /// Omitted by older hosts; absence does not authorize snapshot-restoring TCP.
+ [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
+ public List? Missing { get; init; }
}
/// Subscribe to a URI-identified channel.
@@ -604,6 +614,12 @@ public sealed record SubscribeParams
/// default snapshot. Clients MUST tolerate receiving more state than requested.
[JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
public SubscribeView? View { get; init; }
+
+ /// Atomically create a private child channel and subscribe to it.
+ /// Requires the advertised tcpConnections capability. channel identifies
+ /// the parent session; snapshot.resource identifies the created TCP channel.
+ [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
+ public TcpConnectionSubscription? Create { get; init; }
}
/// Optional client-requested shape for a subscription snapshot.
@@ -643,6 +659,32 @@ public sealed record SubscribeResult
public Snapshot? Snapshot { get; init; }
}
+/// Creates and exclusively subscribes to one TCP connection.
+///
+/// SubscribeParams.channel MUST identify the parent `ahp-session:` channel.
+/// The host returns the new `ahp-tcp:` URI in snapshot.resource, not the parent.
+/// It installs the subscription and sends the response before any TCP actions.
+/// Unknown creation kinds MUST be rejected, never treated as normal subscribe.
+public sealed record TcpConnectionSubscription
+{
+ public required string Type { get; init; }
+
+ /// DNS name or IP literal, not a URL.
+ public required string Host { get; init; }
+
+ /// Destination port.
+ public long Port { get; init; }
+
+ /// Selected from InitializeResult.tcpConnections.encodings.
+ public TcpDataEncoding Encoding { get; init; }
+
+ /// Client receive window in decoded bytes.
+ public long ReceiveWindowBytes { get; init; }
+
+ /// Maximum decoded bytes per output action; MUST NOT exceed receiveWindowBytes.
+ public long MaximumChunkSize { get; init; }
+}
+
// TODO: could not generate SessionForkSource: Error: Interface SessionForkSource not found
/// Creates a new session with the specified agent provider.
diff --git a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Errors.generated.cs b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Errors.generated.cs
index 02a52462b..855ef4182 100644
--- a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Errors.generated.cs
+++ b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/Errors.generated.cs
@@ -45,6 +45,8 @@ public static class AhpErrorCodes
public const int AlreadyExists = -32010;
/// An optimistic-concurrency precondition failed. Returned when a request carries a precondition token that no longer matches the receiver's current state — for example, `resourceWrite` with an `ifMatch` etag that has been superseded by a concurrent write. Callers SHOULD re-read the resource (e.g. via `resourceResolve`) and decide whether to retry the operation with the fresh token or surface the conflict to the user.
public const int Conflict = -32011;
+ /// TCP creation failed; data MUST contain TcpConnectionOpenErrorData.
+ public const int TcpConnectionOpenFailed = -32012;
}
/// Detail payload of an AuthRequired (-32007) error.
diff --git a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/JsonSerializerContext.generated.cs b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/JsonSerializerContext.generated.cs
index 8e496cb27..37774d626 100644
--- a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/JsonSerializerContext.generated.cs
+++ b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/JsonSerializerContext.generated.cs
@@ -221,6 +221,7 @@ namespace Microsoft.AgentHostProtocol;
[JsonSerializable(typeof(FileEditCollection))]
[JsonSerializable(typeof(FileEditDiffStats))]
[JsonSerializable(typeof(FileEditSide))]
+[JsonSerializable(typeof(FlowControlledByteDirectionState))]
[JsonSerializable(typeof(ForkChatSource))]
[JsonSerializable(typeof(HookCustomization))]
[JsonSerializable(typeof(Icon))]
@@ -411,6 +412,26 @@ namespace Microsoft.AgentHostProtocol;
[JsonSerializable(typeof(SubscribeView))]
[JsonSerializable(typeof(SubscriptionDeliveryOptions))]
[JsonSerializable(typeof(SystemNotificationResponsePart))]
+[JsonSerializable(typeof(TcpClientCloseAction))]
+[JsonSerializable(typeof(TcpClientResetAction))]
+[JsonSerializable(typeof(TcpConnectionOpenErrorData))]
+[JsonSerializable(typeof(TcpConnectionOpenFailureReason))]
+[JsonSerializable(typeof(TcpConnectionsCapability))]
+[JsonSerializable(typeof(TcpConnectionState))]
+[JsonSerializable(typeof(TcpConnectionSubscription))]
+[JsonSerializable(typeof(TcpDataAction))]
+[JsonSerializable(typeof(TcpDataConsumedAction))]
+[JsonSerializable(typeof(TcpDataEncoding))]
+[JsonSerializable(typeof(TcpDataEofAction))]
+[JsonSerializable(typeof(TcpEndpoint))]
+[JsonSerializable(typeof(TcpHostCloseAction))]
+[JsonSerializable(typeof(TcpHostResetAction))]
+[JsonSerializable(typeof(TcpInputAction))]
+[JsonSerializable(typeof(TcpInputConsumedAction))]
+[JsonSerializable(typeof(TcpInputEofAction))]
+[JsonSerializable(typeof(TcpResetReason))]
+[JsonSerializable(typeof(TcpResetState))]
+[JsonSerializable(typeof(TcpTarget))]
[JsonSerializable(typeof(TelemetryCapabilities))]
[JsonSerializable(typeof(TerminalClaim))]
[JsonSerializable(typeof(TerminalClaimedAction))]
diff --git a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/State.generated.cs b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/State.generated.cs
index b05b6b10e..cf0d2b95d 100644
--- a/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/State.generated.cs
+++ b/clients/dotnet/src/AgentHostProtocol.Abstractions/Generated/State.generated.cs
@@ -9,6 +9,181 @@ namespace Microsoft.AgentHostProtocol;
// ─── Enums ────────────────────────────────────────────────────────────
+/// Payload encodings advertised by the host.
+[JsonConverter(typeof(TcpDataEncodingConverter))]
+public readonly struct TcpDataEncoding : IEquatable
+{
+ private readonly string? _value;
+
+ /// Wraps a raw wire value — including one this build does not recognize.
+ /// The raw wire string.
+ public TcpDataEncoding(string value)
+ {
+ _value = value;
+ }
+
+ /// The raw wire value.
+ public string Value => _value ?? string.Empty;
+
+ public static readonly TcpDataEncoding Base64 = new TcpDataEncoding("base64");
+
+ ///
+ public bool Equals(TcpDataEncoding other) => string.Equals(Value, other.Value, StringComparison.Ordinal);
+
+ ///
+ public override bool Equals(object? obj) => obj is TcpDataEncoding other && Equals(other);
+
+ ///
+ public override int GetHashCode() => StringComparer.Ordinal.GetHashCode(Value);
+
+ ///
+ public override string ToString() => Value;
+
+ /// Ordinal equality over the raw wire value.
+ public static bool operator ==(TcpDataEncoding left, TcpDataEncoding right) => left.Equals(right);
+
+ /// Ordinal inequality over the raw wire value.
+ public static bool operator !=(TcpDataEncoding left, TcpDataEncoding right) => !left.Equals(right);
+}
+
+/// Reads and writes as its raw wire string, preserving unrecognized values.
+internal sealed class TcpDataEncodingConverter : JsonConverter
+{
+ ///
+ public override TcpDataEncoding Read(ref Utf8JsonReader reader, Type typeToConvert, JsonSerializerOptions options)
+ => new TcpDataEncoding(reader.GetString() ?? throw new JsonException("TcpDataEncoding expects a JSON string."));
+
+ ///
+ public override void Write(Utf8JsonWriter writer, TcpDataEncoding value, JsonSerializerOptions options)
+ => writer.WriteStringValue(value.Value);
+}
+
+/// Endpoint that closes or resets a connection.
+[JsonConverter(typeof(WireEnumConverter))]
+public enum TcpEndpoint
+{
+ [WireValue("client")]
+ Client,
+ [WireValue("host")]
+ Host,
+}
+
+/// Why a connection was aborted.
+[JsonConverter(typeof(TcpResetReasonConverter))]
+public readonly struct TcpResetReason : IEquatable
+{
+ private readonly string? _value;
+
+ /// Wraps a raw wire value — including one this build does not recognize.
+ /// The raw wire string.
+ public TcpResetReason(string value)
+ {
+ _value = value;
+ }
+
+ /// The raw wire value.
+ public string Value => _value ?? string.Empty;
+
+ public static readonly TcpResetReason ConnectionReset = new TcpResetReason("connectionReset");
+
+ public static readonly TcpResetReason ConnectionAborted = new TcpResetReason("connectionAborted");
+
+ public static readonly TcpResetReason ProtocolError = new TcpResetReason("protocolError");
+
+ public static readonly TcpResetReason ReplayUnavailable = new TcpResetReason("replayUnavailable");
+
+ public static readonly TcpResetReason PolicyRevoked = new TcpResetReason("policyRevoked");
+
+ public static readonly TcpResetReason SessionDisposed = new TcpResetReason("sessionDisposed");
+
+ public static readonly TcpResetReason InternalError = new TcpResetReason("internalError");
+
+ ///
+ public bool Equals(TcpResetReason other) => string.Equals(Value, other.Value, StringComparison.Ordinal);
+
+ ///
+ public override bool Equals(object? obj) => obj is TcpResetReason other && Equals(other);
+
+ ///
+ public override int GetHashCode() => StringComparer.Ordinal.GetHashCode(Value);
+
+ ///
+ public override string ToString() => Value;
+
+ /// Ordinal equality over the raw wire value.
+ public static bool operator ==(TcpResetReason left, TcpResetReason right) => left.Equals(right);
+
+ /// Ordinal inequality over the raw wire value.
+ public static bool operator !=(TcpResetReason left, TcpResetReason right) => !left.Equals(right);
+}
+
+/// Reads and writes as its raw wire string, preserving unrecognized values.
+internal sealed class TcpResetReasonConverter : JsonConverter
+{
+ ///
+ public override TcpResetReason Read(ref Utf8JsonReader reader, Type typeToConvert, JsonSerializerOptions options)
+ => new TcpResetReason(reader.GetString() ?? throw new JsonException("TcpResetReason expects a JSON string."));
+
+ ///
+ public override void Write(Utf8JsonWriter writer, TcpResetReason value, JsonSerializerOptions options)
+ => writer.WriteStringValue(value.Value);
+}
+
+/// Expected connection establishment failures.
+[JsonConverter(typeof(TcpConnectionOpenFailureReasonConverter))]
+public readonly struct TcpConnectionOpenFailureReason : IEquatable
+{
+ private readonly string? _value;
+
+ /// Wraps a raw wire value — including one this build does not recognize.
+ /// The raw wire string.
+ public TcpConnectionOpenFailureReason(string value)
+ {
+ _value = value;
+ }
+
+ /// The raw wire value.
+ public string Value => _value ?? string.Empty;
+
+ public static readonly TcpConnectionOpenFailureReason ConnectionFailed = new TcpConnectionOpenFailureReason("connectionFailed");
+
+ public static readonly TcpConnectionOpenFailureReason NameResolutionFailed = new TcpConnectionOpenFailureReason("nameResolutionFailed");
+
+ public static readonly TcpConnectionOpenFailureReason ResourceShortage = new TcpConnectionOpenFailureReason("resourceShortage");
+
+ public static readonly TcpConnectionOpenFailureReason SessionNotReady = new TcpConnectionOpenFailureReason("sessionNotReady");
+
+ ///
+ public bool Equals(TcpConnectionOpenFailureReason other) => string.Equals(Value, other.Value, StringComparison.Ordinal);
+
+ ///
+ public override bool Equals(object? obj) => obj is TcpConnectionOpenFailureReason other && Equals(other);
+
+ ///
+ public override int GetHashCode() => StringComparer.Ordinal.GetHashCode(Value);
+
+ ///
+ public override string ToString() => Value;
+
+ /// Ordinal equality over the raw wire value.
+ public static bool operator ==(TcpConnectionOpenFailureReason left, TcpConnectionOpenFailureReason right) => left.Equals(right);
+
+ /// Ordinal inequality over the raw wire value.
+ public static bool operator !=(TcpConnectionOpenFailureReason left, TcpConnectionOpenFailureReason right) => !left.Equals(right);
+}
+
+/// Reads and writes as its raw wire string, preserving unrecognized values.
+internal sealed class TcpConnectionOpenFailureReasonConverter : JsonConverter
+{
+ ///
+ public override TcpConnectionOpenFailureReason Read(ref Utf8JsonReader reader, Type typeToConvert, JsonSerializerOptions options)
+ => new TcpConnectionOpenFailureReason(reader.GetString() ?? throw new JsonException("TcpConnectionOpenFailureReason expects a JSON string."));
+
+ ///
+ public override void Write(Utf8JsonWriter writer, TcpConnectionOpenFailureReason value, JsonSerializerOptions options)
+ => writer.WriteStringValue(value.Value);
+}
+
/// Policy configuration state for a model.
[JsonConverter(typeof(WireEnumConverter))]
public enum PolicyState
@@ -6158,6 +6333,99 @@ public sealed record ResourceWatchState
public JsonElement? Includes { get; init; }
}
+/// State of one host-assigned `ahp-tcp:` channel.
+///
+/// Payload is never stored in this state. Only the creating authenticated
+/// logical client may observe or dispatch to the channel. Reconnect requires
+/// the original sockets, local stream state, and complete action replay;
+/// a snapshot cannot restore this channel.
+///
+/// Close flags record the two-sided handshake. Either flag means closing;
+/// both mean closed. A present reset terminates the connection immediately,
+/// independently of the close history.
+public sealed record TcpConnectionState
+{
+ public required string Session { get; init; }
+
+ public required TcpTarget Target { get; init; }
+
+ public TcpDataEncoding Encoding { get; init; }
+
+ /// Client to destination socket.
+ public required FlowControlledByteDirectionState Input { get; init; }
+
+ /// Destination socket to client.
+ public required FlowControlledByteDirectionState Output { get; init; }
+
+ public bool ClientClosed { get; init; }
+
+ public bool HostClosed { get; init; }
+
+ [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
+ public TcpResetState? Reset { get; init; }
+}
+
+public sealed record TcpTarget
+{
+ /// DNS name or IP literal, resolved and connected in the host endpoint's network.
+ public required string Host { get; init; }
+
+ /// Destination port.
+ public long Port { get; init; }
+}
+
+public sealed record TcpResetState
+{
+ public TcpEndpoint Source { get; init; }
+
+ public TcpResetReason Reason { get; init; }
+}
+
+/// Bounded byte credit in one direction of a stream.
+/// All counters are nonnegative safe integers (at most 2^53 - 1).
+/// 0 <= consumedBytes <= receivedBytes and
+/// receivedBytes - consumedBytes <= windowBytes.
+public sealed record FlowControlledByteDirectionState
+{
+ /// Maximum accepted-but-not-consumed decoded bytes.
+ public long WindowBytes { get; init; }
+
+ /// Maximum decoded bytes per chunk; MUST NOT exceed windowBytes.
+ public long MaximumChunkSize { get; init; }
+
+ /// Cumulative accepted bytes.
+ public long ReceivedBytes { get; init; }
+
+ /// Cumulative bytes released by the bounded consumer.
+ public long ConsumedBytes { get; init; }
+
+ /// Present after EOF; equals receivedBytes permanently.
+ [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
+ public long? EofAtBytes { get; init; }
+}
+
+/// Host support for private, session-scoped TCP channels.
+/// Presence on initialize is required before using subscribe.create.
+public sealed record TcpConnectionsCapability
+{
+ /// Supported encodings. The base64 profile MUST be supported.
+ public required List Encodings { get; init; }
+
+ /// Informational limit; runtime policy may impose a lower limit.
+ [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
+ public long? MaximumConnectionsPerClient { get; init; }
+}
+
+/// Required detail for TcpConnectionOpenFailed (-32012).
+/// Policy denial and malformed requests use PermissionDenied and InvalidParams.
+public sealed record TcpConnectionOpenErrorData
+{
+ public TcpConnectionOpenFailureReason Reason { get; init; }
+
+ [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)]
+ public bool? Retryable { get; init; }
+}
+
/// A single change observed by a resource watcher.
public sealed record ResourceChange
{
@@ -7726,14 +7994,18 @@ public override void Write(Utf8JsonWriter writer, ToolInput value, JsonSerialize
///
/// SnapshotState is the state payload of a snapshot — root, session,
- /// chat, terminal, changeset, resource-watch, annotations, automation catalogue,
- /// or automation-run state. Read
+/// chat, terminal, changeset, resource-watch, annotations, automation catalogue,
+/// automation-run, or TCP state. Read
/// probes for distinctive fields in an order where no probe shadows another
-/// (chat → session → terminal → changeset → resource-watch → annotations → root).
+/// (tcp → automationRun → automations → session → chat → terminal → changeset →
+/// resource-watch → annotations → root).
///
[JsonConverter(typeof(SnapshotStateConverter))]
public sealed class SnapshotState
{
+ /// Private TCP channel state variant, when populated.
+ public TcpConnectionState? Tcp { get; set; }
+
/// Root state variant, when populated.
public RootState? Root { get; set; }
@@ -7770,7 +8042,13 @@ public override SnapshotState Read(ref Utf8JsonReader reader, Type typeToConvert
using var doc = JsonDocument.ParseValue(ref reader);
var root = doc.RootElement;
var result = new SnapshotState();
- if (root.TryGetProperty("automation", out _) &&
+ if (root.TryGetProperty("input", out _) &&
+ root.TryGetProperty("output", out _) &&
+ root.TryGetProperty("target", out _))
+ {
+ result.Tcp = root.Deserialize(AhpJsonTypeInfo.Get(options));
+ }
+ else if (root.TryGetProperty("automation", out _) &&
root.TryGetProperty("origin", out _) &&
root.TryGetProperty("sessions", out _))
{
@@ -7816,6 +8094,7 @@ public override SnapshotState Read(ref Utf8JsonReader reader, Type typeToConvert
public override void Write(Utf8JsonWriter writer, SnapshotState value, JsonSerializerOptions options)
{
+ if (value.Tcp is not null) { JsonSerializer.Serialize(writer, value.Tcp, AhpJsonTypeInfo.Get(options)); return; }
if (value.AutomationRun is not null) { JsonSerializer.Serialize(writer, value.AutomationRun, AhpJsonTypeInfo.Get(options)); return; }
if (value.Automations is not null) { JsonSerializer.Serialize(writer, value.Automations, AhpJsonTypeInfo.Get(options)); return; }
if (value.Chat is not null) { JsonSerializer.Serialize(writer, value.Chat, AhpJsonTypeInfo.Get(options)); return; }
diff --git a/clients/dotnet/src/AgentHostProtocol/AhpClient.cs b/clients/dotnet/src/AgentHostProtocol/AhpClient.cs
index e90eb47d0..f3a2645ac 100644
--- a/clients/dotnet/src/AgentHostProtocol/AhpClient.cs
+++ b/clients/dotnet/src/AgentHostProtocol/AhpClient.cs
@@ -156,7 +156,7 @@ public enum ConnectionState
/// All public methods are safe to call from multiple threads.
///
///
-public sealed class AhpClient : IAhpClient
+public sealed partial class AhpClient : IAhpClient
{
// ── State that lives for the client lifetime ──────────────────────────
@@ -169,11 +169,13 @@ public sealed class AhpClient : IAhpClient
// In-flight request correlation keyed by JSON-RPC id.
private readonly ConcurrentDictionary> _pending = new();
+ private readonly ConcurrentDictionary> _resultHandlers = new();
// Per-URI subscription fan-out.
private readonly object _subsLock = new();
private readonly Dictionary> _subscriptions = new();
private readonly List _eventListeners = new();
+ private readonly Dictionary> _resourceEventListeners = new(StringComparer.Ordinal);
// Multicast connection-state fan-out. Guarded by `_subsLock` (same lock as the
// event listeners — every fan-out path already takes it).
@@ -223,7 +225,18 @@ internal int SubscriptionCount
}
}
- internal int EventListenerCount { get { lock (_subsLock) { return _eventListeners.Count; } } }
+ internal int EventListenerCount
+ {
+ get
+ {
+ lock (_subsLock)
+ {
+ int count = _eventListeners.Count;
+ foreach (var list in _resourceEventListeners.Values) count += list.Count;
+ return count;
+ }
+ }
+ }
internal int StateListenerCount { get { lock (_subsLock) { return _stateListeners.Count; } } }
@@ -311,7 +324,7 @@ public static AhpClient Connect(
///
/// A that completes once the client begins teardown (either
- /// via or a transport failure).
+ /// via or a transport failure).
///
public Task Completion => _doneTcs.Task;
@@ -327,11 +340,22 @@ public static AhpClient Connect(
/// closed. The underlying transport is closed too.
/// Safe to call multiple times.
///
- public async Task ShutdownAsync(CancellationToken cancellationToken = default)
+ public Task ShutdownAsync(CancellationToken cancellationToken = default)
+ => ShutdownAsync(preserveTcpConnections: false, cancellationToken);
+
+ /// Shuts down transport, optionally retaining TCP handles for explicit reconnect.
+ public async Task ShutdownAsync(bool preserveTcpConnections, CancellationToken cancellationToken = default)
{
- await ShutdownWithErrorAsync(null).ConfigureAwait(false);
- // Wait for both background tasks to exit.
- await Task.WhenAll(_readerTask, _writerTask).WaitAsync(cancellationToken).ConfigureAwait(false);
+ try
+ {
+ if (!preserveTcpConnections) await DisposeTcpConnectionsAsync().ConfigureAwait(false);
+ }
+ finally
+ {
+ await ShutdownWithErrorAsync(null).ConfigureAwait(false);
+ // Wait for both background tasks to exit.
+ await Task.WhenAll(_readerTask, _writerTask).WaitAsync(cancellationToken).ConfigureAwait(false);
+ }
}
///
@@ -408,6 +432,8 @@ private async Task ShutdownWithErrorAsync(Exception? cause, bool fromKeepAlive =
tcs.TrySetException(shutdownEx);
}
}
+ foreach (var kv in _tcpCreationResponses)
+ if (_tcpCreationResponses.TryRemove(kv.Key, out var response)) response.TrySetException(shutdownEx);
// Close every subscription and listener.
List allSubs;
@@ -419,6 +445,8 @@ private async Task ShutdownWithErrorAsync(Exception? cause, bool fromKeepAlive =
allSubs.AddRange(list);
allListeners = new List(_eventListeners);
+ foreach (var list in _resourceEventListeners.Values) allListeners.AddRange(list);
+ _resourceEventListeners.Clear();
_eventListeners.Clear();
}
// Each Close() runs the subscription's detach hook, which removes it from the
@@ -615,9 +643,10 @@ private async Task RunReaderAsync()
{
msg = _serializer.DecodeMessage(frame);
}
- catch
+ catch (Exception ex)
{
- // Skip malformed frames; protocol resync is the server's responsibility.
+ FailStrictEvents("ahp: malformed inbound JSON-RPC frame", ex);
+ // Ordinary receivers retain their existing skip behavior.
AhpTelemetry.MalformedFrames.Add(1);
continue;
}
@@ -656,16 +685,25 @@ private void Dispatch(JsonRpcMessage msg)
// requests" limitation; mirrors the TS client's handleServerRequest.)
_ = HandleServerRequestAsync(msg.Request);
}
+ else
+ {
+ FailStrictEvents("ahp: inbound frame has no JSON-RPC message");
+ }
}
private void Deliver(ulong id, JsonElement result, AhpRpcException? rpcError)
{
- if (_pending.TryRemove(id, out var tcs))
+ if (_pending.TryRemove(id, out var tcs) || _tcpCreationResponses.TryRemove(id, out tcs))
{
+ _resultHandlers.TryRemove(id, out var onResult);
if (rpcError is not null)
tcs.TrySetException(rpcError);
else
+ {
+ try { onResult?.Invoke(result); }
+ catch (Exception ex) { tcs.TrySetException(ex); return; }
tcs.TrySetResult(result);
+ }
}
}
@@ -740,7 +778,13 @@ private async Task EnqueueReplyAsync(JsonRpcMessage msg)
private void HandleNotification(JsonRpcNotification n)
{
- if (n.Params is null) return;
+ if (n.Params is null)
+ {
+ if (n.Method is "action" or "root/sessionAdded" or "root/sessionRemoved"
+ or "root/sessionSummaryChanged" or "root/progress" or "auth/required")
+ FailStrictEvents("ahp: inbound subscription notification is missing params");
+ return;
+ }
var paramsEl = n.Params.Value;
switch (n.Method)
@@ -749,7 +793,7 @@ private void HandleNotification(JsonRpcNotification n)
{
ActionEnvelope env;
try { env = _serializer.Deserialize(paramsEl); }
- catch { return; }
+ catch (Exception ex) { FailStrictEvents("ahp: failed to decode inbound action envelope", ex); return; }
FanOut(env.Channel, new SubscriptionEventAction(env));
break;
}
@@ -757,7 +801,7 @@ private void HandleNotification(JsonRpcNotification n)
{
SessionAddedParams p;
try { p = _serializer.Deserialize(paramsEl); }
- catch { return; }
+ catch (Exception ex) { FailStrictEvents("ahp: failed to decode inbound session-added notification", ex); return; }
FanOut(p.Channel, new SubscriptionEventSessionAdded(p));
break;
}
@@ -765,7 +809,7 @@ private void HandleNotification(JsonRpcNotification n)
{
SessionRemovedParams p;
try { p = _serializer.Deserialize(paramsEl); }
- catch { return; }
+ catch (Exception ex) { FailStrictEvents("ahp: failed to decode inbound session-removed notification", ex); return; }
FanOut(p.Channel, new SubscriptionEventSessionRemoved(p));
break;
}
@@ -773,7 +817,7 @@ private void HandleNotification(JsonRpcNotification n)
{
SessionSummaryChangedParams p;
try { p = _serializer.Deserialize(paramsEl); }
- catch { return; }
+ catch (Exception ex) { FailStrictEvents("ahp: failed to decode inbound session-summary notification", ex); return; }
FanOut(p.Channel, new SubscriptionEventSessionSummaryChanged(p));
break;
}
@@ -781,7 +825,7 @@ private void HandleNotification(JsonRpcNotification n)
{
ProgressParams p;
try { p = _serializer.Deserialize(paramsEl); }
- catch { return; }
+ catch (Exception ex) { FailStrictEvents("ahp: failed to decode inbound progress notification", ex); return; }
FanOut(p.Channel, new SubscriptionEventProgress(p));
break;
}
@@ -789,13 +833,25 @@ private void HandleNotification(JsonRpcNotification n)
{
AuthRequiredParams p;
try { p = _serializer.Deserialize(paramsEl); }
- catch { return; }
+ catch (Exception ex) { FailStrictEvents("ahp: failed to decode inbound auth notification", ex); return; }
FanOut(p.Channel, new SubscriptionEventAuthRequired(p));
break;
}
}
}
+ private void FailStrictEvents(string message, Exception? inner = null)
+ {
+ List listeners;
+ lock (_subsLock)
+ {
+ listeners = new List(_eventListeners);
+ foreach (var list in _resourceEventListeners.Values) listeners.AddRange(list);
+ }
+ var error = new AhpTransportException("protocol", message, inner);
+ foreach (var listener in listeners) listener.FailIfStrict(error);
+ }
+
private void FanOut(string channel, SubscriptionEvent ev)
{
// Snapshot each bucket under the lock (so TrySend runs outside it), but
@@ -812,6 +868,12 @@ private void FanOut(string channel, SubscriptionEvent ev)
listeners = _eventListeners.Count > 0
? new List(_eventListeners)
: Array.Empty();
+ if (_resourceEventListeners.TryGetValue(channel, out var resourceListeners))
+ {
+ var matching = new List(listeners);
+ matching.AddRange(resourceListeners);
+ listeners = matching;
+ }
}
for (var i = 0; i < subs.Count; i++) subs[i].TrySend(ev);
@@ -838,10 +900,15 @@ private void FanOut(string channel, SubscriptionEvent ev)
/// would be non-null.
///
///
- public async Task RequestAsync(
+ public Task RequestAsync(
string method,
TParams parameters,
CancellationToken cancellationToken = default)
+ => RequestCoreAsync(method, parameters, cancellationToken);
+
+ private async Task RequestCoreAsync(
+ string method, TParams parameters, CancellationToken cancellationToken,
+ bool ownsTcpCreation = false, Action? onResult = null)
{
Guard.ThrowIfNull(method, nameof(method));
@@ -879,7 +946,11 @@ private void FanOut(string channel, SubscriptionEvent ev)
ulong id = (ulong)(Interlocked.Increment(ref _nextId) - 1);
var tcs = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously);
+ if (ownsTcpCreation) _tcpCreationResponses[id] = tcs;
+ if (onResult is not null) _resultHandlers[id] = onResult;
_pending[id] = tcs;
+ bool abandonedTcpCreation = false;
+ bool requestMayHaveBeenSent = false;
activity?.SetTag(AhpTelemetryNames.AttrRequestId, id);
AhpTelemetry.InflightRequests.Add(1);
@@ -907,6 +978,7 @@ private void FanOut(string channel, SubscriptionEvent ev)
try
{
+ requestMayHaveBeenSent = true;
await SendMessageAsync(req, requestCts.Token).ConfigureAwait(false);
}
catch
@@ -936,6 +1008,15 @@ private void FanOut(string channel, SubscriptionEvent ev)
// Distinguish caller cancellation (the caller's token) from a request
// timeout (the configured default-timeout fired its linked token) from a
// genuine error — the metric and the span status differ for each.
+ if (ownsTcpCreation)
+ {
+ _pending.TryRemove(id, out _);
+ if (requestMayHaveBeenSent)
+ {
+ abandonedTcpCreation = true;
+ _ = ReleaseAbandonedTcpCreationAsync(id, tcs.Task);
+ }
+ }
bool callerCancelled = ex is OperationCanceledException && cancellationToken.IsCancellationRequested;
outcome = ex switch
{
@@ -949,6 +1030,8 @@ private void FanOut(string channel, SubscriptionEvent ev)
}
finally
{
+ _resultHandlers.TryRemove(id, out _);
+ if (ownsTcpCreation && !abandonedTcpCreation) _tcpCreationResponses.TryRemove(id, out _);
AhpTelemetry.InflightRequests.Add(-1);
AhpTelemetry.RequestDuration.Record(
Compatibility.GetElapsedTime(startTimestamp).TotalMilliseconds,
@@ -1045,6 +1128,8 @@ public async Task InitializeAsync(
"protocol",
$"ahp: server selected unoffered protocol version '{result.ProtocolVersion}'");
}
+ _tcpClientId = clientId;
+ _tcpCapability = result.TcpConnections;
return result;
}
@@ -1399,17 +1484,71 @@ T Decode(JsonElement? el) => el is { } e
/// streams may exist concurrently.
///
public EventStream CreateEventStream()
+ => CreateEventStream(failOnOverflow: false);
+
+ ///
+ /// Registers a bounded global receiver before returning. With
+ /// enabled, overflow preserves the buffered
+ /// prefix and permanently faults the receiver with .
+ /// Discarded malformed inbound frames or notification payloads instead fault
+ /// strict receivers with of kind protocol.
+ /// Other receivers are unaffected. The caller owns reset/unsubscribe and reconnect handling.
+ ///
+ public EventStream CreateEventStream(bool failOnOverflow)
{
- var stream = new EventStream(_cfg.SubscriptionBufferCapacity);
+ var stream = new EventStream(_cfg.SubscriptionBufferCapacity, failOnOverflow);
+ stream.OnClose(() => { lock (_subsLock) { _eventListeners.Remove(stream); } });
lock (_subsLock)
{
_eventListeners.Add(stream);
}
- // Detach on dispose so an abandoned stream is removed from the fan-out
- // list rather than receiving (dropped) events for the client's lifetime.
- stream.OnClose(() => { lock (_subsLock) { _eventListeners.Remove(stream); } });
return stream;
}
+
+ private EventStream CreateResourceEventStream(string resource)
+ {
+ var stream = new EventStream(_cfg.SubscriptionBufferCapacity, failOnOverflow: true) { Resource = resource };
+ // Install cleanup before publishing the receiver: it may fail as soon
+ // as the receive loop sees it.
+ stream.OnClose(() => { lock (_subsLock) { RemoveEventStream(stream); } });
+ lock (_subsLock)
+ {
+ if (Volatile.Read(ref _shutdownStarted) == 1) stream.Close();
+ else AddEventStream(stream);
+ }
+ return stream;
+ }
+
+ private void AddEventStream(EventStream stream)
+ {
+ if (stream.Resource is null) _eventListeners.Add(stream);
+ else
+ {
+ if (!_resourceEventListeners.TryGetValue(stream.Resource, out var list))
+ _resourceEventListeners[stream.Resource] = list = new List();
+ list.Add(stream);
+ }
+ }
+
+ private void RemoveEventStream(EventStream stream)
+ {
+ if (stream.Resource is null) _eventListeners.Remove(stream);
+ else if (_resourceEventListeners.TryGetValue(stream.Resource, out var list))
+ {
+ list.Remove(stream);
+ if (list.Count == 0) _resourceEventListeners.Remove(stream.Resource);
+ }
+ }
+
+ private void BindEventStream(EventStream stream, string resource)
+ {
+ lock (_subsLock)
+ {
+ RemoveEventStream(stream);
+ stream.Resource = resource;
+ if (!stream.IsClosed) AddEventStream(stream);
+ }
+ }
}
// ─── Connection-state stream ────────────────────────────────────────────────────
diff --git a/clients/dotnet/src/AgentHostProtocol/Errors.cs b/clients/dotnet/src/AgentHostProtocol/Errors.cs
index 78e89c18f..5a648afed 100644
--- a/clients/dotnet/src/AgentHostProtocol/Errors.cs
+++ b/clients/dotnet/src/AgentHostProtocol/Errors.cs
@@ -20,6 +20,20 @@ protected AhpException(string message) : base(message) { }
protected AhpException(string message, Exception? inner) : base(message, inner) { }
}
+/// A strict event receiver overflowed and permanently terminated.
+public sealed class SubscriptionLagException : AhpException
+{
+ /// The receiver's maximum buffered event count.
+ public int Capacity { get; }
+
+ /// Creates a terminal receiver-lag exception.
+ public SubscriptionLagException(int capacity)
+ : base($"ahp: event receiver exceeded its capacity of {capacity}; the receiver is permanently terminated")
+ {
+ Capacity = capacity;
+ }
+}
+
///
/// Thrown by implementations when the underlying
/// connection experiences a transport-level fault.
diff --git a/clients/dotnet/src/AgentHostProtocol/Generated/ActionMetadata.generated.cs b/clients/dotnet/src/AgentHostProtocol/Generated/ActionMetadata.generated.cs
index e20e88d36..6c4d9ac69 100644
--- a/clients/dotnet/src/AgentHostProtocol/Generated/ActionMetadata.generated.cs
+++ b/clients/dotnet/src/AgentHostProtocol/Generated/ActionMetadata.generated.cs
@@ -346,6 +346,36 @@ public static bool TryGetActionType(object action, out ActionType actionType)
case SessionWorkingDirectorySetAction value:
actionType = value.Type;
return true;
+ case TcpClientCloseAction value:
+ actionType = value.Type;
+ return true;
+ case TcpClientResetAction value:
+ actionType = value.Type;
+ return true;
+ case TcpDataAction value:
+ actionType = value.Type;
+ return true;
+ case TcpDataConsumedAction value:
+ actionType = value.Type;
+ return true;
+ case TcpDataEofAction value:
+ actionType = value.Type;
+ return true;
+ case TcpHostCloseAction value:
+ actionType = value.Type;
+ return true;
+ case TcpHostResetAction value:
+ actionType = value.Type;
+ return true;
+ case TcpInputAction value:
+ actionType = value.Type;
+ return true;
+ case TcpInputConsumedAction value:
+ actionType = value.Type;
+ return true;
+ case TcpInputEofAction value:
+ actionType = value.Type;
+ return true;
case TerminalClaimedAction value:
actionType = value.Type;
return true;
diff --git a/clients/dotnet/src/AgentHostProtocol/Hosts/HostClientHandle.cs b/clients/dotnet/src/AgentHostProtocol/Hosts/HostClientHandle.cs
index 733e50902..797b93ec4 100644
--- a/clients/dotnet/src/AgentHostProtocol/Hosts/HostClientHandle.cs
+++ b/clients/dotnet/src/AgentHostProtocol/Hosts/HostClientHandle.cs
@@ -16,6 +16,8 @@ namespace Microsoft.AgentHostProtocol.Hosts;
/// Port of Swift's HostClientHandle (Swift surfaces the reconnect case as
/// hostReconnected; the .NET typed-error set folds that into "not the
/// connection you held — reacquire").
+/// Owned TCP streams opened through this handle survive replay reconnects even
+/// though the handle itself becomes stale. Removal/shutdown terminates them.
///
public sealed class HostClientHandle
{
@@ -57,6 +59,26 @@ private AhpClient CheckAlive()
///
public void CheckAliveOrThrow() => CheckAlive();
+ /// Opens a TCP stream retained and replayed by this host's supervisor.
+ public async Task OpenTcpConnectionAsync(
+ string session, TcpConnectionSubscription create, CancellationToken cancellationToken = default)
+ {
+ CheckAlive();
+ var entry = _owner.TryGetEntry(HostId) ?? throw new HostShutDownException(HostId);
+ using var linked = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken, entry.LifetimeCts.Token);
+ await entry.ConnectionGate.WaitAsync(linked.Token).ConfigureAwait(false);
+ try
+ {
+ var client = CheckAlive();
+ using var creation = CancellationTokenSource.CreateLinkedTokenSource(linked.Token, entry.TcpCreations.Token);
+ var connection = await client.OpenTcpConnectionAsync(session, create, creation.Token).ConfigureAwait(false);
+ entry.TcpConnections[connection.Resource] = connection;
+ connection.OnRelease(() => entry.TcpConnections.TryRemove(connection.Resource, out _));
+ return connection;
+ }
+ finally { entry.ConnectionGate.Release(); }
+ }
+
///
/// Dispatches an action through this connection on ,
/// refusing (throwing) if the host was removed or the connection has been
diff --git a/clients/dotnet/src/AgentHostProtocol/Hosts/MultiHostClient.cs b/clients/dotnet/src/AgentHostProtocol/Hosts/MultiHostClient.cs
index 9e7ce14bf..8eb2b4582 100644
--- a/clients/dotnet/src/AgentHostProtocol/Hosts/MultiHostClient.cs
+++ b/clients/dotnet/src/AgentHostProtocol/Hosts/MultiHostClient.cs
@@ -5,6 +5,7 @@
using System;
using System.Collections.Concurrent;
using System.Collections.Generic;
+using System.Linq;
using System.Security.Cryptography;
using System.Text.Json;
using System.Threading;
@@ -156,6 +157,7 @@ public void Dispose()
{
if (Interlocked.Exchange(ref _disposed, 1) != 0) return;
LifetimeCts.Dispose();
+ TcpCreations.Dispose();
ConnectionGate.Dispose();
_manualReconnect.Dispose();
// EndAttempt normally disposes the per-attempt CTS, but a teardown that
@@ -186,6 +188,19 @@ public HostEntry(HostId id, HostConfig config, string clientId)
/// to read one published reference.
///
public AhpClient? CurrentClient => _client;
+ internal AhpClient? PreviousClient { get; set; }
+ internal ConcurrentDictionary TcpConnections { get; } = new();
+ internal CancellationTokenSource TcpCreations { get; private set; } = new();
+
+ internal async Task CloseTcpAsync()
+ {
+ try
+ {
+ await Task.WhenAll(TcpConnections.Values.Select(c => c.DisposeAsync().AsTask())).ConfigureAwait(false);
+ return null;
+ }
+ catch (Exception error) { return error; }
+ }
public void SetClient(AhpClient? client, string protoVer)
{
@@ -202,6 +217,8 @@ public void SetClient(AhpClient? client, string protoVer)
}
else
{
+ TcpCreations.Dispose();
+ TcpCreations = new CancellationTokenSource();
_clientReady.TrySetResult(client);
}
}
@@ -658,9 +675,11 @@ public async Task RemoveHostAsync(HostId id, CancellationToken cancellationToken
FinishPerHostListeners(id.ToString());
entry!.LifetimeCts.Cancel();
+ Exception? tcpCleanupError = null;
await entry.ConnectionGate.WaitAsync(CancellationToken.None).ConfigureAwait(false);
try
{
+ tcpCleanupError = await entry.CloseTcpAsync().ConfigureAwait(false);
var client = entry.CurrentClient;
if (client is not null)
{
@@ -680,6 +699,7 @@ public async Task RemoveHostAsync(HostId id, CancellationToken cancellationToken
// after teardown so a consumer that reacts to the removed event observes
// a host that is already gone (Host(id) == null).
BroadcastHostEvent(HostEvent.Removed(id));
+ if (tcpCleanupError is not null) throw tcpCleanupError;
}
// ── Event channels ────────────────────────────────────────────────────
@@ -724,6 +744,7 @@ public async Task ShutdownAsync(CancellationToken cancellationToken = default)
_rootCts.Cancel();
var entries = new List(_hosts.Values);
+ var tcpCleanupErrors = new List();
_hosts.Clear();
// Finish per-host listener streams for every host so their consumers'
@@ -741,6 +762,7 @@ public async Task ShutdownAsync(CancellationToken cancellationToken = default)
await entry.ConnectionGate.WaitAsync(CancellationToken.None).ConfigureAwait(false);
try
{
+ if (await entry.CloseTcpAsync().ConfigureAwait(false) is { } error) tcpCleanupErrors.Add(error);
var client = entry.CurrentClient;
if (client is not null)
{
@@ -772,6 +794,7 @@ public async Task ShutdownAsync(CancellationToken cancellationToken = default)
}
_rootCts.Dispose();
+ if (tcpCleanupErrors.Count > 0) throw new AggregateException(tcpCleanupErrors);
}
///
@@ -840,6 +863,7 @@ private async Task OpenHostCoreAsync(
// Register before the first handshake request so notifications that
// race initialize/reconnect are buffered rather than discarded.
var stream = client.CreateEventStream();
+ if (entry.PreviousClient is { } previous) client.InheritTcpClient(previous);
// On a reconnect with a known serverSeq, issue the AHP `reconnect` command
// (clientId + lastSeenServerSeq) so the host REPLAYS the actions missed
@@ -851,24 +875,38 @@ private async Task OpenHostCoreAsync(
{
var snap = entry.Snapshot();
subscriptions = snap.Subscriptions;
+ var tcpConnections = entry.TcpConnections.Values.ToArray();
ReconnectResult? reconnectResult = null;
try
{
- reconnectResult = await client.ReconnectAsync(
- snap.ClientId, snap.ServerSeq, subscriptions, cancellationToken)
+ reconnectResult = tcpConnections.Length == 0
+ ? await client.ReconnectAsync(snap.ClientId, snap.ServerSeq, subscriptions, cancellationToken).ConfigureAwait(false)
+ : await client.ReconnectTcpConnectionsAsync(
+ new ReconnectParams
+ {
+ Channel = ProtocolVersion.RootResourceUri,
+ ClientId = snap.ClientId,
+ LastSeenServerSeq = snap.ServerSeq,
+ Subscriptions = subscriptions.ToList(),
+ }, tcpConnections, cancellationToken)
.ConfigureAwait(false);
}
- catch (Exception) when (!cancellationToken.IsCancellationRequested)
+ catch (Exception reconnectError) when (!cancellationToken.IsCancellationRequested
+ && (tcpConnections.Length == 0 || reconnectError is AhpRpcException))
{
// Host does not support `reconnect` (or it errored) — fall through
// to a fresh `initialize` on the still-live client below. A
// cancellation (shutdown/dispose) is NOT swallowed: it propagates
// so the supervisor tears down promptly instead of blocking on a
// fallback initialize.
+ if (await entry.CloseTcpAsync().ConfigureAwait(false) is { } error) throw error;
}
if (reconnectResult?.Value is ReconnectReplayResult replay)
{
+ // TCP may need an older checkpoint than the ordinary host mirror.
+ if (tcpConnections.Length > 0)
+ replay = replay with { Actions = replay.Actions.Where(action => action.ServerSeq > snap.ServerSeq).ToList() };
var summaries = await FetchSessionSummariesAsync(entry, client, cancellationToken).ConfigureAwait(false);
cancellationToken.ThrowIfCancellationRequested();
await entry.ConnectionGate.WaitAsync(cancellationToken).ConfigureAwait(false);
@@ -914,6 +952,7 @@ private async Task OpenHostCoreAsync(
}
}
+ subscriptions = subscriptions?.Where(uri => !uri.StartsWith("ahp-tcp:", StringComparison.Ordinal)).ToArray();
var result = await client.InitializeAsync(
entry.ClientId,
entry.Config.ProtocolVersions,
@@ -951,7 +990,7 @@ private async Task OpenHostCoreAsync(
{
if (client is not null)
{
- try { await client.ShutdownAsync(CancellationToken.None).ConfigureAwait(false); } catch { }
+ try { await client.ShutdownAsync(preserveTcpConnections: true, CancellationToken.None).ConfigureAwait(false); } catch { }
}
else if (transport is not null)
{
@@ -983,6 +1022,7 @@ private void InstallOpenHost(
{
applyHandshake();
CompleteOpenHost(entry, stream);
+ entry.PreviousClient = null;
}
catch
{
@@ -1237,6 +1277,7 @@ private async Task SuperviseAsync(HostEntry entry)
// or we're forcing a manual reconnect). Serialize replacement with
// subscribe/unsubscribe, then drain the old event pump so reconnect
// snapshots the final sequence observed on that connection.
+ entry.TcpCreations.Cancel();
try
{
await entry.ConnectionGate.WaitAsync(ct).ConfigureAwait(false);
@@ -1248,12 +1289,13 @@ private async Task SuperviseAsync(HostEntry entry)
try
{
var oldPump = entry.PumpTask;
+ entry.PreviousClient = client;
BeginReconnect(entry, new HostState
{
Kind = HostStateKind.Reconnecting,
Attempt = 1,
});
- try { await client.ShutdownAsync(CancellationToken.None).ConfigureAwait(false); } catch { }
+ try { await client.ShutdownAsync(preserveTcpConnections: true, CancellationToken.None).ConfigureAwait(false); } catch { }
try { await oldPump.ConfigureAwait(false); } catch (OperationCanceledException) { } catch { }
}
finally
diff --git a/clients/dotnet/src/AgentHostProtocol/Reducers.cs b/clients/dotnet/src/AgentHostProtocol/Reducers.cs
index af977f375..0ee30b7d0 100644
--- a/clients/dotnet/src/AgentHostProtocol/Reducers.cs
+++ b/clients/dotnet/src/AgentHostProtocol/Reducers.cs
@@ -30,6 +30,118 @@ public enum ReduceOutcome
///
public static class Reducers
{
+ ///
+ /// Returns TCP accounting state without retaining payloads or restoring streams.
+ /// Invalid actions throw before changing state. Adapters must reset/close the
+ /// channel on failure and only write data when ReceivedBytes advances.
+ ///
+ /// The TCP action violates the stream contract.
+ public static TcpConnectionState TcpReducer(TcpConnectionState state, StateAction action)
+ {
+ Guard.ThrowIfNull(state, nameof(state));
+ Guard.ThrowIfNull(action, nameof(action));
+ if (state.Reset is not null) return state;
+ var input = state.Input;
+ var output = state.Output;
+ switch (action.Value)
+ {
+ case TcpInputAction a:
+ input = TcpReceive(input, a.Offset, a.Data, state.ClientClosed);
+ break;
+ case TcpDataAction a:
+ output = TcpReceive(output, a.Offset, a.Data, state.HostClosed);
+ break;
+ case TcpInputConsumedAction a:
+ input = TcpConsume(input, a.ConsumedBytes);
+ break;
+ case TcpDataConsumedAction a:
+ output = TcpConsume(output, a.ConsumedBytes);
+ break;
+ case TcpInputEofAction a:
+ input = TcpEof(input, a.FinalOffset, state.ClientClosed);
+ break;
+ case TcpDataEofAction a:
+ output = TcpEof(output, a.FinalOffset, state.HostClosed);
+ break;
+ case TcpClientCloseAction:
+ return state.ClientClosed ? state : state with { ClientClosed = true };
+ case TcpHostCloseAction:
+ return state.HostClosed ? state : state with { HostClosed = true };
+ case TcpClientResetAction a:
+ return state with { Reset = new TcpResetState { Source = TcpEndpoint.Client, Reason = a.Reason } };
+ case TcpHostResetAction a:
+ return state with { Reset = new TcpResetState { Source = TcpEndpoint.Host, Reason = a.Reason } };
+ default:
+ return state;
+ }
+ return ReferenceEquals(input, state.Input) && ReferenceEquals(output, state.Output)
+ ? state : state with { Input = input, Output = output };
+ }
+
+ private static void RequireTcp(bool condition, string message)
+ {
+ if (!condition) throw new InvalidOperationException($"Invalid TCP action: {message}");
+ }
+
+ private static void RequireTcpOffset(long value) =>
+ RequireTcp(value >= 0 && value <= 9007199254740991L, "offset must be a nonnegative safe integer");
+
+ private static int TcpBase64Value(char c) => c switch
+ {
+ >= 'A' and <= 'Z' => c - 'A',
+ >= 'a' and <= 'z' => c - 'a' + 26,
+ >= '0' and <= '9' => c - '0' + 52,
+ '+' => 62,
+ '/' => 63,
+ _ => -1,
+ };
+
+ private static long TcpPayloadLength(string data, long maximumChunkSize)
+ {
+ RequireTcp(data.Length > 0 && data.Length <= 4 * (maximumChunkSize / 3 + (maximumChunkSize % 3 > 0 ? 1 : 0)), "chunk size");
+ int padding = data.EndsWith("==", StringComparison.Ordinal) ? 2 : data[data.Length - 1] == '=' ? 1 : 0;
+ RequireTcp(data.Length % 4 == 0, "base64 encoding");
+ int last = 0;
+ for (int i = 0; i < data.Length - padding; i++)
+ {
+ last = TcpBase64Value(data[i]);
+ RequireTcp(last >= 0, "base64 encoding");
+ }
+ if (padding > 0)
+ RequireTcp(last % (padding == 2 ? 16 : 4) == 0, "noncanonical base64 padding bits");
+ long length = (long)data.Length / 4 * 3 - padding;
+ RequireTcp(length <= maximumChunkSize, "chunk size");
+ return length;
+ }
+
+ private static FlowControlledByteDirectionState TcpReceive(FlowControlledByteDirectionState direction, long offset, string data, bool senderClosed)
+ {
+ RequireTcpOffset(offset);
+ long end = offset + TcpPayloadLength(data, direction.MaximumChunkSize);
+ RequireTcpOffset(end);
+ if (end <= direction.ReceivedBytes) return direction;
+ RequireTcp(offset == direction.ReceivedBytes, "gap or overlapping byte range");
+ RequireTcp(!senderClosed && direction.EofAtBytes is null, "data after EOF or sender close");
+ RequireTcp(end - direction.ConsumedBytes <= direction.WindowBytes, "receive window exceeded");
+ return direction with { ReceivedBytes = end };
+ }
+
+ private static FlowControlledByteDirectionState TcpConsume(FlowControlledByteDirectionState direction, long consumedBytes)
+ {
+ RequireTcpOffset(consumedBytes);
+ RequireTcp(consumedBytes <= direction.ReceivedBytes, "consuming bytes not received");
+ return consumedBytes <= direction.ConsumedBytes ? direction : direction with { ConsumedBytes = consumedBytes };
+ }
+
+ private static FlowControlledByteDirectionState TcpEof(FlowControlledByteDirectionState direction, long finalOffset, bool senderClosed)
+ {
+ RequireTcpOffset(finalOffset);
+ RequireTcp(finalOffset == direction.ReceivedBytes, "EOF offset");
+ if (direction.EofAtBytes == finalOffset) return direction;
+ RequireTcp(!senderClosed, "EOF after sender close");
+ return direction with { EofAtBytes = finalOffset };
+ }
+
// ─── Injectable timestamp ──────────────────────────────────────────────
private static volatile Func s_now = () => DateTimeOffset.UtcNow.ToUnixTimeMilliseconds();
diff --git a/clients/dotnet/src/AgentHostProtocol/Subscription.cs b/clients/dotnet/src/AgentHostProtocol/Subscription.cs
index 684321c77..14e9a6d2a 100644
--- a/clients/dotnet/src/AgentHostProtocol/Subscription.cs
+++ b/clients/dotnet/src/AgentHostProtocol/Subscription.cs
@@ -93,7 +93,7 @@ public sealed class SubscriptionEventAuthRequired : SubscriptionEvent
///
/// A tagged with the channel URI it was
-/// scoped to. Returned by .
+/// scoped to. Returned by .
///
public sealed class ClientEvent
{
@@ -117,19 +117,26 @@ public ClientEvent(string channel, SubscriptionEvent @event)
/// Internal bounded drop-oldest channel shared by the three public stream
/// wrappers (, , and
/// StateChangeStream). Encapsulates the creation,
-/// the idempotent close lifecycle, and the drop-oldest delivery so each public
+/// the idempotent close lifecycle, and the default drop-oldest delivery so each public
/// wrapper stays a thin, sealed, domain-named handle.
+/// Strict event receivers opt into terminal overflow instead.
///
internal sealed class BoundedDropOldestChannel
{
private readonly Channel _channel;
+ private readonly object? _strictWriteLock;
+ private readonly int _capacity;
+ private readonly Action? _onDropped;
private int _closed;
- internal BoundedDropOldestChannel(int bufferCapacity, Action? onDropped = null)
+ internal BoundedDropOldestChannel(int bufferCapacity, Action? onDropped = null, bool failOnOverflow = false)
{
+ _capacity = bufferCapacity;
+ _onDropped = onDropped;
+ _strictWriteLock = failOnOverflow ? new object() : null;
var options = new BoundedChannelOptions(bufferCapacity)
{
- FullMode = BoundedChannelFullMode.DropOldest,
+ FullMode = failOnOverflow ? BoundedChannelFullMode.Wait : BoundedChannelFullMode.DropOldest,
SingleReader = false,
SingleWriter = false,
};
@@ -145,24 +152,56 @@ internal BoundedDropOldestChannel(int bufferCapacity, Action? onDropped = nul
/// Completes the channel. Safe to call multiple times.
internal void Close()
+ {
+ if (_strictWriteLock is not null)
+ {
+ lock (_strictWriteLock) { Complete(null); }
+ return;
+ }
+ Complete(null);
+ }
+
+ private void Complete(Exception? error)
{
if (Interlocked.CompareExchange(ref _closed, 1, 0) == 0)
{
- _channel.Writer.TryComplete();
+ _channel.Writer.TryComplete(error);
}
}
+ internal bool FailIfStrict(Exception error)
+ {
+ if (_strictWriteLock is null) return false;
+ lock (_strictWriteLock) { Complete(error); }
+ return true;
+ }
+
///
- /// Delivers the item, evicting the oldest buffered item if the channel is full
+ /// By default delivers the item, evicting the oldest buffered item if the channel is full
/// () — the newest item is always
/// accepted, so a slow consumer loses the stalest items rather than the latest.
/// Each eviction is reported via the onDropped callback supplied at
- /// construction. Mirrors the Go trySend.
+ /// construction. Strict receivers instead preserve the prefix and fault
+ /// permanently on the first rejected write.
///
- internal void TrySend(T item)
+ internal bool TrySend(T item)
{
- if (Volatile.Read(ref _closed) == 1) return;
- _channel.Writer.TryWrite(item);
+ if (_strictWriteLock is not null)
+ {
+ lock (_strictWriteLock)
+ {
+ if (Volatile.Read(ref _closed) == 1) return false;
+ if (!_channel.Writer.TryWrite(item))
+ {
+ Complete(new SubscriptionLagException(_capacity));
+ _onDropped?.Invoke(item);
+ return false;
+ }
+ return true;
+ }
+ }
+ if (Volatile.Read(ref _closed) == 1) return false;
+ return _channel.Writer.TryWrite(item);
}
}
@@ -172,7 +211,7 @@ internal void TrySend(T item)
/// Per-URI fan-out handle returned by and
/// . Drop the handle by calling
/// (or ) or let
-/// tear it down.
+/// tear it down.
///
public sealed class Subscription : IDisposable
{
@@ -227,7 +266,7 @@ public void Close()
///
/// Top-level fan-in receiver over every inbound event from an ,
/// tagged with the channel URI. Multiple streams may exist concurrently.
-/// Returned by .
+/// Returned by .
///
// CA1711: "Stream" here names the AHP event-stream concept (mirroring Go's
// EventStream and Swift's AsyncStream usage), not a System.IO.Stream subclass.
@@ -236,6 +275,8 @@ public void Close()
Justification = "EventStream names the AHP event-stream abstraction (mirrors Go/Swift API), not a System.IO.Stream subclass.")]
public sealed class EventStream : IDisposable
{
+ internal string? Resource { get; set; }
+ internal bool IsClosed => Volatile.Read(ref _closed) != 0;
private readonly BoundedDropOldestChannel _channel;
private Action? _onClose;
private int _closed;
@@ -243,10 +284,10 @@ public sealed class EventStream : IDisposable
private static readonly KeyValuePair DropTag = new(AhpTelemetryNames.AttrStream, AhpTelemetryNames.StreamEvent);
/// Creates a new event stream.
- internal EventStream(int bufferCapacity)
+ internal EventStream(int bufferCapacity, bool failOnOverflow = false)
{
_channel = new BoundedDropOldestChannel(
- bufferCapacity, _ => AhpTelemetry.DroppedEvents.Add(1, DropTag));
+ bufferCapacity, _ => AhpTelemetry.DroppedEvents.Add(1, DropTag), failOnOverflow);
}
///
@@ -272,7 +313,15 @@ public void Close()
///
public void Dispose() => Close();
- internal void TrySend(ClientEvent ev) => _channel.TrySend(ev);
+ internal void TrySend(ClientEvent ev)
+ {
+ if (!_channel.TrySend(ev)) Close();
+ }
+
+ internal void FailIfStrict(Exception error)
+ {
+ if (_channel.FailIfStrict(error)) Close();
+ }
}
///
diff --git a/clients/dotnet/src/AgentHostProtocol/TcpConnection.cs b/clients/dotnet/src/AgentHostProtocol/TcpConnection.cs
new file mode 100644
index 000000000..cbf287c9b
--- /dev/null
+++ b/clients/dotnet/src/AgentHostProtocol/TcpConnection.cs
@@ -0,0 +1,734 @@
+#nullable enable
+
+using System;
+using System.Collections.Generic;
+using System.Collections.Concurrent;
+using System.Linq;
+using System.Threading;
+using System.Threading.Tasks;
+using System.Text.Json;
+
+namespace Microsoft.AgentHostProtocol;
+
+/// An owned TCP byte stream. One reader and one writer may run concurrently.
+public sealed class TcpConnection : IAsyncDisposable
+{
+ private readonly object _gate = new();
+ private readonly Queue _received = new();
+ private readonly SortedDictionary _pending = new();
+ private TaskCompletionSource _changed = Signal();
+ private AhpClient _client;
+ private EventStream? _events;
+ private CancellationTokenSource? _pumpCancellation;
+ private TcpConnectionState _state;
+ private Exception? _error;
+ private long _sentBytes;
+ private long _consumedBytes;
+ private long _checkpoint;
+ private long _lastClientSeq;
+ private int _generation;
+ private bool _suspended;
+ private bool _writing;
+ private bool _reading;
+ private bool _ending;
+ private bool _closing;
+ private bool _closed;
+ private bool _released;
+ private Action? _onRelease;
+
+ internal void OnRelease(Action callback)
+ {
+ lock (_gate)
+ {
+ if (!_released) { _onRelease = callback; return; }
+ }
+ callback();
+ }
+
+ internal TcpConnection(AhpClient client, string clientId, Snapshot snapshot)
+ {
+ _client = client;
+ ClientId = clientId;
+ Resource = snapshot.Resource;
+ _state = snapshot.State.Tcp!;
+ _sentBytes = _state.Input.ReceivedBytes;
+ _consumedBytes = _state.Output.ConsumedBytes;
+ _checkpoint = snapshot.FromSeq;
+ }
+
+ /// The private channel URI, unchanged across replay reconnects.
+ public string Resource { get; }
+ /// The logical client that owns this stream.
+ public string ClientId { get; }
+ /// Last accepted server state; writes are never optimistically reduced.
+ public TcpConnectionState State { get { lock (_gate) return _state; } }
+ /// Last applied checkpoint, including bytes retained in the read buffer.
+ public long AppliedCheckpoint { get { lock (_gate) return _checkpoint; } }
+ /// Whether transport delivery is suspended.
+ public bool IsSuspended { get { lock (_gate) return _suspended; } }
+ internal long LastClientSequence { get { lock (_gate) return _lastClientSeq; } }
+ internal bool IsClosed { get { lock (_gate) return _closed; } }
+ internal AhpClient Owner { get { lock (_gate) return _client; } }
+ internal bool CanRebind { get { lock (_gate) return !_closed && (_suspended || _client.ConnectionState == ConnectionState.Disconnected); } }
+
+ private static TaskCompletionSource Signal() => new(TaskCreationOptions.RunContinuationsAsynchronously);
+ private void Wake()
+ {
+ var previous = _changed;
+ _changed = Signal();
+ previous.TrySetResult(true);
+ }
+
+ private void ThrowIfFailed()
+ {
+ if (_error is not null) throw _error;
+ }
+
+ private (AhpClient Client, long Sequence, StateAction Action)? Enqueue(StateAction action)
+ {
+ long seq = _client.ReserveTcpSequence();
+ _lastClientSeq = seq;
+ _pending.Add(seq, action);
+ return _suspended ? null : (_client, seq, action);
+ }
+
+ private async Task SendAsync((AhpClient Client, long Sequence, StateAction Action)? item)
+ {
+ if (item is not { } send) return;
+ try
+ {
+ await send.Client.DispatchAsync(Resource, send.Action, send.Sequence).ConfigureAwait(false);
+ }
+ catch (AhpException) when (send.Client.ConnectionState == ConnectionState.Disconnected)
+ {
+ lock (_gate) { if (ReferenceEquals(send.Client, _client)) Suspend(); }
+ }
+ }
+
+ /// Reads one chunk, releasing its receive credit. Null means drained EOF.
+ public async Task ReadAsync(CancellationToken cancellationToken = default)
+ {
+ lock (_gate)
+ {
+ ThrowIfFailed();
+ if (_reading) throw new InvalidOperationException("TCP permits one reader at a time");
+ _reading = true;
+ }
+ try
+ {
+ while (true)
+ {
+ Task wait;
+ byte[]? data = null;
+ (AhpClient, long, StateAction)? send = null;
+ lock (_gate)
+ {
+ cancellationToken.ThrowIfCancellationRequested();
+ ThrowIfFailed();
+ if (!_suspended || _closed)
+ {
+ if (_received.Count > 0)
+ {
+ data = _received.Peek();
+ if (!_closed)
+ send = Enqueue(new StateAction(new TcpDataConsumedAction { Type = ActionType.TcpDataConsumed, ConsumedBytes = _consumedBytes + data.Length }));
+ _received.Dequeue();
+ _consumedBytes += data.Length;
+ }
+ else if (_closed || _state.HostClosed || _state.Output.EofAtBytes is not null) return null;
+ }
+ wait = _changed.Task;
+ }
+ if (data is not null)
+ {
+ await SendAsync(send).ConfigureAwait(false);
+ return data;
+ }
+ await wait.WaitAsync(cancellationToken).ConfigureAwait(false);
+ }
+ }
+ finally { lock (_gate) _reading = false; }
+ }
+
+ /// Writes bounded chunks, waiting for destination credit. Concurrent writes are rejected.
+ public async Task WriteAsync(byte[] data, CancellationToken cancellationToken = default)
+ {
+ Guard.ThrowIfNull(data, nameof(data));
+ lock (_gate)
+ {
+ ThrowIfFailed();
+ if (_writing || _ending || _closing || _closed) throw new InvalidOperationException("TCP write requires an open, idle writer");
+ _writing = true;
+ }
+ try
+ {
+ int offset = 0;
+ while (offset < data.Length)
+ {
+ Task wait;
+ (AhpClient, long, StateAction)? send = null;
+ lock (_gate)
+ {
+ cancellationToken.ThrowIfCancellationRequested();
+ ThrowIfFailed();
+ if (_closing || _closed) throw new InvalidOperationException("TCP closed during write");
+ long credit = _state.Input.WindowBytes - (_sentBytes - _state.Input.ConsumedBytes);
+ if (!_suspended && credit > 0)
+ {
+ int length = (int)Math.Min(data.Length - offset, Math.Min(credit, _state.Input.MaximumChunkSize));
+ TcpProtocol.Safe(_sentBytes + length);
+ var action = new StateAction(new TcpInputAction
+ {
+ Type = ActionType.TcpInput,
+ Offset = _sentBytes,
+ Data = Convert.ToBase64String(data, offset, length),
+ });
+ send = Enqueue(action);
+ _sentBytes += length;
+ offset += length;
+ }
+ wait = _changed.Task;
+ }
+ if (send is not null) await SendAsync(send).ConfigureAwait(false);
+ else await wait.WaitAsync(cancellationToken).ConfigureAwait(false);
+ }
+ }
+ finally { lock (_gate) { _writing = false; Wake(); } }
+ }
+
+ /// Waits for all reserved input bytes to be consumed by the destination.
+ public async Task DrainAsync(CancellationToken cancellationToken = default)
+ {
+ while (true)
+ {
+ Task wait;
+ lock (_gate)
+ {
+ ThrowIfFailed();
+ if (_state.Input.ConsumedBytes >= _sentBytes) return;
+ if (_closed) throw new InvalidOperationException("TCP closed before drain completed");
+ wait = _changed.Task;
+ }
+ await wait.WaitAsync(cancellationToken).ConfigureAwait(false);
+ }
+ }
+
+ /// Half-closes input after all writes have finished; output remains readable.
+ public async Task EndAsync(CancellationToken cancellationToken = default)
+ {
+ lock (_gate)
+ {
+ ThrowIfFailed();
+ if (_writing || _closing || _closed) throw new InvalidOperationException("TCP end requires an open, idle writer");
+ if (_ending) return;
+ _ending = true;
+ }
+ bool queued = false;
+ try
+ {
+ while (true)
+ {
+ Task wait;
+ (AhpClient, long, StateAction)? send;
+ lock (_gate)
+ {
+ cancellationToken.ThrowIfCancellationRequested();
+ ThrowIfFailed();
+ if (_closing || _closed) throw new InvalidOperationException("TCP closed before EOF");
+ send = _suspended ? null : Enqueue(new StateAction(new TcpInputEofAction { Type = ActionType.TcpInputEof, FinalOffset = _sentBytes }));
+ queued = send is not null;
+ wait = _changed.Task;
+ }
+ if (send is not null) { await SendAsync(send).ConfigureAwait(false); return; }
+ await wait.WaitAsync(cancellationToken).ConfigureAwait(false);
+ }
+ }
+ catch
+ {
+ if (!queued) lock (_gate) { _ending = false; Wake(); }
+ throw;
+ }
+ }
+
+ /// Stops input; retains crossing output and ownership until both sides close and drain.
+ public async Task CloseAsync()
+ {
+ (AhpClient, long, StateAction)? send;
+ lock (_gate)
+ {
+ if (_closing || _closed) return;
+ _closing = true;
+ send = Enqueue(new StateAction(new TcpClientCloseAction { Type = ActionType.TcpClientClose }));
+ Wake();
+ }
+ try
+ {
+ await SendAsync(send).ConfigureAwait(false);
+ }
+ catch (Exception error)
+ {
+ await FailAsync(error).ConfigureAwait(false);
+ throw;
+ }
+ }
+
+ private async Task FinishCloseAsync()
+ {
+ lock (_gate)
+ {
+ if (_closed || !_closing || !_state.ClientClosed || !_state.HostClosed
+ || _state.Input.ConsumedBytes < _sentBytes
+ || _state.Output.ConsumedBytes < _state.Output.ReceivedBytes
+ || _received.Count != 0 || _pending.Count != 0) return;
+ _closed = true;
+ Wake();
+ }
+ await ReleaseAsync().ConfigureAwait(false);
+ }
+
+ internal void Suspend()
+ {
+ lock (_gate)
+ {
+ if (_closed) return;
+ _suspended = true;
+ _generation++;
+ _pumpCancellation?.Cancel();
+ _pumpCancellation?.Dispose();
+ _pumpCancellation = null;
+ _events?.Dispose();
+ _events = null;
+ Wake();
+ }
+ }
+
+ internal void Bind(AhpClient client)
+ {
+ lock (_gate)
+ {
+ if (_closed) return;
+ if (!_suspended && _client.ConnectionState != ConnectionState.Disconnected)
+ throw new InvalidOperationException("Suspend the previous transport before reconnecting TCP");
+ Suspend();
+ _lastClientSeq = Math.Max(_lastClientSeq, _client.LastAssignedClientSequence);
+ if (!ReferenceEquals(client, _client))
+ {
+ client.TrackTcpConnection(this);
+ _client.ForgetTcpConnection(this);
+ }
+ _client = client;
+ }
+ }
+
+ internal async Task AcceptAsync(ActionEnvelope envelope, int? generation = null)
+ {
+ Exception? failure = null;
+ bool close = false;
+ lock (_gate)
+ {
+ if (_closed || envelope.Channel != Resource || (generation is not null && generation != _generation)) return;
+ try
+ {
+ TcpProtocol.Safe(envelope.ServerSeq);
+ bool clientEcho = envelope.Action.Value is TcpInputAction or TcpDataConsumedAction or TcpInputEofAction
+ or TcpClientCloseAction or TcpClientResetAction;
+ if (clientEcho)
+ {
+ if (envelope.Origin is not { } origin || origin.ClientId != ClientId)
+ throw new InvalidOperationException("TCP client echo requires the owning client origin");
+ TcpProtocol.Safe(origin.ClientSeq);
+ if (origin.ClientSeq > _lastClientSeq)
+ throw new InvalidOperationException("TCP client echo has an unassigned sequence");
+ }
+ if (envelope.ServerSeq <= _checkpoint) return;
+ if (envelope.RejectionReason is not null) throw new InvalidOperationException(envelope.RejectionReason);
+ var previous = _state;
+ var next = Reducers.TcpReducer(previous, envelope.Action);
+ if (clientEcho)
+ {
+ if (_pending.TryGetValue(envelope.Origin!.ClientSeq, out var expected))
+ {
+ if (!Equals(expected.Value, envelope.Action.Value))
+ throw new InvalidOperationException("TCP client echo does not match its pending action");
+ }
+ else if (next != previous)
+ throw new InvalidOperationException("TCP advancing client echo has no pending action");
+ }
+ if (next.Output.ReceivedBytes - _consumedBytes > next.Output.WindowBytes)
+ throw new InvalidOperationException("TCP output exceeds locally released credit");
+ if (envelope.Action.Value is TcpDataAction data && next.Output.ReceivedBytes > previous.Output.ReceivedBytes)
+ _received.Enqueue(Convert.FromBase64String(data.Data));
+ _state = next;
+ _checkpoint = envelope.ServerSeq;
+ if (clientEcho) _pending.Remove(envelope.Origin!.ClientSeq);
+ if (next.Reset is not null) failure = new InvalidOperationException($"TCP reset: {next.Reset.Reason}");
+ else close = next.HostClosed;
+ Wake();
+ }
+ catch (Exception ex) when (ex is InvalidOperationException or FormatException)
+ {
+ failure = ex;
+ }
+ }
+ if (failure is not null) await FailAsync(failure, reset: State.Reset is null).ConfigureAwait(false);
+ else
+ {
+ if (close) await CloseAsync().ConfigureAwait(false);
+ await FinishCloseAsync().ConfigureAwait(false);
+ }
+ }
+
+ internal async Task ResumeAsync(EventStream events)
+ {
+ KeyValuePair[] pending;
+ lock (_gate)
+ {
+ if (_closed) { events.Dispose(); return; }
+ pending = _pending.ToArray();
+ }
+ foreach (var item in pending)
+ await _client.DispatchAsync(Resource, item.Value, item.Key).ConfigureAwait(false);
+ Start(events);
+ }
+
+ internal void Start(EventStream events)
+ {
+ int generation;
+ CancellationToken cancellation;
+ lock (_gate)
+ {
+ if (_closed) { events.Dispose(); return; }
+ _events = events;
+ _suspended = false;
+ _pumpCancellation?.Dispose();
+ _pumpCancellation = new CancellationTokenSource();
+ cancellation = _pumpCancellation.Token;
+ generation = ++_generation;
+ Wake();
+ }
+ _ = PumpAsync(events, generation, cancellation);
+ }
+
+ private async Task PumpAsync(EventStream events, int generation, CancellationToken cancellation)
+ {
+ try
+ {
+ await foreach (var item in events.Events.ReadAllAsync(cancellation).ConfigureAwait(false))
+ if (item.Event is SubscriptionEventAction action)
+ await AcceptAsync(action.Envelope, generation).ConfigureAwait(false);
+ lock (_gate) { if (generation == _generation) Suspend(); }
+ }
+ catch (OperationCanceledException) when (cancellation.IsCancellationRequested) { }
+ catch (Exception ex)
+ {
+ lock (_gate) { if (generation != _generation) return; }
+ try { await FailAsync(ex, reset: true).ConfigureAwait(false); }
+ catch (Exception cleanupError)
+ {
+ lock (_gate) { _error = new AggregateException(ex, cleanupError); Wake(); }
+ }
+ }
+ finally { events.Dispose(); }
+ }
+
+ internal async Task FailAsync(Exception error, bool reset = false)
+ {
+ (AhpClient, long, StateAction)? send = null;
+ lock (_gate)
+ {
+ if (_closed) return;
+ _error = error;
+ _closed = true;
+ if (reset && !_suspended)
+ {
+ try { send = Enqueue(new StateAction(new TcpClientResetAction { Type = ActionType.TcpClientReset, Reason = TcpResetReason.ProtocolError })); }
+ catch (InvalidOperationException sequenceError) { _error = new AggregateException(error, sequenceError); }
+ }
+ _received.Clear();
+ _pending.Clear();
+ Wake();
+ }
+ try { await SendAsync(send).ConfigureAwait(false); }
+ finally { await ReleaseAsync().ConfigureAwait(false); }
+ }
+
+ private async Task ReleaseAsync()
+ {
+ AhpClient client;
+ Action? onRelease;
+ lock (_gate)
+ {
+ if (_released) return;
+ _released = true;
+ _pending.Clear();
+ _generation++;
+ _pumpCancellation?.Cancel();
+ _pumpCancellation?.Dispose();
+ _pumpCancellation = null;
+ _events?.Dispose();
+ _events = null;
+ client = _client;
+ onRelease = _onRelease;
+ _onRelease = null;
+ }
+ client.ForgetTcpConnection(this);
+ onRelease?.Invoke();
+ if (client.ConnectionState != ConnectionState.Disconnected)
+ {
+ try { await client.UnsubscribeAsync(Resource).ConfigureAwait(false); }
+ catch (Exception cleanupError)
+ {
+ lock (_gate) { _error = _error is null ? cleanupError : new AggregateException(_error, cleanupError); Wake(); }
+ throw;
+ }
+ }
+ }
+
+ /// Terminates pending operations and releases this subscription once.
+ public ValueTask DisposeAsync() => new(FailAsync(new ObjectDisposedException(nameof(TcpConnection))));
+
+}
+
+public sealed partial class AhpClient
+{
+ private string? _tcpClientId;
+ private TcpConnectionsCapability? _tcpCapability;
+ private readonly object _tcpOwnershipGate = new();
+ private readonly HashSet _ownedTcpConnections = new();
+ private bool _tcpDisposed;
+
+ internal void TrackTcpConnection(TcpConnection connection)
+ {
+ lock (_tcpOwnershipGate)
+ {
+ if (_tcpDisposed) throw new AhpClientClosedException();
+ _ownedTcpConnections.Add(connection);
+ }
+ }
+
+ internal void ForgetTcpConnection(TcpConnection connection)
+ {
+ lock (_tcpOwnershipGate) _ownedTcpConnections.Remove(connection);
+ }
+
+ private Task DisposeTcpConnectionsAsync()
+ {
+ TcpConnection[] connections;
+ lock (_tcpOwnershipGate)
+ {
+ _tcpDisposed = true;
+ connections = _ownedTcpConnections.ToArray();
+ }
+ return Task.WhenAll(connections.Select(c => c.DisposeAsync().AsTask()));
+ }
+ internal void InheritTcpClient(AhpClient previous)
+ {
+ _tcpClientId = previous._tcpClientId;
+ _tcpCapability = previous._tcpCapability;
+ if (previous.LastAssignedClientSequence >= 0) AdvanceTcpSequence(previous.LastAssignedClientSequence);
+ }
+ internal long LastAssignedClientSequence => Interlocked.Read(ref _nextClientSeq) - 1;
+ private readonly ConcurrentDictionary> _tcpCreationResponses = new();
+
+ private async Task ReleaseAbandonedTcpCreationAsync(ulong id, Task response)
+ {
+ try
+ {
+ JsonElement result;
+ try { result = await response.ConfigureAwait(false); }
+ catch (AhpRpcException) { return; }
+ catch (AhpClientClosedException) { return; }
+ if (result.ValueKind == JsonValueKind.Object
+ && result.TryGetProperty("snapshot", out var snapshot) && snapshot.ValueKind == JsonValueKind.Object
+ && snapshot.TryGetProperty("resource", out var resource) && resource.ValueKind == JsonValueKind.String
+ && resource.GetString() is { } uri && uri.StartsWith("ahp-tcp:", StringComparison.Ordinal)
+ && ConnectionState != ConnectionState.Disconnected)
+ await UnsubscribeAsync(uri).ConfigureAwait(false);
+ }
+ catch (Exception error)
+ {
+ await ShutdownWithErrorAsync(new AhpTransportException("io", "ahp: failed to release abandoned TCP creation", error)).ConfigureAwait(false);
+ }
+ finally { _tcpCreationResponses.TryRemove(id, out _); }
+ }
+
+ internal long ReserveTcpSequence()
+ {
+ while (true)
+ {
+ long current = Interlocked.Read(ref _nextClientSeq);
+ TcpProtocol.Safe(current);
+ if (Interlocked.CompareExchange(ref _nextClientSeq, current + 1, current) == current) return current;
+ }
+ }
+
+ private void AdvanceTcpSequence(long sequence)
+ {
+ TcpProtocol.Safe(sequence);
+ while (true)
+ {
+ long current = Interlocked.Read(ref _nextClientSeq);
+ if (current > sequence) return;
+ if (Interlocked.CompareExchange(ref _nextClientSeq, sequence + 1, current) == current) return;
+ }
+ }
+
+ /// Creates an owned TCP stream, registering its strict child route during reply processing.
+ public async Task OpenTcpConnectionAsync(
+ string session, TcpConnectionSubscription create, CancellationToken cancellationToken = default)
+ {
+ Guard.ThrowIfNull(session, nameof(session));
+ Guard.ThrowIfNull(create, nameof(create));
+ cancellationToken.ThrowIfCancellationRequested();
+ var clientId = _tcpClientId ?? throw new InvalidOperationException("Initialize before opening TCP");
+ TcpProtocol.ValidateRequest(session, create, _tcpCapability);
+ var events = CreateResourceEventStream("");
+ Snapshot? snapshot = null;
+ TcpConnection? connection = null;
+ try
+ {
+ var result = await RequestCoreAsync("subscribe",
+ new SubscribeParams { Channel = session, Create = create }, cancellationToken, ownsTcpCreation: true,
+ onResult: raw =>
+ {
+ if (raw.ValueKind == JsonValueKind.Object && raw.TryGetProperty("snapshot", out var child)
+ && child.ValueKind == JsonValueKind.Object && child.TryGetProperty("resource", out var resource)
+ && resource.ValueKind == JsonValueKind.String)
+ BindEventStream(events, resource.GetString()!);
+ }).ConfigureAwait(false);
+ snapshot = result?.Snapshot;
+ TcpProtocol.ValidateSnapshot(session, create, snapshot);
+ cancellationToken.ThrowIfCancellationRequested();
+ connection = new TcpConnection(this, clientId, snapshot!);
+ TrackTcpConnection(connection);
+ connection.Start(events);
+ lock (_tcpOwnershipGate) { if (_tcpDisposed) throw new AhpClientClosedException(); }
+ return connection;
+ }
+ catch
+ {
+ events.Dispose();
+ if (connection is not null) await connection.DisposeAsync().ConfigureAwait(false);
+ else if (snapshot?.Resource.StartsWith("ahp-tcp:", StringComparison.Ordinal) == true
+ && ConnectionState != ConnectionState.Disconnected)
+ await UnsubscribeAsync(snapshot.Resource, CancellationToken.None).ConfigureAwait(false);
+ throw;
+ }
+ }
+
+ ///
+ /// Rebinds suspended TCP handles to this fresh transport. Reconciles replay before live
+ /// delivery and resends only unacknowledged actions with their original identities.
+ /// Returned replay excludes actions at or below the caller's original checkpoint.
+ /// Snapshot fallback and missing resources close handles; sockets are never recreated.
+ ///
+ public async Task ReconnectTcpConnectionsAsync(
+ ReconnectParams parameters, IReadOnlyList connections,
+ CancellationToken cancellationToken = default)
+ {
+ Guard.ThrowIfNull(parameters, nameof(parameters));
+ Guard.ThrowIfNull(connections, nameof(connections));
+ TcpProtocol.Safe(parameters.LastSeenServerSeq);
+ if (parameters.Channel != ProtocolVersion.RootResourceUri
+ || (_tcpClientId is not null && _tcpClientId != parameters.ClientId)
+ || connections.Any(c => c.ClientId != parameters.ClientId || !c.CanRebind)
+ || connections.Select(c => c.Resource).Distinct(StringComparer.Ordinal).Count() != connections.Count)
+ throw new InvalidOperationException("TCP reconnect requires the same logical client and distinct live handles");
+ long checkpoint = parameters.LastSeenServerSeq;
+ var resources = new HashSet(parameters.Subscriptions, StringComparer.Ordinal);
+ foreach (var connection in connections)
+ {
+ checkpoint = Math.Min(checkpoint, connection.AppliedCheckpoint);
+ resources.Add(connection.Resource);
+ }
+ var receivers = connections.Select(c => CreateResourceEventStream(c.Resource)).ToArray();
+ try
+ {
+ if (connections.Count > 0) InheritTcpClient(connections[0].Owner);
+ foreach (var connection in connections) connection.Bind(this);
+ foreach (var connection in connections) AdvanceTcpSequence(connection.LastClientSequence);
+ var result = await ReconnectAsync(parameters.ClientId, checkpoint, resources.ToArray(), cancellationToken).ConfigureAwait(false);
+ _tcpClientId = parameters.ClientId;
+ if (result.Value is ReconnectReplayResult replay)
+ {
+ long previous = checkpoint;
+ foreach (var envelope in replay.Actions)
+ {
+ TcpProtocol.Safe(envelope.ServerSeq);
+ if (envelope.ServerSeq <= previous) throw new InvalidOperationException("TCP replay is not ordered");
+ previous = envelope.ServerSeq;
+ }
+ foreach (var connection in connections)
+ {
+ if (replay.Missing.Contains(connection.Resource))
+ await connection.FailAsync(new InvalidOperationException("TCP resource is missing on reconnect")).ConfigureAwait(false);
+ else
+ foreach (var envelope in replay.Actions)
+ await connection.AcceptAsync(envelope).ConfigureAwait(false);
+ }
+ for (int i = 0; i < connections.Count; i++)
+ {
+ cancellationToken.ThrowIfCancellationRequested();
+ await connections[i].ResumeAsync(receivers[i]).ConfigureAwait(false);
+ }
+ return new ReconnectResult(replay with
+ {
+ Actions = replay.Actions.Where(action => action.ServerSeq > parameters.LastSeenServerSeq).ToList(),
+ });
+ }
+ else
+ {
+ foreach (var connection in connections)
+ await connection.FailAsync(new InvalidOperationException("TCP cannot be restored from a reconnect snapshot")).ConfigureAwait(false);
+ foreach (var receiver in receivers) receiver.Dispose();
+ }
+ return result;
+ }
+ catch
+ {
+ foreach (var connection in connections) connection.Suspend();
+ foreach (var receiver in receivers) receiver.Dispose();
+ throw;
+ }
+ }
+}
+
+internal static class TcpProtocol
+{
+ internal static void Safe(long value)
+ {
+ if (value < 0 || value > 9007199254740991L) throw new InvalidOperationException("TCP counter must be a nonnegative safe integer");
+ }
+
+ private static void ValidateLimits(long windowBytes, long maximumChunkSize)
+ {
+ if (windowBytes < 1 || windowBytes > uint.MaxValue || maximumChunkSize < 1 || maximumChunkSize > windowBytes)
+ throw new InvalidOperationException("TCP window and chunk limits must be positive UInt32 values, with chunk no larger than window");
+ }
+
+ internal static void ValidateRequest(string session, TcpConnectionSubscription create, TcpConnectionsCapability? capability)
+ {
+ if (!session.StartsWith("ahp-session:", StringComparison.Ordinal) || create.Type != "tcpConnection"
+ || string.IsNullOrWhiteSpace(create.Host) || create.Port < 1 || create.Port > 65535
+ || create.Encoding != TcpDataEncoding.Base64 || capability?.Encodings.Contains(create.Encoding) != true)
+ throw new InvalidOperationException("Invalid or unsupported TCP creation request");
+ ValidateLimits(create.ReceiveWindowBytes, create.MaximumChunkSize);
+ }
+
+ internal static void ValidateSnapshot(string session, TcpConnectionSubscription create, Snapshot? snapshot)
+ {
+ if (snapshot?.State?.Tcp is not { } state || !snapshot.Resource.StartsWith("ahp-tcp:", StringComparison.Ordinal)
+ || state.Session != session || state.Target.Host != create.Host || state.Target.Port != create.Port
+ || state.Encoding != create.Encoding || state.ClientClosed || state.HostClosed || state.Reset is not null)
+ throw new InvalidOperationException("Invalid TCP creation snapshot");
+ Safe(snapshot.FromSeq);
+ foreach (var direction in new[] { state.Input, state.Output })
+ {
+ ValidateLimits(direction.WindowBytes, direction.MaximumChunkSize);
+ if (direction.ReceivedBytes != 0 || direction.ConsumedBytes != 0 || direction.EofAtBytes is not null)
+ throw new InvalidOperationException("TCP creation requires fresh byte directions");
+ }
+ if (state.Output.WindowBytes > create.ReceiveWindowBytes || state.Output.MaximumChunkSize > create.MaximumChunkSize)
+ throw new InvalidOperationException("TCP creation exceeded requested receive limits");
+ }
+}
diff --git a/clients/dotnet/tests/AgentHostProtocol.Tests/ClientTests.cs b/clients/dotnet/tests/AgentHostProtocol.Tests/ClientTests.cs
index cc3beed99..446f3ab01 100644
--- a/clients/dotnet/tests/AgentHostProtocol.Tests/ClientTests.cs
+++ b/clients/dotnet/tests/AgentHostProtocol.Tests/ClientTests.cs
@@ -9,6 +9,7 @@
using System.Threading.Channels;
using System.Threading.Tasks;
using Microsoft.AgentHostProtocol;
+using Microsoft.AgentHostProtocol.Hosts;
using Microsoft.Extensions.Time.Testing;
using Xunit;
@@ -20,21 +21,18 @@ namespace Microsoft.AgentHostProtocol.Tests;
/// Paired in-memory transport. The two ends share linked channels so frames
/// flow from one's outbox directly into the other's inbox, exactly as the Go
/// memTransport helper works.
+/// Graceful close rejects new sends but delivers accepted frames before closed.
///
internal sealed class MemTransport : ITransport
{
private readonly Channel _inbox;
private readonly Channel _outbox;
- private readonly CancellationTokenSource _closeCts;
-
private MemTransport(
Channel inbox,
- Channel outbox,
- CancellationTokenSource closeCts)
+ Channel outbox)
{
_inbox = inbox;
_outbox = outbox;
- _closeCts = closeCts;
}
/// Creates a linked pair. Frames sent to A appear on B's inbox and vice versa.
@@ -42,29 +40,25 @@ public static (MemTransport A, MemTransport B) CreatePair()
{
var a2b = Channel.CreateBounded(new BoundedChannelOptions(16) { FullMode = BoundedChannelFullMode.Wait });
var b2a = Channel.CreateBounded(new BoundedChannelOptions(16) { FullMode = BoundedChannelFullMode.Wait });
- var cts = new CancellationTokenSource(); // shared — closing either side closes both.
- return (new MemTransport(b2a, a2b, cts), new MemTransport(a2b, b2a, cts));
+ return (new MemTransport(b2a, a2b), new MemTransport(a2b, b2a));
}
public async ValueTask SendAsync(TransportMessage message, CancellationToken cancellationToken = default)
{
- using var linked = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken, _closeCts.Token);
- try { await _outbox.Writer.WriteAsync(message, linked.Token).ConfigureAwait(false); }
- catch (OperationCanceledException) when (_closeCts.IsCancellationRequested)
+ try { await _outbox.Writer.WriteAsync(message, cancellationToken).ConfigureAwait(false); }
+ catch (ChannelClosedException)
{ throw new AhpTransportException("closed"); }
}
public async ValueTask ReceiveAsync(CancellationToken cancellationToken = default)
{
- using var linked = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken, _closeCts.Token);
- try { return await _inbox.Reader.ReadAsync(linked.Token).ConfigureAwait(false); }
- catch (OperationCanceledException) when (_closeCts.IsCancellationRequested)
+ try { return await _inbox.Reader.ReadAsync(cancellationToken).ConfigureAwait(false); }
+ catch (ChannelClosedException)
{ throw new AhpTransportException("closed"); }
}
public ValueTask CloseAsync(CancellationToken cancellationToken = default)
{
- _closeCts.Cancel();
_outbox.Writer.TryComplete();
_inbox.Writer.TryComplete();
return ValueTask.CompletedTask;
@@ -109,8 +103,1394 @@ public sealed class ClientTests
{
private static readonly SystemTextJsonAhpSerializer Ser = SystemTextJsonAhpSerializer.Default;
+ private static TcpConnectionSubscription TcpCreation() => new()
+ {
+ Type = "tcpConnection",
+ Host = "localhost",
+ Port = 3000,
+ Encoding = TcpDataEncoding.Base64,
+ ReceiveWindowBytes = 4,
+ MaximumChunkSize = 2,
+ };
+
+ private static Snapshot TcpSnapshot(string resource = "ahp-tcp:/created") => new()
+ {
+ Resource = resource,
+ FromSeq = 0,
+ State = new SnapshotState
+ {
+ Tcp = new TcpConnectionState
+ {
+ Session = "ahp-session:/s1",
+ Target = new TcpTarget { Host = "localhost", Port = 3000 },
+ Encoding = TcpDataEncoding.Base64,
+ Input = new FlowControlledByteDirectionState { WindowBytes = 4, MaximumChunkSize = 2 },
+ Output = new FlowControlledByteDirectionState { WindowBytes = 4, MaximumChunkSize = 2 },
+ },
+ },
+ };
+
+ private static async Task<(MultiHostClient Multi, ChannelReader Servers, MemTransport Server, TcpConnection Connection)>
+ OpenTcpHost(CancellationToken token, bool autoReconnect = false)
+ {
+ var servers = Channel.CreateUnbounded();
+ var multi = new MultiHostClient();
+ var add = multi.AddHostAsync(new HostConfig
+ {
+ Id = new HostId("tcp"),
+ ClientId = "owner",
+ ReconnectPolicy = autoReconnect
+ ? new ReconnectPolicy { InitialBackoff = TimeSpan.FromMilliseconds(1), MaxBackoff = TimeSpan.FromMilliseconds(10) }
+ : ReconnectPolicy.Disabled,
+ TransportFactory = (_, _) =>
+ {
+ var (side, server) = MemTransport.CreatePair();
+ servers.Writer.TryWrite(server);
+ return Task.FromResult(side);
+ },
+ }, token);
+ var initial = await servers.Reader.ReadAsync(token);
+ await TcpResponse(initial, await TcpRequest(initial, "initialize", token), new InitializeResult
+ {
+ ProtocolVersion = ProtocolVersion.Current,
+ Snapshots = new(),
+ TcpConnections = new TcpConnectionsCapability { Encodings = new() { TcpDataEncoding.Base64 } },
+ }, token);
+ await TcpHostSessions(initial, token);
+ await add;
+ var open = multi.ClientFor(new HostId("tcp"))!.OpenTcpConnectionAsync("ahp-session:/s1", TcpCreation(), token);
+ await TcpResponse(initial, await TcpRequest(initial, "subscribe", token), new SubscribeResult { Snapshot = TcpSnapshot() }, token);
+ return (multi, servers.Reader, initial, await open);
+ }
+
+ private static async Task TcpHostSessions(MemTransport server, CancellationToken token)
+ => await TcpResponse(server, await TcpRequest(server, "listSessions", token), new ListSessionsResult { Items = new() }, token);
+
+ [Theory]
+ [InlineData(false)]
+ [InlineData(true)]
+ public async Task TcpHostReconnectRetainsStreamCreditPayloadAndGlobalSequence(bool spontaneous)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(15));
+ var token = timeout.Token;
+ var (multi, servers, oldServer, connection) = await OpenTcpHost(token, spontaneous);
+ await using var cleanup = multi;
+ var id = new HostId("tcp");
+ var handle = multi.ClientFor(id)!;
+ var write = connection.WriteAsync(new byte[] { 1, 2, 3, 4, 5, 6 }, token);
+ var first = await TcpDispatch(oldServer, token);
+ var second = await TcpDispatch(oldServer, token);
+ await TcpPush(oldServer, 1, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" }), token);
+ await TcpPush(oldServer, 2, first.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = first.ClientSeq });
+ await handle.DispatchAsync(new StateAction(new SessionTitleChangedAction { Type = ActionType.SessionTitleChanged, Title = "ordinary" }),
+ "ahp-session:/s1", 1000, token);
+ Assert.Equal(1000, (await TcpDispatch(oldServer, token)).ClientSeq);
+ while (connection.AppliedCheckpoint != 2) await Task.Delay(1, token);
+ Assert.False(write.IsCompleted);
+ await FakeHost.SendNotificationAsync(oldServer, "action", new ActionEnvelope
+ {
+ Channel = ProtocolVersion.RootResourceUri,
+ ServerSeq = 50,
+ Action = new StateAction(new RootActiveSessionsChangedAction { Type = ActionType.RootActiveSessionsChanged, ActiveSessions = 1 }),
+ }, token);
+ while (multi.Host(id)!.ServerSeq != 50) await Task.Delay(1, token);
+
+ if (spontaneous) await oldServer.CloseAsync(token);
+ else await multi.ReconnectAsync(id, token);
+ var server = await servers.ReadAsync(token);
+ var request = await TcpRequest(server, "reconnect", token);
+ var parameters = Ser.Deserialize(request.Params!.Value);
+ Assert.Equal("owner", parameters.ClientId);
+ Assert.Equal(2, parameters.LastSeenServerSeq);
+ Assert.Contains(connection.Resource, parameters.Subscriptions);
+ await TcpResponse(server, request, new ReconnectResult(new ReconnectReplayResult
+ {
+ Type = ReconnectResultType.Replay,
+ Missing = new(),
+ Actions = new()
+ {
+ new ActionEnvelope { Channel = connection.Resource, ServerSeq = 3,
+ Action = new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = 2 }) },
+ new ActionEnvelope { Channel = connection.Resource, ServerSeq = 4,
+ Action = new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" }) },
+ new ActionEnvelope { Channel = ProtocolVersion.RootResourceUri, ServerSeq = 5,
+ Action = new StateAction(new RootActiveSessionsChangedAction { Type = ActionType.RootActiveSessionsChanged, ActiveSessions = 999 }) },
+ },
+ }), token);
+ var resent = await TcpDispatch(server, token);
+ Assert.Equal(second.ClientSeq, resent.ClientSeq);
+ Assert.Equal(second.Action.Value, resent.Action.Value);
+ DispatchActionParams? tail = null;
+ bool sessions = false;
+ while (tail is null || !sessions)
+ {
+ var frame = Ser.DecodeMessage(await server.ReceiveAsync(token));
+ if (frame.Request is { } list)
+ {
+ Assert.Equal("listSessions", list.Method);
+ await TcpResponse(server, list, new ListSessionsResult { Items = new() }, token);
+ sessions = true;
+ }
+ else
+ {
+ tail = Ser.Deserialize(Assert.IsType(frame.Notification).Params!.Value);
+ }
+ }
+ Assert.Equal(4, Assert.IsType(tail.Action.Value).Offset);
+ Assert.True(tail.ClientSeq > 1000);
+ await write.WaitAsync(token);
+ while (multi.Host(id)!.Generation == handle.Generation) await Task.Delay(1, token);
+ Assert.Equal(1, multi.Host(id)!.ActiveSessions);
+ Assert.Throws(() => handle.CheckAliveOrThrow());
+ Assert.Equal(new byte[] { 7, 8 }, await connection.ReadAsync(token));
+ Assert.IsType((await TcpDispatch(server, token)).Action.Value);
+
+ var fresh = multi.ClientFor(id)!;
+ var open = fresh.OpenTcpConnectionAsync("ahp-session:/s1", TcpCreation(), token);
+ await TcpResponse(server, await TcpRequest(server, "subscribe", token), new SubscribeResult { Snapshot = TcpSnapshot("ahp-tcp:/second") }, token);
+ var additional = await open;
+ await additional.DisposeAsync();
+ await TcpUnsubscribe(server, additional.Resource, token);
+ var read = connection.ReadAsync(token);
+ Assert.False(read.IsCompleted); // replayed duplicate data was not enqueued twice
+ var remove = multi.RemoveHostAsync(id, token);
+ await remove;
+ await TcpUnsubscribe(server, connection.Resource, token);
+ await Assert.ThrowsAsync(() => read);
+ }
+
+ [Theory]
+ [InlineData("snapshot")]
+ [InlineData("missing")]
+ [InlineData("initialize")]
+ public async Task TcpHostReconnectFallbackFailsStreamsClosed(string mode)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(15));
+ var token = timeout.Token;
+ var (multi, servers, _, connection) = await OpenTcpHost(token);
+ await using var cleanup = multi;
+ var id = new HostId("tcp");
+ var generation = multi.Host(id)!.Generation;
+ var read = connection.ReadAsync(token);
+ await multi.ReconnectAsync(id, token);
+ var server = await servers.ReadAsync(token);
+ var request = await TcpRequest(server, "reconnect", token);
+ if (mode == "initialize")
+ {
+ await server.SendAsync(Ser.EncodeMessage(new JsonRpcMessage
+ {
+ ErrorResponse = new JsonRpcErrorResponse
+ {
+ Id = request.Id,
+ Error = new JsonRpcErrorObject { Code = -32601, Message = "reconnect unavailable" },
+ },
+ }), token);
+ }
+ else
+ {
+ var result = mode == "snapshot"
+ ? new ReconnectResult(new ReconnectSnapshotResult { Type = ReconnectResultType.Snapshot, Snapshots = new() { TcpSnapshot() } })
+ : new ReconnectResult(new ReconnectReplayResult { Type = ReconnectResultType.Replay, Actions = new(), Missing = new() { connection.Resource } });
+ await TcpResponse(server, request, result, token);
+ }
+ await TcpUnsubscribe(server, connection.Resource, token);
+ if (mode == "initialize")
+ {
+ var initialize = await TcpRequest(server, "initialize", token);
+ var parameters = Ser.Deserialize(initialize.Params!.Value);
+ Assert.NotNull(parameters.InitialSubscriptions);
+ Assert.DoesNotContain(connection.Resource, parameters.InitialSubscriptions);
+ await TcpResponse(server, initialize, new InitializeResult { ProtocolVersion = ProtocolVersion.Current, Snapshots = new() }, token);
+ }
+ await TcpHostSessions(server, token);
+ while (multi.Host(id)!.Generation == generation) await Task.Delay(1, token);
+ await Assert.ThrowsAnyAsync(() => read);
+ await connection.DisposeAsync();
+ await multi.ShutdownAsync(token);
+ }
+
+ [Theory]
+ [InlineData(false)]
+ [InlineData(true)]
+ public async Task TcpHostShutdownTerminatesBlockedOperationsAndPendingCreation(bool disconnected)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(15));
+ var token = timeout.Token;
+ var (multi, _, server, connection) = await OpenTcpHost(token);
+ await using var cleanup = multi;
+ var write = connection.WriteAsync(new byte[6], token);
+ _ = await TcpDispatch(server, token);
+ _ = await TcpDispatch(server, token);
+ var read = connection.ReadAsync(token);
+ var drain = connection.DrainAsync(token);
+ Task? creation = null;
+ if (disconnected)
+ {
+ await server.CloseAsync(token);
+ while (!connection.IsSuspended) await Task.Delay(1, token);
+ }
+ else
+ {
+ creation = multi.ClientFor(new HostId("tcp"))!.OpenTcpConnectionAsync("ahp-session:/s1", TcpCreation(), token);
+ _ = await TcpRequest(server, "subscribe", token);
+ }
+ Assert.False(write.IsCompleted);
+ Assert.False(read.IsCompleted);
+ Assert.False(drain.IsCompleted);
+ var shutdown = multi.ShutdownAsync(token);
+ if (!disconnected) await TcpUnsubscribe(server, connection.Resource, token);
+ await shutdown.WaitAsync(token);
+ if (creation is not null) await Assert.ThrowsAnyAsync(() => creation);
+ await Assert.ThrowsAsync(() => read);
+ await Assert.ThrowsAsync(() => write);
+ await Assert.ThrowsAsync(() => drain);
+ }
+
+ [Fact]
+ public async Task TcpHostShutdownDuringReconnectTerminatesRetainedStream()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (multi, servers, oldServer, connection) = await OpenTcpHost(token, autoReconnect: true);
+ await using var cleanup = multi;
+ var read = connection.ReadAsync(token);
+ await oldServer.CloseAsync(token);
+ var server = await servers.ReadAsync(token);
+ _ = await TcpRequest(server, "reconnect", token);
+ await multi.ShutdownAsync(token).WaitAsync(TimeSpan.FromSeconds(2), token);
+ await Assert.ThrowsAsync(() => read);
+ }
+
+ private static async Task TcpRequest(MemTransport server, string method, CancellationToken token)
+ {
+ var request = Assert.IsType(Ser.DecodeMessage(await server.ReceiveAsync(token)).Request);
+ Assert.Equal(method, request.Method);
+ if (method == "subscribe" && request.Params!.Value.TryGetProperty("create", out var create))
+ Assert.Equal("tcpConnection", create.GetProperty("type").GetString());
+ return request;
+ }
+
+ private static async Task TcpResponse(MemTransport server, JsonRpcRequest request, T result, CancellationToken token)
+ => await server.SendAsync(Ser.EncodeMessage(new JsonRpcMessage
+ {
+ SuccessResponse = new JsonRpcSuccessResponse { Id = request.Id, Result = Ser.SerializeToElement(result) },
+ }), token);
+
+ private static async Task TcpDispatch(MemTransport server, CancellationToken token)
+ {
+ var notification = Assert.IsType(Ser.DecodeMessage(await server.ReceiveAsync(token)).Notification);
+ Assert.Equal("dispatchAction", notification.Method);
+ return Ser.Deserialize(notification.Params!.Value);
+ }
+
+ private static async Task TcpUnsubscribe(MemTransport server, string resource, CancellationToken token)
+ {
+ var notification = Assert.IsType(Ser.DecodeMessage(await server.ReceiveAsync(token)).Notification);
+ Assert.Equal("unsubscribe", notification.Method);
+ Assert.Equal(resource, notification.Params!.Value.GetProperty("channel").GetString());
+ }
+
+ private static async Task TcpPush(MemTransport server, long sequence, StateAction action, CancellationToken token, ActionOrigin? origin = null, string? rejectionReason = null, string channel = "ahp-tcp:/created")
+ => await server.SendAsync(Ser.EncodeMessage(new JsonRpcMessage
+ {
+ Notification = new JsonRpcNotification
+ {
+ Method = "action",
+ Params = Ser.SerializeToElement(new ActionEnvelope
+ {
+ Channel = channel,
+ ServerSeq = sequence,
+ Action = action,
+ Origin = origin,
+ RejectionReason = rejectionReason,
+ })
+ },
+ }), token);
+
+ private static async Task TcpUnrelatedBurst(AhpClient client, MemTransport server, long firstSequence, CancellationToken token)
+ {
+ using var barrier = client.AttachSubscription("ahp-session:/barrier");
+ for (int i = 0; i < 16; i++)
+ {
+ await TcpPush(server, firstSequence + i * 2, new StateAction(new SessionTitleChangedAction
+ { Type = ActionType.SessionTitleChanged, Title = "busy" }), token, channel: "ahp-session:/other");
+ await TcpPush(server, firstSequence + i * 2 + 1, new StateAction(new TcpDataAction
+ { Type = ActionType.TcpData, Offset = i, Data = "AA==" }), token, channel: "ahp-tcp:/other");
+ }
+ await TcpPush(server, firstSequence + 32, new StateAction(new SessionTitleChangedAction
+ { Type = ActionType.SessionTitleChanged, Title = "barrier" }), token, channel: barrier.Uri);
+ await barrier.Events.ReadAsync(token);
+ }
+
+ [Fact]
+ public async Task TcpScopedCreationAndActiveTraffic()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side, new ClientConfig { SubscriptionBufferCapacity = 2 });
+ var initialize = client.InitializeAsync("owner", cancellationToken: token);
+ await TcpResponse(server, await TcpRequest(server, "initialize", token), new InitializeResult
+ {
+ ProtocolVersion = ProtocolVersion.Current,
+ Snapshots = new(),
+ TcpConnections = new TcpConnectionsCapability { Encodings = new() { TcpDataEncoding.Base64 } },
+ }, token);
+ await initialize;
+ var opening = client.OpenTcpConnectionAsync("ahp-session:/s1", TcpCreation(), token);
+ var request = await TcpRequest(server, "subscribe", token);
+ await TcpUnrelatedBurst(client, server, 1, token);
+ await TcpResponse(server, request, new SubscribeResult { Snapshot = TcpSnapshot() with { FromSeq = 33 } }, token);
+ await TcpPush(server, 34, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bw==" }), token);
+ var connection = await opening.WaitAsync(token);
+ await TcpUnrelatedBurst(client, server, 35, token);
+ Assert.Equal(new byte[] { 7 }, await connection.ReadAsync(token));
+ Assert.IsType((await TcpDispatch(server, token)).Action.Value);
+ await TcpPush(server, 68, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 1, Data = "CA==" }), token);
+ Assert.Equal(new byte[] { 8 }, await connection.ReadAsync(token));
+ Assert.IsType((await TcpDispatch(server, token)).Action.Value);
+ await CloseTcp(connection, server, token);
+ Assert.Equal(0, client.EventListenerCount);
+ }
+
+ [Theory]
+ [InlineData(false)]
+ [InlineData(true)]
+ public async Task TcpScopedReconnectIsolatesTrafficAndReportsOwnedOverflow(bool overflow)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (oldSide, oldServer) = MemTransport.CreatePair();
+ await using var old = AhpClient.Connect(oldSide);
+ var connection = await OpenTcp(old, oldServer, token);
+ await old.ShutdownAsync(preserveTcpConnections: true, cancellationToken: token);
+ var (side, server) = MemTransport.CreatePair();
+ await using var fresh = AhpClient.Connect(side, new ClientConfig { SubscriptionBufferCapacity = 2 });
+ var reconnect = fresh.ReconnectTcpConnectionsAsync(new ReconnectParams
+ {
+ Channel = ProtocolVersion.RootResourceUri,
+ ClientId = "owner",
+ LastSeenServerSeq = 0,
+ Subscriptions = new(),
+ }, new[] { connection }, token);
+ var request = await TcpRequest(server, "reconnect", token);
+ await TcpUnrelatedBurst(fresh, server, 2, token);
+ if (overflow)
+ {
+ for (int i = 0; i < 3; i++)
+ await TcpPush(server, 35 + i, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = i, Data = "AA==" }), token);
+ await TcpUnrelatedBurst(fresh, server, 38, token);
+ }
+ await TcpResponse(server, request, new ReconnectResult(new ReconnectReplayResult
+ {
+ Type = ReconnectResultType.Replay,
+ Missing = new(),
+ Actions = overflow ? new() : new()
+ {
+ new ActionEnvelope { Channel = connection.Resource, ServerSeq = 1,
+ Action = new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bw==" }) },
+ },
+ }), token);
+ if (!overflow)
+ await TcpPush(server, 35, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 1, Data = "CA==" }), token);
+ await reconnect.WaitAsync(token);
+ if (overflow)
+ {
+ Assert.IsType((await TcpDispatch(server, token)).Action.Value);
+ await TcpUnsubscribe(server, connection.Resource, token);
+ await Assert.ThrowsAsync(() => connection.ReadAsync(token));
+ }
+ else
+ {
+ await TcpUnrelatedBurst(fresh, server, 36, token);
+ Assert.Equal(new byte[] { 7 }, await connection.ReadAsync(token));
+ Assert.Equal(new byte[] { 8 }, await connection.ReadAsync(token));
+ _ = await TcpDispatch(server, token);
+ _ = await TcpDispatch(server, token);
+ await CloseTcp(connection, server, token);
+ }
+ Assert.Equal(0, fresh.EventListenerCount);
+ }
+
+ private static async Task OpenTcp(AhpClient client, MemTransport server, CancellationToken token, bool firstAction = false, bool invalidSnapshot = false, int maximumChunkSize = 2)
+ {
+ var initialize = client.InitializeAsync("owner", cancellationToken: token);
+ await TcpResponse(server, await TcpRequest(server, "initialize", token), new InitializeResult
+ {
+ ProtocolVersion = ProtocolVersion.Current,
+ Snapshots = new(),
+ TcpConnections = new TcpConnectionsCapability { Encodings = new() { TcpDataEncoding.Base64 } },
+ }, token);
+ await initialize;
+ var open = client.OpenTcpConnectionAsync("ahp-session:/s1", new TcpConnectionSubscription
+ {
+ Type = "tcpConnection",
+ Host = "localhost",
+ Port = 3000,
+ Encoding = TcpDataEncoding.Base64,
+ ReceiveWindowBytes = Math.Max(4, maximumChunkSize),
+ MaximumChunkSize = maximumChunkSize,
+ }, token);
+ var direction = new FlowControlledByteDirectionState { WindowBytes = Math.Max(4, maximumChunkSize), MaximumChunkSize = maximumChunkSize, ReceivedBytes = invalidSnapshot ? 1 : 0 };
+ await TcpResponse(server, await TcpRequest(server, "subscribe", token), new SubscribeResult
+ {
+ Snapshot = new Snapshot
+ {
+ Resource = "ahp-tcp:/created",
+ FromSeq = 0,
+ State = new SnapshotState
+ {
+ Tcp = new TcpConnectionState
+ {
+ Session = "ahp-session:/s1",
+ Target = new TcpTarget { Host = "localhost", Port = 3000 },
+ Encoding = TcpDataEncoding.Base64,
+ Input = direction,
+ Output = direction,
+ },
+ }
+ },
+ }, token);
+ if (firstAction)
+ await TcpPush(server, 1, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" }), token);
+ return await open;
+ }
+
+ private static async Task CloseTcp(TcpConnection connection, MemTransport server, CancellationToken token)
+ {
+ var close = connection.CloseAsync();
+ Assert.IsType((await TcpDispatch(server, token)).Action.Value);
+ await close;
+ await connection.DisposeAsync();
+ await TcpUnsubscribe(server, connection.Resource, token);
+ }
+
+ [Fact]
+ public void TcpCreationRequiresCanonicalDiscriminator()
+ {
+ var create = TcpCreation() with { Type = "tcpConnection" };
+ var capability = new TcpConnectionsCapability { Encodings = new() { TcpDataEncoding.Base64 } };
+ TcpProtocol.ValidateRequest("ahp-session:/s1", create, capability);
+ Assert.Equal("tcpConnection", Ser.SerializeToElement(new SubscribeParams
+ {
+ Channel = "ahp-session:/s1",
+ Create = create,
+ }).GetProperty("create").GetProperty("type").GetString());
+ Assert.Throws(() =>
+ TcpProtocol.ValidateRequest("ahp-session:/s1", create with { Type = "tcp" }, capability));
+ }
+
+ [Theory]
+ [InlineData(1L)]
+ [InlineData(4294967295L)]
+ [InlineData(0L)]
+ [InlineData(-1L)]
+ [InlineData(4294967296L)]
+ [InlineData(9007199254740991L)]
+ public void TcpCreationAndSnapshotLimitsUseUInt32Range(long limit)
+ {
+ var capability = new TcpConnectionsCapability { Encodings = new() { TcpDataEncoding.Base64 } };
+ var valid = limit >= 1 && limit <= 4294967295L;
+ foreach (bool chunk in new[] { false, true })
+ {
+ var create = TcpCreation() with { ReceiveWindowBytes = limit, MaximumChunkSize = chunk ? limit : 1 };
+ if (valid) TcpProtocol.ValidateRequest("ahp-session:/s1", create, capability);
+ else Assert.Throws(() => TcpProtocol.ValidateRequest("ahp-session:/s1", create, capability));
+ var request = TcpCreation() with { ReceiveWindowBytes = 4294967295L, MaximumChunkSize = 4294967295L };
+ foreach (bool input in new[] { false, true })
+ {
+ var snapshot = TcpSnapshot();
+ var direction = new FlowControlledByteDirectionState { WindowBytes = limit, MaximumChunkSize = chunk ? limit : 1 };
+ var state = snapshot.State.Tcp!;
+ snapshot = snapshot with { State = new SnapshotState { Tcp = input ? state with { Input = direction } : state with { Output = direction } } };
+ if (valid) TcpProtocol.ValidateSnapshot("ahp-session:/s1", request, snapshot);
+ else Assert.Throws(() => TcpProtocol.ValidateSnapshot("ahp-session:/s1", request, snapshot));
+ }
+ }
+ }
+
+ [Fact]
+ public async Task TcpSingleClientReconnectFiltersReturnedReplayAtCallerCheckpoint()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (oldSide, oldServer) = MemTransport.CreatePair();
+ await using var old = AhpClient.Connect(oldSide);
+ var connection = await OpenTcp(old, oldServer, token);
+ await connection.AcceptAsync(new ActionEnvelope
+ {
+ Channel = connection.Resource,
+ ServerSeq = 10,
+ Action = new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = 0 })
+ });
+ await old.ShutdownAsync(preserveTcpConnections: true, cancellationToken: token);
+ var (side, server) = MemTransport.CreatePair();
+ await using var fresh = AhpClient.Connect(side);
+ var parameters = new ReconnectParams
+ {
+ Channel = ProtocolVersion.RootResourceUri,
+ ClientId = "owner",
+ LastSeenServerSeq = 100,
+ Subscriptions = new() { "ahp-session:/s1" },
+ };
+ var reconnect = fresh.ReconnectTcpConnectionsAsync(parameters, new[] { connection }, token);
+ var request = await TcpRequest(server, "reconnect", token);
+ Assert.Equal(10, Ser.Deserialize(request.Params!.Value).LastSeenServerSeq);
+ var actions = new System.Collections.Generic.List();
+ for (long sequence = 11; sequence <= 101; sequence++)
+ actions.Add(new ActionEnvelope
+ {
+ Channel = sequence == 50 ? connection.Resource : "ahp-session:/s1",
+ ServerSeq = sequence,
+ Action = sequence == 50
+ ? new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" })
+ : new StateAction(new SessionTitleChangedAction { Type = ActionType.SessionTitleChanged, Title = $"title-{sequence}" })
+ });
+ actions.Add(new ActionEnvelope
+ {
+ Channel = connection.Resource,
+ ServerSeq = 102,
+ Action = new StateAction(new TcpDataEofAction { Type = ActionType.TcpDataEof, FinalOffset = 2 })
+ });
+ await TcpResponse(server, request, new ReconnectResult(new ReconnectReplayResult
+ {
+ Type = ReconnectResultType.Replay,
+ Actions = actions,
+ Missing = new() { "ahp-session:/missing" },
+ }), token);
+ var returned = Assert.IsType((await reconnect.WaitAsync(token)).Value);
+ Assert.Equal(100, parameters.LastSeenServerSeq);
+ Assert.Equal(102, connection.AppliedCheckpoint);
+ Assert.Equal(2, connection.State.Output.ReceivedBytes);
+ Assert.Equal(2, connection.State.Output.EofAtBytes);
+ Assert.Equal(new byte[] { 7, 8 }, await connection.ReadAsync(token));
+ Assert.IsType((await TcpDispatch(server, token)).Action.Value);
+ Assert.Equal(new long[] { 101, 102 }, returned.Actions.ConvertAll(action => action.ServerSeq));
+ Assert.Equal(new[] { "ahp-session:/missing" }, returned.Missing);
+ await connection.DisposeAsync();
+ await TcpUnsubscribe(server, connection.Resource, token);
+ }
+
+ [Fact]
+ public async Task TcpPeerCloseRespondsWithoutWaitingForCreditOrUnreadOutput()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side);
+ var connection = await OpenTcp(client, server, token);
+ var write = connection.WriteAsync(new byte[5], token);
+ var first = await TcpDispatch(server, token);
+ var second = await TcpDispatch(server, token);
+ var drain = connection.DrainAsync(token);
+ await TcpPush(server, 1, first.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = first.ClientSeq });
+ await TcpPush(server, 2, second.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = second.ClientSeq });
+ await TcpPush(server, 3, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" }), token);
+ await TcpPush(server, 4, new StateAction(new TcpHostCloseAction { Type = ActionType.TcpHostClose }), token);
+ var close = await TcpDispatch(server, token);
+ Assert.IsType(close.Action.Value);
+ Assert.False(drain.IsCompleted);
+ Assert.Equal(0, connection.State.Input.ConsumedBytes);
+ Assert.Equal(0, connection.State.Output.ConsumedBytes);
+ Assert.Equal(1, client.EventListenerCount);
+ await Assert.ThrowsAsync(() => write.WaitAsync(token));
+ await TcpPush(server, 5, close.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = close.ClientSeq });
+ await TcpPush(server, 6, new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = 4 }), token);
+ await drain.WaitAsync(token);
+ Assert.Equal(new byte[] { 7, 8 }, await connection.ReadAsync(token));
+ var credit = await TcpDispatch(server, token);
+ Assert.IsType(credit.Action.Value);
+ await TcpPush(server, 7, credit.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = credit.ClientSeq });
+ await TcpUnsubscribe(server, connection.Resource, token);
+ Assert.Null(await connection.ReadAsync(token));
+ }
+
+ [Fact]
+ public async Task TcpLocalCloseRetainsCrossingTrafficUntilBothDirectionsDrain()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side);
+ var connection = await OpenTcp(client, server, token);
+ await connection.WriteAsync(new byte[] { 1, 2 }, token);
+ var input = await TcpDispatch(server, token);
+ var drain = connection.DrainAsync(token);
+ await connection.CloseAsync();
+ var close = await TcpDispatch(server, token);
+ Assert.IsType(close.Action.Value);
+ Assert.False(connection.IsClosed);
+ Assert.Equal(1, client.EventListenerCount);
+ var read = connection.ReadAsync(token);
+ Assert.False(read.IsCompleted);
+ await TcpPush(server, 1, input.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = input.ClientSeq });
+ await TcpPush(server, 2, close.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = close.ClientSeq });
+ await TcpPush(server, 3, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" }), token);
+ Assert.Equal(new byte[] { 7, 8 }, await read.WaitAsync(token));
+ var credit = await TcpDispatch(server, token);
+ Assert.Equal(2, Assert.IsType(credit.Action.Value).ConsumedBytes);
+ await connection.AcceptAsync(new ActionEnvelope
+ {
+ Channel = connection.Resource,
+ ServerSeq = 4,
+ Action = new StateAction(new TcpHostCloseAction { Type = ActionType.TcpHostClose })
+ });
+ Assert.False(connection.IsClosed);
+ Assert.False(drain.IsCompleted);
+ Assert.Equal(1, client.EventListenerCount);
+ await connection.AcceptAsync(new ActionEnvelope
+ {
+ Channel = connection.Resource,
+ ServerSeq = 5,
+ Action = credit.Action,
+ Origin = new ActionOrigin { ClientId = "owner", ClientSeq = credit.ClientSeq }
+ });
+ Assert.False(connection.IsClosed);
+ await TcpPush(server, 6, new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = 2 }), token);
+ await TcpUnsubscribe(server, connection.Resource, token);
+ await drain.WaitAsync(token);
+ Assert.Null(await connection.ReadAsync(token));
+ Assert.True(connection.IsClosed);
+ Assert.Equal(0, client.EventListenerCount);
+ await connection.CloseAsync();
+ await connection.DisposeAsync();
+ }
+
+ [Fact]
+ public async Task TcpAdapterRejectsStaleCreationAndDetachesCancelledSetup()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side);
+ await Assert.ThrowsAsync(() => OpenTcp(client, server, token, invalidSnapshot: true));
+ await TcpUnsubscribe(server, "ahp-tcp:/created", token);
+ Assert.Equal(0, client.EventListenerCount);
+ using var cancellation = new CancellationTokenSource();
+ var open = client.OpenTcpConnectionAsync("ahp-session:/s1", new TcpConnectionSubscription
+ {
+ Type = "tcpConnection",
+ Host = "localhost",
+ Port = 3000,
+ Encoding = TcpDataEncoding.Base64,
+ ReceiveWindowBytes = 4,
+ MaximumChunkSize = 2,
+ }, cancellation.Token);
+ _ = await TcpRequest(server, "subscribe", token);
+ cancellation.Cancel();
+ await Assert.ThrowsAnyAsync(() => open);
+ Assert.Equal(0, client.EventListenerCount);
+ }
+
+ [Theory]
+ [InlineData(false)]
+ [InlineData(true)]
+ public async Task TcpResetOrDisposeTerminatesClosingStream(bool reset)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side);
+ var connection = await OpenTcp(client, server, token);
+ var write = connection.WriteAsync(new byte[5], token);
+ _ = await TcpDispatch(server, token);
+ _ = await TcpDispatch(server, token);
+ var drain = connection.DrainAsync(token);
+ await connection.CloseAsync();
+ _ = await TcpDispatch(server, token);
+ await Assert.ThrowsAsync(() => write.WaitAsync(token));
+ Assert.False(drain.IsCompleted);
+ if (reset)
+ {
+ await connection.AcceptAsync(new ActionEnvelope
+ {
+ Channel = connection.Resource,
+ ServerSeq = 1,
+ Action = new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" })
+ });
+ await TcpPush(server, 2, new StateAction(new TcpHostResetAction { Type = ActionType.TcpHostReset, Reason = TcpResetReason.ProtocolError }), token);
+ await TcpUnsubscribe(server, connection.Resource, token);
+ await Assert.ThrowsAsync(() => connection.ReadAsync(token));
+ }
+ else
+ {
+ var read = connection.ReadAsync(token);
+ Assert.False(read.IsCompleted);
+ await connection.DisposeAsync();
+ await TcpUnsubscribe(server, connection.Resource, token);
+ await Assert.ThrowsAsync(() => read);
+ }
+ await Assert.ThrowsAnyAsync(() => drain.WaitAsync(token));
+ Assert.Equal(0, client.EventListenerCount);
+ await connection.DisposeAsync();
+ await connection.CloseAsync();
+ }
+
+ [Fact]
+ public async Task TcpCloseWhileSuspendedReplaysAndDrainsBeforeRelease()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (oldSide, oldServer) = MemTransport.CreatePair();
+ await using var old = AhpClient.Connect(oldSide);
+ var connection = await OpenTcp(old, oldServer, token);
+ await old.ShutdownAsync(preserveTcpConnections: true, cancellationToken: token);
+ await connection.CloseAsync();
+ Assert.False(connection.IsClosed);
+ var read = connection.ReadAsync(token);
+ Assert.False(read.IsCompleted);
+ var (side, server) = MemTransport.CreatePair();
+ await using var fresh = AhpClient.Connect(side);
+ var reconnect = fresh.ReconnectTcpConnectionsAsync(new ReconnectParams
+ {
+ Channel = ProtocolVersion.RootResourceUri,
+ ClientId = "owner",
+ LastSeenServerSeq = 0,
+ Subscriptions = new(),
+ }, new[] { connection }, token);
+ var request = await TcpRequest(server, "reconnect", token);
+ Assert.Contains(connection.Resource, Ser.Deserialize(request.Params!.Value).Subscriptions);
+ await TcpResponse(server, request, new ReconnectResult(new ReconnectReplayResult
+ {
+ Type = ReconnectResultType.Replay,
+ Missing = new(),
+ Actions = new()
+ {
+ new() { Channel = connection.Resource, ServerSeq = 1, Action = new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" }) },
+ new() { Channel = connection.Resource, ServerSeq = 2, Action = new StateAction(new TcpHostCloseAction { Type = ActionType.TcpHostClose }) },
+ },
+ }), token);
+ await reconnect.WaitAsync(token);
+ var close = await TcpDispatch(server, token);
+ Assert.IsType(close.Action.Value);
+ Assert.Equal(new byte[] { 7, 8 }, await read.WaitAsync(token));
+ var credit = await TcpDispatch(server, token);
+ Assert.IsType(credit.Action.Value);
+ Assert.Equal(1, fresh.EventListenerCount);
+ await TcpPush(server, 3, close.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = close.ClientSeq });
+ await TcpPush(server, 4, credit.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = credit.ClientSeq });
+ await TcpUnsubscribe(server, connection.Resource, token);
+ Assert.Null(await connection.ReadAsync(token));
+ Assert.Equal(0, fresh.EventListenerCount);
+ }
+
+ [Theory]
+ [InlineData(false, "ahp-tcp:/late")]
+ [InlineData(true, "ahp-tcp:/late")]
+ [InlineData(true, "ahp-session:/s1")]
+ public async Task TcpAdapterReleasesLateCreationWithoutUnsubscribingParent(bool cancel, string resource)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var clock = new FakeTimeProvider();
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side, new ClientConfig { TimeProvider = clock, DefaultRequestTimeout = TimeSpan.FromMinutes(1) });
+ var initial = await OpenTcp(client, server, token);
+ await CloseTcp(initial, server, token);
+ using var cancellation = new CancellationTokenSource();
+ var open = client.OpenTcpConnectionAsync("ahp-session:/s1", new TcpConnectionSubscription
+ {
+ Type = "tcpConnection",
+ Host = "localhost",
+ Port = 3000,
+ Encoding = TcpDataEncoding.Base64,
+ ReceiveWindowBytes = 4,
+ MaximumChunkSize = 2,
+ }, cancellation.Token);
+ var request = await TcpRequest(server, "subscribe", token);
+ if (cancel) cancellation.Cancel();
+ else clock.Advance(TimeSpan.FromMinutes(1));
+ await Assert.ThrowsAnyAsync(() => open.WaitAsync(token));
+ Assert.Equal(0, client.EventListenerCount);
+ Assert.Equal(0, client.PendingRequestCount);
+ await TcpResponse(server, request, new { snapshot = new { resource } }, token);
+ if (resource.StartsWith("ahp-tcp:", StringComparison.Ordinal)) await TcpUnsubscribe(server, resource, token);
+ using var barrier = client.AttachSubscription("ahp-session:/barrier");
+ await TcpResponse(server, request, new { snapshot = new { resource } }, token);
+ await server.SendAsync(BuildActionNotification("ahp-session:/barrier", 99, "barrier"), token);
+ _ = await barrier.Events.ReadAsync(token);
+ var probe = client.RequestAsync("probe", new SubscribeParams { Channel = "ahp-session:/s1" }, token);
+ await TcpResponse(server, await TcpRequest(server, "probe", token), new SubscribeResult(), token);
+ await probe.WaitAsync(token);
+ Assert.Equal(ConnectionState.Connected, client.ConnectionState);
+ }
+
+ [Theory]
+ [InlineData(0)]
+ [InlineData(1)]
+ [InlineData(2)]
+ public async Task TcpAdapterResetCloseAndDisposeWakeAllBlockedOperations(int terminal)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side);
+ var connection = await OpenTcp(client, server, token);
+ var read = connection.ReadAsync(token);
+ var write = connection.WriteAsync(new byte[5], token);
+ _ = await TcpDispatch(server, token);
+ _ = await TcpDispatch(server, token);
+ var drain = connection.DrainAsync(token);
+ if (terminal == 1)
+ await TcpPush(server, 1, new StateAction(new TcpHostResetAction { Type = ActionType.TcpHostReset, Reason = TcpResetReason.ProtocolError }), token);
+ else if (terminal == 2)
+ await CloseTcp(connection, server, token);
+ else
+ await connection.DisposeAsync();
+ if (terminal != 2) await TcpUnsubscribe(server, connection.Resource, token);
+ await Assert.ThrowsAnyAsync(() => read.WaitAsync(token));
+ foreach (var operation in new Task[] { write, drain })
+ await Assert.ThrowsAnyAsync(() => operation.WaitAsync(token));
+ await connection.DisposeAsync();
+ Assert.Equal(0, client.EventListenerCount);
+ }
+
+ [Fact]
+ public async Task TcpAdapterReconnectContinuesBlockedWriterAfterReplayCredit()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (oldSide, oldServer) = MemTransport.CreatePair();
+ await using var oldClient = AhpClient.Connect(oldSide);
+ var connection = await OpenTcp(oldClient, oldServer, token);
+ var write = connection.WriteAsync(new byte[6], token);
+ var first = await TcpDispatch(oldServer, token);
+ var second = await TcpDispatch(oldServer, token);
+ await TcpPush(oldServer, 1, first.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = first.ClientSeq });
+ await TcpPush(oldServer, 2, second.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = second.ClientSeq });
+ await TcpPush(oldServer, 3, new StateAction(new TcpDataEofAction { Type = ActionType.TcpDataEof, FinalOffset = 0 }), token);
+ Assert.Null(await connection.ReadAsync(token));
+ Assert.False(write.IsCompleted);
+ await oldClient.DispatchAsync("ahp-session:/s1",
+ new StateAction(new SessionTitleChangedAction { Type = ActionType.SessionTitleChanged, Title = "ordinary action" }), 100, token);
+ Assert.Equal(100, (await TcpDispatch(oldServer, token)).ClientSeq);
+ await oldClient.ShutdownAsync(preserveTcpConnections: true, cancellationToken: token);
+ var (side, server) = MemTransport.CreatePair();
+ await using var fresh = AhpClient.Connect(side);
+ var reconnect = fresh.ReconnectTcpConnectionsAsync(new ReconnectParams
+ {
+ Channel = ProtocolVersion.RootResourceUri,
+ ClientId = "owner",
+ Subscriptions = new(),
+ LastSeenServerSeq = 20,
+ }, new[] { connection }, token);
+ var request = await TcpRequest(server, "reconnect", token);
+ Assert.Equal(3, Ser.Deserialize(request.Params!.Value).LastSeenServerSeq);
+ await TcpResponse(server, request, new ReconnectResult(new ReconnectReplayResult
+ {
+ Type = ReconnectResultType.Replay,
+ Missing = new(),
+ Actions = new()
+ {
+ new() { Channel = connection.Resource, ServerSeq = 4, Action = new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = 2 }) },
+ },
+ }), token);
+ await reconnect.WaitAsync(token);
+ var tail = await TcpDispatch(server, token);
+ Assert.Equal(4, Assert.IsType(tail.Action.Value).Offset);
+ Assert.True(tail.ClientSeq > 100);
+ await write.WaitAsync(token);
+ await CloseTcp(connection, server, token);
+ }
+
+ [Fact]
+ public async Task TcpAdapterReservesCreditChunksReadsDuplicatesAndHalfCloses()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (clientSide, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(clientSide);
+ var connection = await OpenTcp(client, server, token);
+ var write = connection.WriteAsync(new byte[] { 1, 2, 3, 4, 5 }, token);
+ var first = await TcpDispatch(server, token);
+ var second = await TcpDispatch(server, token);
+ Assert.Equal(2, Convert.FromBase64String(Assert.IsType(first.Action.Value).Data).Length);
+ Assert.Equal(2, Assert.IsType(second.Action.Value).Offset);
+ Assert.False(write.IsCompleted);
+ await Assert.ThrowsAsync(() => connection.WriteAsync(new byte[] { 9 }, token));
+ await TcpPush(server, 1, first.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = first.ClientSeq });
+ await TcpPush(server, 2, second.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = second.ClientSeq });
+ await TcpPush(server, 3, new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = 2 }), token);
+ var third = await TcpDispatch(server, token);
+ await write.WaitAsync(token);
+ Assert.Equal(4, Assert.IsType(third.Action.Value).Offset);
+ var drain = connection.DrainAsync(token);
+ Assert.False(drain.IsCompleted);
+ await TcpPush(server, 4, third.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = third.ClientSeq });
+ await TcpPush(server, 5, new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = 5 }), token);
+ await drain.WaitAsync(token);
+ var data = new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" });
+ await TcpPush(server, 6, data, token);
+ await TcpPush(server, 7, data, token);
+ await TcpPush(server, 8, new StateAction(new TcpDataEofAction { Type = ActionType.TcpDataEof, FinalOffset = 2 }), token);
+ Assert.Equal(new byte[] { 7, 8 }, await connection.ReadAsync(token));
+ Assert.Equal(2, Assert.IsType((await TcpDispatch(server, token)).Action.Value).ConsumedBytes);
+ Assert.Null(await connection.ReadAsync(token));
+ await connection.EndAsync(token);
+ Assert.Equal(5, Assert.IsType((await TcpDispatch(server, token)).Action.Value).FinalOffset);
+ await CloseTcp(connection, server, token);
+ Assert.Equal(0, client.EventListenerCount);
+ }
+
+ [Fact]
+ public async Task TcpAdapterPreservesFirstActionAndStrictLossWakesWaiters()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (clientSide, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(clientSide);
+ var connection = await OpenTcp(client, server, token, firstAction: true);
+ Assert.Equal(new byte[] { 7, 8 }, await connection.ReadAsync(token));
+ _ = await TcpDispatch(server, token);
+ var read = connection.ReadAsync(token);
+ var write = connection.WriteAsync(new byte[5], token);
+ _ = await TcpDispatch(server, token);
+ _ = await TcpDispatch(server, token);
+ var drain = connection.DrainAsync(token);
+ await server.SendAsync(TransportMessage.FromText("{"), token);
+ await Assert.ThrowsAsync(() => read);
+ await Assert.ThrowsAsync(() => write);
+ await Assert.ThrowsAsync(() => drain);
+ Assert.IsType((await TcpDispatch(server, token)).Action.Value);
+ await TcpUnsubscribe(server, connection.Resource, token);
+ await connection.DisposeAsync();
+ Assert.Equal(0, client.EventListenerCount);
+ }
+
+ [Theory]
+ [InlineData("missing")]
+ [InlineData("owner")]
+ [InlineData("negative")]
+ [InlineData("unsafe")]
+ [InlineData("unassigned")]
+ [InlineData("wrong-pending")]
+ [InlineData("reused")]
+ [InlineData("payload")]
+ [InlineData("eof")]
+ [InlineData("credit")]
+ [InlineData("close")]
+ [InlineData("reset")]
+ [InlineData("rejected-empty")]
+ public async Task TcpAdapterRejectsMalformedClientEchoWithoutAdvancingState(string malformed)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side);
+ var connection = await OpenTcp(client, server, token);
+ var read = connection.ReadAsync(token);
+ var write = connection.WriteAsync(new byte[] { 1, 2, 3, 4, 5 }, token);
+ var first = await TcpDispatch(server, token);
+ var second = await TcpDispatch(server, token);
+ var drain = connection.DrainAsync(token);
+ StateAction action = first.Action;
+ ActionOrigin? origin = new() { ClientId = "owner", ClientSeq = first.ClientSeq };
+ switch (malformed)
+ {
+ case "missing": origin = null; break;
+ case "owner": origin = origin with { ClientId = "other" }; break;
+ case "negative": origin = origin with { ClientSeq = -1 }; break;
+ case "unsafe": origin = origin with { ClientSeq = 9007199254740992 }; break;
+ case "unassigned": origin = origin with { ClientSeq = second.ClientSeq + 1 }; break;
+ case "wrong-pending": origin = origin with { ClientSeq = second.ClientSeq }; break;
+ case "reused":
+ await TcpPush(server, 1, first.Action, token, origin);
+ action = second.Action;
+ break;
+ case "payload": action = new StateAction(new TcpInputAction { Type = ActionType.TcpInput, Offset = 0, Data = "AgE=" }); break;
+ case "eof": origin = null; action = new StateAction(new TcpInputEofAction { Type = ActionType.TcpInputEof, FinalOffset = 0 }); break;
+ case "credit": origin = null; action = new StateAction(new TcpDataConsumedAction { Type = ActionType.TcpDataConsumed, ConsumedBytes = 0 }); break;
+ case "close": origin = null; action = new StateAction(new TcpClientCloseAction { Type = ActionType.TcpClientClose }); break;
+ case "reset": origin = null; action = new StateAction(new TcpClientResetAction { Type = ActionType.TcpClientReset, Reason = TcpResetReason.ProtocolError }); break;
+ }
+ await TcpPush(server, 2, action, token, origin, malformed == "rejected-empty" ? "" : null);
+ foreach (var pending in new Task[] { read, write, drain })
+ await Assert.ThrowsAsync(() => pending.WaitAsync(token));
+ Assert.Equal(malformed == "reused" ? 2 : 0, connection.State.Input.ReceivedBytes);
+ Assert.Equal(0, connection.State.Input.ConsumedBytes);
+ Assert.IsType((await TcpDispatch(server, token)).Action.Value);
+ await TcpUnsubscribe(server, connection.Resource, token);
+ Assert.Equal(0, client.EventListenerCount);
+ await connection.DisposeAsync();
+ }
+
+ [Theory]
+ [InlineData(false)]
+ [InlineData(true)]
+ public async Task TcpAdapterReconnectRetainsReadersAndResendsOnlyUnacknowledgedActions(bool acknowledged)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (oldSide, oldServer) = MemTransport.CreatePair();
+ await using var oldClient = AhpClient.Connect(oldSide);
+ var connection = await OpenTcp(oldClient, oldServer, token);
+ await connection.WriteAsync(new byte[] { 1, 2 }, token);
+ var original = await TcpDispatch(oldServer, token);
+ var read = connection.ReadAsync(token);
+ await oldClient.ShutdownAsync(preserveTcpConnections: true, cancellationToken: token);
+ var (freshSide, server) = MemTransport.CreatePair();
+ await using var fresh = AhpClient.Connect(freshSide);
+ await Assert.ThrowsAsync(() => fresh.ReconnectTcpConnectionsAsync(
+ new ReconnectParams { Channel = ProtocolVersion.RootResourceUri, ClientId = "other", LastSeenServerSeq = 20, Subscriptions = new() },
+ new[] { connection }, token));
+ var reconnect = fresh.ReconnectTcpConnectionsAsync(new ReconnectParams
+ {
+ Channel = ProtocolVersion.RootResourceUri,
+ ClientId = "owner",
+ LastSeenServerSeq = 20,
+ Subscriptions = new() { "ahp-session:/s1" },
+ }, new[] { connection }, token);
+ var request = await TcpRequest(server, "reconnect", token);
+ var parameters = Ser.Deserialize(request.Params!.Value);
+ Assert.Equal(0, parameters.LastSeenServerSeq);
+ Assert.Contains(connection.Resource, parameters.Subscriptions);
+ var actions = new System.Collections.Generic.List
+ {
+ new() { Channel = connection.Resource, ServerSeq = 1, Action = new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" }),
+ Origin = new ActionOrigin { ClientId = "owner", ClientSeq = original.ClientSeq } },
+ };
+ if (acknowledged)
+ actions.Add(new ActionEnvelope
+ {
+ Channel = connection.Resource,
+ ServerSeq = 2,
+ Action = original.Action,
+ Origin = new ActionOrigin { ClientId = "owner", ClientSeq = original.ClientSeq }
+ });
+ await TcpResponse(server, request, new ReconnectResult(new ReconnectReplayResult
+ {
+ Type = ReconnectResultType.Replay,
+ Actions = actions,
+ Missing = new(),
+ }), token);
+ await TcpPush(server, 3, original.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = original.ClientSeq });
+ await TcpPush(server, 4, new StateAction(new TcpDataEofAction { Type = ActionType.TcpDataEof, FinalOffset = 2 }), token);
+ await reconnect.WaitAsync(token);
+ Assert.Equal(new byte[] { 7, 8 }, await read.WaitAsync(token));
+ if (!acknowledged)
+ {
+ var resent = await TcpDispatch(server, token);
+ Assert.Equal(original.ClientSeq, resent.ClientSeq);
+ Assert.Equal(original.Action.Value, resent.Action.Value);
+ }
+ var credit = await TcpDispatch(server, token);
+ Assert.IsType(credit.Action.Value);
+ Assert.True(credit.ClientSeq > original.ClientSeq);
+ Assert.Null(await connection.ReadAsync(token));
+ await CloseTcp(connection, server, token);
+ }
+
+ [Theory]
+ [InlineData(false)]
+ [InlineData(true)]
+ public async Task TcpAdapterReconnectSnapshotOrMissingFailsClosed(bool snapshot)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (oldSide, oldServer) = MemTransport.CreatePair();
+ await using var oldClient = AhpClient.Connect(oldSide);
+ var connection = await OpenTcp(oldClient, oldServer, token);
+ var read = connection.ReadAsync(token);
+ await oldClient.ShutdownAsync(preserveTcpConnections: true, cancellationToken: token);
+ var (freshSide, server) = MemTransport.CreatePair();
+ await using var fresh = AhpClient.Connect(freshSide);
+ var reconnect = fresh.ReconnectTcpConnectionsAsync(new ReconnectParams
+ {
+ Channel = ProtocolVersion.RootResourceUri,
+ ClientId = "owner",
+ Subscriptions = new(),
+ LastSeenServerSeq = 0,
+ }, new[] { connection }, token);
+ var request = await TcpRequest(server, "reconnect", token);
+ var result = snapshot
+ ? new ReconnectResult(new ReconnectSnapshotResult { Type = ReconnectResultType.Snapshot, Snapshots = new() })
+ : new ReconnectResult(new ReconnectReplayResult { Type = ReconnectResultType.Replay, Actions = new(), Missing = new() { connection.Resource } });
+ await TcpResponse(server, request, result, token);
+ await TcpUnsubscribe(server, connection.Resource, token);
+ await reconnect.WaitAsync(token);
+ await Assert.ThrowsAsync(() => read);
+ Assert.Equal(0, fresh.EventListenerCount);
+ await connection.DisposeAsync();
+ var open = fresh.OpenTcpConnectionAsync("ahp-session:/s1", TcpCreation(), token);
+ await TcpResponse(server, await TcpRequest(server, "subscribe", token),
+ new SubscribeResult { Snapshot = TcpSnapshot("ahp-tcp:/replacement") }, token);
+ var replacement = await open;
+ await replacement.DisposeAsync();
+ await TcpUnsubscribe(server, replacement.Resource, token);
+ }
+
+ [Theory]
+ [InlineData(false)]
+ [InlineData(true)]
+ public async Task TcpClientDefaultShutdownDisposesEvenAfterPreservedTransportShutdown(bool preserved)
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side);
+ var connection = await OpenTcp(client, server, token);
+ var read = connection.ReadAsync(token);
+ var write = connection.WriteAsync(new byte[6], token);
+ _ = await TcpDispatch(server, token);
+ _ = await TcpDispatch(server, token);
+ var drain = connection.DrainAsync(token);
+ if (preserved)
+ {
+ await client.ShutdownAsync(preserveTcpConnections: true, cancellationToken: token);
+ Assert.False(read.IsCompleted);
+ Assert.False(write.IsCompleted);
+ Assert.False(drain.IsCompleted);
+ }
+ var shutdown = client.ShutdownAsync(token);
+ if (!preserved) await TcpUnsubscribe(server, connection.Resource, token);
+ await shutdown.WaitAsync(token);
+ await Assert.ThrowsAsync(() => read);
+ await Assert.ThrowsAsync(() => write);
+ await Assert.ThrowsAsync(() => drain);
+ Assert.Equal(0, client.EventListenerCount);
+ }
+
// ── Request round-trip ────────────────────────────────────────────────
+ [Fact]
+ public async Task TcpAdapterLargeEncodingAndFinalClosePreserveCompletedDrainAndBufferedReads()
+ {
+ using var timeout = new CancellationTokenSource(TimeSpan.FromSeconds(10));
+ var token = timeout.Token;
+ var (side, server) = MemTransport.CreatePair();
+ await using var client = AhpClient.Connect(side);
+ var bytes = new byte[4 * 1024 * 1024];
+ bytes[0] = 1;
+ bytes[bytes.Length - 1] = 255;
+ var connection = await OpenTcp(client, server, token, maximumChunkSize: bytes.Length);
+ await connection.WriteAsync(bytes, token);
+ var sent = await TcpDispatch(server, token);
+ var input = Assert.IsType(sent.Action.Value);
+ Assert.Equal(0, input.Offset);
+ Assert.Equal(bytes, Convert.FromBase64String(input.Data));
+ var drain = connection.DrainAsync(token);
+ Assert.False(drain.IsCompleted);
+ await TcpPush(server, 1, sent.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = sent.ClientSeq });
+ await TcpPush(server, 2, new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = bytes.Length }), token);
+ await TcpPush(server, 3, new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "Bwg=" }), token);
+ await TcpPush(server, 4, new StateAction(new TcpHostCloseAction { Type = ActionType.TcpHostClose }), token);
+ var close = await TcpDispatch(server, token);
+ Assert.IsType(close.Action.Value);
+ Assert.Equal(1, client.EventListenerCount);
+ await TcpPush(server, 5, close.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = close.ClientSeq });
+ await drain.WaitAsync(token);
+ await connection.DrainAsync(token);
+ Assert.Equal(new byte[] { 7, 8 }, await connection.ReadAsync(token));
+ var credit = await TcpDispatch(server, token);
+ Assert.IsType(credit.Action.Value);
+ await TcpPush(server, 6, credit.Action, token, new ActionOrigin { ClientId = "owner", ClientSeq = credit.ClientSeq });
+ await TcpUnsubscribe(server, connection.Resource, token);
+ Assert.Null(await connection.ReadAsync(token));
+ await connection.DisposeAsync();
+ }
+
+ [Theory]
+ [InlineData("{", false)]
+ [InlineData("{", true)]
+ [InlineData("""{"jsonrpc":"2.0","id":1}""", false)]
+ [InlineData("""{"jsonrpc":"2.0","method":"action"}""", false)]
+ [InlineData("""{"jsonrpc":"2.0","method":"action","params":null}""", false)]
+ [InlineData("""{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-tcp:/child","serverSeq":2,"action":{"type":"tcp/dataEof","finalOffset":0.5}}}""", false)]
+ [InlineData("""{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-tcp:/child","serverSeq":2,"action":{"type":"tcp/inputConsumed","consumedBytes":0.5}}}""", false)]
+ [InlineData("""{"jsonrpc":"2.0","method":"root/sessionAdded","params":[]}""", false)]
+ [InlineData("""{"jsonrpc":"2.0","method":"root/sessionAdded"}""", false)]
+ public async Task StrictEventsFailOnMalformedFramesAndNotificationPayloads(string wire, bool binary)
+ {
+ var (clientSide, serverSide) = MemTransport.CreatePair();
+ using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(5));
+ await using var client = AhpClient.Connect(clientSide);
+ using var strict = client.CreateEventStream(failOnOverflow: true);
+ using var ordinary = client.CreateEventStream();
+ using var barrier = client.AttachSubscription("ahp-session:/barrier");
+ await serverSide.SendAsync(BuildActionNotification("ahp-session:/s1", 1, "prefix"), cts.Token);
+ await serverSide.SendAsync(binary
+ ? TransportMessage.FromBinary(System.Text.Encoding.UTF8.GetBytes(wire))
+ : TransportMessage.FromText(wire), cts.Token);
+ await serverSide.SendAsync(BuildActionNotification("ahp-session:/s1", 2, "later"), cts.Token);
+ await serverSide.SendAsync(BuildActionNotification("ahp-session:/barrier", 99, "barrier"), cts.Token);
+ _ = await barrier.Events.ReadAsync(cts.Token);
+ Assert.Equal(1, client.EventListenerCount);
+ var prefix = await strict.Events.ReadAsync(cts.Token);
+ Assert.Equal(1, Assert.IsType(prefix.Event).Envelope.ServerSeq);
+ var error = await Assert.ThrowsAsync(async () =>
+ {
+ await foreach (var item in strict.Events.ReadAllAsync(cts.Token))
+ Assert.Fail($"unexpected event after decode loss: {item.Channel}");
+ });
+ Assert.Equal("protocol", error.Kind);
+ Assert.False(strict.Events.TryRead(out _), "a decode-failed receiver must never resume");
+ foreach (long expected in new[] { 1L, 2L, 99L })
+ {
+ var item = await ordinary.Events.ReadAsync(cts.Token);
+ Assert.Equal(expected, Assert.IsType(item.Event).Envelope.ServerSeq);
+ }
+ }
+
+ [Fact]
+ public async Task StrictEventsAllowUnknownNotificationsAndActions()
+ {
+ var (clientSide, serverSide) = MemTransport.CreatePair();
+ using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(5));
+ await using var client = AhpClient.Connect(clientSide);
+ using var events = client.CreateEventStream(failOnOverflow: true);
+ await serverSide.SendAsync(TransportMessage.FromText("""{"jsonrpc":"2.0","method":"future/notification"}"""), cts.Token);
+ await serverSide.SendAsync(TransportMessage.FromText(
+ """{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-tcp:/child","serverSeq":1,"action":{"type":"tcp/future"}}}"""), cts.Token);
+ var item = await events.Events.ReadAsync(cts.Token);
+ var envelope = Assert.IsType(item.Event).Envelope;
+ Assert.Equal(1, envelope.ServerSeq);
+ Assert.Equal("ahp-tcp:/child", envelope.Channel);
+ }
+
+ [Fact]
+ public async Task StrictDecodeFailureWakesBlockedReaderAndUnregisters()
+ {
+ var (clientSide, serverSide) = MemTransport.CreatePair();
+ using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(5));
+ await using var client = AhpClient.Connect(clientSide);
+ using var strict = client.CreateEventStream(failOnOverflow: true);
+ using var ordinary = client.CreateEventStream();
+ var pending = strict.Events.WaitToReadAsync(cts.Token).AsTask();
+ Assert.False(pending.IsCompleted);
+ await serverSide.SendAsync(TransportMessage.FromText("{"), cts.Token);
+ var error = await Assert.ThrowsAsync(async () => await pending);
+ Assert.Equal("protocol", error.Kind);
+ await serverSide.SendAsync(BuildActionNotification("ahp-session:/s1", 1, "still connected"), cts.Token);
+ var item = await ordinary.Events.ReadAsync(cts.Token);
+ Assert.Equal(1, Assert.IsType(item.Event).Envelope.ServerSeq);
+ Assert.Equal(1, client.EventListenerCount);
+ Assert.False(strict.Events.TryRead(out _));
+ await Assert.ThrowsAsync(async () => await strict.Events.Completion);
+ }
+
+ [Fact]
+ public async Task StrictEventsPreserveFirstTcpActionBeforeCreateReturns()
+ {
+ var (clientSide, serverSide) = MemTransport.CreatePair();
+ using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(5));
+ await using var client = AhpClient.Connect(clientSide);
+ using var events = client.CreateEventStream(failOnOverflow: true);
+ using var barrier = client.AttachSubscription("ahp-session:/barrier");
+ var server = Task.Run(async () =>
+ {
+ var message = Ser.DecodeMessage(await serverSide.ReceiveAsync(cts.Token));
+ var request = Assert.IsType(message.Request);
+ Assert.Equal("subscribe", request.Method);
+ Assert.Equal("tcpConnection", request.Params!.Value.GetProperty("create").GetProperty("type").GetString());
+ var parameters = Ser.Deserialize(request.Params!.Value);
+ Assert.Equal("ahp-session:/s1", parameters.Channel);
+ Assert.Equal("localhost", parameters.Create!.Host);
+ var direction = new FlowControlledByteDirectionState { WindowBytes = 8, MaximumChunkSize = 8 };
+ var result = new SubscribeResult
+ {
+ Snapshot = new Snapshot
+ {
+ Resource = "ahp-tcp:/created",
+ State = new SnapshotState
+ {
+ Tcp = new TcpConnectionState
+ {
+ Session = parameters.Channel,
+ Target = new TcpTarget { Host = "localhost", Port = 3000 },
+ Encoding = TcpDataEncoding.Base64,
+ Input = direction,
+ Output = direction,
+ }
+ },
+ FromSeq = 0,
+ },
+ };
+ await serverSide.SendAsync(Ser.EncodeMessage(new JsonRpcMessage
+ {
+ SuccessResponse = new JsonRpcSuccessResponse { Id = request.Id, Result = Ser.SerializeToElement(result) },
+ }), cts.Token);
+ await serverSide.SendAsync(Ser.EncodeMessage(new JsonRpcMessage
+ {
+ Notification = new JsonRpcNotification
+ {
+ Method = "action",
+ Params = Ser.SerializeToElement(new ActionEnvelope
+ {
+ Channel = "ahp-tcp:/created",
+ ServerSeq = 1,
+ Action = new StateAction(new TcpDataAction { Type = ActionType.TcpData, Offset = 0, Data = "AA==" }),
+ }),
+ },
+ }), cts.Token);
+ await serverSide.SendAsync(BuildActionNotification("ahp-session:/barrier", 2, "barrier"), cts.Token);
+ }, cts.Token);
+ var result = await client.RequestAsync("subscribe", new SubscribeParams
+ {
+ Channel = "ahp-session:/s1",
+ Create = new TcpConnectionSubscription
+ {
+ Type = "tcpConnection",
+ Host = "localhost",
+ Port = 3000,
+ Encoding = TcpDataEncoding.Base64,
+ ReceiveWindowBytes = 8,
+ MaximumChunkSize = 8,
+ },
+ }, cts.Token);
+ await server;
+ _ = await barrier.Events.ReadAsync(cts.Token);
+ var snapshot = Assert.IsType(result!.Snapshot);
+ var first = await events.Events.ReadAsync(cts.Token);
+ Assert.Equal(snapshot.Resource, first.Channel);
+ var envelope = Assert.IsType(first.Event).Envelope;
+ var initial = Assert.IsType(snapshot.State.Tcp);
+ Assert.Equal(1, Reducers.TcpReducer(initial, envelope.Action).Output.ReceivedBytes);
+ }
+
+ [Fact]
+ public async Task StrictEventsOverflowIsTerminalAndOrdinaryEventsStillDropOldest()
+ {
+ var (clientSide, serverSide) = MemTransport.CreatePair();
+ using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(5));
+ await using var client = AhpClient.Connect(clientSide, new ClientConfig { SubscriptionBufferCapacity = 2 });
+ using var strict = client.CreateEventStream(failOnOverflow: true);
+ using var ordinary = client.CreateEventStream();
+ using var barrier = client.AttachSubscription("ahp-session:/barrier");
+ for (long seq = 1; seq <= 3; seq++)
+ await serverSide.SendAsync(BuildActionNotification("ahp-session:/s1", seq, $"e{seq}"), cts.Token);
+ await serverSide.SendAsync(BuildActionNotification("ahp-session:/barrier", 99, "barrier"), cts.Token);
+ _ = await barrier.Events.ReadAsync(cts.Token);
+ Assert.Equal(1, client.EventListenerCount);
+ for (long expected = 1; expected <= 2; expected++)
+ {
+ var item = await strict.Events.ReadAsync(cts.Token);
+ Assert.Equal(expected, Assert.IsType(item.Event).Envelope.ServerSeq);
+ }
+ var error = await Assert.ThrowsAsync(async () =>
+ {
+ await foreach (var item in strict.Events.ReadAllAsync(cts.Token))
+ Assert.Fail($"unexpected event after overflow: {item.Channel}");
+ });
+ Assert.Equal(2, error.Capacity);
+ var closed = await Assert.ThrowsAsync(
+ async () => await strict.Events.ReadAsync(cts.Token));
+ Assert.IsType(closed.InnerException);
+ foreach (long expected in new[] { 3L, 99L })
+ {
+ var item = await ordinary.Events.ReadAsync(cts.Token);
+ Assert.Equal(expected, Assert.IsType(item.Event).Envelope.ServerSeq);
+ }
+ using var healthy = client.CreateEventStream(failOnOverflow: true);
+ await serverSide.SendAsync(BuildActionNotification("ahp-session:/barrier", 100, "later"), cts.Token);
+ _ = await barrier.Events.ReadAsync(cts.Token);
+ Assert.False(strict.Events.TryRead(out _), "an overflowed receiver must never resume");
+ await Assert.ThrowsAsync(async () => await strict.Events.Completion);
+ var later = await healthy.Events.ReadAsync(cts.Token);
+ Assert.Equal(100, Assert.IsType(later.Event).Envelope.ServerSeq);
+ healthy.Dispose();
+ Assert.False(await healthy.Events.WaitToReadAsync(cts.Token));
+ }
+
[Fact]
public async Task RequestRoundTrip_InitializeReturnsProtocolVersion()
{
diff --git a/clients/dotnet/tests/AgentHostProtocol.Tests/FixtureDrivenReducerTests.cs b/clients/dotnet/tests/AgentHostProtocol.Tests/FixtureDrivenReducerTests.cs
index 4bf9390e4..3b3caaed5 100644
--- a/clients/dotnet/tests/AgentHostProtocol.Tests/FixtureDrivenReducerTests.cs
+++ b/clients/dotnet/tests/AgentHostProtocol.Tests/FixtureDrivenReducerTests.cs
@@ -49,40 +49,38 @@ public void ReducerMatchesFixture(string name, string path)
string reducer = root.GetProperty("reducer").GetString()!;
JsonElement initial = root.GetProperty("initial");
JsonElement expected = root.GetProperty("expected");
- var actions = new List();
- foreach (JsonElement actionElement in root.GetProperty("actions").EnumerateArray())
- {
- actions.Add(actionElement.Deserialize(Options)!);
- }
+ JsonElement actions = root.GetProperty("actions");
+ string? expectedError = root.TryGetProperty("expectedError", out var error) ? error.GetString() : null;
switch (reducer)
{
case "root":
- RunFixture(initial, expected, actions, Reducers.ApplyToRoot);
+ RunFixture(initial, expected, actions, expectedError, Reducers.ApplyToRoot);
break;
case "session":
- RunFixture(initial, expected, actions, Reducers.ApplyToSession);
+ RunFixture(initial, expected, actions, expectedError, Reducers.ApplyToSession);
break;
case "terminal":
- RunFixture(initial, expected, actions, Reducers.ApplyToTerminal);
+ RunFixture(initial, expected, actions, expectedError, Reducers.ApplyToTerminal);
break;
case "changeset":
- RunFixture(initial, expected, actions, Reducers.ApplyToChangeset);
+ RunFixture(initial, expected, actions, expectedError, Reducers.ApplyToChangeset);
break;
case "resourceWatch":
- RunFixture(initial, expected, actions, Reducers.ApplyToResourceWatch);
+ RunFixture(initial, expected, actions, expectedError, Reducers.ApplyToResourceWatch);
break;
case "annotations":
- RunFixture(initial, expected, actions, Reducers.ApplyToAnnotations);
+ RunFixture(initial, expected, actions, expectedError, Reducers.ApplyToAnnotations);
break;
case "chat":
- RunFixture(initial, expected, actions, Reducers.ApplyToChat);
+ RunFixture(initial, expected, actions, expectedError, Reducers.ApplyToChat);
break;
case "automation":
RunFixture(
initial,
expected,
actions,
+ expectedError,
Reducers.ApplyToAutomation);
break;
case "automationRun":
@@ -90,8 +88,12 @@ public void ReducerMatchesFixture(string name, string path)
initial,
expected,
actions,
+ expectedError,
Reducers.ApplyToAutomationRun);
break;
+ case "tcp":
+ RunFixture(initial, expected, actions, expectedError, Reducers.TcpReducer);
+ break;
default:
throw new Xunit.Sdk.XunitException($"unknown reducer kind '{reducer}'");
}
@@ -105,9 +107,23 @@ public void ReducerMatchesFixture(string name, string path)
private static void RunFixture(
JsonElement initial,
JsonElement expected,
- List actions,
+ JsonElement actions,
+ string? expectedError,
Func apply)
where T : class
+ => RunFixture(initial, expected, actions, expectedError, (state, action) =>
+ {
+ apply(state, action);
+ return state;
+ });
+
+ private static void RunFixture(
+ JsonElement initial,
+ JsonElement expected,
+ JsonElement actions,
+ string? expectedError,
+ Func apply)
+ where T : class
{
T state = initial.Deserialize(Options)!;
@@ -123,9 +139,33 @@ private static void RunFixture(
$"initial state did not survive round-trip:\nre-serialized: {actual}\noriginal: {original}");
}
- foreach (StateAction action in actions)
+ if (expectedError is not null) Assert.True(actions.GetArrayLength() > 0, "expectedError requires a final action");
+ int index = 0;
+ foreach (JsonElement raw in actions.EnumerateArray())
{
- apply(state, action);
+ bool mustFail = expectedError is not null && index++ == actions.GetArrayLength() - 1;
+ string before = JsonSerializer.Serialize(state, Options);
+ if (mustFail && typeof(T) == typeof(TcpConnectionState)
+ && raw.TryGetProperty("offset", out var offset)
+ && offset.GetDouble() % 1 != 0)
+ {
+ Assert.Equal("Invalid TCP action: offset must be a nonnegative safe integer", expectedError);
+ Assert.Throws(() => raw.Deserialize(Options));
+ }
+ else
+ {
+ StateAction action = raw.Deserialize(Options)!;
+ if (mustFail)
+ {
+ var failure = Assert.Throws(() => apply(state, action));
+ Assert.Equal(expectedError, failure.Message);
+ }
+ else
+ {
+ state = apply(state, action);
+ }
+ }
+ if (mustFail) Assert.Equal(before, JsonSerializer.Serialize(state, Options));
}
string got = Canon(JsonSerializer.SerializeToElement(state, Options));
diff --git a/clients/dotnet/tests/AgentHostProtocol.Tests/NativeReducerTests.cs b/clients/dotnet/tests/AgentHostProtocol.Tests/NativeReducerTests.cs
index ad32be82f..5edcae7da 100644
--- a/clients/dotnet/tests/AgentHostProtocol.Tests/NativeReducerTests.cs
+++ b/clients/dotnet/tests/AgentHostProtocol.Tests/NativeReducerTests.cs
@@ -10,6 +10,7 @@
// generated union + serializer's [WireValue] mapping, not a hand-typed literal.
#nullable enable
+using System;
using System.Collections.Generic;
using System.Linq;
using Microsoft.AgentHostProtocol;
@@ -19,6 +20,80 @@ namespace Microsoft.AgentHostProtocol.Tests;
public sealed class NativeReducerTests
{
+ private static TcpConnectionState TcpState(long size = 8) => new()
+ {
+ Session = "ahp-session:/test",
+ Target = new TcpTarget { Host = "localhost", Port = 3000 },
+ Encoding = TcpDataEncoding.Base64,
+ Input = new FlowControlledByteDirectionState { WindowBytes = size, MaximumChunkSize = size },
+ Output = new FlowControlledByteDirectionState { WindowBytes = size, MaximumChunkSize = size },
+ };
+
+ [Fact]
+ public void TcpFourMiBChunkDoesNotAllocateDecodedPayload()
+ {
+ const int size = 4 * 1024 * 1024;
+ string data = new string('A', size / 3 * 4) + "AA==";
+ var before = TcpState(size);
+ var action = new StateAction(new TcpInputAction { Type = ActionType.TcpInput, Offset = 0, Data = data });
+ _ = Reducers.TcpReducer(before, action);
+ long allocated = GC.GetAllocatedBytesForCurrentThread();
+ var after = Reducers.TcpReducer(before, action);
+ allocated = GC.GetAllocatedBytesForCurrentThread() - allocated;
+ Assert.True(allocated < 64 * 1024, $"Reducer allocated {allocated} bytes for a 4 MiB chunk");
+ Assert.Equal(size, after.Input.ReceivedBytes);
+ Assert.Equal(0, before.Input.ReceivedBytes);
+ Assert.Same(after, Reducers.TcpReducer(after, action));
+ Assert.Same(before.Output, after.Output);
+ var error = Assert.Throws(() => Reducers.TcpReducer(after,
+ new StateAction(new TcpInputAction { Type = ActionType.TcpInput, Offset = size, Data = "AA==" })));
+ Assert.Equal("Invalid TCP action: receive window exceeded", error.Message);
+ Assert.Equal(size, after.Input.ReceivedBytes);
+ }
+
+ [Fact]
+ public void TcpSafeIntegerCounters()
+ {
+ const long max = 9007199254740991L;
+ var before = TcpState();
+ before = before with { Input = before.Input with { ReceivedBytes = max - 1, ConsumedBytes = max - 1 } };
+ var state = Reducers.TcpReducer(before, new StateAction(new TcpInputAction { Type = ActionType.TcpInput, Offset = max - 1, Data = "AA==" }));
+ state = Reducers.TcpReducer(state, new StateAction(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = max }));
+ state = Reducers.TcpReducer(state, new StateAction(new TcpInputEofAction { Type = ActionType.TcpInputEof, FinalOffset = max }));
+ Assert.Equal(max, state.Input.ReceivedBytes);
+ Assert.Equal(max, state.Input.ConsumedBytes);
+ Assert.Equal(max, state.Input.EofAtBytes);
+ foreach (long value in new[] { -1L, max + 1, long.MaxValue })
+ {
+ StateAction[] actions =
+ {
+ new(new TcpInputAction { Type = ActionType.TcpInput, Offset = value, Data = "AA==" }),
+ new(new TcpDataAction { Type = ActionType.TcpData, Offset = value, Data = "AA==" }),
+ new(new TcpInputConsumedAction { Type = ActionType.TcpInputConsumed, ConsumedBytes = value }),
+ new(new TcpDataConsumedAction { Type = ActionType.TcpDataConsumed, ConsumedBytes = value }),
+ new(new TcpInputEofAction { Type = ActionType.TcpInputEof, FinalOffset = value }),
+ new(new TcpDataEofAction { Type = ActionType.TcpDataEof, FinalOffset = value }),
+ };
+ foreach (StateAction action in actions)
+ {
+ var error = Assert.Throws(() => Reducers.TcpReducer(state, action));
+ Assert.Equal("Invalid TCP action: offset must be a nonnegative safe integer", error.Message);
+ Assert.Equal(max, state.Input.ReceivedBytes);
+ }
+ }
+ }
+
+ [Fact]
+ public void TcpRejectsWhitespaceAndUnicodeBase64()
+ {
+ foreach (string data in new[] { "AAA\n", "AAA\r", "AAA\t", "AAA ", "AAA\u00e9", "AA\U0001f600" })
+ {
+ var error = Assert.Throws(() => Reducers.TcpReducer(TcpState(),
+ new StateAction(new TcpInputAction { Type = ActionType.TcpInput, Offset = 0, Data = data })));
+ Assert.Equal("Invalid TCP action: base64 encoding", error.Message);
+ }
+ }
+
// #338: an open request is an UNRESOLVED InputRequestResponsePart living in the
// active turn's response stream — there is no longer a separate live surface, so
// without an active turn there is nowhere for an open request to exist at all.
diff --git a/clients/dotnet/tests/AgentHostProtocol.Tests/TransportTests.cs b/clients/dotnet/tests/AgentHostProtocol.Tests/TransportTests.cs
index 9aefa4c93..ed20b2f5c 100644
--- a/clients/dotnet/tests/AgentHostProtocol.Tests/TransportTests.cs
+++ b/clients/dotnet/tests/AgentHostProtocol.Tests/TransportTests.cs
@@ -72,6 +72,46 @@ await Assert.ThrowsAsync(
async () => await b.ReceiveAsync(cts.Token));
}
+ // Graceful close preserves frames whose sends have completed, even if the
+ // receiver does not start reading until after close.
+ [Theory]
+ [InlineData(false)]
+ [InlineData(true)]
+ public async Task InMemoryTransport_Close_DrainsAlreadySentFrames(bool closePeer)
+ {
+ var (a, b) = MemTransport.CreatePair();
+ using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(5));
+ for (int i = 0; i < 2; i++)
+ {
+ await a.SendAsync(TransportMessage.FromText($"a-{i}"), cts.Token);
+ await b.SendAsync(TransportMessage.FromText($"b-{i}"), cts.Token);
+ }
+ await (closePeer ? b : a).CloseAsync(cts.Token);
+ for (int i = 0; i < 2; i++)
+ {
+ Assert.Equal($"a-{i}", (await b.ReceiveAsync(cts.Token)).Text);
+ Assert.Equal($"b-{i}", (await a.ReceiveAsync(cts.Token)).Text);
+ }
+ await Assert.ThrowsAsync(async () => await a.ReceiveAsync(cts.Token));
+ await Assert.ThrowsAsync(async () => await b.ReceiveAsync(cts.Token));
+ }
+
+ [Fact]
+ public async Task InMemoryTransport_Close_UnblocksBackpressuredSendWithoutLosingAcceptedFrames()
+ {
+ var (a, b) = MemTransport.CreatePair();
+ using var cts = new CancellationTokenSource(TimeSpan.FromSeconds(5));
+ for (int i = 0; i < 16; i++)
+ await a.SendAsync(TransportMessage.FromText($"frame-{i}"), cts.Token);
+ var blocked = a.SendAsync(TransportMessage.FromText("not-accepted"), cts.Token).AsTask();
+ Assert.False(blocked.IsCompleted);
+ await b.CloseAsync(cts.Token);
+ await Assert.ThrowsAsync(() => blocked);
+ for (int i = 0; i < 16; i++)
+ Assert.Equal($"frame-{i}", (await b.ReceiveAsync(cts.Token)).Text);
+ await Assert.ThrowsAsync(async () => await b.ReceiveAsync(cts.Token));
+ }
+
// ── E: send after close throws ────────────────────────────────────────
[Fact]
public async Task InMemoryTransport_SendAfterClose_Throws()
diff --git a/clients/dotnet/tests/AgentHostProtocol.Tests/TypesRoundTripFixtures.cs b/clients/dotnet/tests/AgentHostProtocol.Tests/TypesRoundTripFixtures.cs
index d0fa9c1d6..91a1b0b87 100644
--- a/clients/dotnet/tests/AgentHostProtocol.Tests/TypesRoundTripFixtures.cs
+++ b/clients/dotnet/tests/AgentHostProtocol.Tests/TypesRoundTripFixtures.cs
@@ -164,6 +164,12 @@ private static (object decoded, string reencoded) DecodeAndReencode(string type,
return Wrap(Ser.Deserialize(inputJson));
case "InitializeResult":
return Wrap(Ser.Deserialize(inputJson));
+ case "SubscribeParams":
+ return Wrap(Ser.Deserialize(inputJson));
+ case "ReconnectResult":
+ return Wrap(Ser.Deserialize(inputJson));
+ case "TcpConnectionOpenErrorData":
+ return Wrap(Ser.Deserialize(inputJson));
case "Snapshot":
return Wrap(Ser.Deserialize(inputJson));
default:
diff --git a/clients/go/README.md b/clients/go/README.md
index 13562b870..8e06632c7 100644
--- a/clients/go/README.md
+++ b/clients/go/README.md
@@ -60,6 +60,72 @@ for evt := range sub.Events() {
}
```
+## Owned TCP byte streams
+
+After `Initialize` advertises `tcpConnections` with `base64` support, create a
+stream atomically through its parent session:
+
+```go
+connection, err := client.OpenTCPConnection(ctx, sessionURI, ahptypes.TcpConnectionSubscription{
+ Type: "tcpConnection", Host: "localhost", Port: 3000,
+ Encoding: ahptypes.TcpDataEncodingBase64,
+ ReceiveWindowBytes: 256 * 1024, MaximumChunkSize: 64 * 1024,
+})
+if err != nil {
+ return err
+}
+defer connection.Dispose(context.Background())
+
+if _, err := connection.Write(ctx, requestBytes); err != nil {
+ return err
+}
+if err := connection.End(ctx); err != nil { // Input EOF; output remains readable.
+ return err
+}
+chunk, err := connection.Read(ctx) // io.EOF only after buffered output is drained.
+```
+
+The SDK owns buffering, flow control, and replay. `Read` releases receive credit;
+`Drain` waits for destination consumption. Only one writer may run at a time;
+`Write` returns the accepted prefix on cancellation. Finish the writer before
+calling `Close`, and keep reading until EOF while the close handshake drains.
+Cancellation stops waiting, not the handshake. `Dispose` aborts without draining.
+
+Transport loss suspends the same handles. For deliberate transport replacement,
+use `ShutdownPreservingTCP(ctx)`; normal shutdown disposes streams. Create a fresh
+`Client` and resume instead of initializing again:
+
+```go
+result, err := freshClient.ReconnectTCPConnections(ctx, ahptypes.ReconnectParams{
+ ClientId: "my-client", // Same ID used by the original Initialize.
+ LastSeenServerSeq: lastSeenServerSeq,
+ Subscriptions: []string{sessionURI},
+}, []*ahp.TCPConnection{connection})
+```
+
+Continue using `connection`; apply the returned result only to non-TCP
+subscriptions. Snapshot fallback or missing resources fail streams rather than
+creating new sockets.
+
+For a managed host, call `hostClientHandle.OpenTCPConnection(ctx, sessionURI, create)`
+instead of opening through its borrowed `Client()`. The manager automatically
+resumes streams across reconnects and disposes them on host removal or shutdown.
+Transport/retry policy, connection limits, and native socket bridges remain
+application-owned. See the [TCP channel contract](../../docs/specification/tcp-channel.md).
+
+## Loss-sensitive raw events (advanced)
+
+Ordinary `client.Events()` drops events when its bounded buffer is full.
+For loss-sensitive consumers, attach `client.EventsStrict()` before sending
+requests. Drain `events.Events()`, check `events.Err()` when it closes, and call
+`events.Close()` when done. Overflow reports `*ahp.SubscriptionLagError`; decode
+loss reports `*ahp.TransportError` (`Kind == "protocol"`). Both terminate the
+receiver rather than skipping events. Capacity uses `Config.SubscriptionBuffer`;
+ordinary receivers are unchanged. These raw receivers are global. The owned TCP
+adapter instead registers a strict child-scoped receiver during creation reply
+processing and reattaches it per child on reconnect. Unrelated traffic cannot
+exhaust a TCP stream's event buffer.
+
## Code generation
The contents of `ahptypes/*.go` (except `common.go`) are auto-generated
diff --git a/clients/go/ahp/client.go b/clients/go/ahp/client.go
index 1938bf369..30d85f1d2 100644
--- a/clients/go/ahp/client.go
+++ b/clients/go/ahp/client.go
@@ -151,6 +151,11 @@ type pendingResult struct {
err *ahptypes.JsonRpcError
}
+type pendingRequest struct {
+ result chan pendingResult
+ onResult func(json.RawMessage)
+}
+
// outboundMsg is the writer goroutine's input queue payload.
type outboundMsg struct {
msg ahptypes.JsonRpcMessage
@@ -159,36 +164,80 @@ type outboundMsg struct {
}
// EventStream is a top-level fan-in receiver over every inbound event
-// from a [Client]. Returned by [Client.Events].
+// from a [Client]. Returned by [Client.Events] or [Client.EventsStrict].
type EventStream struct {
- events chan ClientEvent
- closeMu sync.Mutex
- closed bool
+ events chan ClientEvent
+ closeMu sync.Mutex
+ closed bool
+ strict bool
+ err error
+ onClose func()
+ resource *string
}
// Events returns a receive-only channel of every [ClientEvent].
+// For a strict receiver, check [EventStream.Err] when this channel closes.
func (s *EventStream) Events() <-chan ClientEvent { return s.events }
+// Err returns the terminal error, if any. A strict receiver records a
+// *SubscriptionLagError on overflow or a *TransportError for malformed input.
+// Already-buffered events form a valid prefix and can still be drained,
+// but no events after the gap are delivered.
+func (s *EventStream) Err() error {
+ s.closeMu.Lock()
+ defer s.closeMu.Unlock()
+ return s.err
+}
+
// Close stops the stream. Safe to call multiple times.
func (s *EventStream) Close() {
s.closeMu.Lock()
- defer s.closeMu.Unlock()
+ onClose := s.closeLocked(nil)
+ s.closeMu.Unlock()
+ if onClose != nil {
+ onClose()
+ }
+}
+
+func (s *EventStream) closeLocked(err error) func() {
if s.closed {
- return
+ return nil
}
+ s.err = err
s.closed = true
close(s.events)
+ onClose := s.onClose
+ s.onClose = nil
+ return onClose
}
func (s *EventStream) trySend(ev ClientEvent) {
+ var onClose func()
s.closeMu.Lock()
- defer s.closeMu.Unlock()
- if s.closed {
- return
+ if !s.closed {
+ select {
+ case s.events <- ev:
+ default:
+ if s.strict {
+ onClose = s.closeLocked(&SubscriptionLagError{Capacity: cap(s.events)})
+ }
+ }
}
- select {
- case s.events <- ev:
- default:
+ s.closeMu.Unlock()
+ if onClose != nil {
+ onClose()
+ }
+}
+
+func (s *EventStream) fail(err error) {
+ var onClose func()
+ s.closeMu.Lock()
+ if s.strict {
+ onClose = s.closeLocked(err)
+ }
+ s.closeMu.Unlock()
+ if onClose != nil {
+ onClose()
}
}
@@ -229,19 +278,25 @@ type Client struct {
// pending is the request-correlation map keyed by JSON-RPC id.
pendingMu sync.Mutex
- pending map[uint64]chan pendingResult
+ pending map[uint64]pendingRequest
// subscriptionsMu guards subscriptions and the all-events
// fan-out registry.
- subscriptionsMu sync.Mutex
- subscriptions map[string][]*Subscription
- eventListeners []*EventStream
+ subscriptionsMu sync.Mutex
+ subscriptions map[string][]*Subscription
+ eventListeners []*EventStream
+ resourceEventListeners map[string][]*EventStream
serverRequestMu sync.Mutex
serverRequestHandler ServerRequestHandler
nextID atomic.Uint64
nextClientSeq atomic.Int64
+ tcpMu sync.Mutex
+ tcpClientID string
+ tcpCapability *ahptypes.TcpConnectionsCapability
+ tcpStreams map[*TCPConnection]struct{}
+ tcpDisposed bool
// done closes once the client has begun teardown. Subsequent
// sends fail with [ErrShutdown].
@@ -268,7 +323,7 @@ func Connect(_ context.Context, transport Transport, cfg Config) (*Client, error
cfg: cfg,
transport: transport,
outbound: make(chan outboundMsg, 64),
- pending: make(map[uint64]chan pendingResult),
+ pending: make(map[uint64]pendingRequest),
subscriptions: make(map[string][]*Subscription),
done: make(chan struct{}),
}
@@ -302,6 +357,14 @@ func (c *Client) Err() error {
//
// Safe to call multiple times.
func (c *Client) Shutdown(ctx context.Context) error {
+ err := c.disposeTCPStreams(ctx)
+ return errors.Join(err, c.ShutdownPreservingTCP(ctx))
+}
+
+// ShutdownPreservingTCP closes this transport but retains owned TCP handles for
+// ReconnectTCPConnections. Ordinary Shutdown permanently disposes the handles,
+// including when the transport has already failed.
+func (c *Client) ShutdownPreservingTCP(ctx context.Context) error {
c.shutdownWithError(nil)
doneCh := make(chan struct{})
go func() { c.wg.Wait(); close(doneCh) }()
@@ -334,7 +397,8 @@ func (c *Client) shutdownWithError(err error) {
failErr.Message = fmt.Sprintf("client shut down: %v", err)
}
c.pendingMu.Lock()
- for id, ch := range c.pending {
+ for id, request := range c.pending {
+ ch := request.result
select {
case ch <- pendingResult{err: failErr}:
default:
@@ -348,6 +412,10 @@ func (c *Client) shutdownWithError(err error) {
c.subscriptionsMu.Lock()
subs := c.subscriptions
listeners := c.eventListeners
+ for _, list := range c.resourceEventListeners {
+ listeners = append(listeners, list...)
+ }
+ c.resourceEventListeners = nil
c.subscriptions = map[string][]*Subscription{}
c.eventListeners = nil
c.subscriptionsMu.Unlock()
@@ -412,11 +480,13 @@ func (c *Client) runReader() {
msg, err := c.transport.Recv(ctx)
cancel()
if err != nil {
+ c.failStrictEvents(err)
c.shutdownWithError(fmt.Errorf("ahp: transport recv: %w", err))
return
}
parsed, perr := msg.IntoParsed()
if perr != nil {
+ c.failStrictEvents(perr)
// Skip malformed frames; protocol resync is the server's
// responsibility.
continue
@@ -461,7 +531,7 @@ func (c *Client) dispatch(msg ahptypes.JsonRpcMessage) {
func (c *Client) deliver(id uint64, r pendingResult) {
c.pendingMu.Lock()
- ch, ok := c.pending[id]
+ request, ok := c.pending[id]
if ok {
delete(c.pending, id)
}
@@ -469,6 +539,10 @@ func (c *Client) deliver(id uint64, r pendingResult) {
if !ok {
return
}
+ if r.err == nil && request.onResult != nil {
+ request.onResult(r.value)
+ }
+ ch := request.result
ch <- r
close(ch)
}
@@ -478,40 +552,58 @@ func (c *Client) handleNotification(n ahptypes.JsonRpcNotification) {
case "action":
var env ahptypes.ActionEnvelope
if err := json.Unmarshal(n.Params, &env); err != nil {
+ c.failStrictEvents(&TransportError{Kind: "protocol", Err: err})
return
}
c.fanOut(env.Channel, SubscriptionEventAction{Envelope: env})
case "root/sessionAdded":
var p ahptypes.SessionAddedParams
if err := json.Unmarshal(n.Params, &p); err != nil {
+ c.failStrictEvents(&TransportError{Kind: "protocol", Err: err})
return
}
c.fanOut(p.Channel, SubscriptionEventSessionAdded{Params: p})
case "root/sessionRemoved":
var p ahptypes.SessionRemovedParams
if err := json.Unmarshal(n.Params, &p); err != nil {
+ c.failStrictEvents(&TransportError{Kind: "protocol", Err: err})
return
}
c.fanOut(p.Channel, SubscriptionEventSessionRemoved{Params: p})
case "root/sessionSummaryChanged":
var p ahptypes.SessionSummaryChangedParams
if err := json.Unmarshal(n.Params, &p); err != nil {
+ c.failStrictEvents(&TransportError{Kind: "protocol", Err: err})
return
}
c.fanOut(p.Channel, SubscriptionEventSessionSummaryChanged{Params: p})
case "auth/required":
var p ahptypes.AuthRequiredParams
if err := json.Unmarshal(n.Params, &p); err != nil {
+ c.failStrictEvents(&TransportError{Kind: "protocol", Err: err})
return
}
c.fanOut(p.Channel, SubscriptionEventAuthRequired{Params: p})
}
}
+func (c *Client) failStrictEvents(err error) {
+ c.subscriptionsMu.Lock()
+ listeners := append([]*EventStream(nil), c.eventListeners...)
+ for _, list := range c.resourceEventListeners {
+ listeners = append(listeners, list...)
+ }
+ c.subscriptionsMu.Unlock()
+ for _, listener := range listeners {
+ listener.fail(err)
+ }
+}
+
func (c *Client) fanOut(channel string, ev SubscriptionEvent) {
c.subscriptionsMu.Lock()
subs := append([]*Subscription(nil), c.subscriptions[channel]...)
listeners := append([]*EventStream(nil), c.eventListeners...)
+ listeners = append(listeners, c.resourceEventListeners[channel]...)
c.subscriptionsMu.Unlock()
for _, s := range subs {
s.trySend(ev)
@@ -638,6 +730,10 @@ func (c *Client) handleServerRequest(req ahptypes.JsonRpcRequest) {
// Request sends a JSON-RPC request and decodes the response into out.
// If out is nil, the result is discarded.
func (c *Client) Request(ctx context.Context, method string, params any, out any) error {
+ return c.requestWithLateResult(ctx, method, params, out, nil, nil)
+}
+
+func (c *Client) requestWithLateResult(ctx context.Context, method string, params any, out any, onLate func(json.RawMessage), onResult func(json.RawMessage)) error {
select {
case <-c.done:
return ErrShutdown
@@ -651,7 +747,8 @@ func (c *Client) Request(ctx context.Context, method string, params any, out any
resultCh := make(chan pendingResult, 1)
c.pendingMu.Lock()
- c.pending[id] = resultCh
+ request := pendingRequest{result: resultCh, onResult: onResult}
+ c.pending[id] = request
c.pendingMu.Unlock()
req := ahptypes.JsonRpcMessage{Request: &ahptypes.JsonRpcRequest{
@@ -661,9 +758,11 @@ func (c *Client) Request(ctx context.Context, method string, params any, out any
Params: rawParams,
}}
if err := c.send(ctx, req); err != nil {
- c.pendingMu.Lock()
- delete(c.pending, id)
- c.pendingMu.Unlock()
+ if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) {
+ c.abandonRequest(id, resultCh, onLate)
+ } else {
+ c.abandonRequest(id, resultCh, nil)
+ }
return err
}
@@ -688,15 +787,36 @@ func (c *Client) Request(ctx context.Context, method string, params any, out any
}
return nil
case <-ctx.Done():
- c.pendingMu.Lock()
- delete(c.pending, id)
- c.pendingMu.Unlock()
+ c.abandonRequest(id, resultCh, onLate)
return ctx.Err()
case <-c.done:
return ErrShutdown
}
}
+func (c *Client) abandonRequest(id uint64, resultCh <-chan pendingResult, onLate func(json.RawMessage)) {
+ if onLate == nil {
+ c.pendingMu.Lock()
+ delete(c.pending, id)
+ c.pendingMu.Unlock()
+ return
+ }
+ // Send may already have reached the host when its context expires. Retain
+ // correlation until response or shutdown to release any late-created child.
+ go func() {
+ select {
+ case result, ok := <-resultCh:
+ if ok && result.err == nil {
+ onLate(result.value)
+ }
+ case <-c.done:
+ c.pendingMu.Lock()
+ delete(c.pending, id)
+ c.pendingMu.Unlock()
+ }
+ }()
+}
+
// Notify sends a JSON-RPC notification (fire-and-forget).
func (c *Client) Notify(ctx context.Context, method string, params any) error {
select {
@@ -774,6 +894,15 @@ func (c *Client) Initialize(ctx context.Context, clientID string, protocolVersio
if err := c.Request(ctx, "initialize", params, &out); err != nil {
return nil, err
}
+ c.tcpMu.Lock()
+ c.tcpClientID = clientID
+ c.tcpCapability = nil
+ if out.TcpConnections != nil {
+ capability := *out.TcpConnections
+ capability.Encodings = append([]ahptypes.TcpDataEncoding(nil), capability.Encodings...)
+ c.tcpCapability = &capability
+ }
+ c.tcpMu.Unlock()
return &out, nil
}
@@ -1003,9 +1132,87 @@ func (c *Client) SessionConfigCompletions(ctx context.Context, params ahptypes.S
// inbound event from this client, tagged with the channel URI it was
// scoped to. Multiple streams may exist concurrently.
func (c *Client) Events() *EventStream {
+ return c.events(false)
+}
+
+// EventsStrict attaches a bounded global receiver that permanently closes on
+// overflow. Check its Err after draining Events; a *SubscriptionLagError means
+// affected TCP channels must be reset/unsubscribed, never resumed past the gap.
+// Malformed input terminates strict receivers with a *TransportError instead
+// of silently dropping the frame or notification.
+//
+// Attach before requesting subscribe(create), then retain the receiver and
+// filter by the returned child URI. Reconnect replay and credits remain the
+// caller's responsibility.
+func (c *Client) EventsStrict() *EventStream {
+ return c.events(true)
+}
+
+func (c *Client) events(strict bool, resource ...string) *EventStream {
c.subscriptionsMu.Lock()
defer c.subscriptionsMu.Unlock()
- s := &EventStream{events: make(chan ClientEvent, c.cfg.SubscriptionBuffer)}
- c.eventListeners = append(c.eventListeners, s)
+ s := &EventStream{events: make(chan ClientEvent, c.cfg.SubscriptionBuffer), strict: strict}
+ if len(resource) != 0 {
+ s.resource = &resource[0]
+ }
+ if strict {
+ select {
+ case <-c.done:
+ s.Close()
+ return s
+ default:
+ }
+ s.onClose = func() {
+ c.subscriptionsMu.Lock()
+ defer c.subscriptionsMu.Unlock()
+ c.removeEventStream(s)
+ }
+ }
+ c.addEventStream(s)
return s
}
+
+func (c *Client) addEventStream(s *EventStream) {
+ if s.resource == nil {
+ c.eventListeners = append(c.eventListeners, s)
+ } else {
+ if c.resourceEventListeners == nil {
+ c.resourceEventListeners = make(map[string][]*EventStream)
+ }
+ c.resourceEventListeners[*s.resource] = append(c.resourceEventListeners[*s.resource], s)
+ }
+}
+
+func (c *Client) removeEventStream(s *EventStream) {
+ list := c.eventListeners
+ if s.resource != nil {
+ list = c.resourceEventListeners[*s.resource]
+ }
+ for i, listener := range list {
+ if listener == s {
+ copy(list[i:], list[i+1:])
+ list[len(list)-1] = nil
+ list = list[:len(list)-1]
+ break
+ }
+ }
+ if s.resource == nil {
+ c.eventListeners = list
+ } else if len(list) == 0 {
+ delete(c.resourceEventListeners, *s.resource)
+ } else {
+ c.resourceEventListeners[*s.resource] = list
+ }
+}
+
+func (c *Client) bindEventStream(s *EventStream, resource string) {
+ c.subscriptionsMu.Lock()
+ defer c.subscriptionsMu.Unlock()
+ s.closeMu.Lock()
+ defer s.closeMu.Unlock()
+ c.removeEventStream(s)
+ s.resource = &resource
+ if !s.closed {
+ c.addEventStream(s)
+ }
+}
diff --git a/clients/go/ahp/client_test.go b/clients/go/ahp/client_test.go
index 485adaf12..d6943da3b 100644
--- a/clients/go/ahp/client_test.go
+++ b/clients/go/ahp/client_test.go
@@ -680,6 +680,314 @@ func TestClientAutomationCatalogueAction(t *testing.T) {
}
}
+func TestStrictEventsCaptureAtomicTCPCreateFirstAction(t *testing.T) {
+ clientSide, serverSide := newMemTransportPair()
+ ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancel()
+ client, err := Connect(ctx, clientSide, DefaultConfig())
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer client.Shutdown(context.Background())
+ events := client.EventsStrict()
+ defer events.Close()
+ state := tcpTestState()
+ snapshotResult, err := json.Marshal(ahptypes.SubscribeResult{Snapshot: &ahptypes.Snapshot{
+ Resource: "ahp-tcp:/created", State: ahptypes.SnapshotState{Tcp: &state},
+ }})
+ if err != nil {
+ t.Fatal(err)
+ }
+ serverErr := make(chan error, 1)
+ go func() {
+ serverErr <- func() error {
+ message, err := serverSide.Recv(ctx)
+ if err != nil {
+ return err
+ }
+ parsed, err := message.IntoParsed()
+ if err != nil {
+ return err
+ }
+ if parsed.Request == nil || parsed.Request.Method != "subscribe" {
+ return errors.New("expected subscribe request")
+ }
+ var params ahptypes.SubscribeParams
+ if err := json.Unmarshal(parsed.Request.Params, ¶ms); err != nil {
+ return err
+ }
+ if params.Channel != "ahp-session:/s1" || params.Create == nil || params.Create.Type != "tcpConnection" {
+ return fmt.Errorf("unexpected create parameters: %+v", params)
+ }
+ if err := serverSide.Send(ctx, NewTextMessage(fmt.Sprintf(
+ `{"jsonrpc":"2.0","id":%d,"result":%s}`, parsed.Request.ID, snapshotResult,
+ ))); err != nil {
+ return err
+ }
+ if err := serverSide.Send(ctx, NewTextMessage(`{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-tcp:/created","serverSeq":1,"origin":null,"action":{"type":"tcp/data","offset":0,"data":"AA=="}}}`)); err != nil {
+ return err
+ }
+ message, err = serverSide.Recv(ctx)
+ if err != nil {
+ return err
+ }
+ parsed, err = message.IntoParsed()
+ if err != nil {
+ return err
+ }
+ if parsed.Request == nil || parsed.Request.Method != "ping" {
+ return errors.New("expected ping barrier")
+ }
+ return serverSide.Send(ctx, NewTextMessage(fmt.Sprintf(`{"jsonrpc":"2.0","id":%d,"result":null}`, parsed.Request.ID)))
+ }()
+ }()
+ var result ahptypes.SubscribeResult
+ err = client.Request(ctx, "subscribe", ahptypes.SubscribeParams{
+ Channel: "ahp-session:/s1",
+ Create: &ahptypes.TcpConnectionSubscription{
+ Type: "tcpConnection", Host: "localhost", Port: 3000,
+ Encoding: ahptypes.TcpDataEncodingBase64, ReceiveWindowBytes: 8, MaximumChunkSize: 6,
+ },
+ }, &result)
+ if err != nil {
+ t.Fatal(err)
+ }
+ // Ensure the first action has been processed before draining this receiver.
+ if err := client.Ping(ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case event, ok := <-events.Events():
+ if !ok {
+ t.Fatalf("receiver closed: %v", events.Err())
+ }
+ if result.Snapshot == nil || result.Snapshot.State.Tcp == nil || event.Channel != result.Snapshot.Resource {
+ t.Fatalf("event did not match created TCP snapshot: %+v", event)
+ }
+ action, ok := event.Event.(SubscriptionEventAction)
+ if !ok || action.Envelope.ServerSeq != 1 {
+ t.Fatalf("unexpected first action: %+v", event)
+ }
+ outcome, err := ApplyActionToTCP(result.Snapshot.State.Tcp, action.Envelope.Action)
+ if err != nil || outcome != ReduceOutcomeApplied || result.Snapshot.State.Tcp.Output.ReceivedBytes != 1 {
+ t.Fatalf("first action reduction: %v, %v", outcome, err)
+ }
+ case <-ctx.Done():
+ t.Fatal("missing first action")
+ }
+ if err := <-serverErr; err != nil {
+ t.Fatal(err)
+ }
+}
+
+func TestStrictEventsOverflowIsTerminalAndOtherReceiversContinue(t *testing.T) {
+ clientSide, _ := newMemTransportPair()
+ ctx := context.Background()
+ config := DefaultConfig()
+ config.SubscriptionBuffer = 1
+ client, err := Connect(ctx, clientSide, config)
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer client.Shutdown(ctx)
+ strict := client.EventsStrict()
+ ordinary := client.Events()
+ send := func(seq int64) {
+ client.fanOut("ahp-tcp:/created", SubscriptionEventAction{Envelope: ahptypes.ActionEnvelope{ServerSeq: seq}})
+ }
+ send(1)
+ send(2)
+ var lag *SubscriptionLagError
+ if !errors.As(strict.Err(), &lag) || lag.Capacity != 1 {
+ t.Fatalf("expected typed overflow error, got %v", strict.Err())
+ }
+ if event := <-strict.Events(); event.Event.(SubscriptionEventAction).Envelope.ServerSeq != 1 {
+ t.Fatal("strict receiver lost its valid buffered prefix")
+ }
+ if event := <-ordinary.Events(); event.Event.(SubscriptionEventAction).Envelope.ServerSeq != 1 {
+ t.Fatal("ordinary receiver's drop-newest behavior changed")
+ }
+ send(3)
+ select {
+ case _, open := <-strict.Events():
+ if open {
+ t.Fatal("strict receiver resumed after overflow")
+ }
+ default:
+ t.Fatal("strict receiver did not terminate")
+ }
+ if event := <-ordinary.Events(); event.Event.(SubscriptionEventAction).Envelope.ServerSeq != 3 {
+ t.Fatal("ordinary receiver did not continue")
+ }
+ strict.Close()
+ if !errors.As(strict.Err(), &lag) {
+ t.Fatal("close erased the lag error")
+ }
+ if err := client.Shutdown(ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case _, open := <-client.EventsStrict().Events():
+ if open {
+ t.Fatal("receiver created after shutdown remained open")
+ }
+ default:
+ t.Fatal("receiver created after shutdown did not terminate")
+ }
+}
+
+func TestStrictEventsDecodeLossIsTerminalButFutureActionsAreAllowed(t *testing.T) {
+ for _, test := range []struct {
+ name, wire string
+ fails bool
+ }{
+ {"malformedJSON", "{", true},
+ {"invalidEnvelope", `{"jsonrpc":"2.0","method":"action","params":{"channel":42,"serverSeq":1,"action":{"type":"tcp/dataEof","finalOffset":0}}}`, true},
+ {"invalidEOF", `{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-tcp:/created","serverSeq":1,"action":{"type":"tcp/dataEof","finalOffset":"bad"}}}`, true},
+ {"invalidCredit", `{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-tcp:/created","serverSeq":1,"action":{"type":"tcp/inputConsumed","consumedBytes":"bad"}}}`, true},
+ {"futureAction", `{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-tcp:/created","serverSeq":1,"action":{"type":"tcp/futureControl"}}}`, false},
+ } {
+ t.Run(test.name, func(t *testing.T) {
+ ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancel()
+ clientSide, serverSide := newMemTransportPair()
+ client, err := Connect(ctx, clientSide, DefaultConfig())
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer client.Shutdown(context.Background())
+ strict, ordinary := client.EventsStrict(), client.Events()
+ serverErr := make(chan error, 1)
+ go func() {
+ serverErr <- func() error {
+ message, err := serverSide.Recv(ctx)
+ if err != nil {
+ return err
+ }
+ request, err := message.IntoParsed()
+ if err != nil {
+ return err
+ }
+ if request.Request == nil || request.Request.Method != "ping" {
+ return errors.New("expected ping barrier")
+ }
+ for _, wire := range []string{
+ test.wire,
+ `{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-tcp:/created","serverSeq":2,"action":{"type":"tcp/data","offset":0,"data":"AA=="}}}`,
+ fmt.Sprintf(`{"jsonrpc":"2.0","id":%d,"result":null}`, request.Request.ID),
+ } {
+ if err := serverSide.Send(ctx, NewTextMessage(wire)); err != nil {
+ return err
+ }
+ }
+ return nil
+ }()
+ }()
+ if err := client.Ping(ctx); err != nil {
+ t.Fatal(err)
+ }
+ receive := func(stream *EventStream, seq int64) {
+ t.Helper()
+ select {
+ case event, ok := <-stream.Events():
+ action, isAction := event.Event.(SubscriptionEventAction)
+ if !ok || !isAction || action.Envelope.ServerSeq != seq {
+ t.Fatalf("expected action %d, got %+v (error %v)", seq, event, stream.Err())
+ }
+ case <-ctx.Done():
+ t.Fatal("missing event")
+ }
+ }
+ if test.fails {
+ var protocolError *TransportError
+ if !errors.As(strict.Err(), &protocolError) || protocolError.Kind != "protocol" {
+ t.Fatalf("expected typed protocol error, got %v", strict.Err())
+ }
+ select {
+ case _, open := <-strict.Events():
+ if open {
+ t.Fatal("strict receiver resumed after decode loss")
+ }
+ default:
+ t.Fatal("strict receiver did not terminate")
+ }
+ } else {
+ receive(strict, 1)
+ receive(strict, 2)
+ receive(ordinary, 1)
+ if strict.Err() != nil {
+ t.Fatal(strict.Err())
+ }
+ }
+ receive(ordinary, 2)
+ if err := <-serverErr; err != nil {
+ t.Fatal(err)
+ }
+ })
+ }
+}
+
+func TestStrictEventsUnregisterAndWakePendingReceive(t *testing.T) {
+ clientSide, _ := newMemTransportPair()
+ ctx, cancel := context.WithTimeout(context.Background(), 3*time.Second)
+ defer cancel()
+ config := DefaultConfig()
+ config.SubscriptionBuffer = 1
+ client, err := Connect(ctx, clientSide, config)
+ if err != nil {
+ t.Fatal(err)
+ }
+ defer client.Shutdown(context.Background())
+ ordinary := client.Events()
+ assertOnlyOrdinary := func() {
+ t.Helper()
+ client.subscriptionsMu.Lock()
+ defer client.subscriptionsMu.Unlock()
+ if len(client.eventListeners) != 1 || client.eventListeners[0] != ordinary {
+ t.Fatal("strict receiver remained registered or affected another receiver")
+ }
+ }
+ strict := client.EventsStrict()
+ ready := make(chan struct{})
+ received := make(chan bool, 1)
+ go func() {
+ close(ready)
+ _, open := <-strict.Events()
+ received <- open
+ }()
+ <-ready
+ client.failStrictEvents(&TransportError{Kind: "protocol", Err: errors.New("malformed frame")})
+ assertOnlyOrdinary()
+ select {
+ case open := <-received:
+ if open {
+ t.Fatal("failed receiver delivered an event")
+ }
+ case <-ctx.Done():
+ t.Fatal("failure did not wake the pending receive")
+ }
+ var protocolError *TransportError
+ if !errors.As(strict.Err(), &protocolError) || protocolError.Kind != "protocol" {
+ t.Fatalf("missing protocol error: %v", strict.Err())
+ }
+ closed := client.EventsStrict()
+ closed.Close()
+ closed.Close()
+ assertOnlyOrdinary()
+ overflowed := client.EventsStrict()
+ for i := int64(1); i <= 2; i++ {
+ client.fanOut("ahp-tcp:/created", SubscriptionEventAction{Envelope: ahptypes.ActionEnvelope{ServerSeq: i}})
+ }
+ var lag *SubscriptionLagError
+ if !errors.As(overflowed.Err(), &lag) {
+ t.Fatalf("missing overflow error: %v", overflowed.Err())
+ }
+ assertOnlyOrdinary()
+ if event := <-ordinary.Events(); event.Event.(SubscriptionEventAction).Envelope.ServerSeq != 1 {
+ t.Fatal("ordinary receiver was affected")
+ }
+}
+
// TestClientShutdownFailsInFlightRequest confirms a Shutdown unblocks
// any pending request with ErrShutdown.
func TestClientShutdownFailsInFlightRequest(t *testing.T) {
diff --git a/clients/go/ahp/error.go b/clients/go/ahp/error.go
index ec97c2930..01eeb09f3 100644
--- a/clients/go/ahp/error.go
+++ b/clients/go/ahp/error.go
@@ -54,6 +54,16 @@ var ErrShutdown = errors.New("ahp: client shut down")
// caller should resubscribe to recover.
var ErrSequenceGap = errors.New("ahp: sequence gap detected; resubscribe required")
+// SubscriptionLagError reports overflow of a strict event receiver.
+// That receiver is permanently closed and will never resume past the gap.
+type SubscriptionLagError struct {
+ Capacity int
+}
+
+func (e *SubscriptionLagError) Error() string {
+ return fmt.Sprintf("ahp: subscription lag: event buffer capacity %d exceeded; receiver terminated", e.Capacity)
+}
+
// TransportError wraps any error produced by an underlying [Transport]
// implementation so that callers can distinguish transport faults from
// protocol-level RPC errors via [errors.As].
diff --git a/clients/go/ahp/hosts/hosts.go b/clients/go/ahp/hosts/hosts.go
index 65396c4c9..428f3e44d 100644
--- a/clients/go/ahp/hosts/hosts.go
+++ b/clients/go/ahp/hosts/hosts.go
@@ -19,8 +19,10 @@ import (
"encoding/json"
"errors"
"fmt"
+ "log"
"os"
"path/filepath"
+ "strings"
"sync"
"time"
@@ -366,9 +368,13 @@ func (h *HostClientHandle) HostID() HostID { return h.host.id }
// Client returns the live [ahp.Client] for the current generation, or
// [ErrHostReconnected] if the connection has since been replaced.
+// For managed TCP ownership use OpenTCPConnection on this handle instead.
func (h *HostClientHandle) Client() (*ahp.Client, error) {
h.host.mu.RLock()
defer h.host.mu.RUnlock()
+ if h.host.removed {
+ return nil, ErrUnknownHost
+ }
if h.host.generation != h.generation {
return nil, ErrHostReconnected
}
@@ -378,6 +384,34 @@ func (h *HostClientHandle) Client() (*ahp.Client, error) {
return h.host.client, nil
}
+// OpenTCPConnection creates a stream retained across the host's reconnects.
+// The stream remains usable after this generation-checked handle becomes stale.
+func (h *HostClientHandle) OpenTCPConnection(ctx context.Context, session string, create ahptypes.TcpConnectionSubscription) (*ahp.TCPConnection, error) {
+ client, err := h.Client()
+ if err != nil {
+ return nil, err
+ }
+ connection, err := client.OpenTCPConnection(ctx, session, create)
+ if err != nil {
+ return nil, err
+ }
+ h.host.mu.Lock()
+ if h.host.generation != h.generation || h.host.client != client || h.host.state.Kind != HostStateConnected || h.host.removed {
+ h.host.mu.Unlock()
+ return nil, errors.Join(ErrHostReconnected, connection.Dispose(ctx))
+ }
+ live := h.host.tcpConnections[:0]
+ for _, retained := range h.host.tcpConnections {
+ if !retained.IsClosed() {
+ live = append(live, retained)
+ }
+ }
+ clear(h.host.tcpConnections[len(live):])
+ h.host.tcpConnections = append(live, connection)
+ h.host.mu.Unlock()
+ return connection, nil
+}
+
// ─── Errors ────────────────────────────────────────────────────────────
// ErrHostReconnected is returned by [HostClientHandle.Client] when
@@ -401,22 +435,26 @@ var ErrDuplicateHost = errors.New("hosts: host id already registered")
// hostState is the per-host bookkeeping the multi-host runtime owns.
type hostState struct {
- id HostID
- label string
- cfg HostConfig
- mu sync.RWMutex
- client *ahp.Client
- state HostState
- clientID string
- protoVer string
- automations *ahptypes.AutomationCapabilities
- agents []ahptypes.AgentInfo
- sessions []ahptypes.SessionSummary
- terminals []ahptypes.TerminalInfo
- updatedAt time.Time
- generation uint64
- cancel context.CancelFunc
- supervised sync.WaitGroup
+ id HostID
+ label string
+ cfg HostConfig
+ mu sync.RWMutex
+ client *ahp.Client
+ resumeClient *ahp.Client
+ state HostState
+ clientID string
+ protoVer string
+ automations *ahptypes.AutomationCapabilities
+ agents []ahptypes.AgentInfo
+ sessions []ahptypes.SessionSummary
+ terminals []ahptypes.TerminalInfo
+ updatedAt time.Time
+ generation uint64
+ cancel context.CancelFunc
+ supervised sync.WaitGroup
+ tcpConnections []*ahp.TCPConnection
+ serverSeq int64
+ removed bool
}
// MultiHostClient is the public multi-host registry + reconnect
@@ -565,42 +603,123 @@ func (m *MultiHostClient) openHost(ctx context.Context, hs *hostState) error {
if err != nil {
return fmt.Errorf("hosts: connect: %w", err)
}
+ hs.mu.Lock()
+ previous := hs.resumeClient
+ hs.resumeClient = client
+ hs.mu.Unlock()
+ if previous != nil {
+ if err := client.RestoreReconnectState(previous); err != nil {
+ return errors.Join(err, client.Shutdown(ctx))
+ }
+ }
- result, err := client.Initialize(ctx, hs.clientID, hs.cfg.ProtocolVersions, hs.cfg.InitialSubscriptions)
+ events := client.Events()
+ hs.mu.RLock()
+ retained := append([]*ahp.TCPConnection(nil), hs.tcpConnections...)
+ serverSeq := hs.serverSeq
+ hs.mu.RUnlock()
+ live := retained[:0]
+ for _, connection := range retained {
+ if !connection.IsClosed() {
+ live = append(live, connection)
+ }
+ }
+ subscriptions := make([]string, 0, len(hs.cfg.InitialSubscriptions))
+ for _, resource := range hs.cfg.InitialSubscriptions {
+ if !strings.HasPrefix(resource, "ahp-tcp:") {
+ subscriptions = append(subscriptions, resource)
+ }
+ }
+ var result *ahptypes.InitializeResult
+ var replay *ahptypes.ReconnectResult
+ if len(live) != 0 {
+ replay, err = client.ReconnectTCPConnections(ctx, ahptypes.ReconnectParams{
+ ClientId: hs.clientID, LastSeenServerSeq: serverSeq, Subscriptions: subscriptions,
+ }, live)
+ var rpc *ahp.RPCError
+ transportClosed := false
+ select {
+ case <-client.Done():
+ transportClosed = true
+ default:
+ }
+ if errors.As(err, &rpc) && !transportClosed {
+ for _, connection := range live {
+ if cleanup := connection.Dispose(ctx); cleanup != nil {
+ log.Printf("hosts: TCP fallback cleanup: %v", cleanup)
+ }
+ }
+ result, err = client.Initialize(ctx, hs.clientID, hs.cfg.ProtocolVersions, subscriptions)
+ }
+ } else {
+ result, err = client.Initialize(ctx, hs.clientID, hs.cfg.ProtocolVersions, subscriptions)
+ }
if err != nil {
- _ = client.Shutdown(ctx)
- return fmt.Errorf("hosts: initialize: %w", err)
+ events.Close()
+ return errors.Join(fmt.Errorf("hosts: handshake: %w", err), client.ShutdownPreservingTCP(ctx))
}
hs.mu.Lock()
+ if hs.removed {
+ hs.mu.Unlock()
+ events.Close()
+ return errors.Join(ErrUnknownHost, client.Shutdown(ctx))
+ }
hs.client = client
- hs.protoVer = result.ProtocolVersion
- hs.automations = cloneAutomationCapabilities(result.Automations)
+ if result != nil {
+ hs.protoVer = result.ProtocolVersion
+ hs.automations = cloneAutomationCapabilities(result.Automations)
+ hs.serverSeq = result.ServerSeq
+ }
hs.generation++
hs.mu.Unlock()
+ if replay != nil {
+ if actions, ok := replay.Value.(*ahptypes.ReconnectReplayResult); ok {
+ for _, envelope := range actions.Actions {
+ if !strings.HasPrefix(envelope.Channel, "ahp-tcp:") && envelope.ServerSeq <= serverSeq {
+ continue
+ }
+ m.publishHostEvent(hs, ahp.ClientEvent{Channel: envelope.Channel, Event: ahp.SubscriptionEventAction{Envelope: envelope}})
+ }
+ } else if snapshots, ok := replay.Value.(*ahptypes.ReconnectSnapshotResult); ok {
+ hs.mu.Lock()
+ for _, snapshot := range snapshots.Snapshots {
+ hs.serverSeq = max(hs.serverSeq, snapshot.FromSeq)
+ }
+ hs.mu.Unlock()
+ }
+ }
m.setHostState(hs, HostState{Kind: HostStateConnected})
// Fan inbound events out to subscribers.
- go m.pumpEvents(hs, client)
+ go m.pumpEvents(hs, events)
return nil
}
// pumpEvents drains the per-host [ahp.Client.Events] stream and
// re-emits each event tagged with the host id.
-func (m *MultiHostClient) pumpEvents(hs *hostState, client *ahp.Client) {
- stream := client.Events()
+func (m *MultiHostClient) pumpEvents(hs *hostState, stream *ahp.EventStream) {
defer stream.Close()
for ev := range stream.Events() {
- m.subMu.Lock()
- subs := append([]chan HostSubscriptionEvent(nil), m.subs...)
- m.subMu.Unlock()
- out := HostSubscriptionEvent{HostID: hs.id, Channel: ev.Channel, Event: ev.Event}
- for _, ch := range subs {
- select {
- case ch <- out:
- default:
- }
+ m.publishHostEvent(hs, ev)
+ }
+}
+
+func (m *MultiHostClient) publishHostEvent(hs *hostState, ev ahp.ClientEvent) {
+ if action, ok := ev.Event.(ahp.SubscriptionEventAction); ok {
+ hs.mu.Lock()
+ hs.serverSeq = max(hs.serverSeq, action.Envelope.ServerSeq)
+ hs.mu.Unlock()
+ }
+ m.subMu.Lock()
+ subs := append([]chan HostSubscriptionEvent(nil), m.subs...)
+ m.subMu.Unlock()
+ out := HostSubscriptionEvent{HostID: hs.id, Channel: ev.Channel, Event: ev.Event}
+ for _, ch := range subs {
+ select {
+ case ch <- out:
+ default:
}
}
}
@@ -633,6 +752,7 @@ func (m *MultiHostClient) supervise(ctx context.Context, hs *hostState) {
default:
}
if policy.IsDisabled() {
+ m.disposeTCPConnections(ctx, hs)
m.setHostState(hs, HostState{Kind: HostStateFailed, Err: errors.New("hosts: transport closed and reconnect disabled")})
return
}
@@ -640,7 +760,7 @@ func (m *MultiHostClient) supervise(ctx context.Context, hs *hostState) {
// Ensure the old client is fully torn down before opening a
// replacement (a no-op if Done fired from Shutdown already).
shutdownCtx, shutdownCancel := context.WithTimeout(context.Background(), 2*time.Second)
- _ = client.Shutdown(shutdownCtx)
+ _ = client.ShutdownPreservingTCP(shutdownCtx)
shutdownCancel()
var attempt uint32 = 1
@@ -661,9 +781,23 @@ func (m *MultiHostClient) supervise(ctx context.Context, hs *hostState) {
}
attempt++
if policy.MaxAttempts > 0 && attempt > policy.MaxAttempts {
+ m.disposeTCPConnections(ctx, hs)
m.setHostState(hs, HostState{Kind: HostStateFailed, Err: fmt.Errorf("hosts: exceeded %d reconnect attempts", policy.MaxAttempts)})
return
}
+
+ }
+ }
+}
+
+func (m *MultiHostClient) disposeTCPConnections(ctx context.Context, hs *hostState) {
+ hs.mu.Lock()
+ connections := hs.tcpConnections
+ hs.tcpConnections = nil
+ hs.mu.Unlock()
+ for _, connection := range connections {
+ if err := connection.Dispose(ctx); err != nil {
+ log.Printf("hosts: TCP disposal: %v", err)
}
}
}
@@ -782,6 +916,10 @@ func (m *MultiHostClient) RemoveHost(ctx context.Context, id HostID) error {
return ErrUnknownHost
}
hs.cancel()
+ hs.mu.Lock()
+ hs.removed = true
+ hs.mu.Unlock()
+ m.disposeTCPConnections(ctx, hs)
hs.mu.RLock()
client := hs.client
hs.mu.RUnlock()
@@ -833,6 +971,10 @@ func (m *MultiHostClient) Shutdown(ctx context.Context) error {
m.mu.Unlock()
for _, hs := range hosts {
hs.cancel()
+ hs.mu.Lock()
+ hs.removed = true
+ hs.mu.Unlock()
+ m.disposeTCPConnections(ctx, hs)
hs.mu.RLock()
client := hs.client
hs.mu.RUnlock()
diff --git a/clients/go/ahp/hosts/hosts_test.go b/clients/go/ahp/hosts/hosts_test.go
index 8b35fd92f..2f8e29d52 100644
--- a/clients/go/ahp/hosts/hosts_test.go
+++ b/clients/go/ahp/hosts/hosts_test.go
@@ -3,6 +3,8 @@ package hosts
import (
"context"
"encoding/json"
+ "io"
+ "strings"
"sync"
"testing"
"time"
@@ -100,6 +102,280 @@ func runFakeServerWithInitializeResult(t *testing.T, serverSide *fakeTransport,
}
}
+func TestManagedTCPReconnectAndShutdown(t *testing.T) {
+ for _, mode := range []string{"replay", "retry", "snapshot", "missing", "refused"} {
+ t.Run(mode, func(t *testing.T) {
+ ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
+ defer cancel()
+ multi := NewMultiHostClient()
+ events := multi.Subscriptions()
+ defer multi.Shutdown(context.Background())
+ servers := make(chan *fakeTransport, 4)
+ reconnects := make(chan ahptypes.ReconnectParams, 4)
+ dispatches := make(chan ahptypes.DispatchActionParams, 32)
+ attempt := 0
+ snapshot := func(resource string) map[string]any {
+ direction := map[string]any{"windowBytes": 4, "maximumChunkSize": 3, "receivedBytes": 0, "consumedBytes": 0, "eofAtBytes": nil}
+ return map[string]any{"snapshot": map[string]any{"resource": resource, "fromSeq": 10,
+ "state": map[string]any{"type": "tcp", "session": "ahp-session:/s1", "target": map[string]any{"host": "localhost", "port": 3000},
+ "encoding": "base64", "input": direction, "output": direction, "clientClosed": false, "hostClosed": false, "reset": nil}}}
+ }
+ cfg := NewHostConfig("tcp", "TCP", func(_ context.Context, _ HostID) (ahp.Transport, error) {
+ attempt++
+ current := attempt
+ client, server := newFakePair()
+ servers <- server
+ go func() {
+ send := func(value any) {
+ data, err := json.Marshal(value)
+ if err != nil {
+ t.Error(err)
+ return
+ }
+ if err := server.Send(ctx, ahp.NewTextMessage(string(data))); err != nil && ctx.Err() == nil {
+ t.Error(err)
+ }
+ }
+ for {
+ frame, err := server.Recv(ctx)
+ if err != nil {
+ return
+ }
+ message, err := frame.IntoParsed()
+ if err != nil {
+ t.Error(err)
+ return
+ }
+ if message.Notification != nil {
+ if message.Notification.Method == "dispatchAction" {
+ var params ahptypes.DispatchActionParams
+ if err := json.Unmarshal(message.Notification.Params, ¶ms); err != nil {
+ t.Error(err)
+ return
+ }
+ dispatches <- params
+ }
+ continue
+ }
+ if message.Request == nil {
+ continue
+ }
+ req := message.Request
+ var result any
+ switch req.Method {
+ case "initialize":
+ var params ahptypes.InitializeParams
+ if err := json.Unmarshal(req.Params, ¶ms); err != nil {
+ t.Error(err)
+ return
+ }
+ for _, resource := range params.InitialSubscriptions {
+ if strings.HasPrefix(resource, "ahp-tcp:") {
+ t.Error("TCP entered initialize fallback")
+ }
+ }
+ result = map[string]any{"protocolVersion": ahptypes.ProtocolVersion, "serverSeq": 10, "snapshots": []any{}, "tcpConnections": map[string]any{"encodings": []string{"base64"}}}
+ case "subscribe":
+ resource := "ahp-tcp:/owned"
+ if current > 1 {
+ resource = "ahp-tcp:/second"
+ }
+ result = snapshot(resource)
+ case "reconnect":
+ var params ahptypes.ReconnectParams
+ if err := json.Unmarshal(req.Params, ¶ms); err != nil {
+ t.Error(err)
+ return
+ }
+ reconnects <- params
+ if mode == "retry" && current == 2 {
+ server.Close(ctx)
+ return
+ }
+ if mode == "refused" {
+ send(map[string]any{"jsonrpc": "2.0", "id": req.ID, "error": map[string]any{"code": -32000, "message": "replay expired"}})
+ continue
+ }
+ result = map[string]any{"type": "replay", "actions": []any{}, "missing": []string{}}
+ if mode == "replay" || mode == "retry" {
+ result = map[string]any{"type": "replay", "missing": []string{}, "actions": []any{
+ map[string]any{"channel": "ahp-terminal:/t", "serverSeq": 20, "action": map[string]any{"type": "terminal/data", "data": "hello"}},
+ map[string]any{"channel": "ahp-terminal:/t", "serverSeq": 21, "action": map[string]any{"type": "terminal/data", "data": "!"}},
+ }}
+ }
+ if mode == "snapshot" {
+ result = map[string]any{"type": "snapshot", "snapshots": []any{}}
+ }
+ if mode == "missing" {
+ result = map[string]any{"type": "replay", "actions": []any{}, "missing": []string{"ahp-tcp:/owned"}}
+ }
+ default:
+ t.Errorf("unexpected request %s", req.Method)
+ return
+ }
+ send(map[string]any{"jsonrpc": "2.0", "id": req.ID, "result": result})
+ if req.Method == "subscribe" && current == 1 {
+ send(map[string]any{"jsonrpc": "2.0", "method": "action", "params": map[string]any{"channel": "ahp-tcp:/owned", "serverSeq": 11, "action": map[string]any{"type": "tcp/data", "offset": 0, "data": "eA=="}}})
+ }
+ if req.Method == "reconnect" && (mode == "replay" || mode == "retry") {
+ send(map[string]any{"jsonrpc": "2.0", "method": "action", "params": map[string]any{"channel": "ahp-tcp:/owned", "serverSeq": 22, "action": map[string]any{"type": "tcp/dataEof", "finalOffset": 1}}})
+ }
+ }
+ }()
+ return client, nil
+ })
+ cfg.ClientID = "owner"
+ cfg.InitialSubscriptions = []string{ahptypes.RootResourceURI, "ahp-tcp:/must-not-initialize"}
+ cfg.ReconnectPolicy = ReconnectPolicy{MaxAttempts: 2, InitialBackoff: time.Millisecond, MaxBackoff: time.Millisecond, BackoffMultiplier: 1, ResetOnSuccess: true}
+ if _, err := multi.AddHost(ctx, cfg); err != nil {
+ t.Fatal(err)
+ }
+ old, err := multi.ClientHandle(cfg.ID)
+ if err != nil {
+ t.Fatal(err)
+ }
+ create := ahptypes.TcpConnectionSubscription{Type: "tcpConnection", Host: "localhost", Port: 3000, Encoding: ahptypes.TcpDataEncodingBase64, ReceiveWindowBytes: 4, MaximumChunkSize: 3}
+ connection, err := old.OpenTCPConnection(ctx, "ahp-session:/s1", create)
+ if err != nil {
+ t.Fatal(err)
+ }
+ receive := func() ahptypes.DispatchActionParams {
+ select {
+ case p := <-dispatches:
+ return p
+ case <-ctx.Done():
+ t.Fatal("dispatch timeout")
+ return ahptypes.DispatchActionParams{}
+ }
+ }
+ if data, err := connection.Read(ctx); err != nil || string(data) != "x" {
+ t.Fatalf("first data %q: %v", data, err)
+ }
+ credit := receive()
+ if _, err := connection.Write(ctx, []byte("ab")); err != nil {
+ t.Fatal(err)
+ }
+ input := receive()
+ raw, err := old.Client()
+ if err != nil {
+ t.Fatal(err)
+ }
+ ordinary, err := raw.Dispatch(ctx, "ahp-session:/s1", ahptypes.StateAction{Value: &ahptypes.SessionTitleChangedAction{Type: ahptypes.ActionTypeSessionTitleChanged, Title: "ordinary"}})
+ if err != nil {
+ t.Fatal(err)
+ }
+ receive()
+ oldServer := <-servers
+ if err := oldServer.Send(ctx, ahp.NewTextMessage(`{"jsonrpc":"2.0","method":"action","params":{"channel":"ahp-terminal:/t","serverSeq":20,"action":{"type":"terminal/data","data":"hello"}}}`)); err != nil {
+ t.Fatal(err)
+ }
+ for {
+ select {
+ case event := <-events:
+ if event.Channel == "ahp-terminal:/t" {
+ goto ordinaryApplied
+ }
+ case <-ctx.Done():
+ t.Fatal("ordinary event not delivered")
+ }
+ }
+ ordinaryApplied:
+ if err := oldServer.Close(ctx); err != nil {
+ t.Fatal(err)
+ }
+ var reconnect ahptypes.ReconnectParams
+ select {
+ case reconnect = <-reconnects:
+ case <-ctx.Done():
+ t.Fatal("missing TCP reconnect")
+ }
+ found := false
+ for _, resource := range reconnect.Subscriptions {
+ found = found || resource == connection.Resource()
+ }
+ if !found || reconnect.ClientId != "owner" || reconnect.LastSeenServerSeq > 11 {
+ t.Fatalf("unsafe reconnect %+v", reconnect)
+ }
+ var fresh *HostClientHandle
+ for {
+ fresh, err = multi.ClientHandle(cfg.ID)
+ if err == nil && fresh.generation > old.generation && multi.Host(cfg.ID).State.Kind == HostStateConnected {
+ break
+ }
+ select {
+ case <-ctx.Done():
+ t.Fatal("reconnect did not finish")
+ case <-time.After(time.Millisecond):
+ }
+ }
+ if mode == "replay" || mode == "retry" {
+ text := "hello"
+ replayed:
+ for {
+ select {
+ case event := <-events:
+ if action, ok := event.Event.(ahp.SubscriptionEventAction); ok {
+ if data, ok := action.Envelope.Action.Value.(*ahptypes.TerminalDataAction); ok {
+ text += data.Data
+ }
+ if action.Envelope.ServerSeq == 22 {
+ break replayed
+ }
+ }
+ case <-ctx.Done():
+ t.Fatal("replay not delivered")
+ }
+ }
+ if text != "hello!" {
+ t.Fatalf("managed replay duplicated terminal output: %q", text)
+ }
+ if got := receive(); got.ClientSeq != credit.ClientSeq {
+ t.Fatal("credit was renumbered")
+ }
+ if got := receive(); got.ClientSeq != input.ClientSeq {
+ t.Fatal("input was renumbered")
+ }
+ if _, err := connection.Read(ctx); err != io.EOF {
+ t.Fatalf("replayed EOF: %v", err)
+ }
+ if _, err := connection.Write(ctx, []byte("c")); err != nil {
+ t.Fatal(err)
+ }
+ if got := receive(); got.ClientSeq <= ordinary.ClientSeq {
+ t.Fatal("ordinary sequence floor lost")
+ }
+ second, err := fresh.OpenTCPConnection(ctx, "ahp-session:/s1", create)
+ if err != nil {
+ t.Fatal(err)
+ }
+ waiter := make(chan error, 1)
+ go func() { _, err := second.Read(ctx); waiter <- err }()
+ if err := multi.Shutdown(ctx); err != nil {
+ t.Fatal(err)
+ }
+ if !connection.IsClosed() || !second.IsClosed() {
+ t.Fatal("permanent shutdown retained live streams")
+ }
+ select {
+ case err := <-waiter:
+ if err == nil {
+ t.Fatal("shutdown did not fail blocked read")
+ }
+ case <-ctx.Done():
+ t.Fatal("shutdown blocked read")
+ }
+ } else {
+ if !connection.IsClosed() {
+ t.Fatal("snapshot/missing/fallback revived TCP")
+ }
+ if _, err := connection.Read(ctx); err == nil {
+ t.Fatal("terminated stream read succeeded")
+ }
+ }
+ })
+ }
+}
+
func TestAutomationCapabilitiesUpdatedAcrossReconnect(t *testing.T) {
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
diff --git a/clients/go/ahp/reducers.go b/clients/go/ahp/reducers.go
index ebef72956..84f2775e4 100644
--- a/clients/go/ahp/reducers.go
+++ b/clients/go/ahp/reducers.go
@@ -25,6 +25,175 @@ const (
ReduceOutcomeOutOfScope
)
+// TcpReduceError reports a rejected TCP action. The reducer leaves state unchanged.
+type TcpReduceError struct {
+ Reason string
+}
+
+func (e *TcpReduceError) Error() string {
+ return "Invalid TCP action: " + e.Reason
+}
+
+const tcpMaxSafeInteger int64 = 1<<53 - 1
+
+func requireTCPOffset(value int64) error {
+ if value < 0 || value > tcpMaxSafeInteger {
+ return &TcpReduceError{Reason: "offset must be a nonnegative safe integer"}
+ }
+ return nil
+}
+
+func tcpPayloadLength(data string, maximumChunkSize int64) (int64, error) {
+ // Match JavaScript string length for the size check, including invalid
+ // non-ASCII payloads, before validating the base64 alphabet.
+ encodedLength := 0
+ for _, char := range data {
+ encodedLength++
+ if char > 0xffff {
+ encodedLength++
+ }
+ }
+ if maximumChunkSize < 0 || len(data) == 0 ||
+ uint64(encodedLength) > 4*((uint64(maximumChunkSize)+2)/3) {
+ return 0, &TcpReduceError{Reason: "chunk size"}
+ }
+ padding := 0
+ if data[len(data)-1] == '=' {
+ padding = 1
+ if len(data) > 1 && data[len(data)-2] == '=' {
+ padding = 2
+ }
+ }
+ if len(data)%4 != 0 {
+ return 0, &TcpReduceError{Reason: "base64 encoding"}
+ }
+ var last byte
+ for i := 0; i < len(data)-padding; i++ {
+ switch c := data[i]; {
+ case c >= 'A' && c <= 'Z':
+ last = c - 'A'
+ case c >= 'a' && c <= 'z':
+ last = c - 'a' + 26
+ case c >= '0' && c <= '9':
+ last = c - '0' + 52
+ case c == '+':
+ last = 62
+ case c == '/':
+ last = 63
+ default:
+ return 0, &TcpReduceError{Reason: "base64 encoding"}
+ }
+ }
+ if (padding == 2 && last%16 != 0) || (padding == 1 && last%4 != 0) {
+ return 0, &TcpReduceError{Reason: "noncanonical base64 padding bits"}
+ }
+ length := int64(len(data)/4)*3 - int64(padding)
+ if length > maximumChunkSize {
+ return 0, &TcpReduceError{Reason: "chunk size"}
+ }
+ return length, nil
+}
+
+func receiveTCP(direction *ahptypes.FlowControlledByteDirectionState, offset int64, data string, senderClosed bool) (ReduceOutcome, error) {
+ if err := requireTCPOffset(offset); err != nil {
+ return ReduceOutcomeNoOp, err
+ }
+ length, err := tcpPayloadLength(data, direction.MaximumChunkSize)
+ if err != nil {
+ return ReduceOutcomeNoOp, err
+ }
+ if length > tcpMaxSafeInteger-offset {
+ return ReduceOutcomeNoOp, &TcpReduceError{Reason: "offset must be a nonnegative safe integer"}
+ }
+ end := offset + length
+ if end <= direction.ReceivedBytes {
+ return ReduceOutcomeNoOp, nil
+ }
+ if offset != direction.ReceivedBytes {
+ return ReduceOutcomeNoOp, &TcpReduceError{Reason: "gap or overlapping byte range"}
+ }
+ if senderClosed || direction.EofAtBytes != nil {
+ return ReduceOutcomeNoOp, &TcpReduceError{Reason: "data after EOF or sender close"}
+ }
+ if end-direction.ConsumedBytes > direction.WindowBytes {
+ return ReduceOutcomeNoOp, &TcpReduceError{Reason: "receive window exceeded"}
+ }
+ direction.ReceivedBytes = end
+ return ReduceOutcomeApplied, nil
+}
+
+func consumeTCP(direction *ahptypes.FlowControlledByteDirectionState, consumedBytes int64) (ReduceOutcome, error) {
+ if err := requireTCPOffset(consumedBytes); err != nil {
+ return ReduceOutcomeNoOp, err
+ }
+ if consumedBytes > direction.ReceivedBytes {
+ return ReduceOutcomeNoOp, &TcpReduceError{Reason: "consuming bytes not received"}
+ }
+ if consumedBytes <= direction.ConsumedBytes {
+ return ReduceOutcomeNoOp, nil
+ }
+ direction.ConsumedBytes = consumedBytes
+ return ReduceOutcomeApplied, nil
+}
+
+func eofTCP(direction *ahptypes.FlowControlledByteDirectionState, finalOffset int64, senderClosed bool) (ReduceOutcome, error) {
+ if err := requireTCPOffset(finalOffset); err != nil {
+ return ReduceOutcomeNoOp, err
+ }
+ if finalOffset != direction.ReceivedBytes {
+ return ReduceOutcomeNoOp, &TcpReduceError{Reason: "EOF offset"}
+ }
+ if direction.EofAtBytes != nil && *direction.EofAtBytes == finalOffset {
+ return ReduceOutcomeNoOp, nil
+ }
+ if senderClosed {
+ return ReduceOutcomeNoOp, &TcpReduceError{Reason: "EOF after sender close"}
+ }
+ direction.EofAtBytes = &finalOffset
+ return ReduceOutcomeApplied, nil
+}
+
+// ApplyActionToTCP applies an action without retaining its payload. Invalid
+// actions return a *TcpReduceError and leave state unchanged. Adapters must
+// reset the channel on error and write bytes only when ReceivedBytes advances.
+// This reducer is not a lossless stream subscription or a snapshot mirror.
+func ApplyActionToTCP(state *ahptypes.TcpConnectionState, action ahptypes.StateAction) (ReduceOutcome, error) {
+ if state.Reset != nil {
+ return ReduceOutcomeNoOp, nil
+ }
+ switch a := action.Value.(type) {
+ case *ahptypes.TcpInputAction:
+ return receiveTCP(&state.Input, a.Offset, a.Data, state.ClientClosed)
+ case *ahptypes.TcpDataAction:
+ return receiveTCP(&state.Output, a.Offset, a.Data, state.HostClosed)
+ case *ahptypes.TcpInputConsumedAction:
+ return consumeTCP(&state.Input, a.ConsumedBytes)
+ case *ahptypes.TcpDataConsumedAction:
+ return consumeTCP(&state.Output, a.ConsumedBytes)
+ case *ahptypes.TcpInputEofAction:
+ return eofTCP(&state.Input, a.FinalOffset, state.ClientClosed)
+ case *ahptypes.TcpDataEofAction:
+ return eofTCP(&state.Output, a.FinalOffset, state.HostClosed)
+ case *ahptypes.TcpClientCloseAction:
+ if state.ClientClosed {
+ return ReduceOutcomeNoOp, nil
+ }
+ state.ClientClosed = true
+ case *ahptypes.TcpHostCloseAction:
+ if state.HostClosed {
+ return ReduceOutcomeNoOp, nil
+ }
+ state.HostClosed = true
+ case *ahptypes.TcpClientResetAction:
+ state.Reset = &ahptypes.TcpResetState{Source: ahptypes.TcpEndpointClient, Reason: a.Reason}
+ case *ahptypes.TcpHostResetAction:
+ state.Reset = &ahptypes.TcpResetState{Source: ahptypes.TcpEndpointHost, Reason: a.Reason}
+ default:
+ return ReduceOutcomeOutOfScope, nil
+ }
+ return ReduceOutcomeApplied, nil
+}
+
func addMillisecondsToTimestamp(timestamp string, duration int64) string {
start, err := time.Parse(time.RFC3339Nano, timestamp)
if err != nil {
diff --git a/clients/go/ahp/reducers_fixture_test.go b/clients/go/ahp/reducers_fixture_test.go
index a1617d4b7..9c5ce9dd5 100644
--- a/clients/go/ahp/reducers_fixture_test.go
+++ b/clients/go/ahp/reducers_fixture_test.go
@@ -2,7 +2,9 @@ package ahp
import (
"encoding/json"
+ "errors"
"fmt"
+ "math"
"os"
"path/filepath"
"reflect"
@@ -83,6 +85,15 @@ var reducerFixturesSkipList = map[string]string{
// Add entries like "123-foo.json": "reason" when needed.
}
+type reducerFixture struct {
+ Description string `json:"description"`
+ Reducer string `json:"reducer"`
+ Initial json.RawMessage `json:"initial"`
+ Actions []json.RawMessage `json:"actions"`
+ Expected json.RawMessage `json:"expected"`
+ ExpectedError string `json:"expectedError"`
+}
+
// TestFixtureDrivenReducerParity loads every fixture under
// types/test-cases/reducers/*.json, applies the actions through the
// matching Go reducer, and compares the resulting state with the
@@ -117,13 +128,7 @@ func TestFixtureDrivenReducerParity(t *testing.T) {
continue
}
- var fixture struct {
- Description string `json:"description"`
- Reducer string `json:"reducer"`
- Initial json.RawMessage `json:"initial"`
- Actions []json.RawMessage `json:"actions"`
- Expected json.RawMessage `json:"expected"`
- }
+ var fixture reducerFixture
if err := json.Unmarshal(raw, &fixture); err != nil {
t.Errorf("%s: parse fixture: %v", name, err)
failed++
@@ -131,32 +136,27 @@ func TestFixtureDrivenReducerParity(t *testing.T) {
}
ok := t.Run(fmt.Sprintf("%s/%s", fixture.Reducer, name), func(tt *testing.T) {
- actions := make([]ahptypes.StateAction, len(fixture.Actions))
- for i, raw := range fixture.Actions {
- if err := json.Unmarshal(raw, &actions[i]); err != nil {
- tt.Fatalf("decode action %d: %v", i, err)
- }
- }
-
switch fixture.Reducer {
case "root":
- runFixture[ahptypes.RootState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToRoot)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToRoot))
case "session":
- runFixture[ahptypes.SessionState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToSession)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToSession))
case "chat":
- runFixture[ahptypes.ChatState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToChat)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToChat))
case "terminal":
- runFixture[ahptypes.TerminalState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToTerminal)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToTerminal))
case "changeset":
- runFixture[ahptypes.ChangesetState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToChangeset)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToChangeset))
case "annotations":
- runFixture[ahptypes.AnnotationsState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToAnnotations)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToAnnotations))
case "resourceWatch":
- runFixture[ahptypes.ResourceWatchState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToResourceWatch)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToResourceWatch))
case "automation":
- runFixture[ahptypes.AutomationState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToAutomation)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToAutomation))
case "automationRun":
- runFixture[ahptypes.AutomationRunState](tt, fixture.Initial, fixture.Expected, actions, ApplyActionToAutomationRun)
+ runFixture(tt, fixture, reducerWithoutError(ApplyActionToAutomationRun))
+ case "tcp":
+ runFixture(tt, fixture, ApplyActionToTCP)
default:
tt.Fatalf("unknown reducer kind %q", fixture.Reducer)
}
@@ -171,28 +171,61 @@ func TestFixtureDrivenReducerParity(t *testing.T) {
t.Logf("Fixture results: %d passed, %d skipped, %d failed (of %d total)", passed, skipped, failed, passed+skipped+failed)
}
-func runFixture[T any](t *testing.T, initial, expected json.RawMessage, actions []ahptypes.StateAction, apply func(*T, ahptypes.StateAction) ReduceOutcome) {
+func reducerWithoutError[T any](apply func(*T, ahptypes.StateAction) ReduceOutcome) func(*T, ahptypes.StateAction) (ReduceOutcome, error) {
+ return func(state *T, action ahptypes.StateAction) (ReduceOutcome, error) {
+ return apply(state, action), nil
+ }
+}
+
+func runFixture[T any](t *testing.T, fixture reducerFixture, apply func(*T, ahptypes.StateAction) (ReduceOutcome, error)) {
t.Helper()
var state T
- if err := json.Unmarshal(initial, &state); err != nil {
+ if err := json.Unmarshal(fixture.Initial, &state); err != nil {
t.Fatalf("decode initial state: %v", err)
}
// Round-trip the initial state through marshal/unmarshal to catch
// any data loss in the generated types before we mutate.
roundTripped := stripNulls(reMarshal(t, &state))
- originalParsed := stripNulls(parseJSON(t, initial))
+ originalParsed := stripNulls(parseJSON(t, fixture.Initial))
if !reflect.DeepEqual(roundTripped, originalParsed) {
t.Fatalf("initial state did not survive round-trip:\nre-serialized: %s\noriginal: %s",
mustPretty(roundTripped), mustPretty(originalParsed))
}
- for i, action := range actions {
- _ = apply(&state, action)
- _ = i
+ if fixture.ExpectedError != "" && len(fixture.Actions) == 0 {
+ t.Fatal("expectedError requires a final action")
+ }
+ for i, raw := range fixture.Actions {
+ expectError := fixture.ExpectedError != "" && i == len(fixture.Actions)-1
+ before := reMarshal(t, &state)
+ var action ahptypes.StateAction
+ if err := json.Unmarshal(raw, &action); err != nil {
+ var typeError *json.UnmarshalTypeError
+ offset, fractional := parseJSON(t, raw).(map[string]any)["offset"].(float64)
+ if !expectError || fixture.Reducer != "tcp" ||
+ fixture.ExpectedError != "Invalid TCP action: offset must be a nonnegative safe integer" ||
+ !fractional || math.Mod(offset, 1) == 0 ||
+ !errors.As(err, &typeError) || typeError.Type.Kind() != reflect.Int64 || typeError.Field != "offset" {
+ t.Fatalf("decode action %d: %v", i, err)
+ }
+ t.Logf("final fractional action rejected by native int64 deserializer: %v", err)
+ } else {
+ _, err := apply(&state, action)
+ if expectError {
+ if err == nil || err.Error() != fixture.ExpectedError {
+ t.Fatalf("action %d: expected error %q, got %v", i, fixture.ExpectedError, err)
+ }
+ } else if err != nil {
+ t.Fatalf("action %d: unexpected reducer error: %v", i, err)
+ }
+ }
+ if expectError && !reflect.DeepEqual(before, reMarshal(t, &state)) {
+ t.Fatalf("rejected action %d mutated state", i)
+ }
}
actual := stripNulls(reMarshal(t, &state))
- want := stripNulls(parseJSON(t, expected))
+ want := stripNulls(parseJSON(t, fixture.Expected))
if !reflect.DeepEqual(actual, want) {
t.Fatalf("state mismatch:\nactual: %s\nexpected: %s",
mustPretty(actual), mustPretty(want))
@@ -215,3 +248,123 @@ func mustPretty(v any) string {
}
return string(b)
}
+
+func tcpTestState() ahptypes.TcpConnectionState {
+ direction := ahptypes.FlowControlledByteDirectionState{WindowBytes: 8, MaximumChunkSize: 6}
+ return ahptypes.TcpConnectionState{
+ Session: "ahp-session:/s1", Target: ahptypes.TcpTarget{Host: "localhost", Port: 3000},
+ Encoding: ahptypes.TcpDataEncodingBase64, Input: direction, Output: direction,
+ }
+}
+
+func TestTCPReducerLargePayload(t *testing.T) {
+ const size = 4 * 1024 * 1024
+ data := strings.Repeat("A", 4*((size+2)/3)-2) + "=="
+ for _, actionType := range []string{"tcp/input", "tcp/data"} {
+ t.Run(actionType, func(t *testing.T) {
+ state := tcpTestState()
+ state.Input.WindowBytes, state.Output.WindowBytes = size, size
+ state.Input.MaximumChunkSize, state.Output.MaximumChunkSize = size, size
+ var action ahptypes.StateAction
+ raw, err := json.Marshal(map[string]any{"type": actionType, "offset": 0, "data": data})
+ if err != nil {
+ t.Fatal(err)
+ }
+ if err := json.Unmarshal(raw, &action); err != nil {
+ t.Fatal(err)
+ }
+ if outcome, err := ApplyActionToTCP(&state, action); outcome != ReduceOutcomeApplied || err != nil {
+ t.Fatalf("large chunk: %v, %v", outcome, err)
+ }
+ direction := state.Input
+ if actionType == "tcp/data" {
+ direction = state.Output
+ }
+ if direction.ReceivedBytes != size {
+ t.Fatalf("received %d bytes, want %d", direction.ReceivedBytes, size)
+ }
+ before := reMarshal(t, state)
+ if outcome, err := ApplyActionToTCP(&state, action); outcome != ReduceOutcomeNoOp || err != nil {
+ t.Fatalf("duplicate: %v, %v", outcome, err)
+ }
+ // Even a fully duplicated range must pass canonical encoding validation.
+ switch a := action.Value.(type) {
+ case *ahptypes.TcpInputAction:
+ a.Data = data[:len(data)-3] + "B=="
+ case *ahptypes.TcpDataAction:
+ a.Data = data[:len(data)-3] + "B=="
+ }
+ _, err = ApplyActionToTCP(&state, action)
+ var tcpError *TcpReduceError
+ if !errors.As(err, &tcpError) || tcpError.Reason != "noncanonical base64 padding bits" {
+ t.Fatalf("invalid padding: %v", err)
+ }
+ if !reflect.DeepEqual(before, reMarshal(t, state)) {
+ t.Fatal("duplicate or rejected chunk mutated state")
+ }
+ })
+ }
+}
+
+func TestTCPReducerNativeIntegerBounds(t *testing.T) {
+ for _, actionType := range []string{"tcp/input", "tcp/data", "tcp/inputConsumed", "tcp/dataConsumed", "tcp/inputEof", "tcp/dataEof"} {
+ t.Run(actionType, func(t *testing.T) {
+ field := "offset"
+ if strings.HasSuffix(actionType, "Consumed") {
+ field = "consumedBytes"
+ } else if strings.HasSuffix(actionType, "Eof") {
+ field = "finalOffset"
+ }
+
+ for _, value := range []int64{math.MinInt64, -1, tcpMaxSafeInteger + 1, math.MaxInt64} {
+ state := tcpTestState()
+ before := reMarshal(t, state)
+ raw, err := json.Marshal(map[string]any{"type": actionType, field: value, "data": "AA=="})
+ if err != nil {
+ t.Fatal(err)
+ }
+ var action ahptypes.StateAction
+ if err := json.Unmarshal(raw, &action); err != nil {
+ t.Fatal(err)
+ }
+ outcome, err := ApplyActionToTCP(&state, action)
+ var tcpError *TcpReduceError
+ if outcome != ReduceOutcomeNoOp || !errors.As(err, &tcpError) ||
+ tcpError.Reason != "offset must be a nonnegative safe integer" {
+ t.Fatalf("%d: %v, %v", value, outcome, err)
+ }
+ if !reflect.DeepEqual(before, reMarshal(t, state)) {
+ t.Fatal("invalid integer mutated state")
+ }
+ }
+ })
+ }
+ state := tcpTestState()
+ state.Input.ReceivedBytes, state.Input.ConsumedBytes = tcpMaxSafeInteger-1, tcpMaxSafeInteger-1
+ action := ahptypes.StateAction{Value: &ahptypes.TcpInputAction{Offset: tcpMaxSafeInteger - 1, Data: "AA=="}}
+ if outcome, err := ApplyActionToTCP(&state, action); outcome != ReduceOutcomeApplied || err != nil || state.Input.ReceivedBytes != tcpMaxSafeInteger {
+ t.Fatalf("last safe byte: %v, %v", outcome, err)
+ }
+}
+
+func TestTCPReducerNonASCIIErrorOrder(t *testing.T) {
+ for _, test := range []struct {
+ data, reason string
+ }{
+ {"\u00e9\u00e9\u00e9", "base64 encoding"},
+ {"\U0001f600\U0001f600", "base64 encoding"},
+ {"\U0001f600\U0001f600A", "chunk size"},
+ } {
+ state := tcpTestState()
+ state.Input.MaximumChunkSize = 1
+ before := reMarshal(t, state)
+ _, err := ApplyActionToTCP(&state, ahptypes.StateAction{Value: &ahptypes.TcpInputAction{Offset: 0, Data: test.data}})
+ var tcpError *TcpReduceError
+ if !errors.As(err, &tcpError) || tcpError.Reason != test.reason {
+ t.Fatalf("%q: got %v, want %s", test.data, err, test.reason)
+ }
+ if !reflect.DeepEqual(before, reMarshal(t, state)) {
+ t.Fatal("invalid encoding mutated state")
+ }
+ }
+}
diff --git a/clients/go/ahp/tcp.go b/clients/go/ahp/tcp.go
new file mode 100644
index 000000000..f096def65
--- /dev/null
+++ b/clients/go/ahp/tcp.go
@@ -0,0 +1,1048 @@
+package ahp
+
+import (
+ "context"
+ "encoding/base64"
+ "encoding/json"
+ "errors"
+ "fmt"
+ "io"
+ "log"
+ "sort"
+ "strings"
+ "sync"
+
+ "github.com/microsoft/agent-host-protocol/clients/go/ahptypes"
+)
+
+const tcpMaxWindowBytes int64 = (1 << 32) - 1
+
+// TCPConnectionError is a terminal stream or invalid-operation error.
+type TCPConnectionError struct {
+ Reason string
+ Cause error
+}
+
+func (e *TCPConnectionError) Error() string {
+ if e.Cause != nil {
+ return fmt.Sprintf("ahp: TCP %s: %v", e.Reason, e.Cause)
+ }
+ return "ahp: TCP " + e.Reason
+}
+
+func (e *TCPConnectionError) Unwrap() error { return e.Cause }
+
+type tcpPending struct {
+ action ahptypes.StateAction
+ sent uint64
+}
+
+// TCPConnection owns one protocol byte stream, not a native socket.
+// Write accepts bytes into bounded protocol credit; Drain waits for consumption
+// at the host. Transport loss suspends the same handle until explicitly resumed.
+type TCPConnection struct {
+ mu sync.Mutex
+ sendMu sync.Mutex
+ resumeMu sync.Mutex
+ client *Client
+ events *EventStream
+ resource string
+ owner string
+ state ahptypes.TcpConnectionState
+ checkpoint int64
+ lastSeq int64
+ epoch uint64
+ online bool
+ resuming bool
+ writing bool
+ ending bool
+ closing bool
+ closeQueued bool
+ terminal bool
+ err error
+ pending map[int64]tcpPending
+ received [][]byte
+ sentBytes int64
+ consumed int64
+ changed chan struct{}
+ cleanup chan struct{}
+ cleanupErr error
+}
+
+// Resource is the host-assigned child channel URI.
+func (t *TCPConnection) Resource() string { return t.resource }
+
+// IsClosed reports final close/reset/disposal, not a resumable transport loss.
+func (t *TCPConnection) IsClosed() bool {
+ t.mu.Lock()
+ defer t.mu.Unlock()
+ return t.terminal
+}
+
+func (c *Client) registerTCPStream(connection *TCPConnection) bool {
+ c.tcpMu.Lock()
+ defer c.tcpMu.Unlock()
+ if c.tcpDisposed {
+ return false
+ }
+ if c.tcpStreams == nil {
+ c.tcpStreams = make(map[*TCPConnection]struct{})
+ }
+ c.tcpStreams[connection] = struct{}{}
+ return true
+}
+
+func (c *Client) disposeTCPStreams(ctx context.Context) error {
+ c.tcpMu.Lock()
+ c.tcpDisposed = true
+ streams := make([]*TCPConnection, 0, len(c.tcpStreams))
+ for connection := range c.tcpStreams {
+ streams = append(streams, connection)
+ }
+ c.tcpMu.Unlock()
+ var err error
+ for _, connection := range streams {
+ if connection.finishForClient(tcpConnectionError("disposed", "client shut down"),
+ tcpResetAction(ahptypes.TcpResetReasonConnectionAborted), false, c) {
+ err = errors.Join(err, connection.waitCleanup(ctx))
+ }
+ }
+ return err
+}
+
+// RestoreReconnectState carries negotiated TCP support and the complete client
+// sequence allocator to a fresh transport. Host runtimes call this before their
+// reconnect handshake; it does not initialize or subscribe the new transport.
+func (c *Client) RestoreReconnectState(previous *Client) error {
+ if previous == nil || previous == c {
+ return tcpConnectionError("resume", "a previous client on a different transport is required")
+ }
+ previous.tcpMu.Lock()
+ owner, capability := previous.tcpClientID, previous.tcpCapability
+ previous.tcpMu.Unlock()
+ c.tcpMu.Lock()
+ if c.tcpClientID != "" && c.tcpClientID != owner {
+ c.tcpMu.Unlock()
+ return tcpConnectionError("resume", "previous transport belongs to a different clientId")
+ }
+ c.tcpClientID, c.tcpCapability = owner, capability
+ c.tcpMu.Unlock()
+ floor := previous.nextClientSeq.Load()
+ for {
+ next := c.nextClientSeq.Load()
+ if next >= floor || c.nextClientSeq.CompareAndSwap(next, floor) {
+ return nil
+ }
+ }
+}
+
+func tcpConnectionError(reason, message string) error {
+ return &TCPConnectionError{Reason: reason, Cause: errors.New(message)}
+}
+
+func tcpResetAction(reason ahptypes.TcpResetReason) ahptypes.StateAction {
+ return ahptypes.StateAction{Value: &ahptypes.TcpClientResetAction{Type: ahptypes.ActionTypeTcpClientReset, Reason: reason}}
+}
+
+func validTCPDirection(d ahptypes.FlowControlledByteDirectionState) bool {
+ return d.WindowBytes > 0 && d.WindowBytes <= tcpMaxWindowBytes &&
+ d.MaximumChunkSize > 0 && d.MaximumChunkSize <= d.WindowBytes &&
+ d.ReceivedBytes == 0 && d.ConsumedBytes == 0 && d.EofAtBytes == nil
+}
+
+func validateTCPOpen(session string, create ahptypes.TcpConnectionSubscription) error {
+ if !strings.HasPrefix(session, "ahp-session:") || len(session) == len("ahp-session:") ||
+ create.Type != "tcpConnection" || create.Host == "" || strings.ContainsAny(create.Host, "/\\\x00 \t\r\n") ||
+ create.Port < 1 || create.Port > 65535 || create.Encoding != ahptypes.TcpDataEncodingBase64 ||
+ create.ReceiveWindowBytes < 1 || create.ReceiveWindowBytes > tcpMaxWindowBytes ||
+ create.MaximumChunkSize < 1 || create.MaximumChunkSize > create.ReceiveWindowBytes {
+ return tcpConnectionError("invalidOpen", "invalid session, target, encoding, or byte limits")
+ }
+ return nil
+}
+
+func tcpContext(c *Client) (context.Context, context.CancelFunc) {
+ if c.cfg.DefaultRequestTimeout > 0 {
+ return context.WithTimeout(context.Background(), c.cfg.DefaultRequestTimeout)
+ }
+ return context.WithCancel(context.Background())
+}
+
+// OpenTCPConnection atomically creates a child channel after Initialize has
+// advertised base64 TCP support. Cancellation still cleans up a late creation
+// response, including after the request timeout; the parent is never unsubscribed.
+// For managed hosts use hosts.HostClientHandle.OpenTCPConnection so the runtime
+// retains the stream across reconnects.
+func (c *Client) OpenTCPConnection(ctx context.Context, session string, create ahptypes.TcpConnectionSubscription) (*TCPConnection, error) {
+ if err := ctx.Err(); err != nil {
+ return nil, err
+ }
+ if err := validateTCPOpen(session, create); err != nil {
+ return nil, err
+ }
+ c.tcpMu.Lock()
+ owner, capability := c.tcpClientID, c.tcpCapability
+ supported := false
+ if capability != nil {
+ for _, encoding := range capability.Encodings {
+ supported = supported || encoding == create.Encoding
+ }
+ }
+ c.tcpMu.Unlock()
+ if owner == "" || !supported {
+ return nil, tcpConnectionError("unsupported", "Initialize must advertise the requested TCP encoding")
+ }
+ events := c.events(true, "")
+ type opened struct {
+ connection *TCPConnection
+ err error
+ }
+ result := make(chan opened)
+ go func() {
+ requestCtx, cancel := tcpContext(c)
+ defer cancel()
+ var response ahptypes.SubscribeResult
+ err := c.requestWithLateResult(requestCtx, "subscribe", ahptypes.SubscribeParams{Channel: session, Create: &create}, &response, c.cleanupLateTCPCreation, func(raw json.RawMessage) {
+ var identity tcpCreationIdentity
+ if json.Unmarshal(raw, &identity) == nil && identity.Snapshot != nil {
+ c.bindEventStream(events, identity.Snapshot.Resource)
+ }
+ })
+ var connection *TCPConnection
+ if err == nil {
+ snapshot := response.Snapshot
+ if snapshot == nil || !strings.HasPrefix(snapshot.Resource, "ahp-tcp:") || snapshot.Resource == "ahp-tcp:" ||
+ snapshot.FromSeq < 0 || snapshot.FromSeq > tcpMaxSafeInteger || snapshot.State.Tcp == nil {
+ err = tcpConnectionError("protocol", "creation did not return a TCP snapshot")
+ } else {
+ state := snapshot.State.Tcp
+ if state.Session != session || state.Target.Host != create.Host || state.Target.Port != create.Port ||
+ state.Encoding != create.Encoding || !validTCPDirection(state.Input) || !validTCPDirection(state.Output) ||
+ state.Output.WindowBytes > create.ReceiveWindowBytes || state.Output.MaximumChunkSize > create.MaximumChunkSize ||
+ state.ClientClosed || state.HostClosed || state.Reset != nil {
+ err = tcpConnectionError("protocol", "creation snapshot is not fresh or does not match the request")
+ } else {
+ connection = &TCPConnection{
+ client: c, events: events, resource: snapshot.Resource, owner: owner, state: *state,
+ checkpoint: snapshot.FromSeq, epoch: 1, online: true,
+ pending: make(map[int64]tcpPending), changed: make(chan struct{}),
+ }
+ }
+ }
+ }
+ if err == nil && !c.registerTCPStream(connection) {
+ err = ErrShutdown
+ }
+ if err == nil && events.Err() != nil {
+ err = events.Err()
+ }
+ if err != nil {
+ events.Close()
+ if connection != nil {
+ connection.finish(err, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ err = errors.Join(err, connection.waitCleanup(requestCtx))
+ connection = nil
+ } else if response.Snapshot != nil && strings.HasPrefix(response.Snapshot.Resource, "ahp-tcp:") && response.Snapshot.Resource != "ahp-tcp:" {
+ err = errors.Join(err, c.Unsubscribe(requestCtx, response.Snapshot.Resource))
+ }
+ }
+ select {
+ case result <- opened{connection, err}:
+ if err == nil {
+ connection.start(1, events)
+ }
+ case <-ctx.Done():
+ events.Close()
+ if connection != nil && err == nil {
+ if err := connection.Dispose(requestCtx); err != nil {
+ log.Printf("ahp: cancelled TCP creation cleanup failed: %v", err)
+ }
+ }
+ }
+ }()
+ select {
+ case out := <-result:
+ return out.connection, out.err
+ case <-ctx.Done():
+ events.Close()
+ return nil, ctx.Err()
+ }
+}
+
+func (t *TCPConnection) wakeLocked() {
+ close(t.changed)
+ t.changed = make(chan struct{})
+}
+
+// Cleanup needs the child identity even when the abandoned state cannot decode.
+type tcpCreationIdentity struct {
+ Snapshot *tcpSnapshotIdentity `json:"snapshot"`
+}
+
+type tcpSnapshotIdentity struct {
+ Resource string `json:"resource"`
+}
+
+func (c *Client) cleanupLateTCPCreation(raw json.RawMessage) {
+ var response tcpCreationIdentity
+ if err := json.Unmarshal(raw, &response); err != nil {
+ log.Printf("ahp: invalid late TCP creation response: %v", err)
+ return
+ }
+ if response.Snapshot == nil || !strings.HasPrefix(response.Snapshot.Resource, "ahp-tcp:") || response.Snapshot.Resource == "ahp-tcp:" {
+ log.Printf("ahp: late TCP creation response omitted a valid child resource")
+ return
+ }
+ ctx, cancel := tcpContext(c)
+ defer cancel()
+ seq := c.nextClientSeq.Add(1) - 1
+ var err error
+ if seq < 1 || seq > tcpMaxSafeInteger {
+ err = tcpConnectionError("protocol", "client sequence exhausted")
+ } else {
+ err = c.Notify(ctx, "dispatchAction", ahptypes.DispatchActionParams{
+ Channel: response.Snapshot.Resource, ClientSeq: seq,
+ Action: tcpResetAction(ahptypes.TcpResetReasonConnectionAborted),
+ })
+ }
+ err = errors.Join(err, c.Unsubscribe(ctx, response.Snapshot.Resource))
+ if err != nil {
+ log.Printf("ahp: late TCP creation cleanup failed: %v", err)
+ }
+}
+
+func (t *TCPConnection) queueLocked(action ahptypes.StateAction) error {
+ seq := t.client.nextClientSeq.Add(1) - 1
+ if seq < 1 || seq > tcpMaxSafeInteger {
+ return tcpConnectionError("protocol", "client sequence exhausted")
+ }
+ t.lastSeq = seq
+ t.pending[seq] = tcpPending{action: action}
+ t.wakeLocked()
+ return nil
+}
+
+func waitTCP(ctx context.Context, changed <-chan struct{}) error {
+ select {
+ case <-changed:
+ return nil
+ case <-ctx.Done():
+ return ctx.Err()
+ }
+}
+
+// Read delivers one decoded chunk. Delivery, not arrival, releases receive credit.
+// EOF is returned only after buffered output has been delivered.
+func (t *TCPConnection) Read(ctx context.Context) ([]byte, error) {
+ for {
+ t.mu.Lock()
+ if t.err != nil {
+ err := t.err
+ t.mu.Unlock()
+ return nil, err
+ }
+ if !t.resuming && len(t.received) > 0 {
+ data := t.received[0]
+ consumed := t.consumed + int64(len(data))
+ var err error
+ if !t.terminal {
+ err = t.queueLocked(ahptypes.StateAction{Value: &ahptypes.TcpDataConsumedAction{Type: ahptypes.ActionTypeTcpDataConsumed, ConsumedBytes: consumed}})
+ }
+ if err == nil {
+ t.received[0] = nil
+ t.received = t.received[1:]
+ t.consumed = consumed
+ }
+ t.mu.Unlock()
+ if err != nil {
+ t.finish(err, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ return nil, err
+ }
+ return data, nil
+ }
+ if t.terminal || (!t.resuming && (t.state.Output.EofAtBytes != nil || t.state.HostClosed)) {
+ t.mu.Unlock()
+ return nil, io.EOF
+ }
+ changed := t.changed
+ t.mu.Unlock()
+ if err := waitTCP(ctx, changed); err != nil {
+ return nil, err
+ }
+ }
+}
+
+// Write accepts bytes into the negotiated input window. One concurrent writer
+// is permitted; partial progress is returned on cancellation. Accepted bytes are
+// retained across disconnect, including actions not yet echoed by the host.
+func (t *TCPConnection) Write(ctx context.Context, data []byte) (int, error) {
+ t.mu.Lock()
+ if t.writing || t.ending || t.terminal {
+ t.mu.Unlock()
+ return 0, tcpConnectionError("closed", "write requires an open, idle writer")
+ }
+ t.writing = true
+ t.mu.Unlock()
+ defer func() {
+ t.mu.Lock()
+ t.writing = false
+ t.wakeLocked()
+ t.mu.Unlock()
+ }()
+ written := 0
+ for written < len(data) {
+ if err := ctx.Err(); err != nil {
+ return written, err
+ }
+ t.mu.Lock()
+ if t.err != nil || t.terminal || t.ending {
+ err := t.err
+ if err == nil {
+ err = tcpConnectionError("closed", "connection closed during write")
+ }
+ t.mu.Unlock()
+ return written, err
+ }
+ credit := t.state.Input.WindowBytes - (t.sentBytes - t.state.Input.ConsumedBytes)
+ if !t.online || t.resuming || credit == 0 {
+ changed := t.changed
+ t.mu.Unlock()
+ if err := waitTCP(ctx, changed); err != nil {
+ return written, err
+ }
+ continue
+ }
+ length := min(int64(len(data)-written), credit, t.state.Input.MaximumChunkSize)
+ if length <= 0 || t.sentBytes > tcpMaxSafeInteger-length {
+ t.mu.Unlock()
+ err := tcpConnectionError("protocol", "invalid input credit or byte offset exhaustion")
+ t.finish(err, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ return written, err
+ }
+ action := ahptypes.StateAction{Value: &ahptypes.TcpInputAction{
+ Type: ahptypes.ActionTypeTcpInput, Offset: t.sentBytes, Data: base64.StdEncoding.EncodeToString(data[written : written+int(length)]),
+ }}
+ if err := t.queueLocked(action); err != nil {
+ t.mu.Unlock()
+ return written, err
+ }
+ t.sentBytes += length
+ written += int(length)
+ t.mu.Unlock()
+ }
+ return written, nil
+}
+
+// Drain waits for all accepted input to be consumed at the host.
+func (t *TCPConnection) Drain(ctx context.Context) error {
+ for {
+ t.mu.Lock()
+ if t.err != nil || (t.terminal && t.state.Input.ConsumedBytes < t.sentBytes) {
+ err := t.err
+ if err == nil {
+ err = tcpConnectionError("closed", "connection closed before drain")
+ }
+ t.mu.Unlock()
+ return err
+ }
+ if t.state.Input.ConsumedBytes >= t.sentBytes {
+ t.mu.Unlock()
+ return nil
+ }
+ changed := t.changed
+ t.mu.Unlock()
+ if err := waitTCP(ctx, changed); err != nil {
+ return err
+ }
+ }
+}
+
+// End half-closes input, preserving output reads. Finish a concurrent Write first.
+func (t *TCPConnection) End(ctx context.Context) error {
+ for {
+ t.mu.Lock()
+ if t.terminal || t.writing {
+ t.mu.Unlock()
+ return tcpConnectionError("closed", "end requires an open, idle writer")
+ }
+ if !t.resuming {
+ var err error
+ if !t.ending {
+ err = t.queueLocked(ahptypes.StateAction{Value: &ahptypes.TcpInputEofAction{Type: ahptypes.ActionTypeTcpInputEof, FinalOffset: t.sentBytes}})
+ if err == nil {
+ t.ending = true
+ }
+ }
+ t.mu.Unlock()
+ return err
+ }
+ changed := t.changed
+ t.mu.Unlock()
+ if err := waitTCP(ctx, changed); err != nil {
+ return err
+ }
+ }
+}
+
+// Close stops writes and awaits both close acknowledgements and consumed bytes.
+// Continue reading concurrently to drain output. Dispose aborts without draining.
+func (t *TCPConnection) Close(ctx context.Context) error {
+ t.mu.Lock()
+ if t.terminal {
+ t.mu.Unlock()
+ return t.waitCleanup(ctx)
+ }
+ if t.writing {
+ t.mu.Unlock()
+ return tcpConnectionError("busy", "close requires an idle writer")
+ }
+ t.ending, t.closing = true, true
+ t.wakeLocked()
+ t.mu.Unlock()
+ t.advanceClose()
+ for {
+ t.mu.Lock()
+ err, terminal, changed := t.err, t.terminal, t.changed
+ t.mu.Unlock()
+ if err != nil {
+ return err
+ }
+ if terminal {
+ return t.waitCleanup(ctx)
+ }
+ if err := waitTCP(ctx, changed); err != nil {
+ return err
+ }
+ }
+}
+
+func (t *TCPConnection) advanceClose() {
+ t.mu.Lock()
+ if t.terminal || t.resuming || !t.closing {
+ t.mu.Unlock()
+ return
+ }
+ var err error
+ if !t.closeQueued && (t.state.HostClosed || t.state.Input.ConsumedBytes == t.sentBytes) {
+ err = t.queueLocked(ahptypes.StateAction{Value: &ahptypes.TcpClientCloseAction{Type: ahptypes.ActionTypeTcpClientClose}})
+ t.closeQueued = err == nil
+ }
+ complete := t.state.ClientClosed && t.state.HostClosed &&
+ t.state.Input.ConsumedBytes == t.sentBytes &&
+ t.consumed == t.state.Output.ReceivedBytes &&
+ t.state.Output.ConsumedBytes == t.consumed && len(t.pending) == 0
+ t.mu.Unlock()
+ if err != nil {
+ t.finish(err, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ } else if complete {
+ t.finish(nil, ahptypes.StateAction{}, true)
+ }
+}
+
+// Dispose aborts the stream and discards buffered output. Cleanup is exactly once.
+func (t *TCPConnection) Dispose(ctx context.Context) error {
+ t.finish(tcpConnectionError("disposed", "connection disposed"), tcpResetAction(ahptypes.TcpResetReasonConnectionAborted), false)
+ return t.waitCleanup(ctx)
+}
+
+func (t *TCPConnection) waitCleanup(ctx context.Context) error {
+ t.mu.Lock()
+ done := t.cleanup
+ t.mu.Unlock()
+ select {
+ case <-done:
+ t.mu.Lock()
+ err := t.cleanupErr
+ t.mu.Unlock()
+ return err
+ case <-ctx.Done():
+ return ctx.Err()
+ }
+}
+
+func (t *TCPConnection) finish(reason error, action ahptypes.StateAction, preserve bool) {
+ t.finishForClient(reason, action, preserve, nil)
+}
+
+func (t *TCPConnection) finishForClient(reason error, action ahptypes.StateAction, preserve bool, owner *Client) bool {
+ t.mu.Lock()
+ if owner != nil && t.client != owner {
+ t.mu.Unlock()
+ return false
+ }
+ if t.terminal {
+ t.mu.Unlock()
+ return true
+ }
+ t.terminal, t.online, t.resuming, t.err = true, false, false, reason
+ t.pending = nil
+ if !preserve {
+ t.received = nil
+ }
+ t.cleanup = make(chan struct{})
+ client, events := t.client, t.events
+ t.wakeLocked()
+ t.mu.Unlock()
+ events.Close()
+ go func() {
+ t.sendMu.Lock()
+ defer t.sendMu.Unlock()
+ ctx, cancel := tcpContext(client)
+ defer cancel()
+ var err error
+ closed := false
+ select {
+ case <-client.Done():
+ closed = true
+ default:
+ }
+ if !closed && action.Value != nil {
+ seq := client.nextClientSeq.Add(1) - 1
+ if seq < 1 || seq > tcpMaxSafeInteger {
+ err = tcpConnectionError("protocol", "client sequence exhausted")
+ } else {
+ err = client.Notify(ctx, "dispatchAction", ahptypes.DispatchActionParams{Channel: t.resource, ClientSeq: seq, Action: action})
+ }
+ }
+ if !closed {
+ err = errors.Join(err, client.Unsubscribe(ctx, t.resource))
+ }
+ t.mu.Lock()
+ t.cleanupErr = err
+ if t.err != nil && err != nil {
+ t.err = errors.Join(t.err, err)
+ }
+ close(t.cleanup)
+ t.mu.Unlock()
+ client.tcpMu.Lock()
+ delete(client.tcpStreams, t)
+ client.tcpMu.Unlock()
+ }()
+ return true
+}
+
+func (t *TCPConnection) suspend(epoch uint64) {
+ t.mu.Lock()
+ if t.epoch != epoch || t.terminal {
+ t.mu.Unlock()
+ return
+ }
+ t.online, t.resuming = false, false
+ events := t.events
+ t.wakeLocked()
+ t.mu.Unlock()
+ events.Close()
+}
+
+func tcpEchoMatches(expected, actual ahptypes.StateAction) bool {
+ switch action := actual.Value.(type) {
+ case *ahptypes.TcpInputAction:
+ pending, ok := expected.Value.(*ahptypes.TcpInputAction)
+ return ok && *pending == *action
+ case *ahptypes.TcpDataConsumedAction:
+ pending, ok := expected.Value.(*ahptypes.TcpDataConsumedAction)
+ return ok && *pending == *action
+ case *ahptypes.TcpInputEofAction:
+ pending, ok := expected.Value.(*ahptypes.TcpInputEofAction)
+ return ok && *pending == *action
+ case *ahptypes.TcpClientCloseAction:
+ pending, ok := expected.Value.(*ahptypes.TcpClientCloseAction)
+ return ok && *pending == *action
+ case *ahptypes.TcpClientResetAction:
+ pending, ok := expected.Value.(*ahptypes.TcpClientResetAction)
+ return ok && *pending == *action
+ default:
+ return false
+ }
+}
+
+func (t *TCPConnection) accept(envelope ahptypes.ActionEnvelope, epoch uint64) {
+ t.mu.Lock()
+ if t.epoch != epoch || t.terminal || envelope.Channel != t.resource {
+ t.mu.Unlock()
+ return
+ }
+ var err error
+ clientEcho := false
+ switch envelope.Action.Value.(type) {
+ case *ahptypes.TcpInputAction, *ahptypes.TcpDataConsumedAction, *ahptypes.TcpInputEofAction,
+ *ahptypes.TcpClientCloseAction, *ahptypes.TcpClientResetAction:
+ clientEcho = true
+ }
+ var pending tcpPending
+ var pendingExists bool
+ if clientEcho {
+ origin := envelope.Origin
+ if origin == nil || origin.ClientId != t.owner || origin.ClientSeq < 1 ||
+ origin.ClientSeq > tcpMaxSafeInteger || origin.ClientSeq > t.lastSeq {
+ err = tcpConnectionError("protocol", "invalid client TCP echo origin")
+ } else {
+ pending, pendingExists = t.pending[origin.ClientSeq]
+ if pendingExists && !tcpEchoMatches(pending.action, envelope.Action) {
+ err = tcpConnectionError("protocol", "TCP echo does not match pending action")
+ }
+ }
+ }
+ if err == nil && envelope.ServerSeq <= t.checkpoint {
+ t.mu.Unlock()
+ return
+ }
+ before := t.state.Output.ReceivedBytes
+ next := t.state
+ if envelope.RejectionReason != nil {
+ err = tcpConnectionError("rejected", *envelope.RejectionReason)
+ } else if envelope.ServerSeq > tcpMaxSafeInteger {
+ err = tcpConnectionError("protocol", "invalid server sequence")
+ } else if err == nil {
+ var outcome ReduceOutcome
+ outcome, err = ApplyActionToTCP(&next, envelope.Action)
+ if err == nil && clientEcho && !pendingExists && outcome == ReduceOutcomeApplied {
+ err = tcpConnectionError("protocol", "unacknowledged TCP state advanced without a matching pending action")
+ }
+ }
+ if err == nil && (next.Input.ReceivedBytes > t.sentBytes || next.Output.ConsumedBytes > t.consumed ||
+ next.Output.ReceivedBytes-t.consumed > next.Output.WindowBytes) {
+ err = tcpConnectionError("protocol", "host exceeded owned byte counters")
+ }
+ if err == nil && next.Output.ReceivedBytes > before {
+ if action, ok := envelope.Action.Value.(*ahptypes.TcpDataAction); ok {
+ var bytes []byte
+ bytes, err = base64.StdEncoding.Strict().DecodeString(action.Data)
+ if err == nil {
+ t.received = append(t.received, bytes)
+ }
+ }
+ }
+ if err == nil {
+ t.state = next
+ t.checkpoint = envelope.ServerSeq
+ if clientEcho {
+ delete(t.pending, envelope.Origin.ClientSeq)
+ }
+ }
+ reset := t.state.Reset
+ if t.state.HostClosed {
+ t.ending, t.closing = true, true
+ }
+ t.wakeLocked()
+ t.mu.Unlock()
+ if err != nil {
+ t.finish(err, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ } else if reset != nil {
+ t.finish(tcpConnectionError("reset", string(reset.Reason)), ahptypes.StateAction{}, false)
+ } else {
+ t.advanceClose()
+ }
+}
+
+func (t *TCPConnection) start(epoch uint64, events *EventStream) {
+ go func() {
+ for event := range events.Events() {
+ t.mu.Lock()
+ current := t.epoch == epoch && !t.terminal
+ t.mu.Unlock()
+ if !current {
+ return
+ }
+ if action, ok := event.Event.(SubscriptionEventAction); ok {
+ t.accept(action.Envelope, epoch)
+ }
+ }
+ t.mu.Lock()
+ current := t.epoch == epoch && !t.terminal
+ t.mu.Unlock()
+ if !current {
+ return
+ }
+ var lag *SubscriptionLagError
+ var protocol *TransportError
+ err := events.Err()
+ if errors.As(err, &lag) || (errors.As(err, &protocol) && protocol.Kind == "protocol") {
+ t.finish(err, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ } else {
+ t.suspend(epoch)
+ }
+ }()
+ go t.sendPending(epoch)
+}
+
+func (t *TCPConnection) sendPending(epoch uint64) {
+ for {
+ t.mu.Lock()
+ if t.epoch != epoch || t.terminal || !t.online || t.resuming {
+ t.mu.Unlock()
+ return
+ }
+ var seq int64
+ for candidate, pending := range t.pending {
+ if pending.sent != epoch && (seq == 0 || candidate < seq) {
+ seq = candidate
+ }
+ }
+ pending, client, changed := t.pending[seq], t.client, t.changed
+ t.mu.Unlock()
+ if seq == 0 {
+ select {
+ case <-changed:
+ continue
+ case <-client.Done():
+ t.suspend(epoch)
+ return
+ }
+ }
+ t.sendMu.Lock()
+ t.mu.Lock()
+ current := t.epoch == epoch && !t.terminal && t.online
+ t.mu.Unlock()
+ var err error
+ if current {
+ ctx, cancel := tcpContext(client)
+ err = client.Notify(ctx, "dispatchAction", ahptypes.DispatchActionParams{Channel: t.resource, ClientSeq: seq, Action: pending.action})
+ cancel()
+ }
+ t.sendMu.Unlock()
+ if !current {
+ return
+ }
+ if err != nil {
+ t.suspend(epoch)
+ return
+ }
+ t.mu.Lock()
+ if existing, ok := t.pending[seq]; ok && t.epoch == epoch {
+ existing.sent = epoch
+ t.pending[seq] = existing
+ }
+ t.mu.Unlock()
+ }
+}
+
+// ReconnectTCPConnections resumes existing handles on a caller-provided fresh
+// transport. It never creates replacements. Replay is applied before queued live
+// events and remaining pending actions are resent with their original sequences.
+func (c *Client) ReconnectTCPConnections(ctx context.Context, params ahptypes.ReconnectParams, connections []*TCPConnection) (*ahptypes.ReconnectResult, error) {
+ seen := make(map[*TCPConnection]bool)
+ params.Subscriptions = append([]string(nil), params.Subscriptions...)
+ c.tcpMu.Lock()
+ wrongIdentity := c.tcpClientID != "" && c.tcpClientID != params.ClientId
+ c.tcpMu.Unlock()
+ if wrongIdentity {
+ return nil, tcpConnectionError("resume", "new transport was initialized for a different clientId")
+ }
+ locked := make([]*TCPConnection, 0, len(connections))
+ defer func() {
+ for _, connection := range locked {
+ connection.resumeMu.Unlock()
+ }
+ }()
+ for _, connection := range connections {
+ if connection == nil || seen[connection] || !connection.resumeMu.TryLock() {
+ return nil, tcpConnectionError("resume", "duplicate, nil, or concurrently resuming handle")
+ }
+ seen[connection] = true
+ locked = append(locked, connection)
+ }
+ for _, connection := range connections {
+ connection.mu.Lock()
+ }
+ unlock := func() {
+ for _, connection := range connections {
+ connection.mu.Unlock()
+ }
+ }
+ if params.ClientId == "" || params.LastSeenServerSeq < 0 || params.LastSeenServerSeq > tcpMaxSafeInteger {
+ unlock()
+ return nil, tcpConnectionError("resume", "invalid identity or checkpoint")
+ }
+ consumerCheckpoint := params.LastSeenServerSeq
+ highest := c.nextClientSeq.Load() - 1
+ var capability *ahptypes.TcpConnectionsCapability
+ for _, connection := range connections {
+ select {
+ case <-connection.client.Done():
+ connection.online = false
+ default:
+ }
+ if connection.owner != params.ClientId || connection.terminal || connection.online || connection.client == c {
+ unlock()
+ return nil, tcpConnectionError("resume", "handles must be suspended, live, and owned by the same clientId")
+ }
+ highest = max(highest, connection.lastSeq, connection.client.nextClientSeq.Load()-1)
+ connection.client.tcpMu.Lock()
+ capability = connection.client.tcpCapability
+ connection.client.tcpMu.Unlock()
+ params.LastSeenServerSeq = min(params.LastSeenServerSeq, connection.checkpoint)
+ found := false
+ for _, resource := range params.Subscriptions {
+ found = found || resource == connection.resource
+ }
+ if !found {
+ params.Subscriptions = append(params.Subscriptions, connection.resource)
+ }
+ }
+ for {
+ next := c.nextClientSeq.Load()
+ if next > highest || c.nextClientSeq.CompareAndSwap(next, highest+1) {
+ break
+ }
+ }
+ if len(connections) != 0 {
+ c.tcpMu.Lock()
+ c.tcpClientID, c.tcpCapability = params.ClientId, capability
+ c.tcpMu.Unlock()
+ }
+ epochs := make(map[*TCPConnection]uint64, len(connections))
+ for _, connection := range connections {
+ connection.events.Close()
+ connection.client.tcpMu.Lock()
+ delete(connection.client.tcpStreams, connection)
+ connection.client.tcpMu.Unlock()
+ connection.client, connection.events = c, c.events(true, connection.resource)
+ connection.epoch++
+ epochs[connection] = connection.epoch
+ connection.resuming, connection.online = true, false
+ connection.wakeLocked()
+ }
+ registrationFailed := false
+ for _, connection := range connections {
+ registrationFailed = !c.registerTCPStream(connection) || registrationFailed
+ }
+ unlock()
+ if registrationFailed {
+ for _, connection := range connections {
+ connection.finishForClient(tcpConnectionError("disposed", "client shut down"), ahptypes.StateAction{}, false, c)
+ }
+ return nil, ErrShutdown
+ }
+ params.Channel = ahptypes.RootResourceURI
+ var result ahptypes.ReconnectResult
+ var raw json.RawMessage
+ err := c.Request(ctx, "reconnect", params, &raw)
+ if err != nil {
+ for _, connection := range connections {
+ var lag *SubscriptionLagError
+ var protocol *TransportError
+ failure := connection.events.Err()
+ if errors.As(failure, &lag) || (errors.As(failure, &protocol) && protocol.Kind == "protocol") {
+ connection.finish(failure, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ } else {
+ connection.suspend(epochs[connection])
+ }
+ }
+ return nil, err
+ }
+ if err := json.Unmarshal(raw, &result); err != nil {
+ protocol := &TransportError{Kind: "protocol", Err: err}
+ for _, connection := range connections {
+ connection.finish(protocol, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ }
+ return nil, protocol
+ }
+ replay, ok := result.Value.(*ahptypes.ReconnectReplayResult)
+ if !ok {
+ for _, connection := range connections {
+ connection.finish(tcpConnectionError("replayUnavailable", "TCP cannot recover from a snapshot"), ahptypes.StateAction{}, false)
+ }
+ return &result, nil
+ }
+ actions := replay.Actions[:0]
+ for _, action := range replay.Actions {
+ if strings.HasPrefix(action.Channel, "ahp-tcp:") || action.ServerSeq > consumerCheckpoint {
+ actions = append(actions, action)
+ }
+ }
+ clear(replay.Actions[len(actions):])
+ replay.Actions = actions
+ for _, connection := range connections {
+ for _, missing := range replay.Missing {
+ if missing == connection.resource {
+ connection.finish(tcpConnectionError("replayUnavailable", "TCP channel is missing"), ahptypes.StateAction{}, false)
+ }
+ }
+ for _, action := range replay.Actions {
+ connection.mu.Lock()
+ epoch := connection.epoch
+ connection.mu.Unlock()
+ connection.accept(action, epoch)
+ }
+ }
+ type resend struct {
+ connection *TCPConnection
+ seq int64
+ action ahptypes.StateAction
+ }
+ var pending []resend
+ for _, connection := range connections {
+ drain:
+ for {
+ select {
+ case event, open := <-connection.events.Events():
+ if !open {
+ if failure := connection.events.Err(); failure != nil {
+ var lag *SubscriptionLagError
+ var protocol *TransportError
+ if errors.As(failure, &lag) || (errors.As(failure, &protocol) && protocol.Kind == "protocol") {
+ connection.finish(failure, tcpResetAction(ahptypes.TcpResetReasonProtocolError), false)
+ } else {
+ connection.suspend(epochs[connection])
+ }
+ } else {
+ connection.suspend(epochs[connection])
+ }
+ break drain
+ }
+ if action, ok := event.Event.(SubscriptionEventAction); ok {
+ connection.mu.Lock()
+ epoch := connection.epoch
+ connection.mu.Unlock()
+ connection.accept(action.Envelope, epoch)
+ }
+ default:
+ break drain
+ }
+ }
+ connection.mu.Lock()
+ if !connection.terminal && connection.resuming {
+ for seq, action := range connection.pending {
+ pending = append(pending, resend{connection, seq, action.action})
+ }
+ }
+ connection.mu.Unlock()
+ }
+ sort.Slice(pending, func(i, j int) bool { return pending[i].seq < pending[j].seq })
+ for _, item := range pending {
+ item.connection.sendMu.Lock()
+ item.connection.mu.Lock()
+ terminal := item.connection.terminal
+ item.connection.mu.Unlock()
+ if terminal {
+ item.connection.sendMu.Unlock()
+ continue
+ }
+ err := c.Notify(ctx, "dispatchAction", ahptypes.DispatchActionParams{Channel: item.connection.resource, ClientSeq: item.seq, Action: item.action})
+ item.connection.sendMu.Unlock()
+ if err != nil {
+ for _, connection := range connections {
+ connection.suspend(epochs[connection])
+ }
+ return nil, err
+ }
+ item.connection.mu.Lock()
+ if action, exists := item.connection.pending[item.seq]; exists {
+ action.sent = item.connection.epoch
+ item.connection.pending[item.seq] = action
+ }
+ item.connection.mu.Unlock()
+ }
+ for _, connection := range connections {
+ connection.mu.Lock()
+ if !connection.terminal && connection.resuming {
+ connection.online, connection.resuming = true, false
+ connection.wakeLocked()
+ connection.start(connection.epoch, connection.events)
+ }
+ connection.mu.Unlock()
+ connection.advanceClose()
+ }
+ return &result, nil
+}
diff --git a/clients/go/ahp/tcp_test.go b/clients/go/ahp/tcp_test.go
new file mode 100644
index 000000000..20c8f54dc
--- /dev/null
+++ b/clients/go/ahp/tcp_test.go
@@ -0,0 +1,1376 @@
+package ahp
+
+import (
+ "bytes"
+ "context"
+ "encoding/base64"
+ "encoding/json"
+ "errors"
+ "io"
+ "runtime"
+ "sync"
+ "testing"
+ "time"
+
+ "github.com/microsoft/agent-host-protocol/clients/go/ahptypes"
+)
+
+type tcpTestHost struct {
+ t *testing.T
+ client *Client
+ server *memTransport
+ ctx context.Context
+ requests chan ahptypes.JsonRpcRequest
+ notifications chan ahptypes.JsonRpcNotification
+ mu sync.Mutex
+ seq int64
+ first []byte
+ holdCreate bool
+}
+
+type tcpHeldSubscribeTransport struct {
+ Transport
+ release <-chan struct{}
+}
+
+func (t tcpHeldSubscribeTransport) Send(ctx context.Context, message TransportMessage) error {
+ if err := t.Transport.Send(ctx, message); err != nil {
+ return err
+ }
+ parsed, err := message.IntoParsed()
+ if err != nil {
+ return err
+ }
+ if parsed.Request != nil && parsed.Request.Method == "subscribe" {
+ select {
+ case <-t.release:
+ case <-ctx.Done():
+ return ctx.Err()
+ }
+ }
+ return nil
+}
+
+func newTCPTestHost(t *testing.T, initialize bool, first []byte, holdCreate bool, sendGate ...<-chan struct{}) *tcpTestHost {
+ t.Helper()
+ ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
+ a, b := newMemTransportPair()
+ var transport Transport = a
+ if len(sendGate) != 0 {
+ transport = tcpHeldSubscribeTransport{Transport: a, release: sendGate[0]}
+ }
+ client, err := Connect(ctx, transport, DefaultConfig())
+ if err != nil {
+ t.Fatal(err)
+ }
+ h := &tcpTestHost{t: t, client: client, server: b, ctx: ctx, requests: make(chan ahptypes.JsonRpcRequest, 8), notifications: make(chan ahptypes.JsonRpcNotification, 64), first: first, holdCreate: holdCreate}
+ t.Cleanup(func() { client.Shutdown(context.Background()); cancel() })
+ go func() {
+ for {
+ message, err := b.Recv(ctx)
+ if err != nil {
+ return
+ }
+ parsed, err := message.IntoParsed()
+ if err != nil {
+ t.Error(err)
+ return
+ }
+ if parsed.Notification != nil {
+ select {
+ case h.notifications <- *parsed.Notification:
+ case <-ctx.Done():
+ return
+ }
+ } else if parsed.Request != nil {
+ request := *parsed.Request
+ switch request.Method {
+ case "initialize":
+ h.reply(request, map[string]any{"protocolVersion": ahptypes.ProtocolVersion, "serverSeq": 0, "snapshots": []any{}, "tcpConnections": map[string]any{"encodings": []string{"base64"}}})
+ case "subscribe":
+ if holdCreate {
+ h.requests <- request
+ continue
+ }
+ h.reply(request, h.snapshot())
+ if first != nil {
+ h.emit(&ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Offset: 0, Data: base64.StdEncoding.EncodeToString(first)}, nil, nil)
+ }
+ case "ping":
+ h.reply(request, nil)
+ default:
+ h.requests <- request
+ }
+ }
+ }
+ }()
+ if initialize {
+ if _, err := client.Initialize(ctx, "owner", ahptypes.SupportedProtocolVersions(), nil); err != nil {
+ t.Fatal(err)
+ }
+ }
+ return h
+}
+
+func (h *tcpTestHost) snapshot() ahptypes.SubscribeResult {
+ state := tcpTestState()
+ state.Input.WindowBytes, state.Input.MaximumChunkSize = 4, 3
+ state.Output.WindowBytes, state.Output.MaximumChunkSize = 4, 3
+ return ahptypes.SubscribeResult{Snapshot: &ahptypes.Snapshot{Resource: "ahp-tcp:/owned", State: ahptypes.SnapshotState{Tcp: &state}}}
+}
+
+func (h *tcpTestHost) reply(request ahptypes.JsonRpcRequest, result any) {
+ h.t.Helper()
+ raw, err := json.Marshal(result)
+ if err != nil {
+ h.t.Error(err)
+ return
+ }
+ if err := h.server.Send(h.ctx, NewParsedMessage(ahptypes.JsonRpcMessage{SuccessResponse: &ahptypes.JsonRpcSuccessResponse{JsonRpc: ahptypes.JsonRpcV2, ID: request.ID, Result: raw}})); err != nil {
+ h.t.Error(err)
+ }
+}
+
+func (h *tcpTestHost) emit(action any, origin *ahptypes.ActionOrigin, rejected *string) int64 {
+ return h.emitOn("ahp-tcp:/owned", action, origin, rejected)
+}
+
+func (h *tcpTestHost) emitOn(resource string, action any, origin *ahptypes.ActionOrigin, rejected *string) int64 {
+ h.t.Helper()
+ raw, err := json.Marshal(action)
+ if err != nil {
+ h.t.Fatal(err)
+ }
+ var typed ahptypes.StateAction
+ if err := json.Unmarshal(raw, &typed); err != nil {
+ h.t.Fatal(err)
+ }
+ h.mu.Lock()
+ defer h.mu.Unlock()
+ h.seq++
+ body, err := json.Marshal(ahptypes.ActionEnvelope{Channel: resource, ServerSeq: h.seq, Action: typed, Origin: origin, RejectionReason: rejected})
+ if err != nil {
+ h.t.Fatal(err)
+ }
+ if err := h.server.Send(h.ctx, NewParsedMessage(ahptypes.JsonRpcMessage{Notification: &ahptypes.JsonRpcNotification{JsonRpc: ahptypes.JsonRpcV2, Method: "action", Params: body}})); err != nil {
+ h.t.Fatal(err)
+ }
+ return h.seq
+}
+
+func (h *tcpTestHost) unrelatedBurst() {
+ h.t.Helper()
+ barrier := h.client.AttachSubscription("ahp-session:/barrier")
+ defer barrier.Close()
+ for i := 0; i < 16; i++ {
+ h.emitOn("ahp-session:/other", &ahptypes.SessionTitleChangedAction{Type: ahptypes.ActionTypeSessionTitleChanged, Title: "busy"}, nil, nil)
+ h.emitOn("ahp-tcp:/other", &ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Offset: int64(i), Data: "AA=="}, nil, nil)
+ }
+ h.emitOn(barrier.URI(), &ahptypes.SessionTitleChangedAction{Type: ahptypes.ActionTypeSessionTitleChanged, Title: "barrier"}, nil, nil)
+ select {
+ case <-barrier.Events():
+ case <-h.ctx.Done():
+ h.t.Fatal("traffic barrier timed out")
+ }
+}
+
+func TestOwnedTCPScopedCreationAndActiveTraffic(t *testing.T) {
+ release := make(chan struct{})
+ h := newTCPTestHost(t, true, nil, true, release)
+ h.client.cfg.SubscriptionBuffer = 2
+ type opened struct {
+ connection *TCPConnection
+ err error
+ }
+ result := make(chan opened, 1)
+ go func() {
+ connection, err := h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", tcpCreateParams())
+ result <- opened{connection, err}
+ }()
+ request := <-h.requests
+ h.unrelatedBurst()
+ h.reply(request, h.snapshot())
+ h.emit(&ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Offset: 0, Data: "Bw=="}, nil, nil)
+ // The request writer is still blocked: the child route must already exist.
+ h.unrelatedBurst()
+ close(release)
+ out := <-result
+ if out.err != nil {
+ t.Fatal(out.err)
+ }
+ c := out.connection
+ c.mu.Lock()
+ h.unrelatedBurst()
+ c.mu.Unlock()
+ data, err := c.Read(h.ctx)
+ if err != nil || !bytes.Equal(data, []byte{7}) {
+ t.Fatalf("first child action lost: %v, %v", data, err)
+ }
+ credit := h.dispatch()
+ h.emit(credit.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: credit.ClientSeq}, nil)
+ h.emit(&ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Offset: 1, Data: "CA=="}, nil, nil)
+ data, err = c.Read(h.ctx)
+ if err != nil || !bytes.Equal(data, []byte{8}) {
+ t.Fatalf("active stream failed after unrelated traffic: %v, %v", data, err)
+ }
+ h.dispatch()
+ if err := c.Dispose(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ h.dispatch()
+ h.notification("unsubscribe")
+}
+
+func tcpCreateParams() ahptypes.TcpConnectionSubscription {
+ return ahptypes.TcpConnectionSubscription{Type: "tcpConnection", Host: "localhost", Port: 3000, Encoding: ahptypes.TcpDataEncodingBase64, ReceiveWindowBytes: 4, MaximumChunkSize: 3}
+}
+
+func (h *tcpTestHost) open() *TCPConnection {
+ h.t.Helper()
+ connection, err := h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", tcpCreateParams())
+ if err != nil {
+ h.t.Fatal(err)
+ }
+ return connection
+}
+
+func (h *tcpTestHost) notification(method string) ahptypes.JsonRpcNotification {
+ h.t.Helper()
+ select {
+ case notification := <-h.notifications:
+ if notification.Method != method {
+ h.t.Fatalf("got %s, want %s", notification.Method, method)
+ }
+ return notification
+ case <-h.ctx.Done():
+ h.t.Fatal("notification timeout")
+ return ahptypes.JsonRpcNotification{}
+ }
+}
+
+func (h *tcpTestHost) dispatch() ahptypes.DispatchActionParams {
+ h.t.Helper()
+ n := h.notification("dispatchAction")
+ var params ahptypes.DispatchActionParams
+ if err := json.Unmarshal(n.Params, ¶ms); err != nil {
+ h.t.Fatal(err)
+ }
+ return params
+}
+
+func (h *tcpTestHost) closeGracefully(c *TCPConnection) {
+ h.t.Helper()
+ done := make(chan error, 1)
+ go func() { done <- c.Close(h.ctx) }()
+ action := h.dispatch()
+ if _, ok := action.Action.Value.(*ahptypes.TcpClientCloseAction); !ok {
+ h.t.Fatal("missing client close")
+ }
+ h.emit(action.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: action.ClientSeq}, nil)
+ h.emit(&ahptypes.TcpHostCloseAction{Type: ahptypes.ActionTypeTcpHostClose}, nil, nil)
+ if err := <-done; err != nil {
+ h.t.Fatal(err)
+ }
+ h.notification("unsubscribe")
+}
+
+func waitTCPState(t *testing.T, ctx context.Context, c *TCPConnection, predicate func(*TCPConnection) bool) {
+ t.Helper()
+ for {
+ c.mu.Lock()
+ ok := predicate(c)
+ changed := c.changed
+ c.mu.Unlock()
+ if ok {
+ return
+ }
+ if err := waitTCP(ctx, changed); err != nil {
+ t.Fatal(err)
+ }
+ }
+}
+
+func TestOwnedTCPFlowControlAndHalfClose(t *testing.T) {
+ h := newTCPTestHost(t, true, []byte("abc"), false)
+ c := h.open()
+ waitTCPState(t, h.ctx, c, func(c *TCPConnection) bool { return c.state.Output.ReceivedBytes == 3 })
+ select {
+ case event := <-h.notifications:
+ t.Fatalf("receipt released credit: %s", event.Method)
+ default:
+ }
+ data, err := c.Read(h.ctx)
+ if err != nil || string(data) != "abc" {
+ t.Fatalf("read %q %v", data, err)
+ }
+ credit := h.dispatch()
+ if a, ok := credit.Action.Value.(*ahptypes.TcpDataConsumedAction); !ok || a.ConsumedBytes != 3 {
+ t.Fatalf("bad delivered credit: %+v", credit)
+ }
+ done := make(chan error, 1)
+ go func() {
+ n, err := c.Write(h.ctx, []byte("abcdef"))
+ if err == nil && n != 6 {
+ err = errors.New("short write")
+ }
+ done <- err
+ }()
+ first, second := h.dispatch(), h.dispatch()
+ a, b := first.Action.Value.(*ahptypes.TcpInputAction), second.Action.Value.(*ahptypes.TcpInputAction)
+ if a.Offset != 0 || a.Data != "YWJj" || b.Offset != 3 || b.Data != "ZA==" {
+ t.Fatalf("chunks: %+v %+v", a, b)
+ }
+ c.mu.Lock()
+ received := c.state.Input.ReceivedBytes
+ c.mu.Unlock()
+ if received != 0 {
+ t.Fatal("optimistic reducer mutation")
+ }
+ select {
+ case err := <-done:
+ t.Fatalf("write did not block: %v", err)
+ default:
+ }
+ hostAction := h.emit(&ahptypes.TcpInputConsumedAction{Type: ahptypes.ActionTypeTcpInputConsumed, ConsumedBytes: 0},
+ &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: first.ClientSeq}, nil)
+ waitTCPState(t, h.ctx, c, func(c *TCPConnection) bool { return c.checkpoint >= hostAction })
+ c.mu.Lock()
+ _, retained := c.pending[first.ClientSeq]
+ c.mu.Unlock()
+ if !retained {
+ t.Fatal("host action cleared pending client action")
+ }
+ h.emit(a, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: first.ClientSeq}, nil)
+ h.emit(a, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: first.ClientSeq}, nil)
+ h.emit(b, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: second.ClientSeq}, nil)
+ h.emit(&ahptypes.TcpInputConsumedAction{Type: ahptypes.ActionTypeTcpInputConsumed, ConsumedBytes: 4}, nil, nil)
+ last := h.dispatch()
+ lastInput := last.Action.Value.(*ahptypes.TcpInputAction)
+ if lastInput.Offset != 4 || lastInput.Data != "ZWY=" {
+ t.Fatalf("resumed write %+v", lastInput)
+ }
+ if err := <-done; err != nil {
+ t.Fatal(err)
+ }
+ drained := make(chan error, 1)
+ go func() { drained <- c.Drain(h.ctx) }()
+ h.emit(lastInput, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: last.ClientSeq}, nil)
+ h.emit(&ahptypes.TcpInputConsumedAction{Type: ahptypes.ActionTypeTcpInputConsumed, ConsumedBytes: 6}, nil, nil)
+ if err := <-drained; err != nil {
+ t.Fatal(err)
+ }
+ h.emit(&ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Offset: 0, Data: "YWJj"}, nil, nil)
+ h.emit(&ahptypes.TcpDataEofAction{Type: ahptypes.ActionTypeTcpDataEof, FinalOffset: 3}, nil, nil)
+ if _, err := c.Read(h.ctx); !errors.Is(err, io.EOF) {
+ t.Fatalf("duplicate data delivered or missing EOF: %v", err)
+ }
+ if err := c.End(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ end := h.dispatch()
+ if a, ok := end.Action.Value.(*ahptypes.TcpInputEofAction); !ok || a.FinalOffset != 6 {
+ t.Fatalf("bad EOF %+v", end)
+ }
+ h.emit(credit.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: credit.ClientSeq}, nil)
+ h.emit(end.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: end.ClientSeq}, nil)
+ h.closeGracefully(c)
+ if err := c.Dispose(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ if err := h.client.Ping(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case n := <-h.notifications:
+ t.Fatalf("duplicate cleanup %s", n.Method)
+ default:
+ }
+}
+
+func TestOwnedTCPNegotiatedOutputLimits(t *testing.T) {
+ for _, limits := range []struct {
+ window, chunk int64
+ valid bool
+ }{
+ {2, 1, true}, {4, 2, true}, {0, 1, false}, {2, 0, false},
+ {5, 1, false}, {4, 3, false}, {1, 2, false},
+ } {
+ h := newTCPTestHost(t, true, nil, true)
+ create := tcpCreateParams()
+ create.MaximumChunkSize = 2
+ var connection *TCPConnection
+ var err error
+ done := make(chan struct{})
+ go func() { connection, err = h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", create); close(done) }()
+ req := <-h.requests
+ snapshot := h.snapshot()
+ snapshot.Snapshot.State.Tcp.Output.WindowBytes = limits.window
+ snapshot.Snapshot.State.Tcp.Output.MaximumChunkSize = limits.chunk
+ h.reply(req, snapshot)
+ <-done
+ if (err == nil) != limits.valid {
+ t.Fatalf("limits %d/%d valid=%v: %v", limits.window, limits.chunk, limits.valid, err)
+ }
+ if connection != nil {
+ if connection.state.Output.WindowBytes != limits.window || connection.state.Output.MaximumChunkSize != limits.chunk {
+ t.Fatal("negotiated limits not retained")
+ }
+ if err := connection.Dispose(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ }
+ }
+}
+
+func TestOwnedTCPLimitsUInt32(t *testing.T) {
+ const limit int64 = 4_294_967_295
+ for _, field := range []string{"boundary", "requestWindow", "requestChunk", "inputWindow", "inputChunk", "outputWindow", "outputChunk"} {
+ t.Run(field, func(t *testing.T) {
+ h := newTCPTestHost(t, true, nil, true)
+ create := tcpCreateParams()
+ create.ReceiveWindowBytes, create.MaximumChunkSize = limit, limit
+ snapshot := h.snapshot()
+ state := snapshot.Snapshot.State.Tcp
+ state.Input.WindowBytes, state.Input.MaximumChunkSize = limit, limit
+ state.Output.WindowBytes, state.Output.MaximumChunkSize = limit, limit
+ switch field {
+ case "requestWindow":
+ create.ReceiveWindowBytes++
+ case "requestChunk":
+ create.MaximumChunkSize++
+ case "inputWindow":
+ state.Input.WindowBytes++
+ case "inputChunk":
+ state.Input.MaximumChunkSize++
+ case "outputWindow":
+ state.Output.WindowBytes++
+ case "outputChunk":
+ state.Output.MaximumChunkSize++
+ }
+ var connection *TCPConnection
+ var err error
+ done := make(chan struct{})
+ go func() { connection, err = h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", create); close(done) }()
+ select {
+ case req := <-h.requests:
+ if field == "requestWindow" || field == "requestChunk" {
+ t.Error("out-of-UInt32 request was sent")
+ }
+ h.reply(req, snapshot)
+ <-done
+ case <-done:
+ case <-h.ctx.Done():
+ t.Fatal("creation hung")
+ }
+ if (err == nil) != (field == "boundary") {
+ t.Errorf("UInt32 validation result: %v", err)
+ }
+ if connection != nil {
+ if err := connection.Dispose(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ }
+ })
+ }
+}
+
+func TestOwnedTCPReconnectFiltersOrdinaryReplay(t *testing.T) {
+ old := newTCPTestHost(t, true, nil, false)
+ c := old.open()
+ if err := old.client.ShutdownPreservingTCP(old.ctx); err != nil {
+ t.Fatal(err)
+ }
+ fresh := newTCPTestHost(t, false, nil, false)
+ done := make(chan *ahptypes.ReconnectResult, 1)
+ go func() {
+ result, err := fresh.client.ReconnectTCPConnections(fresh.ctx, ahptypes.ReconnectParams{ClientId: "owner", LastSeenServerSeq: 20}, []*TCPConnection{c})
+ if err != nil {
+ t.Error(err)
+ }
+ done <- result
+ }()
+ req := <-fresh.requests
+ var params ahptypes.ReconnectParams
+ if err := json.Unmarshal(req.Params, ¶ms); err != nil {
+ t.Fatal(err)
+ }
+ if params.LastSeenServerSeq != 0 {
+ t.Fatal("wire checkpoint not clamped for TCP")
+ }
+ fresh.reply(req, &ahptypes.ReconnectReplayResult{Missing: []string{}, Actions: []ahptypes.ActionEnvelope{
+ {Channel: c.Resource(), ServerSeq: 11, Action: ahptypes.StateAction{Value: &ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Data: "eA=="}}},
+ {Channel: "ahp-terminal:/t", ServerSeq: 15, Action: ahptypes.StateAction{Value: &ahptypes.TerminalDataAction{Type: ahptypes.ActionTypeTerminalData, Data: "hello"}}},
+ {Channel: "ahp-terminal:/t", ServerSeq: 21, Action: ahptypes.StateAction{Value: &ahptypes.TerminalDataAction{Type: ahptypes.ActionTypeTerminalData, Data: "!"}}},
+ }})
+ result := <-done
+ if result == nil {
+ t.Fatal("reconnect failed")
+ }
+ text := "hello"
+ for _, action := range result.Value.(*ahptypes.ReconnectReplayResult).Actions {
+ if data, ok := action.Action.Value.(*ahptypes.TerminalDataAction); ok {
+ text += data.Data
+ }
+ }
+ if text != "hello!" {
+ t.Fatalf("ordinary replay duplicated output: %q", text)
+ }
+ if data, err := c.Read(fresh.ctx); err != nil || string(data) != "x" {
+ t.Fatalf("TCP replay lost: %q %v", data, err)
+ }
+ if err := c.Dispose(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+}
+
+func TestOwnedTCPCloseRetainsCrossingDataUntilAcknowledged(t *testing.T) {
+ h := newTCPTestHost(t, true, nil, false)
+ c := h.open()
+ done := make(chan error, 1)
+ go func() { done <- c.Close(h.ctx) }()
+ closeAction := h.dispatch()
+ if _, ok := closeAction.Action.Value.(*ahptypes.TcpClientCloseAction); !ok {
+ t.Fatal("missing client close")
+ }
+ if c.IsClosed() {
+ t.Fatal("local close released ownership before host close")
+ }
+ h.emit(closeAction.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: closeAction.ClientSeq}, nil)
+ h.emit(&ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Offset: 0, Data: "YWI="}, nil, nil)
+ seq := h.emit(&ahptypes.TcpHostCloseAction{Type: ahptypes.ActionTypeTcpHostClose}, nil, nil)
+ waitTCPState(t, h.ctx, c, func(c *TCPConnection) bool { return c.checkpoint == seq })
+ if err := h.client.Ping(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case n := <-h.notifications:
+ t.Fatalf("cleanup before output drain: %s", n.Method)
+ default:
+ }
+ data, err := c.Read(h.ctx)
+ if err != nil || string(data) != "ab" {
+ t.Fatalf("crossing data %q: %v", data, err)
+ }
+ credit := h.dispatch()
+ if a, ok := credit.Action.Value.(*ahptypes.TcpDataConsumedAction); !ok || a.ConsumedBytes != 2 {
+ t.Fatal("missing read credit")
+ }
+ if _, err := c.Read(h.ctx); !errors.Is(err, io.EOF) {
+ t.Fatal("host close did not finish output")
+ }
+ if c.IsClosed() {
+ t.Fatal("released before credit acknowledgement")
+ }
+ h.emit(credit.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: credit.ClientSeq}, nil)
+ if err := <-done; err != nil {
+ t.Fatal(err)
+ }
+ h.notification("unsubscribe")
+ if !c.IsClosed() {
+ t.Fatal("handshake did not release ownership")
+ }
+}
+
+func TestOwnedTCPClosingResumesAndResetWakesClose(t *testing.T) {
+ old := newTCPTestHost(t, true, nil, false)
+ c := old.open()
+ closed := make(chan error, 1)
+ go func() { closed <- c.Close(old.ctx) }()
+ original := old.dispatch()
+ if err := old.client.ShutdownPreservingTCP(old.ctx); err != nil {
+ t.Fatal(err)
+ }
+ fresh := newTCPTestHost(t, false, nil, false)
+ resumed := make(chan error, 1)
+ go func() {
+ _, err := fresh.client.ReconnectTCPConnections(fresh.ctx, ahptypes.ReconnectParams{ClientId: "owner"}, []*TCPConnection{c})
+ resumed <- err
+ }()
+ req := <-fresh.requests
+ fresh.reply(req, &ahptypes.ReconnectReplayResult{Actions: []ahptypes.ActionEnvelope{}, Missing: []string{}})
+ if err := <-resumed; err != nil {
+ t.Fatal(err)
+ }
+ action := fresh.dispatch()
+ if _, ok := action.Action.Value.(*ahptypes.TcpClientCloseAction); !ok || action.ClientSeq != original.ClientSeq {
+ t.Fatal("closing resume replaced or renumbered client close")
+ }
+ select {
+ case err := <-closed:
+ t.Fatalf("close completed before handshake: %v", err)
+ default:
+ }
+ fresh.emit(&ahptypes.TcpHostResetAction{Type: ahptypes.ActionTypeTcpHostReset, Reason: ahptypes.TcpResetReasonConnectionReset}, nil, nil)
+ select {
+ case err := <-closed:
+ var tcpErr *TCPConnectionError
+ if !errors.As(err, &tcpErr) || tcpErr.Reason != "reset" {
+ t.Fatalf("reset close result: %v", err)
+ }
+ case <-fresh.ctx.Done():
+ t.Fatal("reset left close blocked")
+ }
+ fresh.notification("unsubscribe")
+ if err := c.Dispose(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+ if err := fresh.client.Ping(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case n := <-fresh.notifications:
+ t.Fatalf("duplicate cleanup: %s", n.Method)
+ default:
+ }
+}
+
+func TestOwnedTCPResetAndRejectionWakeWaiters(t *testing.T) {
+ for _, reject := range []bool{false, true} {
+ t.Run(map[bool]string{false: "reset", true: "rejection"}[reject], func(t *testing.T) {
+ h := newTCPTestHost(t, true, nil, false)
+ c := h.open()
+ if _, err := c.Write(h.ctx, []byte("abcd")); err != nil {
+ t.Fatal(err)
+ }
+ first := h.dispatch()
+ h.dispatch()
+ results := make(chan error, 3)
+ go func() { _, err := c.Read(h.ctx); results <- err }()
+ go func() { _, err := c.Write(h.ctx, []byte("x")); results <- err }()
+ go func() { results <- c.Drain(h.ctx) }()
+ if reject {
+ reason := ""
+ h.emit(first.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: first.ClientSeq}, &reason)
+ } else {
+ h.emit(&ahptypes.TcpHostResetAction{Type: ahptypes.ActionTypeTcpHostReset, Reason: ahptypes.TcpResetReasonConnectionReset}, nil, nil)
+ }
+ for i := 0; i < 3; i++ {
+ select {
+ case err := <-results:
+ if err == nil {
+ t.Fatal("waiter succeeded after terminal failure")
+ }
+ case <-h.ctx.Done():
+ t.Fatal("waiter hung")
+ }
+ }
+ if reject {
+ if a, ok := h.dispatch().Action.Value.(*ahptypes.TcpClientResetAction); !ok || a.Reason != ahptypes.TcpResetReasonProtocolError {
+ t.Fatal("missing protocol reset")
+ }
+ }
+ h.notification("unsubscribe")
+ })
+ }
+}
+
+func TestOwnedTCPReconnectRetainsHandleAndPendingSequences(t *testing.T) {
+ for _, acknowledged := range []bool{false, true} {
+ t.Run(map[bool]string{false: "resend", true: "acknowledged"}[acknowledged], func(t *testing.T) {
+ old := newTCPTestHost(t, true, nil, false)
+ c := old.open()
+ if _, err := c.Write(old.ctx, []byte("ab")); err != nil {
+ t.Fatal(err)
+ }
+ original := old.dispatch()
+ seq := old.emit(&ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Offset: 0, Data: "eHk="}, nil, nil)
+ waitTCPState(t, old.ctx, c, func(c *TCPConnection) bool { return c.checkpoint == seq })
+ if err := old.client.ShutdownPreservingTCP(old.ctx); err != nil {
+ t.Fatal(err)
+ }
+ fresh := newTCPTestHost(t, false, nil, false)
+ done := make(chan error, 1)
+ go func() {
+ _, err := fresh.client.ReconnectTCPConnections(fresh.ctx, ahptypes.ReconnectParams{ClientId: "owner", LastSeenServerSeq: 999, Subscriptions: []string{"ahp-session:/s1"}}, []*TCPConnection{c})
+ done <- err
+ }()
+ var request ahptypes.JsonRpcRequest
+ select {
+ case request = <-fresh.requests:
+ case <-fresh.ctx.Done():
+ t.Fatal("missing reconnect")
+ }
+ var params ahptypes.ReconnectParams
+ if err := json.Unmarshal(request.Params, ¶ms); err != nil {
+ t.Fatal(err)
+ }
+ if request.Method != "reconnect" || params.LastSeenServerSeq != seq || len(params.Subscriptions) != 2 {
+ t.Fatalf("unsafe reconnect %+v", params)
+ }
+ actions := []ahptypes.ActionEnvelope{}
+ if acknowledged {
+ actions = append(actions, ahptypes.ActionEnvelope{Channel: c.Resource(), ServerSeq: 2, Action: original.Action, Origin: &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: original.ClientSeq}})
+ }
+ fresh.reply(request, &ahptypes.ReconnectReplayResult{Actions: actions, Missing: []string{}})
+ if err := <-done; err != nil {
+ t.Fatal(err)
+ }
+ if err := old.client.Shutdown(old.ctx); err != nil {
+ t.Fatal(err)
+ }
+ if c.IsClosed() {
+ t.Fatal("old client shutdown disposed rebound stream")
+ }
+ if !acknowledged {
+ replayed := fresh.dispatch()
+ if replayed.ClientSeq != original.ClientSeq || replayed.Channel != original.Channel {
+ t.Fatal("pending action was renumbered")
+ }
+ if a := replayed.Action.Value.(*ahptypes.TcpInputAction); a.Offset != 0 || a.Data != "YWI=" {
+ t.Fatal("pending bytes changed")
+ }
+ }
+ data, err := c.Read(fresh.ctx)
+ if err != nil || string(data) != "xy" {
+ t.Fatalf("retained read %q %v", data, err)
+ }
+ credit := fresh.dispatch()
+ if _, ok := credit.Action.Value.(*ahptypes.TcpDataConsumedAction); !ok || credit.ClientSeq <= original.ClientSeq {
+ t.Fatalf("ack replayed or sequence reused: %+v", credit)
+ }
+ fresh.client.tcpMu.Lock()
+ retainedIdentity := fresh.client.tcpClientID == "owner" && fresh.client.tcpCapability != nil
+ fresh.client.tcpMu.Unlock()
+ if !retainedIdentity {
+ t.Fatal("resume lost negotiated TCP capability")
+ }
+ if err := c.Dispose(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+ fresh.dispatch()
+ fresh.notification("unsubscribe")
+ })
+ }
+}
+
+func TestOwnedTCPReconnectAppliesReplayBeforeQueuedLiveData(t *testing.T) {
+ old := newTCPTestHost(t, true, nil, false)
+ c := old.open()
+ if _, err := c.Write(old.ctx, []byte("xy")); err != nil {
+ t.Fatal(err)
+ }
+ original := old.dispatch()
+ if err := old.client.ShutdownPreservingTCP(old.ctx); err != nil {
+ t.Fatal(err)
+ }
+ fresh := newTCPTestHost(t, false, nil, false)
+ fresh.client.cfg.SubscriptionBuffer = 2
+ done := make(chan error, 1)
+ go func() {
+ _, err := fresh.client.ReconnectTCPConnections(fresh.ctx,
+ ahptypes.ReconnectParams{ClientId: "owner", LastSeenServerSeq: 999}, []*TCPConnection{c})
+ done <- err
+ }()
+ request := <-fresh.requests
+ fresh.mu.Lock()
+ fresh.seq = 1
+ fresh.mu.Unlock()
+ fresh.unrelatedBurst()
+ func() {
+ // Hold application of replay while the receive loop queues later live frames.
+ c.mu.Lock()
+ defer c.mu.Unlock()
+ if !c.resuming || c.online {
+ t.Fatal("stream operations were released before replay")
+ }
+ fresh.reply(request, &ahptypes.ReconnectReplayResult{
+ Actions: []ahptypes.ActionEnvelope{{
+ Channel: c.resource, ServerSeq: 1,
+ Action: ahptypes.StateAction{Value: &ahptypes.TcpDataAction{
+ Type: ahptypes.ActionTypeTcpData, Offset: 0, Data: "YWI=",
+ }},
+ }},
+ Missing: []string{},
+ })
+ fresh.unrelatedBurst()
+ fresh.emit(&ahptypes.TcpDataAction{Type: ahptypes.ActionTypeTcpData, Offset: 2, Data: "Y2Q="}, nil, nil)
+ fresh.emit(&ahptypes.TcpDataEofAction{Type: ahptypes.ActionTypeTcpDataEof, FinalOffset: 4}, nil, nil)
+ if err := fresh.client.Ping(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case n := <-fresh.notifications:
+ t.Fatalf("outgoing action before replay applied: %s", n.Method)
+ default:
+ }
+ }()
+ if err := <-done; err != nil {
+ t.Fatal(err)
+ }
+ resend := fresh.dispatch()
+ if resend.ClientSeq != original.ClientSeq {
+ t.Fatal("resume renumbered pending input")
+ }
+ if input, ok := resend.Action.Value.(*ahptypes.TcpInputAction); !ok || input.Offset != 0 || input.Data != "eHk=" {
+ t.Fatalf("incorrect resumed input: %+v", resend)
+ }
+ if err := fresh.client.Ping(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case n := <-fresh.notifications:
+ t.Fatalf("receipt released credit or duplicated input: %s", n.Method)
+ default:
+ }
+ for i, expected := range []string{"ab", "cd"} {
+ data, err := c.Read(fresh.ctx)
+ if err != nil || string(data) != expected {
+ t.Fatalf("read %q: %v", data, err)
+ }
+ credit := fresh.dispatch()
+ if action, ok := credit.Action.Value.(*ahptypes.TcpDataConsumedAction); !ok ||
+ action.ConsumedBytes != int64((i+1)*2) || credit.ClientSeq <= original.ClientSeq {
+ t.Fatalf("incorrect resumed credit: %+v", credit)
+ }
+ }
+ if _, err := c.Read(fresh.ctx); !errors.Is(err, io.EOF) {
+ t.Fatalf("missing live EOF: %v", err)
+ }
+ if err := c.Dispose(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+ fresh.dispatch()
+ fresh.notification("unsubscribe")
+}
+
+func TestOwnedTCPReconnectSnapshotAndMissingAreTerminal(t *testing.T) {
+ for _, snapshot := range []bool{false, true} {
+ t.Run(map[bool]string{false: "missing", true: "snapshot"}[snapshot], func(t *testing.T) {
+ old := newTCPTestHost(t, true, nil, false)
+ c := old.open()
+ old.client.ShutdownPreservingTCP(old.ctx)
+ fresh := newTCPTestHost(t, false, nil, false)
+ done := make(chan error, 1)
+ go func() {
+ _, err := fresh.client.ReconnectTCPConnections(fresh.ctx, ahptypes.ReconnectParams{ClientId: "owner"}, []*TCPConnection{c})
+ done <- err
+ }()
+ request := <-fresh.requests
+ if snapshot {
+ fresh.reply(request, &ahptypes.ReconnectSnapshotResult{Snapshots: []ahptypes.Snapshot{*fresh.snapshot().Snapshot}})
+ } else {
+ fresh.reply(request, &ahptypes.ReconnectReplayResult{Actions: []ahptypes.ActionEnvelope{}, Missing: []string{c.Resource()}})
+ }
+ if err := <-done; err != nil {
+ t.Fatal(err)
+ }
+ _, err := c.Read(fresh.ctx)
+ var tcpErr *TCPConnectionError
+ if !errors.As(err, &tcpErr) || tcpErr.Reason != "replayUnavailable" {
+ t.Fatalf("restored invalid stream: %v", err)
+ }
+ fresh.notification("unsubscribe")
+ next := fresh.open()
+ if err := next.Dispose(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+ })
+ }
+}
+
+func TestOwnedTCPCancelledCreationCleansKnownChild(t *testing.T) {
+ h := newTCPTestHost(t, true, nil, true)
+ ctx, cancel := context.WithCancel(h.ctx)
+ done := make(chan error, 1)
+ go func() { _, err := h.client.OpenTCPConnection(ctx, "ahp-session:/s1", tcpCreateParams()); done <- err }()
+ request := <-h.requests
+ cancel()
+ if err := <-done; !errors.Is(err, context.Canceled) {
+ t.Fatal(err)
+ }
+ h.reply(request, h.snapshot())
+ h.dispatch()
+ n := h.notification("unsubscribe")
+ var params ahptypes.UnsubscribeParams
+ if err := json.Unmarshal(n.Params, ¶ms); err != nil {
+ t.Fatal(err)
+ }
+ if params.Channel != "ahp-tcp:/owned" {
+ t.Fatalf("unsubscribed parent: %s", params.Channel)
+ }
+}
+
+func TestOwnedTCPClientShutdownDisposesLiveAndSuspendedStreams(t *testing.T) {
+ for _, suspended := range []bool{false, true} {
+ t.Run(map[bool]string{false: "live", true: "suspended"}[suspended], func(t *testing.T) {
+ h := newTCPTestHost(t, true, nil, false)
+ c := h.open()
+ if _, err := c.Write(h.ctx, []byte("abcd")); err != nil {
+ t.Fatal(err)
+ }
+ h.dispatch()
+ h.dispatch()
+ results := make(chan error, 3)
+ go func() { _, err := c.Read(h.ctx); results <- err }()
+ go func() { _, err := c.Write(h.ctx, []byte("x")); results <- err }()
+ go func() { results <- c.Drain(h.ctx) }()
+ for {
+ c.mu.Lock()
+ waiting := c.writing
+ c.mu.Unlock()
+ if waiting {
+ break
+ }
+ select {
+ case <-h.ctx.Done():
+ t.Fatal("writer never blocked")
+ default:
+ runtime.Gosched()
+ }
+ }
+ if suspended {
+ if err := h.client.ShutdownPreservingTCP(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ if c.IsClosed() {
+ t.Fatal("preserving shutdown disposed stream")
+ }
+ }
+ if err := h.client.Shutdown(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ if !c.IsClosed() {
+ t.Fatal("shutdown left stream suspended")
+ }
+ if h.client.registerTCPStream(c) {
+ t.Fatal("shutdown permitted new stream registration")
+ }
+ for i := 0; i < 3; i++ {
+ select {
+ case err := <-results:
+ var tcpErr *TCPConnectionError
+ if !errors.As(err, &tcpErr) || tcpErr.Reason != "disposed" {
+ t.Fatalf("waiter: %v", err)
+ }
+ case <-h.ctx.Done():
+ t.Fatal("shutdown left waiter blocked")
+ }
+ }
+ if err := h.client.Shutdown(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ })
+ }
+}
+
+func TestOwnedTCPCloseWaitsForAcceptedInput(t *testing.T) {
+ h := newTCPTestHost(t, true, nil, false)
+ c := h.open()
+ if _, err := c.Write(h.ctx, []byte("ab")); err != nil {
+ t.Fatal(err)
+ }
+ input := h.dispatch()
+ done := make(chan error, 1)
+ go func() { done <- c.Close(h.ctx) }()
+ waitTCPState(t, h.ctx, c, func(c *TCPConnection) bool { return c.ending })
+ select {
+ case err := <-done:
+ t.Fatalf("closed before drain: %v", err)
+ default:
+ }
+ select {
+ case n := <-h.notifications:
+ t.Fatalf("premature %s", n.Method)
+ default:
+ }
+ h.emit(input.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: input.ClientSeq}, nil)
+ h.emit(&ahptypes.TcpInputConsumedAction{Type: ahptypes.ActionTypeTcpInputConsumed, ConsumedBytes: 2}, nil, nil)
+ closeAction := h.dispatch()
+ if _, ok := closeAction.Action.Value.(*ahptypes.TcpClientCloseAction); !ok {
+ t.Fatal("missing final close")
+ }
+ h.emit(closeAction.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: closeAction.ClientSeq}, nil)
+ h.emit(&ahptypes.TcpHostCloseAction{Type: ahptypes.ActionTypeTcpHostClose}, nil, nil)
+ if err := <-done; err != nil {
+ t.Fatal(err)
+ }
+ h.notification("unsubscribe")
+}
+
+func TestOwnedTCPLargePayloadEncoding(t *testing.T) {
+ h := newTCPTestHost(t, true, nil, true)
+ const size = 4 * 1024 * 1024
+ params := tcpCreateParams()
+ params.ReceiveWindowBytes, params.MaximumChunkSize = size, size
+ var c *TCPConnection
+ var err error
+ done := make(chan struct{})
+ go func() { c, err = h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", params); close(done) }()
+ req := <-h.requests
+ snapshot := h.snapshot()
+ snapshot.Snapshot.State.Tcp.Input.WindowBytes, snapshot.Snapshot.State.Tcp.Input.MaximumChunkSize = size, size
+ snapshot.Snapshot.State.Tcp.Output.WindowBytes, snapshot.Snapshot.State.Tcp.Output.MaximumChunkSize = size, size
+ h.reply(req, snapshot)
+ <-done
+ if err != nil {
+ t.Fatal(err)
+ }
+ payload := bytes.Repeat([]byte{0xab}, size)
+ if n, err := c.Write(h.ctx, payload); err != nil || n != size {
+ t.Fatalf("write %d: %v", n, err)
+ }
+ input := h.dispatch()
+ action := input.Action.Value.(*ahptypes.TcpInputAction)
+ decoded, err := base64.StdEncoding.Strict().DecodeString(action.Data)
+ if err != nil || !bytes.Equal(decoded, payload) || action.Offset != 0 {
+ t.Fatalf("incorrect large payload: %v", err)
+ }
+ h.emit(action, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: input.ClientSeq}, nil)
+ h.emit(&ahptypes.TcpInputConsumedAction{Type: ahptypes.ActionTypeTcpInputConsumed, ConsumedBytes: size}, nil, nil)
+ h.closeGracefully(c)
+}
+func TestOwnedTCPStrictLossWakesReader(t *testing.T) {
+ for _, overflow := range []bool{false, true} {
+ t.Run(map[bool]string{false: "decode", true: "overflow"}[overflow], func(t *testing.T) {
+ h := newTCPTestHost(t, true, nil, false)
+ c := h.open()
+ done := make(chan error, 1)
+ go func() { _, err := c.Read(h.ctx); done <- err }()
+ if overflow {
+ c.mu.Lock()
+ for i := 0; i < h.client.cfg.SubscriptionBuffer+2; i++ {
+ h.emit(map[string]any{"type": "future/tcp"}, nil, nil)
+ }
+ err := h.client.Ping(h.ctx)
+ c.mu.Unlock()
+ if err != nil {
+ t.Fatal(err)
+ }
+ } else if err := h.server.Send(h.ctx, NewTextMessage("{")); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case err := <-done:
+ var lag *SubscriptionLagError
+ var protocol *TransportError
+ if overflow && !errors.As(err, &lag) || !overflow && (!errors.As(err, &protocol) || protocol.Kind != "protocol") {
+ t.Fatalf("loss was not surfaced: %v", err)
+ }
+ case <-h.ctx.Done():
+ t.Fatal("loss did not wake reader")
+ }
+ if _, ok := h.dispatch().Action.Value.(*ahptypes.TcpClientResetAction); !ok {
+ t.Fatal("missing loss reset")
+ }
+ h.notification("unsubscribe")
+ if _, err := c.Read(h.ctx); err == nil {
+ t.Fatal("failed stream resumed")
+ }
+ })
+ }
+}
+
+func TestOwnedTCPValidationAndInvalidCreationCleanup(t *testing.T) {
+ h := newTCPTestHost(t, false, nil, true)
+ if c, err := h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", tcpCreateParams()); err == nil || c != nil {
+ t.Fatal("uninitialized creation accepted")
+ }
+ if _, err := h.client.Initialize(h.ctx, "owner", ahptypes.SupportedProtocolVersions(), nil); err != nil {
+ t.Fatal(err)
+ }
+ for _, change := range []func(*ahptypes.TcpConnectionSubscription){
+ func(p *ahptypes.TcpConnectionSubscription) { p.Port = 0 },
+ func(p *ahptypes.TcpConnectionSubscription) { p.Host = "http://localhost" },
+ func(p *ahptypes.TcpConnectionSubscription) { p.MaximumChunkSize = 5 },
+ func(p *ahptypes.TcpConnectionSubscription) { p.ReceiveWindowBytes = tcpMaxSafeInteger + 1 },
+ func(p *ahptypes.TcpConnectionSubscription) { p.Encoding = "future" },
+ } {
+ params := tcpCreateParams()
+ change(¶ms)
+ if c, err := h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", params); err == nil || c != nil {
+ t.Fatal("invalid creation accepted")
+ }
+ }
+ done := make(chan error, 1)
+ go func() {
+ c, err := h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", tcpCreateParams())
+ if c != nil {
+ t.Error("invalid creation returned a handle")
+ }
+ done <- err
+ }()
+ request := <-h.requests
+ response := h.snapshot()
+ response.Snapshot.State.Tcp.Input.ReceivedBytes = 1
+ h.reply(request, response)
+ if err := <-done; err == nil {
+ t.Fatal("non-fresh snapshot accepted")
+ }
+ h.notification("unsubscribe")
+}
+
+func TestOwnedTCPTimeoutReleasesLateChildOnLiveTransport(t *testing.T) {
+ for _, mode := range []string{"valid", "malformedState", "heldSend"} {
+ t.Run(mode, func(t *testing.T) {
+ var gate chan struct{}
+ var gates []<-chan struct{}
+ if mode == "heldSend" {
+ gate = make(chan struct{})
+ gates = append(gates, gate)
+ }
+ h := newTCPTestHost(t, true, nil, true, gates...)
+ h.client.cfg.DefaultRequestTimeout = 100 * time.Millisecond
+ done := make(chan error, 1)
+ go func() {
+ c, err := h.client.OpenTCPConnection(h.ctx, "ahp-session:/s1", tcpCreateParams())
+ if c != nil {
+ t.Error("timeout returned a connection")
+ }
+ done <- err
+ }()
+ request := <-h.requests
+ if err := <-done; !errors.Is(err, context.DeadlineExceeded) {
+ t.Fatalf("expected timeout: %v", err)
+ }
+ if gate != nil {
+ close(gate)
+ }
+ if mode == "malformedState" {
+ h.reply(request, json.RawMessage(`{"snapshot":{"state":{"type":"tcp","input":{"windowBytes":0.5}},"resource":"ahp-tcp:/owned","fromSeq":0}}`))
+ } else {
+ h.reply(request, h.snapshot())
+ }
+ reset := h.dispatch()
+ if reset.Channel != "ahp-tcp:/owned" {
+ t.Fatal("reset targeted parent")
+ }
+ if action, ok := reset.Action.Value.(*ahptypes.TcpClientResetAction); !ok || action.Reason != ahptypes.TcpResetReasonConnectionAborted {
+ t.Fatal("missing late creation reset")
+ }
+ n := h.notification("unsubscribe")
+ var params ahptypes.UnsubscribeParams
+ if err := json.Unmarshal(n.Params, ¶ms); err != nil {
+ t.Fatal(err)
+ }
+ if params.Channel != "ahp-tcp:/owned" {
+ t.Fatal("unsubscribed parent")
+ }
+ h.reply(request, h.snapshot())
+ if err := h.client.Ping(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case n := <-h.notifications:
+ t.Fatalf("duplicate cleanup: %s", n.Method)
+ default:
+ }
+ })
+ }
+}
+
+func TestOwnedTCPResumeDoesNotReuseFullyAckedSequence(t *testing.T) {
+ old := newTCPTestHost(t, true, nil, false)
+ c := old.open()
+ if _, err := c.Write(old.ctx, []byte("ab")); err != nil {
+ t.Fatal(err)
+ }
+ original := old.dispatch()
+ old.emit(original.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: original.ClientSeq}, nil)
+ waitTCPState(t, old.ctx, c, func(c *TCPConnection) bool { return len(c.pending) == 0 })
+ ordinary, err := old.client.Dispatch(old.ctx, "ahp-session:/s1", ahptypes.StateAction{
+ Value: &ahptypes.SessionTitleChangedAction{Type: ahptypes.ActionTypeSessionTitleChanged, Title: "ordinary"},
+ })
+ if err != nil {
+ t.Fatal(err)
+ }
+ old.dispatch()
+ old.client.ShutdownPreservingTCP(old.ctx)
+ fresh := newTCPTestHost(t, false, nil, false)
+ done := make(chan error, 1)
+ go func() {
+ _, err := fresh.client.ReconnectTCPConnections(fresh.ctx, ahptypes.ReconnectParams{ClientId: "owner", LastSeenServerSeq: 1}, []*TCPConnection{c})
+ done <- err
+ }()
+ request := <-fresh.requests
+ fresh.reply(request, &ahptypes.ReconnectReplayResult{Actions: []ahptypes.ActionEnvelope{}, Missing: []string{}})
+ if err := <-done; err != nil {
+ t.Fatal(err)
+ }
+ if _, err := c.Write(fresh.ctx, []byte("c")); err != nil {
+ t.Fatal(err)
+ }
+ next := fresh.dispatch()
+ if next.ClientSeq <= ordinary.ClientSeq {
+ t.Fatal("reused original client's ordinary action sequence")
+ }
+ if action, ok := next.Action.Value.(*ahptypes.TcpInputAction); !ok || action.Offset != 2 {
+ t.Fatal("lost retained input offset")
+ }
+ if err := c.Dispose(fresh.ctx); err != nil {
+ t.Fatal(err)
+ }
+ fresh.dispatch()
+ fresh.notification("unsubscribe")
+}
+
+func TestOwnedTCPFinalCloseDrainsBufferedReads(t *testing.T) {
+ h := newTCPTestHost(t, true, []byte("abc"), false)
+ c := h.open()
+ if _, err := c.Write(h.ctx, []byte("abcd")); err != nil {
+ t.Fatal(err)
+ }
+ first, second := h.dispatch(), h.dispatch()
+ write := make(chan error, 1)
+ drain := make(chan error, 1)
+ go func() { _, err := c.Write(h.ctx, []byte("x")); write <- err }()
+ go func() { drain <- c.Drain(h.ctx) }()
+ h.emit(&ahptypes.TcpHostCloseAction{Type: ahptypes.ActionTypeTcpHostClose}, nil, nil)
+ if err := <-write; err == nil {
+ t.Fatal("host close did not stop new writes")
+ }
+ select {
+ case err := <-drain:
+ t.Fatalf("drain finished before input consumption: %v", err)
+ default:
+ }
+ closeAction := h.dispatch()
+ if _, ok := closeAction.Action.Value.(*ahptypes.TcpClientCloseAction); !ok {
+ t.Fatal("missing close response before input consumption")
+ }
+ if c.IsClosed() {
+ t.Fatal("close response disposed unconsumed input")
+ }
+ if data, err := c.Read(h.ctx); string(data) != "abc" || err != nil {
+ t.Fatalf("lost buffered close data %q: %v", data, err)
+ }
+ if _, err := c.Read(h.ctx); !errors.Is(err, io.EOF) {
+ t.Fatal(err)
+ }
+ credit := h.dispatch()
+ if _, ok := credit.Action.Value.(*ahptypes.TcpDataConsumedAction); !ok {
+ t.Fatal("missing read credit")
+ }
+ h.emit(first.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: first.ClientSeq}, nil)
+ h.emit(second.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: second.ClientSeq}, nil)
+ ack := h.emit(closeAction.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: closeAction.ClientSeq}, nil)
+ waitTCPState(t, h.ctx, c, func(c *TCPConnection) bool { return c.checkpoint == ack })
+ if c.IsClosed() {
+ t.Fatal("two-sided close discarded unconsumed bytes")
+ }
+ if err := h.client.Ping(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ select {
+ case n := <-h.notifications:
+ t.Fatalf("cleanup before drain: %s", n.Method)
+ default:
+ }
+ h.emit(&ahptypes.TcpInputConsumedAction{Type: ahptypes.ActionTypeTcpInputConsumed, ConsumedBytes: 4}, nil, nil)
+ if err := <-drain; err != nil {
+ t.Fatal(err)
+ }
+ if c.IsClosed() {
+ t.Fatal("closed before credit ack")
+ }
+ h.emit(credit.Action.Value, &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: credit.ClientSeq}, nil)
+ h.notification("unsubscribe")
+ if err := c.Close(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+}
+
+func TestOwnedTCPInvalidEchoTerminatesWithoutRetainingPayload(t *testing.T) {
+ for _, mode := range []string{"missing", "foreign", "zero", "negative", "unsafe", "unassigned",
+ "mismatched", "reusedAck", "creditMissing", "eofMissing", "closeMissing", "resetMissing"} {
+ t.Run(mode, func(t *testing.T) {
+ h := newTCPTestHost(t, true, []byte("x"), false)
+ c := h.open()
+ if _, err := c.Write(h.ctx, []byte("ab")); err != nil {
+ t.Fatal(err)
+ }
+ original := h.dispatch()
+ action := original.Action.Value
+ origin := &ahptypes.ActionOrigin{ClientId: "owner", ClientSeq: original.ClientSeq}
+ switch mode {
+ case "missing":
+ origin = nil
+ case "foreign":
+ origin.ClientId = "another"
+ case "zero":
+ origin.ClientSeq = 0
+ case "negative":
+ origin.ClientSeq = -1
+ case "unsafe":
+ origin.ClientSeq = tcpMaxSafeInteger + 1
+ case "unassigned":
+ origin.ClientSeq += 100
+ case "mismatched":
+ action = &ahptypes.TcpInputAction{Type: ahptypes.ActionTypeTcpInput, Offset: 0, Data: "eHk="}
+ case "reusedAck":
+ seq := h.emit(original.Action.Value, origin, nil)
+ waitTCPState(t, h.ctx, c, func(c *TCPConnection) bool { return c.checkpoint == seq })
+ if _, err := c.Write(h.ctx, []byte("c")); err != nil {
+ t.Fatal(err)
+ }
+ action = h.dispatch().Action.Value
+ case "creditMissing":
+ if _, err := c.Read(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ action = h.dispatch().Action.Value
+ origin = nil
+ case "eofMissing":
+ if err := c.End(h.ctx); err != nil {
+ t.Fatal(err)
+ }
+ action = h.dispatch().Action.Value
+ origin = nil
+ case "closeMissing":
+ action = &ahptypes.TcpClientCloseAction{Type: ahptypes.ActionTypeTcpClientClose}
+ origin = nil
+ case "resetMissing":
+ action = &ahptypes.TcpClientResetAction{Type: ahptypes.ActionTypeTcpClientReset, Reason: ahptypes.TcpResetReasonConnectionAborted}
+ origin = nil
+ }
+ c.mu.Lock()
+ beforeInput, beforeConsumed := c.state.Input.ReceivedBytes, c.state.Output.ConsumedBytes
+ c.mu.Unlock()
+ h.emit(action, origin, nil)
+ waitTCPState(t, h.ctx, c, func(c *TCPConnection) bool { return c.terminal })
+ c.mu.Lock()
+ unchanged := c.state.Input.ReceivedBytes == beforeInput && c.state.Output.ConsumedBytes == beforeConsumed
+ released := len(c.pending) == 0
+ c.mu.Unlock()
+ if !unchanged || !released {
+ t.Fatal("bad echo changed counters or retained payload")
+ }
+ if _, err := c.Read(h.ctx); err == nil {
+ t.Fatal("bad echo did not terminate reader")
+ }
+ if _, err := c.Write(h.ctx, []byte("z")); err == nil {
+ t.Fatal("bad echo did not terminate writer")
+ }
+ if reset, ok := h.dispatch().Action.Value.(*ahptypes.TcpClientResetAction); !ok || reset.Reason != ahptypes.TcpResetReasonProtocolError {
+ t.Fatal("missing protocol reset")
+ }
+ h.notification("unsubscribe")
+ })
+ }
+}
+
+func TestOwnedTCPMalformedReplayIsTerminal(t *testing.T) {
+ old := newTCPTestHost(t, true, nil, false)
+ c := old.open()
+ old.client.ShutdownPreservingTCP(old.ctx)
+ fresh := newTCPTestHost(t, false, nil, false)
+ if _, err := fresh.client.ReconnectTCPConnections(fresh.ctx, ahptypes.ReconnectParams{ClientId: "different"}, []*TCPConnection{c}); err == nil {
+ t.Fatal("changed owner accepted")
+ }
+ done := make(chan error, 1)
+ go func() {
+ _, err := fresh.client.ReconnectTCPConnections(fresh.ctx, ahptypes.ReconnectParams{ClientId: "owner"}, []*TCPConnection{c})
+ done <- err
+ }()
+ request := <-fresh.requests
+ fresh.reply(request, map[string]any{"type": "replay", "missing": []any{}, "actions": []any{
+ map[string]any{"channel": c.Resource(), "serverSeq": 1, "action": map[string]any{"type": "tcp/dataEof", "finalOffset": 0.5}},
+ }})
+ if err := <-done; err == nil {
+ t.Fatal("malformed replay succeeded")
+ }
+ if _, err := c.Read(fresh.ctx); err == nil {
+ t.Fatal("malformed replay did not terminate stream")
+ }
+ fresh.dispatch()
+ fresh.notification("unsubscribe")
+}
diff --git a/clients/go/ahptypes/actions.generated.go b/clients/go/ahptypes/actions.generated.go
index dc5e6f386..da5b8e221 100644
--- a/clients/go/ahptypes/actions.generated.go
+++ b/clients/go/ahptypes/actions.generated.go
@@ -122,6 +122,16 @@ const (
ActionTypeAutomationRunSessionRemoved ActionType = "automationRun/sessionRemoved"
ActionTypeAutomationRunPrimarySessionChanged ActionType = "automationRun/primarySessionChanged"
ActionTypeAutomationRunCancelRequested ActionType = "automationRun/cancelRequested"
+ ActionTypeTcpInput ActionType = "tcp/input"
+ ActionTypeTcpData ActionType = "tcp/data"
+ ActionTypeTcpInputConsumed ActionType = "tcp/inputConsumed"
+ ActionTypeTcpDataConsumed ActionType = "tcp/dataConsumed"
+ ActionTypeTcpInputEof ActionType = "tcp/inputEof"
+ ActionTypeTcpDataEof ActionType = "tcp/dataEof"
+ ActionTypeTcpClientClose ActionType = "tcp/clientClose"
+ ActionTypeTcpHostClose ActionType = "tcp/hostClose"
+ ActionTypeTcpClientReset ActionType = "tcp/clientReset"
+ ActionTypeTcpHostReset ActionType = "tcp/hostReset"
)
// ─── Action Envelope ─────────────────────────────────────────────────
@@ -1634,6 +1644,72 @@ type ResourceWatchChangedAction struct {
Changes json.RawMessage `json:"changes"`
}
+// Client bytes. Never apply optimistically to the authoritative reducer.
+// Write to the destination only when accepted input.receivedBytes advances.
+type TcpInputAction struct {
+ Type ActionType `json:"type"`
+ // Absolute decoded-byte offset.
+ Offset int64 `json:"offset"`
+ // Nonempty canonical padded RFC 4648 base64; no whitespace.
+ Data string `json:"data"`
+}
+
+// Host bytes. Deliver once, only when output.receivedBytes advances.
+type TcpDataAction struct {
+ Type ActionType `json:"type"`
+ // Absolute decoded-byte offset.
+ Offset int64 `json:"offset"`
+ // Nonempty canonical padded RFC 4648 base64; no whitespace.
+ Data string `json:"data"`
+}
+
+// Cumulative input bytes released from the host's bounded write buffer.
+// Not an acknowledgment that the destination application processed the bytes.
+type TcpInputConsumedAction struct {
+ Type ActionType `json:"type"`
+ ConsumedBytes int64 `json:"consumedBytes"`
+}
+
+// Cumulative output bytes released by the client's bounded stream consumer.
+type TcpDataConsumedAction struct {
+ Type ActionType `json:"type"`
+ ConsumedBytes int64 `json:"consumedBytes"`
+}
+
+// Half-close client input after all preceding input bytes have been written.
+type TcpInputEofAction struct {
+ Type ActionType `json:"type"`
+ FinalOffset int64 `json:"finalOffset"`
+}
+
+// Half-close host output after all preceding output bytes have been delivered.
+type TcpDataEofAction struct {
+ Type ActionType `json:"type"`
+ FinalOffset int64 `json:"finalOffset"`
+}
+
+// Client's final close. Respond with hostClose if not already sent.
+type TcpClientCloseAction struct {
+ Type ActionType `json:"type"`
+}
+
+// Host's final close. Respond with clientClose if not already sent.
+type TcpHostCloseAction struct {
+ Type ActionType `json:"type"`
+}
+
+// Abort both directions and discard buffered payload.
+type TcpClientResetAction struct {
+ Type ActionType `json:"type"`
+ Reason TcpResetReason `json:"reason"`
+}
+
+// Abort both directions and discard buffered payload.
+type TcpHostResetAction struct {
+ Type ActionType `json:"type"`
+ Reason TcpResetReason `json:"reason"`
+}
+
// Ask the host to create a durable automation at a client-chosen resource.
//
// Clients may dispatch this action only when the host advertises its `create`
@@ -1857,6 +1933,16 @@ func (*TerminalCommandDetectionAvailableAction) isStateAction() {}
func (*TerminalCommandExecutedAction) isStateAction() {}
func (*TerminalCommandFinishedAction) isStateAction() {}
func (*ResourceWatchChangedAction) isStateAction() {}
+func (*TcpInputAction) isStateAction() {}
+func (*TcpDataAction) isStateAction() {}
+func (*TcpInputConsumedAction) isStateAction() {}
+func (*TcpDataConsumedAction) isStateAction() {}
+func (*TcpInputEofAction) isStateAction() {}
+func (*TcpDataEofAction) isStateAction() {}
+func (*TcpClientCloseAction) isStateAction() {}
+func (*TcpHostCloseAction) isStateAction() {}
+func (*TcpClientResetAction) isStateAction() {}
+func (*TcpHostResetAction) isStateAction() {}
func (*AutomationCreateRequestedAction) isStateAction() {}
func (*AutomationUpdateRequestedAction) isStateAction() {}
func (*AutomationSetAction) isStateAction() {}
@@ -2445,6 +2531,66 @@ func (u *StateAction) UnmarshalJSON(data []byte) error {
return err
}
u.Value = &value
+ case "tcp/input":
+ var value TcpInputAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/data":
+ var value TcpDataAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/inputConsumed":
+ var value TcpInputConsumedAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/dataConsumed":
+ var value TcpDataConsumedAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/inputEof":
+ var value TcpInputEofAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/dataEof":
+ var value TcpDataEofAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/clientClose":
+ var value TcpClientCloseAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/hostClose":
+ var value TcpHostCloseAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/clientReset":
+ var value TcpClientResetAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
+ case "tcp/hostReset":
+ var value TcpHostResetAction
+ if err := json.Unmarshal(data, &value); err != nil {
+ return err
+ }
+ u.Value = &value
case "automation/createRequested":
var value AutomationCreateRequestedAction
if err := json.Unmarshal(data, &value); err != nil {
diff --git a/clients/go/ahptypes/commands.generated.go b/clients/go/ahptypes/commands.generated.go
index cd3b2a6ad..a1214a9ec 100644
--- a/clients/go/ahptypes/commands.generated.go
+++ b/clients/go/ahptypes/commands.generated.go
@@ -190,6 +190,8 @@ type InitializeResult struct {
// `ahp-automations://` for {@link AutomationState}; absence means the
// host does not expose an automation catalogue or automation commands.
Automations *AutomationCapabilities `json:"automations,omitempty"`
+ // Enables atomic creation of session-scoped, replay-only TCP channels.
+ TcpConnections *TcpConnectionsCapability `json:"tcpConnections,omitempty"`
}
// Optional capabilities a client declares during `initialize`.
@@ -317,6 +319,8 @@ type ReconnectParams struct {
//
// The server MUST include all replayed data in the response.
type ReconnectReplayResult struct {
+ // Discriminant
+ Type ReconnectResultType `json:"type"`
// Missed action envelopes since `lastSeenServerSeq`
Actions []ActionEnvelope `json:"actions"`
// URIs from `ReconnectParams.subscriptions` that the server cannot resume.
@@ -328,8 +332,14 @@ type ReconnectReplayResult struct {
// Reconnect result when the gap exceeds the replay buffer.
type ReconnectSnapshotResult struct {
+ // Discriminant
+ Type ReconnectResultType `json:"type"`
// Fresh snapshots for each subscription
Snapshots []Snapshot `json:"snapshots"`
+ // Subscriptions that cannot be restored. Hosts supporting TCP MUST list all
+ // requested TCP channels here and dispose their sockets on snapshot fallback.
+ // Omitted by older hosts; absence does not authorize snapshot-restoring TCP.
+ Missing []URI `json:"missing,omitempty"`
}
// Subscribe to a URI-identified channel.
@@ -355,6 +365,10 @@ type SubscribeParams struct {
// Servers that do not understand a requested view ignore it and return their
// default snapshot. Clients MUST tolerate receiving more state than requested.
View *SubscribeView `json:"view,omitempty"`
+ // Atomically create a private child channel and subscribe to it.
+ // Requires the advertised tcpConnections capability. channel identifies
+ // the parent session; snapshot.resource identifies the created TCP channel.
+ Create *TcpConnectionSubscription `json:"create,omitempty"`
}
// Optional client-requested shape for a subscription snapshot.
@@ -388,6 +402,26 @@ type SubscribeResult struct {
Snapshot *Snapshot `json:"snapshot,omitempty"`
}
+// Creates and exclusively subscribes to one TCP connection.
+//
+// SubscribeParams.channel MUST identify the parent `ahp-session:` channel.
+// The host returns the new `ahp-tcp:` URI in snapshot.resource, not the parent.
+// It installs the subscription and sends the response before any TCP actions.
+// Unknown creation kinds MUST be rejected, never treated as normal subscribe.
+type TcpConnectionSubscription struct {
+ Type string `json:"type"`
+ // DNS name or IP literal, not a URL.
+ Host string `json:"host"`
+ // Destination port.
+ Port int64 `json:"port"`
+ // Selected from InitializeResult.tcpConnections.encodings.
+ Encoding TcpDataEncoding `json:"encoding"`
+ // Client receive window in decoded bytes.
+ ReceiveWindowBytes int64 `json:"receiveWindowBytes"`
+ // Maximum decoded bytes per output action; MUST NOT exceed receiveWindowBytes.
+ MaximumChunkSize int64 `json:"maximumChunkSize"`
+}
+
// Creates a new session with the specified agent provider.
//
// If the session URI already exists, the server MUST return an error with code
@@ -1512,6 +1546,62 @@ func (v ChatMoveToNewSessionDestination) MarshalJSON() ([]byte, error) {
return json.Marshal(raw)
}
+func (v *ReconnectReplayResult) UnmarshalJSON(data []byte) error {
+ disc, ok, err := readDiscriminator(data, "type")
+ if err != nil {
+ return err
+ }
+ if !ok {
+ return missingDiscriminatorError("ReconnectReplayResult", "type")
+ }
+ if disc != "replay" {
+ return unknownDiscriminatorError("ReconnectReplayResult", "type", disc)
+ }
+ type wire ReconnectReplayResult
+ var raw wire
+ if err := json.Unmarshal(data, &raw); err != nil {
+ return err
+ }
+ *v = ReconnectReplayResult(raw)
+ v.Type = ReconnectResultTypeReplay
+ return nil
+}
+
+func (v ReconnectReplayResult) MarshalJSON() ([]byte, error) {
+ type wire ReconnectReplayResult
+ raw := wire(v)
+ raw.Type = ReconnectResultTypeReplay
+ return json.Marshal(raw)
+}
+
+func (v *ReconnectSnapshotResult) UnmarshalJSON(data []byte) error {
+ disc, ok, err := readDiscriminator(data, "type")
+ if err != nil {
+ return err
+ }
+ if !ok {
+ return missingDiscriminatorError("ReconnectSnapshotResult", "type")
+ }
+ if disc != "snapshot" {
+ return unknownDiscriminatorError("ReconnectSnapshotResult", "type", disc)
+ }
+ type wire ReconnectSnapshotResult
+ var raw wire
+ if err := json.Unmarshal(data, &raw); err != nil {
+ return err
+ }
+ *v = ReconnectSnapshotResult(raw)
+ v.Type = ReconnectResultTypeSnapshot
+ return nil
+}
+
+func (v ReconnectSnapshotResult) MarshalJSON() ([]byte, error) {
+ type wire ReconnectSnapshotResult
+ raw := wire(v)
+ raw.Type = ReconnectResultTypeSnapshot
+ return json.Marshal(raw)
+}
+
// ─── ChatSource Union ─────────────────────────────────────────────────
// ChatSource identifies how a new chat uses a source chat.
diff --git a/clients/go/ahptypes/errors.generated.go b/clients/go/ahptypes/errors.generated.go
index e6a4fd395..f1d1463ba 100644
--- a/clients/go/ahptypes/errors.generated.go
+++ b/clients/go/ahptypes/errors.generated.go
@@ -32,10 +32,11 @@ const (
ErrorCodeTurnInProgress int32 = -32004
ErrorCodeUnsupportedProtocolVersion int32 = -32005
// -32006 is intentionally reserved and unassigned.
- ErrorCodeAuthRequired int32 = -32007
- ErrorCodeNotFound int32 = -32008
- ErrorCodePermissionDenied int32 = -32009
- ErrorCodeAlreadyExists int32 = -32010
+ ErrorCodeAuthRequired int32 = -32007
+ ErrorCodeNotFound int32 = -32008
+ ErrorCodePermissionDenied int32 = -32009
+ ErrorCodeAlreadyExists int32 = -32010
+ ErrorCodeTcpConnectionOpenFailed int32 = -32012
)
// AhpErrorCode is the type alias used by AHP application error codes.
diff --git a/clients/go/ahptypes/roundtrip_fixture_test.go b/clients/go/ahptypes/roundtrip_fixture_test.go
index 0f7181385..b96ee7795 100644
--- a/clients/go/ahptypes/roundtrip_fixture_test.go
+++ b/clients/go/ahptypes/roundtrip_fixture_test.go
@@ -242,6 +242,18 @@ func decodeAndReencode(t *testing.T, name, typ, inputJSON string) string {
var v InitializeResult
dec(&v)
return enc(&v)
+ case "SubscribeParams":
+ var v SubscribeParams
+ dec(&v)
+ return enc(&v)
+ case "ReconnectResult":
+ var v ReconnectResult
+ dec(&v)
+ return enc(&v)
+ case "TcpConnectionOpenErrorData":
+ var v TcpConnectionOpenErrorData
+ dec(&v)
+ return enc(&v)
case "ChatSource":
var v ChatSource
dec(&v)
diff --git a/clients/go/ahptypes/state.generated.go b/clients/go/ahptypes/state.generated.go
index 4cc70c783..681827353 100644
--- a/clients/go/ahptypes/state.generated.go
+++ b/clients/go/ahptypes/state.generated.go
@@ -15,6 +15,44 @@ var _ = json.RawMessage(nil)
// ─── Enums ────────────────────────────────────────────────────────────
+// Payload encodings advertised by the host.
+type TcpDataEncoding string
+
+const (
+ TcpDataEncodingBase64 TcpDataEncoding = "base64"
+)
+
+// Endpoint that closes or resets a connection.
+type TcpEndpoint string
+
+const (
+ TcpEndpointClient TcpEndpoint = "client"
+ TcpEndpointHost TcpEndpoint = "host"
+)
+
+// Why a connection was aborted.
+type TcpResetReason string
+
+const (
+ TcpResetReasonConnectionReset TcpResetReason = "connectionReset"
+ TcpResetReasonConnectionAborted TcpResetReason = "connectionAborted"
+ TcpResetReasonProtocolError TcpResetReason = "protocolError"
+ TcpResetReasonReplayUnavailable TcpResetReason = "replayUnavailable"
+ TcpResetReasonPolicyRevoked TcpResetReason = "policyRevoked"
+ TcpResetReasonSessionDisposed TcpResetReason = "sessionDisposed"
+ TcpResetReasonInternalError TcpResetReason = "internalError"
+)
+
+// Expected connection establishment failures.
+type TcpConnectionOpenFailureReason string
+
+const (
+ TcpConnectionOpenFailureReasonConnectionFailed TcpConnectionOpenFailureReason = "connectionFailed"
+ TcpConnectionOpenFailureReasonNameResolutionFailed TcpConnectionOpenFailureReason = "nameResolutionFailed"
+ TcpConnectionOpenFailureReasonResourceShortage TcpConnectionOpenFailureReason = "resourceShortage"
+ TcpConnectionOpenFailureReasonSessionNotReady TcpConnectionOpenFailureReason = "sessionNotReady"
+)
+
// Policy configuration state for a model.
type PolicyState string
@@ -3923,6 +3961,74 @@ type ResourceWatchState struct {
Includes *json.RawMessage `json:"includes,omitempty"`
}
+// State of one host-assigned `ahp-tcp:` channel.
+//
+// Payload is never stored in this state. Only the creating authenticated
+// logical client may observe or dispatch to the channel. Reconnect requires
+// the original sockets, local stream state, and complete action replay;
+// a snapshot cannot restore this channel.
+//
+// Close flags record the two-sided handshake. Either flag means closing;
+// both mean closed. A present reset terminates the connection immediately,
+// independently of the close history.
+type TcpConnectionState struct {
+ Session URI `json:"session"`
+ Target TcpTarget `json:"target"`
+ Encoding TcpDataEncoding `json:"encoding"`
+ // Client to destination socket.
+ Input FlowControlledByteDirectionState `json:"input"`
+ // Destination socket to client.
+ Output FlowControlledByteDirectionState `json:"output"`
+ ClientClosed bool `json:"clientClosed"`
+ HostClosed bool `json:"hostClosed"`
+ Reset *TcpResetState `json:"reset,omitempty"`
+}
+
+type TcpTarget struct {
+ // DNS name or IP literal, resolved and connected in the host endpoint's network.
+ Host string `json:"host"`
+ // Destination port.
+ Port int64 `json:"port"`
+}
+
+type TcpResetState struct {
+ Source TcpEndpoint `json:"source"`
+ Reason TcpResetReason `json:"reason"`
+}
+
+// Bounded byte credit in one direction of a stream.
+// All counters are nonnegative safe integers (at most 2^53 - 1).
+// 0 <= consumedBytes <= receivedBytes and
+// receivedBytes - consumedBytes <= windowBytes.
+type FlowControlledByteDirectionState struct {
+ // Maximum accepted-but-not-consumed decoded bytes.
+ WindowBytes int64 `json:"windowBytes"`
+ // Maximum decoded bytes per chunk; MUST NOT exceed windowBytes.
+ MaximumChunkSize int64 `json:"maximumChunkSize"`
+ // Cumulative accepted bytes.
+ ReceivedBytes int64 `json:"receivedBytes"`
+ // Cumulative bytes released by the bounded consumer.
+ ConsumedBytes int64 `json:"consumedBytes"`
+ // Present after EOF; equals receivedBytes permanently.
+ EofAtBytes *int64 `json:"eofAtBytes,omitempty"`
+}
+
+// Host support for private, session-scoped TCP channels.
+// Presence on initialize is required before using subscribe.create.
+type TcpConnectionsCapability struct {
+ // Supported encodings. The base64 profile MUST be supported.
+ Encodings []TcpDataEncoding `json:"encodings"`
+ // Informational limit; runtime policy may impose a lower limit.
+ MaximumConnectionsPerClient *int64 `json:"maximumConnectionsPerClient,omitempty"`
+}
+
+// Required detail for TcpConnectionOpenFailed (-32012).
+// Policy denial and malformed requests use PermissionDenied and InvalidParams.
+type TcpConnectionOpenErrorData struct {
+ Reason TcpConnectionOpenFailureReason `json:"reason"`
+ Retryable *bool `json:"retryable,omitempty"`
+}
+
// A single change observed by a resource watcher.
type ResourceChange struct {
// The URI of the resource that changed.
@@ -6261,12 +6367,13 @@ func (o ChatOrigin) MarshalJSON() ([]byte, error) {
// SnapshotState is the state payload of a snapshot — root, session,
// chat, terminal, changeset, resource-watch, annotations, automation catalogue,
-// or automation-run state. The active
+// automation-run, or TCP state. The active
// variant is chosen by which pointer field is non-nil; UnmarshalJSON probes
// for required fields in the canonical order
-// (automationRun → automations → session → chat → terminal → changeset →
+// (tcp → automationRun → automations → session → chat → terminal → changeset →
// resourceWatch → annotations → root).
type SnapshotState struct {
+ Tcp *TcpConnectionState `json:"-"`
Root *RootState `json:"-"`
Session *SessionState `json:"-"`
Chat *ChatState `json:"-"`
@@ -6281,6 +6388,8 @@ type SnapshotState struct {
// MarshalJSON encodes whichever variant is currently populated.
func (s SnapshotState) MarshalJSON() ([]byte, error) {
switch {
+ case s.Tcp != nil:
+ return json.Marshal(s.Tcp)
case s.AutomationRun != nil:
return json.Marshal(s.AutomationRun)
case s.Automations != nil:
@@ -6313,6 +6422,12 @@ func (s *SnapshotState) UnmarshalJSON(data []byte) error {
return err
}
switch {
+ case containsAll(probe, "input", "output", "target"):
+ var v TcpConnectionState
+ if err := json.Unmarshal(data, &v); err != nil {
+ return err
+ }
+ s.Tcp = &v
case containsAll(probe, "automation", "origin", "sessions"):
var v AutomationRunState
if err := json.Unmarshal(data, &v); err != nil {
diff --git a/clients/kotlin/README.md b/clients/kotlin/README.md
index cae975d5c..4bbdd0ca4 100644
--- a/clients/kotlin/README.md
+++ b/clients/kotlin/README.md
@@ -110,6 +110,12 @@ behavior. This is a client API migration, not a new JSON protocol.
state from the current state and an applied action. Behavior parity with
the canonical TypeScript reducers is verified against the shared
`types/test-cases/reducers/` fixture corpus.
+- **TCP accounting** — `tcpReducer(state, action)` and `TcpReducer.reduce`
+ return the next `TcpConnectionState` without retaining payloads. Invalid
+ actions throw `IllegalArgumentException` and leave the original state unchanged.
+- **Portable TCP streams** — `TcpConnection` owns buffering, flow control, and
+ reconnect replay. It uses JVM
+ `CompletableFuture` operations without adding a coroutine or network runtime.
- **Channel-scoped notification params** — `SessionAddedParams`,
`SessionRemovedParams`, `SessionSummaryChangedParams`, `AuthRequiredParams`,
`OtlpExportLogsParams`, etc. Notifications are routed by their JSON-RPC
@@ -124,6 +130,48 @@ behavior. This is a client API migration, not a new JSON protocol.
- An example Android client — see the Swift `AHPClient` example for the architecture
pattern; a Kotlin/Android equivalent is planned for a follow-up release.
+### Portable TCP adapter
+
+Inject `TcpConnectionTransport(clientId, nextSequence, advanceSequencePast, send,
+unsubscribe, lastAssignedSequence)`. The identity must be the one used to initialize/reconnect.
+Callbacks enqueue without blocking; the sequence allocator is shared with all
+actions sent by that logical client. `send` receives the child resource, original
+client sequence, and typed action. A send exception suspends the handle and is
+available as `lastTransportFailure`; retained actions reconcile on resume.
+`lastAssignedSequence` reads the global allocator, including acknowledged and
+non-TCP actions (use -1 if nothing has been assigned). This prevents a replacement
+allocator from reusing identities even when the TCP pending set is empty.
+1. Construct `TcpConnectionCreation(session, create, initialized, transport)`.
+ Set `create.type` to `"tcpConnection"`. Its `parameters` validate the
+ request before the transport sends `subscribe`.
+2. Attach a strict bounded event receiver **before** that request. Buffer events
+ while the child URI is unknown; overflow and decode loss must be explicit.
+3. Forward the definitive subscribe reply to `creation.accept(result)`, and
+ request errors/timeouts to `creation.fail(error)`. Await `creation.completion`
+ for the connection, then forward buffered/live envelopes to its `accept`.
+ The transport must still forward a late reply after cancellation/timeout:
+ the creation helper releases its child once, never the parent. Cleanup
+ failures throw to the transport callback.
+4. Use `write(ByteArray)`, `read()`, `drain()`, and `end()` futures. One reader and
+ one writer may run concurrently; overlapping writers are rejected. Do not
+ mutate write input before completion. Reading releases cumulative credit;
+ `read()` returns null only after buffered EOF drains.
+5. Call `suspend(cause)` on disconnect. Before the replacement transport's
+ reconnect request, call `TcpConnection.reconnectParameters(params, handles)`.
+ Buffer live delivery until `TcpConnection.resume(adjustedParams, result,
+ handles, replacementTransport)` returns. This helper verifies identity,
+ reconciles replay, and resends pending actions. Snapshot/missing resources
+ fail closed.
+6. Forward strict receiver lag/decode loss to `fail(error)`, never `suspend`.
+ `close()` stops new input but retains crossing host data and buffered reads.
+ Keep reading during graceful close; `dispose()` aborts without draining.
+
+The application supplies typed transport/event delivery and chooses reconnect
+timing, connection-count limits, and native socket bridges. It does not implement
+TCP reduction, credit, acknowledgement bookkeeping, or replay. TCP streams do not
+belong in ordinary state mirrors.
+See the [TCP channel contract](../../docs/specification/tcp-channel.md).
+
## Protocol version mapping
Two constants in `com.microsoft.agenthostprotocol.generated` track which
diff --git a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/Reducers.kt b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/Reducers.kt
index 8fd22a745..c4031d9bb 100644
--- a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/Reducers.kt
+++ b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/Reducers.kt
@@ -85,6 +85,121 @@ public object AutomationRunReducer : Reducer {
automationRunReducer(state, action)
}
+/** Pure TCP reducer. Invalid actions throw before changing state; payloads are never retained. */
+public object TcpReducer : Reducer {
+ override fun reduce(state: TcpConnectionState, action: StateAction): TcpConnectionState =
+ tcpReducer(state, action)
+}
+
+private const val TCP_MAX_SAFE_INTEGER = 9007199254740991L
+
+private fun requireTcp(condition: Boolean, message: String) {
+ require(condition) { "Invalid TCP action: $message" }
+}
+
+private fun requireTcpOffset(value: Long) {
+ requireTcp(value in 0..TCP_MAX_SAFE_INTEGER, "offset must be a nonnegative safe integer")
+}
+
+private fun tcpBase64Value(char: Char): Int = when (char) {
+ in 'A'..'Z' -> char - 'A'
+ in 'a'..'z' -> char - 'a' + 26
+ in '0'..'9' -> char - '0' + 52
+ '+' -> 62
+ '/' -> 63
+ else -> -1
+}
+
+private fun tcpPayloadLength(data: String, maximumChunkSize: Long): Long {
+ requireTcp(
+ data.isNotEmpty() && data.length.toLong() <= 4 * (maximumChunkSize / 3 + if (maximumChunkSize % 3 > 0) 1 else 0),
+ "chunk size",
+ )
+ val padding = if (data.endsWith("==")) 2 else if (data.endsWith("=")) 1 else 0
+ requireTcp(data.length % 4 == 0, "base64 encoding")
+ var last = 0
+ for (index in 0 until data.length - padding) {
+ last = tcpBase64Value(data[index])
+ requireTcp(last >= 0, "base64 encoding")
+ }
+ if (padding > 0) {
+ requireTcp(last % (if (padding == 2) 16 else 4) == 0, "noncanonical base64 padding bits")
+ }
+ val length = data.length.toLong() / 4 * 3 - padding
+ requireTcp(length <= maximumChunkSize, "chunk size")
+ return length
+}
+
+private fun tcpReceive(
+ direction: FlowControlledByteDirectionState, offset: Long, data: String, senderClosed: Boolean,
+): FlowControlledByteDirectionState {
+ requireTcpOffset(offset)
+ val end = offset + tcpPayloadLength(data, direction.maximumChunkSize)
+ requireTcpOffset(end)
+ if (end <= direction.receivedBytes) return direction
+ requireTcp(offset == direction.receivedBytes, "gap or overlapping byte range")
+ requireTcp(!senderClosed && direction.eofAtBytes == null, "data after EOF or sender close")
+ requireTcp(end - direction.consumedBytes <= direction.windowBytes, "receive window exceeded")
+ return direction.copy(receivedBytes = end)
+}
+
+private fun tcpConsume(direction: FlowControlledByteDirectionState, consumedBytes: Long): FlowControlledByteDirectionState {
+ requireTcpOffset(consumedBytes)
+ requireTcp(consumedBytes <= direction.receivedBytes, "consuming bytes not received")
+ return if (consumedBytes <= direction.consumedBytes) direction else direction.copy(consumedBytes = consumedBytes)
+}
+
+private fun tcpEof(direction: FlowControlledByteDirectionState, finalOffset: Long, senderClosed: Boolean): FlowControlledByteDirectionState {
+ requireTcpOffset(finalOffset)
+ requireTcp(finalOffset == direction.receivedBytes, "EOF offset")
+ if (direction.eofAtBytes == finalOffset) return direction
+ requireTcp(!senderClosed, "EOF after sender close")
+ return direction.copy(eofAtBytes = finalOffset)
+}
+
+/**
+ * Reduces TCP accounting only, without restoring streams or performing socket I/O.
+ * Callers must reset/close the channel on [IllegalArgumentException] and must not
+ * write rejected or duplicate data (only write when receivedBytes advances).
+ */
+public fun tcpReducer(state: TcpConnectionState, action: StateAction): TcpConnectionState {
+ if (state.reset != null) return state
+ val input: FlowControlledByteDirectionState
+ val output: FlowControlledByteDirectionState
+ when (action) {
+ is StateActionTcpInput -> {
+ input = tcpReceive(state.input, action.value.offset, action.value.data, state.clientClosed)
+ output = state.output
+ }
+ is StateActionTcpData -> {
+ input = state.input
+ output = tcpReceive(state.output, action.value.offset, action.value.data, state.hostClosed)
+ }
+ is StateActionTcpInputConsumed -> {
+ input = tcpConsume(state.input, action.value.consumedBytes)
+ output = state.output
+ }
+ is StateActionTcpDataConsumed -> {
+ input = state.input
+ output = tcpConsume(state.output, action.value.consumedBytes)
+ }
+ is StateActionTcpInputEof -> {
+ input = tcpEof(state.input, action.value.finalOffset, state.clientClosed)
+ output = state.output
+ }
+ is StateActionTcpDataEof -> {
+ input = state.input
+ output = tcpEof(state.output, action.value.finalOffset, state.hostClosed)
+ }
+ is StateActionTcpClientClose -> return if (state.clientClosed) state else state.copy(clientClosed = true)
+ is StateActionTcpHostClose -> return if (state.hostClosed) state else state.copy(hostClosed = true)
+ is StateActionTcpClientReset -> return state.copy(reset = TcpResetState(TcpEndpoint.CLIENT, action.value.reason))
+ is StateActionTcpHostReset -> return state.copy(reset = TcpResetState(TcpEndpoint.HOST, action.value.reason))
+ else -> return state
+ }
+ return if (input === state.input && output === state.output) state else state.copy(input = input, output = output)
+}
+
private val isoTimestampFormatter = DateTimeFormatterBuilder().appendInstant(3).toFormatter()
private fun addMillisecondsToTimestamp(timestamp: String, duration: Long): String =
diff --git a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/TcpConnection.kt b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/TcpConnection.kt
new file mode 100644
index 000000000..afedecd7f
--- /dev/null
+++ b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/TcpConnection.kt
@@ -0,0 +1,433 @@
+package com.microsoft.agenthostprotocol
+
+import com.microsoft.agenthostprotocol.generated.*
+import java.util.ArrayDeque
+import java.util.Base64
+import java.util.concurrent.CompletableFuture
+
+/**
+ * Transport boundary for the portable TCP adapter. Callbacks must enqueue without
+ * blocking. Sequence allocation is shared with all actions of this logical client.
+ * Throwing from send suspends the handle; retained actions are reconciled on resume.
+ */
+public class TcpConnectionTransport(
+ public val clientId: String,
+ public val nextSequence: () -> Long,
+ public val advanceSequencePast: (Long) -> Unit,
+ public val send: (String, Long, StateAction) -> Unit,
+ public val unsubscribe: (String) -> Unit,
+ public val lastAssignedSequence: () -> Long,
+)
+
+/**
+ * Portable subscribe/create lifecycle. The transport sends [parameters] and
+ * forwards its definitive reply to [accept], even after a timeout/cancellation.
+ * [fail] reports request failure; a late private child is released by the SDK.
+ */
+public class TcpConnectionCreation(
+ private val session: String,
+ private val create: TcpConnectionSubscription,
+ private val initialized: InitializeResult,
+ private val transport: TcpConnectionTransport,
+) {
+ public val parameters: SubscribeParams = TcpConnection.creationParameters(session, create, initialized)
+ public val completion: CompletableFuture = CompletableFuture()
+ private val lock = Any()
+ private var received = false
+
+ public fun fail(error: Throwable) { completion.completeExceptionally(error) }
+
+ /** Cleanup failures throw to the transport caller instead of disappearing in a detached task. */
+ public fun accept(result: SubscribeResult): Unit = synchronized(lock) {
+ if (received) return@synchronized
+ received = true
+ if (completion.isDone) {
+ result.snapshot?.resource?.takeIf { it.startsWith("ahp-tcp:") }?.let(transport.unsubscribe)
+ return@synchronized
+ }
+ try {
+ val connection = TcpConnection.open(session, create, initialized, result, transport)
+ if (!completion.complete(connection)) connection.dispose()
+ } catch (error: Exception) {
+ if (!completion.completeExceptionally(error)) throw error
+ }
+ }
+}
+
+/**
+ * A transport-independent, bounded TCP consumer. No worker threads or network
+ * runtime are created. Feed accepted envelopes through [accept], call [suspend]
+ * on disconnect, and [fail] on strict receiver lag or decode loss.
+ */
+public class TcpConnection private constructor(
+ public val resource: String,
+ initial: TcpConnectionState,
+ checkpoint: Long,
+ private var transport: TcpConnectionTransport,
+) : AutoCloseable {
+ public val clientId: String = transport.clientId
+ private val lock = Any()
+ private var current = initial
+ private var checkpoint = checkpoint
+ private val received = ArrayDeque()
+ private val pending = sortedMapOf()
+ private var sentBytes = initial.input.receivedBytes
+ private var consumedBytes = initial.output.consumedBytes
+ private var lastSequence = -1L
+ private var suspended = false
+ private var closing = false
+ private var closed = false
+ private var released = false
+ private var ending = false
+ private var failure: Throwable? = null
+ private var transportFailure: Throwable? = null
+ private var reader: CompletableFuture? = null
+ private var writer: Write? = null
+ private var endResult: CompletableFuture? = null
+ private val drains = mutableListOf>()
+ private var progressing = false
+ private var dirty = false
+
+ private class Write(val data: ByteArray, val result: CompletableFuture, var offset: Int = 0)
+
+ public val state: TcpConnectionState get() = synchronized(lock) { current }
+ public val appliedCheckpoint: Long get() = synchronized(lock) { checkpoint }
+ public val isSuspended: Boolean get() = synchronized(lock) { suspended }
+ public val isClosed: Boolean get() = synchronized(lock) { closed }
+ public val lastTransportFailure: Throwable? get() = synchronized(lock) { transportFailure }
+
+ /** Pulls a chunk and releases its credit. Null is returned only after buffered EOF drains. */
+ public fun read(): CompletableFuture = synchronized(lock) {
+ check(reader == null) { "TCP permits one reader at a time" }
+ val result = CompletableFuture()
+ reader = result
+ result.whenComplete { _, _ -> synchronized(lock) { if (reader === result && result.isCancelled) reader = null } }
+ progress()
+ result
+ }
+
+ /** One writer at a time. The caller must not mutate data until the result completes. */
+ public fun write(data: ByteArray): CompletableFuture = synchronized(lock) {
+ check(writer == null && !ending && !closing && !closed) { "TCP write requires an open, idle writer" }
+ val result = CompletableFuture()
+ val operation = Write(data, result)
+ writer = operation
+ result.whenComplete { _, _ -> synchronized(lock) { if (writer === operation && result.isCancelled) writer = null } }
+ progress()
+ result
+ }
+
+ /** Waits until all reserved input bytes are consumed by the destination. */
+ public fun drain(): CompletableFuture = synchronized(lock) {
+ val result = CompletableFuture()
+ drains.add(result)
+ result.whenComplete { _, _ -> synchronized(lock) { drains.remove(result) } }
+ progress()
+ result
+ }
+
+ /** Half-closes input. Finish an active write first. */
+ public fun end(): CompletableFuture = synchronized(lock) {
+ check(writer == null && !closing && !closed) { "TCP end requires an open, idle writer" }
+ if (ending) return@synchronized endResult ?: CompletableFuture.completedFuture(Unit)
+ ending = true
+ val result = CompletableFuture()
+ endResult = result
+ result.whenComplete { _, _ -> synchronized(lock) {
+ if (endResult === result && result.isCancelled) { endResult = null; ending = false }
+ } }
+ progress()
+ result
+ }
+
+ /** Stops input; retains crossing output and ownership until both sides close and drain. */
+ override fun close(): Unit = synchronized(lock) {
+ if (closing || closed) return@synchronized
+ closing = true
+ dispatch(StateActionTcpClientClose(TcpClientCloseAction(ActionType.TCP_CLIENT_CLOSE)))
+ progress()
+ failure?.let { throw it }
+ }
+
+ /** Disposal terminates every pending operation, discarding unread bytes. */
+ public fun dispose(): Unit = fail(IllegalStateException("TCP connection disposed"))
+
+ public fun suspend(cause: Throwable? = null): Unit = synchronized(lock) {
+ if (!closed) {
+ suspended = true
+ transportFailure = cause
+ }
+ }
+
+ /** Lag/decode loss is terminal, not a resumable transport interruption. */
+ public fun fail(error: Throwable): Unit = synchronized(lock) { failLocked(error, reset = true) }
+
+ public fun accept(envelope: ActionEnvelope): Unit = synchronized(lock) {
+ if (!suspended) apply(envelope)
+ }
+
+ private fun apply(envelope: ActionEnvelope) {
+ if (closed || envelope.channel != resource) return
+ try {
+ safe(envelope.serverSeq)
+ val clientEcho = when (envelope.action) {
+ is StateActionTcpInput, is StateActionTcpDataConsumed, is StateActionTcpInputEof,
+ is StateActionTcpClientClose, is StateActionTcpClientReset -> true
+ else -> false
+ }
+ val sequence = if (clientEcho) {
+ val origin = envelope.origin
+ check(origin != null && origin.clientId == clientId) { "TCP client echo requires the owning client origin" }
+ safe(origin.clientSeq)
+ check(origin.clientSeq <= lastSequence) { "TCP client echo has an unassigned sequence" }
+ origin.clientSeq
+ } else null
+ if (envelope.serverSeq <= checkpoint) return
+ check(envelope.rejectionReason == null) { envelope.rejectionReason ?: "Rejected TCP action" }
+ val next = tcpReducer(current, envelope.action)
+ if (sequence != null) {
+ val expected = pending[sequence]
+ if (expected != null) check(expected == envelope.action) { "TCP client echo does not match its pending action" }
+ else check(next == current) { "TCP advancing client echo has no pending action" }
+ }
+ check(next.output.receivedBytes - consumedBytes <= next.output.windowBytes) { "TCP output exceeds locally released credit" }
+ val action = envelope.action
+ if (action is StateActionTcpData && next.output.receivedBytes > current.output.receivedBytes) {
+ received.addLast(Base64.getDecoder().decode(action.value.data))
+ }
+ current = next
+ checkpoint = envelope.serverSeq
+ if (sequence != null) pending.remove(sequence)
+ } catch (error: IllegalArgumentException) {
+ failLocked(error, reset = true)
+ return
+ } catch (error: IllegalStateException) {
+ failLocked(error, reset = true)
+ return
+ }
+ if (current.reset != null) failLocked(IllegalStateException("TCP reset: ${current.reset?.reason}"), reset = false)
+ else if (current.hostClosed) close()
+ progress()
+ }
+
+ private fun dispatch(action: StateAction) {
+ val sequence = transport.nextSequence()
+ safe(sequence)
+ check(sequence > lastSequence) { "TCP sequence allocator did not advance" }
+ lastSequence = sequence
+ pending[sequence] = action
+ if (!suspended) send(sequence, action)
+ }
+
+ private fun send(sequence: Long, action: StateAction) {
+ try {
+ transport.send(resource, sequence, action)
+ } catch (error: Exception) {
+ transportFailure = error
+ suspended = true
+ if (closed) {
+ val previous = failure
+ if (previous == null) failure = error else previous.addSuppressed(error)
+ }
+ }
+ }
+
+ private fun failLocked(error: Throwable, reset: Boolean) {
+ if (closed) return
+ failure = error
+ closed = true
+ try {
+ if (reset && !suspended) dispatch(StateActionTcpClientReset(TcpClientResetAction(ActionType.TCP_CLIENT_RESET, TcpResetReason.PROTOCOL_ERROR)))
+ } catch (sendError: Exception) { error.addSuppressed(sendError) }
+ received.clear()
+ release()
+ progress()
+ }
+
+ private fun release() {
+ if (released) return
+ released = true
+ pending.clear()
+ try {
+ transport.unsubscribe(resource)
+ } catch (error: Exception) {
+ val previous = failure
+ if (previous == null) failure = error else previous.addSuppressed(error)
+ }
+ }
+
+ private fun progress() {
+ dirty = true
+ if (progressing) return
+ progressing = true
+ try {
+ while (dirty) {
+ dirty = false
+ val error = failure
+ if (error != null) {
+ val read = reader; reader = null; read?.completeExceptionally(error)
+ val write = writer; writer = null; write?.result?.completeExceptionally(error)
+ endResult?.completeExceptionally(error); endResult = null
+ for (drain in drains.toList()) drain.completeExceptionally(error)
+ continue
+ }
+ if (!suspended || closed) {
+ reader?.let { result ->
+ if (received.isNotEmpty()) {
+ val bytes = received.removeFirst()
+ reader = null
+ consumedBytes += bytes.size
+ if (!closed) dispatch(StateActionTcpDataConsumed(TcpDataConsumedAction(ActionType.TCP_DATA_CONSUMED, consumedBytes)))
+ result.complete(bytes)
+ } else if (closed || current.hostClosed || current.output.eofAtBytes != null) {
+ reader = null
+ result.complete(null)
+ }
+ }
+ }
+ while (!suspended && !closing && !closed) {
+ val write = writer ?: break
+ if (write.offset == write.data.size) {
+ writer = null
+ write.result.complete(Unit)
+ break
+ }
+ val credit = current.input.windowBytes - (sentBytes - current.input.consumedBytes)
+ if (credit == 0L) break
+ val count = minOf(credit, current.input.maximumChunkSize, (write.data.size - write.offset).toLong()).toInt()
+ safe(sentBytes + count)
+ val action = StateActionTcpInput(TcpInputAction(ActionType.TCP_INPUT, sentBytes,
+ Base64.getEncoder().encodeToString(write.data.copyOfRange(write.offset, write.offset + count))))
+ sentBytes += count
+ write.offset += count
+ dispatch(action)
+ }
+ if (!suspended && !closing && !closed) {
+ endResult?.let { result ->
+ endResult = null
+ dispatch(StateActionTcpInputEof(TcpInputEofAction(ActionType.TCP_INPUT_EOF, sentBytes)))
+ result.complete(Unit)
+ }
+ }
+ if (!closed && closing && current.clientClosed && current.hostClosed
+ && current.input.consumedBytes >= sentBytes
+ && current.output.consumedBytes >= current.output.receivedBytes
+ && received.isEmpty() && pending.isEmpty()) {
+ closed = true
+ release()
+ dirty = true
+ }
+ if (closing || closed) {
+ val errorClosed = IllegalStateException("TCP connection closed")
+ val write = writer; writer = null; write?.result?.completeExceptionally(errorClosed)
+ endResult?.completeExceptionally(errorClosed); endResult = null
+ }
+ if (current.input.consumedBytes >= sentBytes) {
+ for (drain in drains.toList()) drain.complete(Unit)
+ }
+ }
+ } catch (error: Exception) {
+ failLocked(error, reset = false)
+ } finally {
+ progressing = false
+ }
+ if (dirty) progress()
+ }
+
+ public companion object {
+ private fun safe(value: Long) { require(value in 0..9007199254740991L) { "TCP counter must be a nonnegative safe integer" } }
+
+ private fun validateLimits(windowBytes: Long, maximumChunkSize: Long) {
+ require(windowBytes in 1..4294967295L && maximumChunkSize in 1..windowBytes) {
+ "TCP window and chunk limits must be positive UInt32 values, with chunk no larger than window"
+ }
+ }
+
+ /** Validates capability and flow limits before the injected transport sends subscribe. */
+ public fun creationParameters(session: String, create: TcpConnectionSubscription, initialized: InitializeResult): SubscribeParams {
+ require(session.startsWith("ahp-session:") && create.type == "tcpConnection" && create.host.isNotBlank()
+ && create.port in 1..65535 && create.encoding == TcpDataEncoding.BASE64
+ && initialized.tcpConnections?.encodings?.contains(create.encoding) == true) { "Invalid or unsupported TCP creation request" }
+ validateLimits(create.receiveWindowBytes, create.maximumChunkSize)
+ return SubscribeParams(channel = session, create = create)
+ }
+
+ /** Seeds only a validated, freshly created stream; never a reconnect snapshot. */
+ public fun open(session: String, create: TcpConnectionSubscription, initialized: InitializeResult,
+ result: SubscribeResult, transport: TcpConnectionTransport): TcpConnection {
+ val snapshot = requireNotNull(result.snapshot) { "Missing TCP creation snapshot" }
+ try {
+ creationParameters(session, create, initialized)
+ val state = (snapshot.state as? SnapshotState.Tcp)?.value
+ ?: throw IllegalArgumentException("Invalid TCP creation snapshot")
+ require(snapshot.resource.startsWith("ahp-tcp:") && state.session == session && state.target.host == create.host
+ && state.target.port == create.port && state.encoding == create.encoding
+ && !state.clientClosed && !state.hostClosed && state.reset == null) { "Invalid TCP creation snapshot" }
+ safe(snapshot.fromSeq)
+ for (direction in listOf(state.input, state.output)) {
+ validateLimits(direction.windowBytes, direction.maximumChunkSize)
+ require(direction.receivedBytes == 0L && direction.consumedBytes == 0L && direction.eofAtBytes == null) { "TCP creation requires fresh directions" }
+ }
+ require(state.output.windowBytes <= create.receiveWindowBytes && state.output.maximumChunkSize <= create.maximumChunkSize)
+ return TcpConnection(snapshot.resource, state, snapshot.fromSeq, transport)
+ } catch (error: IllegalArgumentException) {
+ if (snapshot.resource.startsWith("ahp-tcp:")) transport.unsubscribe(snapshot.resource)
+ throw error
+ }
+ }
+
+ /** Call before sending reconnect; the transport must buffer subsequent live events until resume returns. */
+ public fun reconnectParameters(parameters: ReconnectParams, connections: List): ReconnectParams {
+ safe(parameters.lastSeenServerSeq)
+ require(parameters.channel == "ahp-root://")
+ require(connections.map { it.resource }.distinct().size == connections.size)
+ var checkpoint = parameters.lastSeenServerSeq
+ for (connection in connections) synchronized(connection.lock) {
+ require(connection.clientId == parameters.clientId && connection.suspended && !connection.closed) { "TCP reconnect requires suspended handles owned by the same client" }
+ checkpoint = minOf(checkpoint, connection.checkpoint)
+ }
+ return parameters.copy(lastSeenServerSeq = checkpoint, subscriptions = (parameters.subscriptions + connections.map { it.resource }).distinct())
+ }
+
+ /** Rebinds, applies replay, cleans acknowledgements, then resends original pending identities. */
+ public fun resume(parameters: ReconnectParams, result: ReconnectResult, connections: List, transport: TcpConnectionTransport) {
+ require(transport.clientId == parameters.clientId)
+ val checked = reconnectParameters(parameters, connections)
+ require(checked.lastSeenServerSeq == parameters.lastSeenServerSeq && parameters.subscriptions.containsAll(checked.subscriptions)) {
+ "Use reconnectParameters before sending reconnect"
+ }
+ val lastSequence = connections.maxOfOrNull { synchronized(it.lock) {
+ val assigned = it.transport.lastAssignedSequence()
+ if (assigned != -1L) safe(assigned)
+ it.lastSequence = maxOf(it.lastSequence, assigned)
+ it.lastSequence
+ } }
+ if (lastSequence != null && lastSequence >= 0) transport.advanceSequencePast(lastSequence)
+ if (result is ReconnectResultReplay) {
+ var previous = parameters.lastSeenServerSeq
+ for (envelope in result.value.actions) {
+ safe(envelope.serverSeq)
+ require(envelope.serverSeq > previous) { "TCP replay is not ordered" }
+ previous = envelope.serverSeq
+ }
+ }
+ for (connection in connections) synchronized(connection.lock) {
+ connection.transport = transport
+ if (result !is ReconnectResultReplay || connection.resource in result.value.missing) {
+ connection.failLocked(IllegalStateException("TCP cannot resume missing resources or snapshot fallback"), reset = false)
+ } else {
+ for (envelope in result.value.actions) connection.apply(envelope)
+ if (!connection.closed) {
+ connection.suspended = false
+ connection.transportFailure = null
+ for ((sequence, action) in connection.pending.toMap()) {
+ if (connection.suspended) break
+ connection.send(sequence, action)
+ }
+ connection.progress()
+ }
+ }
+ }
+ }
+ }
+}
diff --git a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Actions.generated.kt b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Actions.generated.kt
index 3607b366b..dc4c6e8ca 100644
--- a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Actions.generated.kt
+++ b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Actions.generated.kt
@@ -132,6 +132,16 @@ value class ActionType(val rawValue: String) {
val AUTOMATION_RUN_SESSION_REMOVED: ActionType = ActionType("automationRun/sessionRemoved")
val AUTOMATION_RUN_PRIMARY_SESSION_CHANGED: ActionType = ActionType("automationRun/primarySessionChanged")
val AUTOMATION_RUN_CANCEL_REQUESTED: ActionType = ActionType("automationRun/cancelRequested")
+ val TCP_INPUT: ActionType = ActionType("tcp/input")
+ val TCP_DATA: ActionType = ActionType("tcp/data")
+ val TCP_INPUT_CONSUMED: ActionType = ActionType("tcp/inputConsumed")
+ val TCP_DATA_CONSUMED: ActionType = ActionType("tcp/dataConsumed")
+ val TCP_INPUT_EOF: ActionType = ActionType("tcp/inputEof")
+ val TCP_DATA_EOF: ActionType = ActionType("tcp/dataEof")
+ val TCP_CLIENT_CLOSE: ActionType = ActionType("tcp/clientClose")
+ val TCP_HOST_CLOSE: ActionType = ActionType("tcp/hostClose")
+ val TCP_CLIENT_RESET: ActionType = ActionType("tcp/clientReset")
+ val TCP_HOST_RESET: ActionType = ActionType("tcp/hostReset")
}
}
@@ -1507,6 +1517,78 @@ data class ResourceWatchChangedAction(
val changes: JsonElement
)
+@Serializable
+data class TcpInputAction(
+ val type: ActionType,
+ /**
+ * Absolute decoded-byte offset.
+ */
+ val offset: Long,
+ /**
+ * Nonempty canonical padded RFC 4648 base64; no whitespace.
+ */
+ val data: String
+)
+
+@Serializable
+data class TcpDataAction(
+ val type: ActionType,
+ /**
+ * Absolute decoded-byte offset.
+ */
+ val offset: Long,
+ /**
+ * Nonempty canonical padded RFC 4648 base64; no whitespace.
+ */
+ val data: String
+)
+
+@Serializable
+data class TcpInputConsumedAction(
+ val type: ActionType,
+ val consumedBytes: Long
+)
+
+@Serializable
+data class TcpDataConsumedAction(
+ val type: ActionType,
+ val consumedBytes: Long
+)
+
+@Serializable
+data class TcpInputEofAction(
+ val type: ActionType,
+ val finalOffset: Long
+)
+
+@Serializable
+data class TcpDataEofAction(
+ val type: ActionType,
+ val finalOffset: Long
+)
+
+@Serializable
+data class TcpClientCloseAction(
+ val type: ActionType
+)
+
+@Serializable
+data class TcpHostCloseAction(
+ val type: ActionType
+)
+
+@Serializable
+data class TcpClientResetAction(
+ val type: ActionType,
+ val reason: TcpResetReason
+)
+
+@Serializable
+data class TcpHostResetAction(
+ val type: ActionType,
+ val reason: TcpResetReason
+)
+
@Serializable
data class AutomationCreateRequestedAction(
val type: ActionType,
@@ -1756,6 +1838,16 @@ sealed interface StateAction
@JvmInline value class StateActionTerminalCommandExecuted(val value: TerminalCommandExecutedAction) : StateAction
@JvmInline value class StateActionTerminalCommandFinished(val value: TerminalCommandFinishedAction) : StateAction
@JvmInline value class StateActionResourceWatchChanged(val value: ResourceWatchChangedAction) : StateAction
+@JvmInline value class StateActionTcpInput(val value: TcpInputAction) : StateAction
+@JvmInline value class StateActionTcpData(val value: TcpDataAction) : StateAction
+@JvmInline value class StateActionTcpInputConsumed(val value: TcpInputConsumedAction) : StateAction
+@JvmInline value class StateActionTcpDataConsumed(val value: TcpDataConsumedAction) : StateAction
+@JvmInline value class StateActionTcpInputEof(val value: TcpInputEofAction) : StateAction
+@JvmInline value class StateActionTcpDataEof(val value: TcpDataEofAction) : StateAction
+@JvmInline value class StateActionTcpClientClose(val value: TcpClientCloseAction) : StateAction
+@JvmInline value class StateActionTcpHostClose(val value: TcpHostCloseAction) : StateAction
+@JvmInline value class StateActionTcpClientReset(val value: TcpClientResetAction) : StateAction
+@JvmInline value class StateActionTcpHostReset(val value: TcpHostResetAction) : StateAction
@JvmInline value class StateActionAutomationCreateRequested(val value: AutomationCreateRequestedAction) : StateAction
@JvmInline value class StateActionAutomationUpdateRequested(val value: AutomationUpdateRequestedAction) : StateAction
@JvmInline value class StateActionAutomationSet(val value: AutomationSetAction) : StateAction
@@ -1874,6 +1966,16 @@ internal object StateActionSerializer : KSerializer {
"terminal/commandExecuted" -> StateActionTerminalCommandExecuted(input.json.decodeFromJsonElement(TerminalCommandExecutedAction.serializer(), element))
"terminal/commandFinished" -> StateActionTerminalCommandFinished(input.json.decodeFromJsonElement(TerminalCommandFinishedAction.serializer(), element))
"resourceWatch/changed" -> StateActionResourceWatchChanged(input.json.decodeFromJsonElement(ResourceWatchChangedAction.serializer(), element))
+ "tcp/input" -> StateActionTcpInput(input.json.decodeFromJsonElement(TcpInputAction.serializer(), element))
+ "tcp/data" -> StateActionTcpData(input.json.decodeFromJsonElement(TcpDataAction.serializer(), element))
+ "tcp/inputConsumed" -> StateActionTcpInputConsumed(input.json.decodeFromJsonElement(TcpInputConsumedAction.serializer(), element))
+ "tcp/dataConsumed" -> StateActionTcpDataConsumed(input.json.decodeFromJsonElement(TcpDataConsumedAction.serializer(), element))
+ "tcp/inputEof" -> StateActionTcpInputEof(input.json.decodeFromJsonElement(TcpInputEofAction.serializer(), element))
+ "tcp/dataEof" -> StateActionTcpDataEof(input.json.decodeFromJsonElement(TcpDataEofAction.serializer(), element))
+ "tcp/clientClose" -> StateActionTcpClientClose(input.json.decodeFromJsonElement(TcpClientCloseAction.serializer(), element))
+ "tcp/hostClose" -> StateActionTcpHostClose(input.json.decodeFromJsonElement(TcpHostCloseAction.serializer(), element))
+ "tcp/clientReset" -> StateActionTcpClientReset(input.json.decodeFromJsonElement(TcpClientResetAction.serializer(), element))
+ "tcp/hostReset" -> StateActionTcpHostReset(input.json.decodeFromJsonElement(TcpHostResetAction.serializer(), element))
"automation/createRequested" -> StateActionAutomationCreateRequested(input.json.decodeFromJsonElement(AutomationCreateRequestedAction.serializer(), element))
"automation/updateRequested" -> StateActionAutomationUpdateRequested(input.json.decodeFromJsonElement(AutomationUpdateRequestedAction.serializer(), element))
"automation/set" -> StateActionAutomationSet(input.json.decodeFromJsonElement(AutomationSetAction.serializer(), element))
@@ -1985,6 +2087,16 @@ internal object StateActionSerializer : KSerializer {
is StateActionTerminalCommandExecuted -> output.json.encodeToJsonElement(TerminalCommandExecutedAction.serializer(), value.value)
is StateActionTerminalCommandFinished -> output.json.encodeToJsonElement(TerminalCommandFinishedAction.serializer(), value.value)
is StateActionResourceWatchChanged -> output.json.encodeToJsonElement(ResourceWatchChangedAction.serializer(), value.value)
+ is StateActionTcpInput -> output.json.encodeToJsonElement(TcpInputAction.serializer(), value.value)
+ is StateActionTcpData -> output.json.encodeToJsonElement(TcpDataAction.serializer(), value.value)
+ is StateActionTcpInputConsumed -> output.json.encodeToJsonElement(TcpInputConsumedAction.serializer(), value.value)
+ is StateActionTcpDataConsumed -> output.json.encodeToJsonElement(TcpDataConsumedAction.serializer(), value.value)
+ is StateActionTcpInputEof -> output.json.encodeToJsonElement(TcpInputEofAction.serializer(), value.value)
+ is StateActionTcpDataEof -> output.json.encodeToJsonElement(TcpDataEofAction.serializer(), value.value)
+ is StateActionTcpClientClose -> output.json.encodeToJsonElement(TcpClientCloseAction.serializer(), value.value)
+ is StateActionTcpHostClose -> output.json.encodeToJsonElement(TcpHostCloseAction.serializer(), value.value)
+ is StateActionTcpClientReset -> output.json.encodeToJsonElement(TcpClientResetAction.serializer(), value.value)
+ is StateActionTcpHostReset -> output.json.encodeToJsonElement(TcpHostResetAction.serializer(), value.value)
is StateActionAutomationCreateRequested -> output.json.encodeToJsonElement(AutomationCreateRequestedAction.serializer(), value.value)
is StateActionAutomationUpdateRequested -> output.json.encodeToJsonElement(AutomationUpdateRequestedAction.serializer(), value.value)
is StateActionAutomationSet -> output.json.encodeToJsonElement(AutomationSetAction.serializer(), value.value)
diff --git a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Commands.generated.kt b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Commands.generated.kt
index 48ed5f7b7..e85766d4d 100644
--- a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Commands.generated.kt
+++ b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Commands.generated.kt
@@ -418,7 +418,11 @@ data class InitializeResult(
* `ahp-automations://` for {@link AutomationState}; absence means the
* host does not expose an automation catalogue or automation commands.
*/
- val automations: AutomationCapabilities? = null
+ val automations: AutomationCapabilities? = null,
+ /**
+ * Enables atomic creation of session-scoped, replay-only TCP channels.
+ */
+ val tcpConnections: TcpConnectionsCapability? = null
)
@Serializable
@@ -557,7 +561,13 @@ data class ReconnectSnapshotResult(
/**
* Fresh snapshots for each subscription
*/
- val snapshots: List
+ val snapshots: List,
+ /**
+ * Subscriptions that cannot be restored. Hosts supporting TCP MUST list all
+ * requested TCP channels here and dispose their sockets on snapshot fallback.
+ * Omitted by older hosts; absence does not authorize snapshot-restoring TCP.
+ */
+ val missing: List? = null
)
@Serializable
@@ -586,7 +596,13 @@ data class SubscribeParams(
* Servers that do not understand a requested view ignore it and return their
* default snapshot. Clients MUST tolerate receiving more state than requested.
*/
- val view: SubscribeView? = null
+ val view: SubscribeView? = null,
+ /**
+ * Atomically create a private child channel and subscribe to it.
+ * Requires the advertised tcpConnections capability. channel identifies
+ * the parent session; snapshot.resource identifies the created TCP channel.
+ */
+ val create: TcpConnectionSubscription? = null
)
@Serializable
@@ -623,6 +639,31 @@ data class SubscribeResult(
val snapshot: Snapshot? = null
)
+@Serializable
+data class TcpConnectionSubscription(
+ val type: String,
+ /**
+ * DNS name or IP literal, not a URL.
+ */
+ val host: String,
+ /**
+ * Destination port.
+ */
+ val port: Long,
+ /**
+ * Selected from InitializeResult.tcpConnections.encodings.
+ */
+ val encoding: TcpDataEncoding,
+ /**
+ * Client receive window in decoded bytes.
+ */
+ val receiveWindowBytes: Long,
+ /**
+ * Maximum decoded bytes per output action; MUST NOT exceed receiveWindowBytes.
+ */
+ val maximumChunkSize: Long
+)
+
@Serializable
data class CreateSessionParams(
/**
diff --git a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Errors.generated.kt b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Errors.generated.kt
index c1ba5c42b..ffe6fa9d4 100644
--- a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Errors.generated.kt
+++ b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/Errors.generated.kt
@@ -57,6 +57,7 @@ object AhpErrorCodes {
const val PERMISSION_DENIED: Int = -32009
/** The target resource already exists and the operation does not allow overwriting */
const val ALREADY_EXISTS: Int = -32010
+ const val TCP_CONNECTION_OPEN_FAILED: Int = -32012
}
// ─── Error Detail Payloads ──────────────────────────────────────────────────
diff --git a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/State.generated.kt b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/State.generated.kt
index 87ba426a8..3f7cdbc74 100644
--- a/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/State.generated.kt
+++ b/clients/kotlin/src/main/kotlin/com/microsoft/agenthostprotocol/generated/State.generated.kt
@@ -72,6 +72,89 @@ internal object StringOrMarkdownSerializer : KSerializer {
// ─── Enums ──────────────────────────────────────────────────────────────────
+/**
+ * Payload encodings advertised by the host.
+ */
+@Serializable(with = TcpDataEncodingSerializer::class)
+@JvmInline
+value class TcpDataEncoding(val rawValue: String) {
+ companion object {
+ val BASE64: TcpDataEncoding = TcpDataEncoding("base64")
+ }
+}
+
+internal object TcpDataEncodingSerializer : KSerializer {
+ override val descriptor: SerialDescriptor =
+ PrimitiveSerialDescriptor("TcpDataEncoding", PrimitiveKind.STRING)
+ override fun serialize(encoder: Encoder, value: TcpDataEncoding) {
+ encoder.encodeString(value.rawValue)
+ }
+ override fun deserialize(decoder: Decoder): TcpDataEncoding =
+ TcpDataEncoding(decoder.decodeString())
+}
+
+/**
+ * Endpoint that closes or resets a connection.
+ */
+@Serializable
+enum class TcpEndpoint {
+ @SerialName("client")
+ CLIENT,
+ @SerialName("host")
+ HOST
+}
+
+/**
+ * Why a connection was aborted.
+ */
+@Serializable(with = TcpResetReasonSerializer::class)
+@JvmInline
+value class TcpResetReason(val rawValue: String) {
+ companion object {
+ val CONNECTION_RESET: TcpResetReason = TcpResetReason("connectionReset")
+ val CONNECTION_ABORTED: TcpResetReason = TcpResetReason("connectionAborted")
+ val PROTOCOL_ERROR: TcpResetReason = TcpResetReason("protocolError")
+ val REPLAY_UNAVAILABLE: TcpResetReason = TcpResetReason("replayUnavailable")
+ val POLICY_REVOKED: TcpResetReason = TcpResetReason("policyRevoked")
+ val SESSION_DISPOSED: TcpResetReason = TcpResetReason("sessionDisposed")
+ val INTERNAL_ERROR: TcpResetReason = TcpResetReason("internalError")
+ }
+}
+
+internal object TcpResetReasonSerializer : KSerializer {
+ override val descriptor: SerialDescriptor =
+ PrimitiveSerialDescriptor("TcpResetReason", PrimitiveKind.STRING)
+ override fun serialize(encoder: Encoder, value: TcpResetReason) {
+ encoder.encodeString(value.rawValue)
+ }
+ override fun deserialize(decoder: Decoder): TcpResetReason =
+ TcpResetReason(decoder.decodeString())
+}
+
+/**
+ * Expected connection establishment failures.
+ */
+@Serializable(with = TcpConnectionOpenFailureReasonSerializer::class)
+@JvmInline
+value class TcpConnectionOpenFailureReason(val rawValue: String) {
+ companion object {
+ val CONNECTION_FAILED: TcpConnectionOpenFailureReason = TcpConnectionOpenFailureReason("connectionFailed")
+ val NAME_RESOLUTION_FAILED: TcpConnectionOpenFailureReason = TcpConnectionOpenFailureReason("nameResolutionFailed")
+ val RESOURCE_SHORTAGE: TcpConnectionOpenFailureReason = TcpConnectionOpenFailureReason("resourceShortage")
+ val SESSION_NOT_READY: TcpConnectionOpenFailureReason = TcpConnectionOpenFailureReason("sessionNotReady")
+ }
+}
+
+internal object TcpConnectionOpenFailureReasonSerializer : KSerializer {
+ override val descriptor: SerialDescriptor =
+ PrimitiveSerialDescriptor("TcpConnectionOpenFailureReason", PrimitiveKind.STRING)
+ override fun serialize(encoder: Encoder, value: TcpConnectionOpenFailureReason) {
+ encoder.encodeString(value.rawValue)
+ }
+ override fun deserialize(decoder: Decoder): TcpConnectionOpenFailureReason =
+ TcpConnectionOpenFailureReason(decoder.decodeString())
+}
+
/**
* Policy configuration state for a model.
*/
@@ -5427,6 +5510,84 @@ data class ResourceChange(
val type: ResourceChangeType
)
+@Serializable
+data class TcpConnectionState(
+ val session: String,
+ val target: TcpTarget,
+ val encoding: TcpDataEncoding,
+ /**
+ * Client to destination socket.
+ */
+ val input: FlowControlledByteDirectionState,
+ /**
+ * Destination socket to client.
+ */
+ val output: FlowControlledByteDirectionState,
+ val clientClosed: Boolean,
+ val hostClosed: Boolean,
+ val reset: TcpResetState? = null
+)
+
+@Serializable
+data class TcpTarget(
+ /**
+ * DNS name or IP literal, resolved and connected in the host endpoint's network.
+ */
+ val host: String,
+ /**
+ * Destination port.
+ */
+ val port: Long
+)
+
+@Serializable
+data class TcpResetState(
+ val source: TcpEndpoint,
+ val reason: TcpResetReason
+)
+
+@Serializable
+data class FlowControlledByteDirectionState(
+ /**
+ * Maximum accepted-but-not-consumed decoded bytes.
+ */
+ val windowBytes: Long,
+ /**
+ * Maximum decoded bytes per chunk; MUST NOT exceed windowBytes.
+ */
+ val maximumChunkSize: Long,
+ /**
+ * Cumulative accepted bytes.
+ */
+ val receivedBytes: Long,
+ /**
+ * Cumulative bytes released by the bounded consumer.
+ */
+ val consumedBytes: Long,
+ /**
+ * Present after EOF; equals receivedBytes permanently.
+ */
+ val eofAtBytes: Long? = null
+)
+
+@Serializable
+data class TcpConnectionsCapability(
+ /**
+ * Supported encodings. The base64 profile MUST be supported.
+ */
+ val encodings: List,
+ /**
+ * Informational limit; runtime policy may impose a lower limit.
+ */
+ val maximumConnectionsPerClient: Long? = null
+)
+
+@Serializable
+data class TcpConnectionOpenErrorData(
+ val reason: TcpConnectionOpenFailureReason,
+ val retryable: Boolean? = null
+)
+
@Serializable
data class AutomationSessionOrigin(
val kind: SessionOriginKind,
@@ -7420,6 +7581,7 @@ internal object ToolResultContentSerializer : KSerializer {
*/
@Serializable(with = SnapshotStateSerializer::class)
sealed interface SnapshotState {
+ @JvmInline value class Tcp(val value: TcpConnectionState) : SnapshotState
@JvmInline value class Root(val value: RootState) : SnapshotState
@JvmInline value class Session(val value: SessionState) : SnapshotState
@JvmInline value class Chat(val value: ChatState) : SnapshotState
@@ -7441,7 +7603,8 @@ internal object SnapshotStateSerializer : KSerializer {
val element = input.decodeJsonElement()
val obj = element as? JsonObject
?: error("Expected JsonObject for SnapshotState")
- // Try the most distinctive shape first. AutomationRunState has required
+ // Try the most distinctive shape first. TcpConnectionState has required
+ // `input`, `output`, and `target`; AutomationRunState has required
// `automation`, `origin`, and `sessions`; AutomationState has
// required `entries`; SessionState has required
// `lifecycle`; ChatState has required `turns`; ChangesetState has
@@ -7451,6 +7614,8 @@ internal object SnapshotStateSerializer : KSerializer {
// key); TerminalState has required `content`; RootState is the
// catch-all.
return when {
+ obj.containsKey("input") && obj.containsKey("output") && obj.containsKey("target") ->
+ SnapshotState.Tcp(input.json.decodeFromJsonElement(TcpConnectionState.serializer(), element))
obj.containsKey("automation") && obj.containsKey("origin") && obj.containsKey("sessions") ->
SnapshotState.AutomationRun(input.json.decodeFromJsonElement(AutomationRunState.serializer(), element))
obj.containsKey("entries") ->
@@ -7473,6 +7638,7 @@ internal object SnapshotStateSerializer : KSerializer {
val output = encoder as? JsonEncoder
?: error("SnapshotState can only be serialized to JSON")
val element: JsonElement = when (value) {
+ is SnapshotState.Tcp -> output.json.encodeToJsonElement(TcpConnectionState.serializer(), value.value)
is SnapshotState.Root -> output.json.encodeToJsonElement(RootState.serializer(), value.value)
is SnapshotState.Session -> output.json.encodeToJsonElement(SessionState.serializer(), value.value)
is SnapshotState.Chat -> output.json.encodeToJsonElement(ChatState.serializer(), value.value)
diff --git a/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/FixtureDrivenReducerTest.kt b/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/FixtureDrivenReducerTest.kt
index c3f4675eb..1d286b55a 100644
--- a/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/FixtureDrivenReducerTest.kt
+++ b/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/FixtureDrivenReducerTest.kt
@@ -10,10 +10,14 @@ import com.microsoft.agenthostprotocol.generated.RootState
import com.microsoft.agenthostprotocol.generated.SessionState
import com.microsoft.agenthostprotocol.generated.StateAction
import com.microsoft.agenthostprotocol.generated.TerminalState
+import com.microsoft.agenthostprotocol.generated.TcpConnectionState
import java.io.File
-import kotlinx.serialization.builtins.ListSerializer
+import kotlinx.serialization.SerializationException
import kotlinx.serialization.json.JsonElement
import kotlinx.serialization.json.JsonObject
+import kotlinx.serialization.json.jsonArray
+import kotlinx.serialization.json.jsonPrimitive
+import kotlinx.serialization.json.doubleOrNull
import kotlinx.serialization.json.jsonObject
import org.junit.jupiter.api.DynamicTest
import org.junit.jupiter.api.Test
@@ -21,6 +25,8 @@ import org.junit.jupiter.api.TestFactory
import org.junit.jupiter.api.assertAll
import org.junit.jupiter.api.fail
import kotlin.test.assertTrue
+import kotlin.test.assertEquals
+import kotlin.test.assertFailsWith
/**
* JSON-fixture-driven reducer tests for cross-language parity.
@@ -110,10 +116,33 @@ class FixtureDrivenReducerTest {
)
}
- val actions = Ahp.json.decodeFromJsonElement(
- ListSerializer(StateAction.serializer()),
- actionsArr,
- )
+ val actions = actionsArr.jsonArray
+ val expectedError = fixture["expectedError"]?.jsonPrimitiveContent()
+ if (expectedError != null) assertTrue(actions.isNotEmpty(), "expectedError requires a final action")
+
+ fun runActions(initialState: S, reduce: (S, StateAction) -> S): S {
+ var state = initialState
+ for ((index, raw) in actions.withIndex()) {
+ val mustFail = expectedError != null && index == actions.lastIndex
+ val fractionalOffset = raw.jsonObject["offset"]?.jsonPrimitive?.doubleOrNull
+ ?.let { it % 1.0 != 0.0 } == true
+ if (mustFail && reducer == "tcp" && fractionalOffset) {
+ assertEquals("Invalid TCP action: offset must be a nonnegative safe integer", expectedError)
+ assertFailsWith {
+ Ahp.json.decodeFromJsonElement(StateAction.serializer(), raw)
+ }
+ } else {
+ val action = Ahp.json.decodeFromJsonElement(StateAction.serializer(), raw)
+ if (mustFail) {
+ val error = assertFailsWith { reduce(state, action) }
+ assertEquals(expectedError, error.message, file.name)
+ } else {
+ state = reduce(state, action)
+ }
+ }
+ }
+ return state
+ }
when (reducer) {
"root" -> compareFixture(
@@ -122,9 +151,7 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = RootState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = rootReducer(s, action)
- s
+ runActions(state, ::rootReducer)
},
)
@@ -134,9 +161,7 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = SessionState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = sessionReducer(s, action)
- s
+ runActions(state, ::sessionReducer)
},
)
@@ -146,9 +171,7 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = ChatState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = chatReducer(s, action)
- s
+ runActions(state, ::chatReducer)
},
)
@@ -158,9 +181,7 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = TerminalState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = terminalReducer(s, action)
- s
+ runActions(state, ::terminalReducer)
},
)
@@ -170,9 +191,7 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = ChangesetState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = changesetReducer(s, action)
- s
+ runActions(state, ::changesetReducer)
},
)
@@ -182,9 +201,7 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = AnnotationsState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = annotationsReducer(s, action)
- s
+ runActions(state, ::annotationsReducer)
},
)
@@ -194,9 +211,7 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = ResourceWatchState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = resourceWatchReducer(s, action)
- s
+ runActions(state, ::resourceWatchReducer)
},
)
@@ -206,9 +221,7 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = AutomationState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = automationReducer(s, action)
- s
+ runActions(state, ::automationReducer)
},
)
@@ -218,12 +231,18 @@ class FixtureDrivenReducerTest {
expected = expected,
serializer = AutomationRunState.serializer(),
run = { state ->
- var s = state
- for (action in actions) s = automationRunReducer(s, action)
- s
+ runActions(state, ::automationRunReducer)
},
)
+ "tcp" -> compareFixture(
+ file = file,
+ initial = initial,
+ expected = expected,
+ serializer = TcpConnectionState.serializer(),
+ run = { state -> runActions(state, ::tcpReducer) },
+ )
+
else -> fail("${file.name}: unsupported reducer '$reducer'")
}
}
diff --git a/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/GeneratedStructsTest.kt b/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/GeneratedStructsTest.kt
index f792a4ff2..954eb9cef 100644
--- a/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/GeneratedStructsTest.kt
+++ b/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/GeneratedStructsTest.kt
@@ -9,6 +9,9 @@ import com.microsoft.agenthostprotocol.generated.ProtectedResourceMetadata
import com.microsoft.agenthostprotocol.generated.SessionAddedParams
import com.microsoft.agenthostprotocol.generated.SessionModelInfo
import com.microsoft.agenthostprotocol.generated.SessionStatus
+import com.microsoft.agenthostprotocol.generated.SubscribeParams
+import com.microsoft.agenthostprotocol.generated.SubscribeView
+import com.microsoft.agenthostprotocol.generated.SubscriptionDeliveryOptions
import kotlinx.serialization.json.Json
import kotlinx.serialization.json.JsonObject
import kotlinx.serialization.json.JsonPrimitive
@@ -32,6 +35,22 @@ import kotlin.test.assertTrue
class GeneratedStructsTest {
private val json: Json = Ahp.json
+ @Test
+ fun `subscribe preserves existing positional constructor copy and component arguments`() {
+ val delivery = SubscriptionDeliveryOptions(0L)
+ val view = SubscribeView(10L)
+ val params = SubscribeParams("ahp-session:/s1", null, delivery, view)
+ assertEquals(delivery, params.delivery)
+ assertEquals(view, params.view)
+ val copy = params.copy("ahp-session:/s2", null, delivery, view)
+ val (channel, meta, copiedDelivery, copiedView) = copy
+ assertEquals("ahp-session:/s2", channel)
+ assertEquals(null, meta)
+ assertEquals(delivery, copiedDelivery)
+ assertEquals(view, copiedView)
+ assertEquals(null, copy.create)
+ }
+
@Test
fun `plain enum encodes wire string and decodes back`() {
val encoded = json.encodeToString(PolicyState.serializer(), PolicyState.UNCONFIGURED)
diff --git a/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/ReducersTest.kt b/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/ReducersTest.kt
index 78f73c2a5..d59f29fb4 100644
--- a/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/ReducersTest.kt
+++ b/clients/kotlin/src/test/kotlin/com/microsoft/agenthostprotocol/ReducersTest.kt
@@ -11,6 +11,7 @@ import kotlin.test.assertIs
import kotlin.test.assertNull
import kotlin.test.assertSame
import kotlin.test.assertTrue
+import kotlin.test.assertFailsWith
/**
* Focused unit tests covering the reducer module's public surface, the
@@ -21,6 +22,471 @@ import kotlin.test.assertTrue
*/
class ReducersTest {
+ @Test
+ fun `TCP creation requires canonical discriminator`() {
+ val create = TcpConnectionSubscription("tcpConnection", "localhost", 3000, TcpDataEncoding.BASE64, 4, 2)
+ val initialized = InitializeResult(protocolVersion = "1", snapshots = emptyList(), serverSeq = 0,
+ tcpConnections = TcpConnectionsCapability(listOf(TcpDataEncoding.BASE64)))
+ val creation = TcpConnectionCreation("ahp-session:/test", create, initialized, TcpHarness().transport)
+ val wire = Ahp.json.encodeToJsonElement(SubscribeParams.serializer(), creation.parameters) as JsonObject
+ assertEquals(JsonPrimitive("tcpConnection"), (wire["create"] as JsonObject)["type"])
+ assertFailsWith {
+ TcpConnection.creationParameters("ahp-session:/test", create.copy(type = "tcp"), initialized)
+ }
+ }
+
+ @Test
+ fun `TCP local close retains crossing traffic until both directions drain`() {
+ val harness = TcpHarness()
+ val connection = harness.open()
+ connection.write(byteArrayOf(1, 2)).join()
+ val input = harness.sent.last()
+ val drain = connection.drain()
+ connection.close()
+ val close = harness.sent.last()
+ assertIs(close.second)
+ assertTrue(!connection.isClosed)
+ assertTrue(harness.unsubscribed.isEmpty())
+ val read = connection.read()
+ assertTrue(!read.isDone)
+ connection.accept(tcpEnvelope(1, input.second, ActionOrigin("owner", input.first)))
+ connection.accept(tcpEnvelope(2, close.second, ActionOrigin("owner", close.first)))
+ connection.accept(tcpEnvelope(3, StateActionTcpData(TcpDataAction(ActionType.TCP_DATA, 0, "Bwg="))))
+ assertTrue(byteArrayOf(7, 8).contentEquals(read.join()))
+ val credit = harness.sent.last()
+ assertEquals(2, assertIs(credit.second).value.consumedBytes)
+ connection.accept(tcpEnvelope(4, StateActionTcpHostClose(TcpHostCloseAction(ActionType.TCP_HOST_CLOSE))))
+ assertTrue(!connection.isClosed && !drain.isDone && harness.unsubscribed.isEmpty())
+ connection.accept(tcpEnvelope(5, credit.second, ActionOrigin("owner", credit.first)))
+ assertTrue(!connection.isClosed && harness.unsubscribed.isEmpty())
+ connection.accept(tcpEnvelope(6, StateActionTcpInputConsumed(TcpInputConsumedAction(ActionType.TCP_INPUT_CONSUMED, 2))))
+ drain.join()
+ assertNull(connection.read().join())
+ assertTrue(connection.isClosed)
+ connection.close()
+ connection.dispose()
+ assertEquals(listOf(connection.resource), harness.unsubscribed)
+ }
+
+ @Test
+ fun `TCP creation and snapshot limits use UInt32 range`() {
+ val initialized = InitializeResult(protocolVersion = "1", snapshots = emptyList(), serverSeq = 0,
+ tcpConnections = TcpConnectionsCapability(listOf(TcpDataEncoding.BASE64)))
+ val request = TcpConnectionSubscription("tcpConnection", "localhost", 3000, TcpDataEncoding.BASE64, 4294967295L, 4294967295L)
+ for (limit in listOf(1L, 4294967295L, 0L, -1L, 4294967296L, 9007199254740991L)) {
+ val valid = limit in 1..4294967295L
+ for (chunk in listOf(false, true)) {
+ val create = request.copy(receiveWindowBytes = limit, maximumChunkSize = if (chunk) limit else 1)
+ if (valid) TcpConnection.creationParameters("ahp-session:/test", create, initialized)
+ else assertFailsWith("creation $limit") {
+ TcpConnection.creationParameters("ahp-session:/test", create, initialized)
+ }
+ for (input in listOf(false, true)) {
+ val direction = FlowControlledByteDirectionState(limit, if (chunk) limit else 1, 0, 0)
+ val state = if (input) tcpState().copy(input = direction) else tcpState().copy(output = direction)
+ val result = SubscribeResult(Snapshot("ahp-tcp:/created", SnapshotState.Tcp(state), 0))
+ val harness = TcpHarness()
+ if (valid) {
+ val connection = TcpConnection.open("ahp-session:/test", request, initialized, result, harness.transport)
+ connection.dispose()
+ } else assertFailsWith("snapshot $limit, input $input") {
+ TcpConnection.open("ahp-session:/test", request, initialized, result, harness.transport)
+ }
+ assertEquals(listOf("ahp-tcp:/created"), harness.unsubscribed)
+ }
+ }
+ }
+ }
+
+ @Test
+ fun `TCP peer close responds without waiting for credit or unread output`() {
+ val harness = TcpHarness()
+ val connection = harness.open()
+ val write = connection.write(ByteArray(5))
+ val drain = connection.drain()
+ harness.sent.toList().forEachIndexed { index, item ->
+ connection.accept(tcpEnvelope(index + 1L, item.second, ActionOrigin("owner", item.first)))
+ }
+ connection.accept(tcpEnvelope(3, StateActionTcpData(TcpDataAction(ActionType.TCP_DATA, 0, "Bwg="))))
+ connection.accept(tcpEnvelope(4, StateActionTcpHostClose(TcpHostCloseAction(ActionType.TCP_HOST_CLOSE))))
+ val close = harness.sent.last()
+ assertIs(close.second)
+ assertEquals(0, connection.state.input.consumedBytes)
+ assertEquals(0, connection.state.output.consumedBytes)
+ assertTrue(write.isCompletedExceptionally && !drain.isDone && harness.unsubscribed.isEmpty())
+ connection.accept(tcpEnvelope(5, close.second, ActionOrigin("owner", close.first)))
+ connection.accept(tcpEnvelope(6, StateActionTcpInputConsumed(TcpInputConsumedAction(ActionType.TCP_INPUT_CONSUMED, 4))))
+ drain.join()
+ assertTrue(byteArrayOf(7, 8).contentEquals(connection.read().join()))
+ val credit = harness.sent.last()
+ assertIs(credit.second)
+ connection.accept(tcpEnvelope(7, credit.second, ActionOrigin("owner", credit.first)))
+ assertNull(connection.read().join())
+ assertEquals(listOf(connection.resource), harness.unsubscribed)
+ }
+
+ @Test
+ fun `TCP creation releases late children after timeout or cancellation but never the parent`() {
+ for (cancel in listOf(false, true)) for (resource in listOf("ahp-tcp:/late", "ahp-session:/test")) {
+ val harness = TcpHarness()
+ val creation = TcpConnectionCreation("ahp-session:/test",
+ TcpConnectionSubscription("tcpConnection", "localhost", 3000, TcpDataEncoding.BASE64, 4, 4),
+ InitializeResult(protocolVersion = "1", snapshots = emptyList(), serverSeq = 0,
+ tcpConnections = TcpConnectionsCapability(listOf(TcpDataEncoding.BASE64))), harness.transport)
+ if (cancel) creation.completion.cancel(false)
+ else creation.fail(java.util.concurrent.TimeoutException("create timeout"))
+ val result = SubscribeResult(Snapshot(resource, SnapshotState.Tcp(tcpState(4)), 0))
+ creation.accept(result)
+ creation.accept(result)
+ assertTrue(creation.completion.isCompletedExceptionally)
+ assertEquals(if (resource.startsWith("ahp-tcp:")) listOf(resource) else emptyList(), harness.unsubscribed)
+ assertTrue(harness.sent.isEmpty())
+ }
+ }
+
+ private class TcpHarness {
+ val sent = mutableListOf>()
+ val unsubscribed = mutableListOf()
+ var sequence = 1L
+ val transport = TcpConnectionTransport("owner", { sequence++ }, { sequence = maxOf(sequence, it + 1) },
+ { _, seq, action -> sent.add(seq to action) }, { unsubscribed.add(it) }, { sequence - 1 })
+ fun open(maximumChunkSize: Long = 2): TcpConnection {
+ val window = maxOf(4, maximumChunkSize)
+ val direction = FlowControlledByteDirectionState(window, maximumChunkSize, 0, 0)
+ val state = TcpConnectionState("ahp-session:/test", TcpTarget("localhost", 3000), TcpDataEncoding.BASE64,
+ direction, direction, false, false)
+ val creation = TcpConnectionCreation("ahp-session:/test",
+ TcpConnectionSubscription("tcpConnection", "localhost", 3000, TcpDataEncoding.BASE64, window, maximumChunkSize),
+ InitializeResult(protocolVersion = "1", snapshots = emptyList(), serverSeq = 0, tcpConnections = TcpConnectionsCapability(listOf(TcpDataEncoding.BASE64))),
+ transport)
+ val wire = Ahp.json.encodeToJsonElement(SubscribeParams.serializer(), creation.parameters) as JsonObject
+ assertEquals(JsonPrimitive("tcpConnection"), (wire["create"] as JsonObject)["type"])
+ creation.accept(SubscribeResult(Snapshot("ahp-tcp:/created", SnapshotState.Tcp(state), 0)))
+ return creation.completion.join()
+ }
+ }
+
+ @Test
+ fun `TCP reset or dispose terminates a closing stream`() {
+ for (reset in listOf(false, true)) {
+ val harness = TcpHarness()
+ val connection = harness.open()
+ val write = connection.write(ByteArray(5))
+ val drain = connection.drain()
+ connection.close()
+ assertTrue(write.isCompletedExceptionally)
+ assertTrue(!drain.isDone)
+ if (reset) {
+ connection.accept(tcpEnvelope(1, StateActionTcpData(TcpDataAction(ActionType.TCP_DATA, 0, "Bwg="))))
+ connection.accept(tcpEnvelope(2, StateActionTcpHostReset(TcpHostResetAction(ActionType.TCP_HOST_RESET, TcpResetReason.PROTOCOL_ERROR))))
+ assertTrue(connection.read().isCompletedExceptionally)
+ } else {
+ val read = connection.read()
+ assertTrue(!read.isDone)
+ connection.dispose()
+ assertTrue(read.isCompletedExceptionally)
+ }
+ assertTrue(drain.isCompletedExceptionally)
+ connection.dispose()
+ connection.close()
+ assertEquals(listOf(connection.resource), harness.unsubscribed)
+ }
+ }
+
+ @Test
+ fun `TCP close while suspended replays and drains before release`() {
+ val old = TcpHarness()
+ val connection = old.open()
+ connection.suspend()
+ connection.close()
+ assertTrue(!connection.isClosed && old.sent.isEmpty() && old.unsubscribed.isEmpty())
+ val read = connection.read()
+ assertTrue(!read.isDone)
+ val request = TcpConnection.reconnectParameters(ReconnectParams("ahp-root://", clientId = "owner",
+ lastSeenServerSeq = 0, subscriptions = emptyList()), listOf(connection))
+ assertTrue(connection.resource in request.subscriptions)
+ val fresh = TcpHarness()
+ TcpConnection.resume(request, ReconnectResultReplay(ReconnectReplayResult(ReconnectResultType.REPLAY, listOf(
+ tcpEnvelope(1, StateActionTcpData(TcpDataAction(ActionType.TCP_DATA, 0, "Bwg="))),
+ tcpEnvelope(2, StateActionTcpHostClose(TcpHostCloseAction(ActionType.TCP_HOST_CLOSE))),
+ ), emptyList())), listOf(connection), fresh.transport)
+ val close = fresh.sent.first()
+ assertIs(close.second)
+ assertTrue(byteArrayOf(7, 8).contentEquals(read.join()))
+ val credit = fresh.sent.last()
+ assertIs(credit.second)
+ connection.accept(tcpEnvelope(3, close.second, ActionOrigin("owner", close.first)))
+ assertTrue(fresh.unsubscribed.isEmpty())
+ connection.accept(tcpEnvelope(4, credit.second, ActionOrigin("owner", credit.first)))
+ assertNull(connection.read().join())
+ assertEquals(listOf(connection.resource), fresh.unsubscribed)
+ assertTrue(old.unsubscribed.isEmpty())
+ }
+
+ private fun tcpEnvelope(sequence: Long, action: StateAction, origin: ActionOrigin? = null) =
+ ActionEnvelope(channel = "ahp-tcp:/created", action = action, serverSeq = sequence, origin = origin)
+
+ @Test
+ fun `TCP adapter encodes 4 MiB and final close preserves drain and buffered reads`() {
+ val harness = TcpHarness()
+ val bytes = ByteArray(4 * 1024 * 1024)
+ bytes[0] = 1
+ bytes[bytes.lastIndex] = -1
+ val connection = harness.open(bytes.size.toLong())
+ connection.write(bytes).join()
+ val sent = harness.sent.single()
+ val input = assertIs(sent.second).value
+ assertEquals(0L, input.offset)
+ assertTrue(bytes.contentEquals(java.util.Base64.getDecoder().decode(input.data)))
+ val drain = connection.drain()
+ assertTrue(!drain.isDone)
+ connection.accept(tcpEnvelope(1, sent.second, ActionOrigin("owner", sent.first)))
+ connection.accept(tcpEnvelope(2, StateActionTcpInputConsumed(TcpInputConsumedAction(ActionType.TCP_INPUT_CONSUMED, bytes.size.toLong()))))
+ connection.accept(tcpEnvelope(3, StateActionTcpData(TcpDataAction(ActionType.TCP_DATA, 0, "Bwg="))))
+ connection.accept(tcpEnvelope(4, StateActionTcpHostClose(TcpHostCloseAction(ActionType.TCP_HOST_CLOSE))))
+ val close = harness.sent.last()
+ assertIs(close.second)
+ connection.accept(tcpEnvelope(5, close.second, ActionOrigin("owner", close.first)))
+ assertTrue(harness.unsubscribed.isEmpty())
+ drain.join()
+ connection.drain().join()
+ assertTrue(byteArrayOf(7, 8).contentEquals(connection.read().join()))
+ val credit = harness.sent.last()
+ assertIs(credit.second)
+ assertTrue(harness.unsubscribed.isEmpty())
+ connection.accept(tcpEnvelope(6, credit.second, ActionOrigin("owner", credit.first)))
+ assertNull(connection.read().join())
+ assertEquals(listOf(connection.resource), harness.unsubscribed)
+ }
+
+ @Test
+ fun `TCP adapter rejects malformed client echoes before advancing or releasing payloads`() {
+ for (malformed in listOf("missing", "owner", "negative", "unsafe", "unassigned", "wrong-pending", "reused", "payload", "eof", "credit", "close", "reset", "rejected-empty")) {
+ val harness = TcpHarness()
+ val connection = harness.open()
+ val read = connection.read()
+ val write = connection.write(byteArrayOf(1, 2, 3, 4, 5))
+ val drain = connection.drain()
+ val first = harness.sent[0]
+ val second = harness.sent[1]
+ var origin: ActionOrigin? = ActionOrigin("owner", first.first)
+ var action = first.second
+ when (malformed) {
+ "missing" -> origin = null
+ "owner" -> origin = ActionOrigin("other", first.first)
+ "negative" -> origin = ActionOrigin("owner", -1)
+ "unsafe" -> origin = ActionOrigin("owner", 9007199254740992)
+ "unassigned" -> origin = ActionOrigin("owner", second.first + 1)
+ "wrong-pending" -> origin = ActionOrigin("owner", second.first)
+ "reused" -> { connection.accept(tcpEnvelope(1, first.second, origin)); action = second.second }
+ "payload" -> action = StateActionTcpInput(TcpInputAction(ActionType.TCP_INPUT, 0, "AgE="))
+ "eof" -> { origin = null; action = StateActionTcpInputEof(TcpInputEofAction(ActionType.TCP_INPUT_EOF, 0)) }
+ "credit" -> { origin = null; action = StateActionTcpDataConsumed(TcpDataConsumedAction(ActionType.TCP_DATA_CONSUMED, 0)) }
+ "close" -> { origin = null; action = StateActionTcpClientClose(TcpClientCloseAction(ActionType.TCP_CLIENT_CLOSE)) }
+ "reset" -> { origin = null; action = StateActionTcpClientReset(TcpClientResetAction(ActionType.TCP_CLIENT_RESET, TcpResetReason.PROTOCOL_ERROR)) }
+ }
+ connection.accept(tcpEnvelope(2, action, origin).copy(rejectionReason = if (malformed == "rejected-empty") "" else null))
+ for (future in listOf(read, write, drain)) assertTrue(future.isCompletedExceptionally, malformed)
+ assertEquals(if (malformed == "reused") 2L else 0L, connection.state.input.receivedBytes)
+ assertEquals(0, connection.state.input.consumedBytes)
+ assertIs(harness.sent.last().second)
+ assertEquals(listOf(connection.resource), harness.unsubscribed)
+ }
+ }
+
+ @Test
+ fun `TCP adapter reserves credit chunks reads duplicates and half closes`() {
+ val harness = TcpHarness()
+ val connection = harness.open()
+ val write = connection.write(byteArrayOf(1, 2, 3, 4, 5))
+ assertTrue(!write.isDone)
+ assertEquals(2, harness.sent.size)
+ assertFailsWith { connection.write(byteArrayOf(9)) }
+ for ((index, item) in harness.sent.toList().withIndex()) {
+ assertEquals(2, java.util.Base64.getDecoder().decode(assertIs(item.second).value.data).size)
+ connection.accept(tcpEnvelope(index + 1L, item.second, ActionOrigin("owner", item.first)))
+ }
+ connection.accept(tcpEnvelope(3, StateActionTcpInputConsumed(TcpInputConsumedAction(ActionType.TCP_INPUT_CONSUMED, 2))))
+ write.get(1, java.util.concurrent.TimeUnit.SECONDS)
+ val last = harness.sent.last()
+ assertEquals(4, assertIs(last.second).value.offset)
+ val drain = connection.drain()
+ assertTrue(!drain.isDone)
+ connection.accept(tcpEnvelope(4, last.second, ActionOrigin("owner", last.first)))
+ connection.accept(tcpEnvelope(5, StateActionTcpInputConsumed(TcpInputConsumedAction(ActionType.TCP_INPUT_CONSUMED, 5))))
+ drain.get(1, java.util.concurrent.TimeUnit.SECONDS)
+ val data = StateActionTcpData(TcpDataAction(ActionType.TCP_DATA, 0, "Bwg="))
+ connection.accept(tcpEnvelope(6, data))
+ connection.accept(tcpEnvelope(7, data))
+ connection.accept(tcpEnvelope(8, StateActionTcpDataEof(TcpDataEofAction(ActionType.TCP_DATA_EOF, 2))))
+ assertTrue(byteArrayOf(7, 8).contentEquals(connection.read().get()))
+ assertEquals(2, assertIs(harness.sent.last().second).value.consumedBytes)
+ assertNull(connection.read().get())
+ connection.end().get()
+ assertEquals(5, assertIs(harness.sent.last().second).value.finalOffset)
+ connection.close()
+ connection.dispose()
+ assertEquals(listOf(connection.resource), harness.unsubscribed)
+ }
+
+ @Test
+ fun `TCP adapter resumes retained readers and only resends unacknowledged identities`() {
+ val first = TcpHarness()
+ val connection = first.open()
+ connection.write(byteArrayOf(1, 2)).get()
+ connection.end().get()
+ connection.accept(tcpEnvelope(1, StateActionTcpData(TcpDataAction(ActionType.TCP_DATA, 0, "Bwg=")), ActionOrigin("owner", first.sent[0].first)))
+ connection.suspend()
+ val reader = connection.read()
+ assertTrue(!reader.isDone)
+ val request = TcpConnection.reconnectParameters(ReconnectParams("ahp-root://", clientId = "owner",
+ lastSeenServerSeq = 20, subscriptions = listOf("ahp-session:/test")), listOf(connection))
+ assertEquals(1, request.lastSeenServerSeq)
+ assertTrue(connection.resource in request.subscriptions)
+ val fresh = TcpHarness()
+ TcpConnection.resume(request, ReconnectResultReplay(ReconnectReplayResult(ReconnectResultType.REPLAY, emptyList(), emptyList())),
+ listOf(connection), fresh.transport)
+ assertEquals(first.sent.take(2), fresh.sent.take(2))
+ assertEquals(3, fresh.sent.last().first)
+ assertTrue(byteArrayOf(7, 8).contentEquals(reader.get()))
+ connection.suspend()
+ val again = TcpHarness()
+ val replay = tcpEnvelope(3, first.sent[0].second, ActionOrigin("owner", first.sent[0].first))
+ TcpConnection.resume(request, ReconnectResultReplay(ReconnectReplayResult(ReconnectResultType.REPLAY, listOf(replay), emptyList())),
+ listOf(connection), again.transport)
+ assertEquals(listOf(2L, 3L), again.sent.map { it.first })
+ connection.accept(tcpEnvelope(4, first.sent[0].second, ActionOrigin("owner", first.sent[0].first)))
+ assertTrue(!connection.isClosed)
+ connection.dispose()
+ }
+
+ @Test
+ fun `TCP adapter continues a blocked writer after replay releases credit`() {
+ val first = TcpHarness()
+ val connection = first.open()
+ val write = connection.write(ByteArray(6))
+ assertTrue(!write.isDone)
+ first.sent.forEachIndexed { index, item -> connection.accept(tcpEnvelope(index + 1L, item.second, ActionOrigin("owner", item.first))) }
+ val unrelatedSequence = first.transport.nextSequence()
+ connection.suspend()
+ val request = TcpConnection.reconnectParameters(ReconnectParams("ahp-root://", clientId = "owner",
+ lastSeenServerSeq = 20, subscriptions = emptyList()), listOf(connection))
+ assertEquals(2, request.lastSeenServerSeq)
+ val replay = listOf(tcpEnvelope(3, StateActionTcpInputConsumed(TcpInputConsumedAction(ActionType.TCP_INPUT_CONSUMED, 2))))
+ val fresh = TcpHarness()
+ TcpConnection.resume(request, ReconnectResultReplay(ReconnectReplayResult(ReconnectResultType.REPLAY, replay, emptyList())),
+ listOf(connection), fresh.transport)
+ write.get(1, java.util.concurrent.TimeUnit.SECONDS)
+ assertEquals(1, fresh.sent.size)
+ assertEquals(4, assertIs(fresh.sent.single().second).value.offset)
+ assertTrue(fresh.sent.single().first > unrelatedSequence)
+ connection.dispose()
+ }
+
+ @Test
+ fun `TCP adapter snapshot missing and strict loss terminate pending operations`() {
+ val results = listOf(
+ ReconnectResultSnapshot(ReconnectSnapshotResult(ReconnectResultType.SNAPSHOT, emptyList())),
+ ReconnectResultReplay(ReconnectReplayResult(ReconnectResultType.REPLAY, emptyList(), listOf("ahp-tcp:/created"))),
+ )
+ for (result in results) {
+ val harness = TcpHarness()
+ val connection = harness.open()
+ val read = connection.read()
+ val write = connection.write(ByteArray(5))
+ val drain = connection.drain()
+ connection.suspend()
+ val request = TcpConnection.reconnectParameters(ReconnectParams("ahp-root://", clientId = "owner", lastSeenServerSeq = 0,
+ subscriptions = emptyList()), listOf(connection))
+ TcpConnection.resume(request, result, listOf(connection), harness.transport)
+ for (future in listOf(read, write, drain)) assertTrue(future.isCompletedExceptionally)
+ assertEquals(listOf(connection.resource), harness.unsubscribed)
+ }
+ val harness = TcpHarness()
+ val connection = harness.open()
+ val read = connection.read()
+ val write = connection.write(ByteArray(5))
+ connection.fail(IllegalStateException("strict decode loss"))
+ assertTrue(read.isCompletedExceptionally && write.isCompletedExceptionally)
+ assertIs(harness.sent.last().second)
+ connection.dispose()
+ assertEquals(1, harness.unsubscribed.size)
+ val closing = TcpHarness().open()
+ val pendingRead = closing.read()
+ val pendingWrite = closing.write(ByteArray(5))
+ val pendingDrain = closing.drain()
+ closing.close()
+ assertTrue(!pendingRead.isDone && !pendingDrain.isDone)
+ assertTrue(pendingWrite.isCompletedExceptionally)
+ closing.dispose()
+ assertTrue(pendingRead.isCompletedExceptionally && pendingDrain.isCompletedExceptionally)
+ }
+
+ private fun tcpState(size: Long = 8): TcpConnectionState = TcpConnectionState(
+ session = "ahp-session:/test",
+ target = TcpTarget(host = "localhost", port = 3000),
+ encoding = TcpDataEncoding.BASE64,
+ input = FlowControlledByteDirectionState(size, size, 0, 0),
+ output = FlowControlledByteDirectionState(size, size, 0, 0),
+ clientClosed = false,
+ hostClosed = false,
+ )
+
+ @Test
+ fun `TCP validates a four MiB chunk without decoding or retaining payload`() {
+ val size = 4 * 1024 * 1024
+ val data = "AAAA".repeat(size / 3) + "AA=="
+ val before = tcpState(size.toLong())
+ val action = StateActionTcpInput(TcpInputAction(ActionType.TCP_INPUT, 0, data))
+ val after = TcpReducer.reduce(before, action)
+ assertEquals(size.toLong(), after.input.receivedBytes)
+ assertEquals(0L, before.input.receivedBytes)
+ assertSame(after, tcpReducer(after, action))
+ assertSame(before.output, after.output)
+ val error = assertFailsWith {
+ tcpReducer(after, StateActionTcpInput(TcpInputAction(ActionType.TCP_INPUT, size.toLong(), "AA==")))
+ }
+ assertEquals("Invalid TCP action: receive window exceeded", error.message)
+ assertEquals(size.toLong(), after.input.receivedBytes)
+ }
+
+ @Test
+ fun `TCP checks all integer counters and accepts the safe boundary`() {
+ val max = 9007199254740991L
+ val before = tcpState().let { it.copy(input = it.input.copy(receivedBytes = max - 1, consumedBytes = max - 1)) }
+ var state = tcpReducer(before, StateActionTcpInput(TcpInputAction(ActionType.TCP_INPUT, max - 1, "AA==")))
+ state = tcpReducer(state, StateActionTcpInputConsumed(TcpInputConsumedAction(ActionType.TCP_INPUT_CONSUMED, max)))
+ state = tcpReducer(state, StateActionTcpInputEof(TcpInputEofAction(ActionType.TCP_INPUT_EOF, max)))
+ assertEquals(max, state.input.receivedBytes)
+ assertEquals(max, state.input.consumedBytes)
+ assertEquals(max, state.input.eofAtBytes)
+ for (value in listOf(-1L, max + 1, Long.MAX_VALUE)) {
+ val actions = listOf(
+ StateActionTcpInput(TcpInputAction(ActionType.TCP_INPUT, value, "AA==")),
+ StateActionTcpData(TcpDataAction(ActionType.TCP_DATA, value, "AA==")),
+ StateActionTcpInputConsumed(TcpInputConsumedAction(ActionType.TCP_INPUT_CONSUMED, value)),
+ StateActionTcpDataConsumed(TcpDataConsumedAction(ActionType.TCP_DATA_CONSUMED, value)),
+ StateActionTcpInputEof(TcpInputEofAction(ActionType.TCP_INPUT_EOF, value)),
+ StateActionTcpDataEof(TcpDataEofAction(ActionType.TCP_DATA_EOF, value)),
+ )
+ for (action in actions) {
+ val error = assertFailsWith