diff --git a/docs/ARCHITECTURE.md b/docs/ARCHITECTURE.md index 962fc9c58..393a727cf 100644 --- a/docs/ARCHITECTURE.md +++ b/docs/ARCHITECTURE.md @@ -185,6 +185,7 @@ leading and trailing pipe. Columns, in order: | cmd-payload-tokenization | authoritative | src/OpenClaw.Shared/ExecApprovals/ExecReusableCommandBinder.cs | parsing a cmd payload into tokens and rewriting its executable token | CmdPayloadTokenizer | ExecReusableCommandBinder.TryTokenizeStaticCmdPayload remains as a delegating wrapper for existing callers and tests | a payload rewrite is built from parsed token spans and is accepted only after re-parsing proves the argument list is unchanged except for the pinned executable | ExecReusableCommandBinderTests.PinnedCarrier_DoesNotRewriteArgumentsThatRepeatTheExecutableText | behavioral | - | | exec-carrier-cwd-ambiguity-check | closed | src/OpenClaw.Shared/ExecApprovals/ExecReusableCommandBinder.cs | deciding whether a carrier payload may be durably approved when the working directory could shadow it | CanonicalCmdCarrier.TryBuildPinnedCarrier (payload executable pinning) | - | the approval-time working-directory check is deleted, not merely bypassed: ExecCommandResolver exposes no HasCurrentDirectoryCandidate, a trusted carrier's payload executable is pinned to its resolved absolute path so cmd has nothing to search for, and a post-approval shadow cannot win | ExecReusableCommandBinderTests.PinnedCarrier_IgnoresShadowInsertedAfterApproval | behavioral | - | | exec-legacy-host-quarantine | authoritative | src/OpenClaw.Shared/ExecApprovals/ExecCommandToken.cs | deciding whether a provenance-less path-only allowlist entry authorizes an interpreter or code host | ExecAllowlistMatcher.MatchInternal via ExecCommandToken.IsLegacyQuarantinedHost | argument binding remains the security boundary for every rule this node generates | an allowlist entry with no source and no argPattern is inert when its resolved target is a command host the previous model refused, is never deleted or migrated, and is superseded only by an explicit allow-always sibling carrying source and argPattern | ExecAllowlistArgBindingTests.LegacyPathOnlyEntryForACommandHost_IsInert | behavioral | - | +| app-ssh-restart-closed | closed | src/OpenClaw.Tray.WinUI/App.xaml.cs and ConnectionPage.xaml.cs | stopping, starting, reconnecting, and declaring success for a user-requested SSH tunnel restart | GatewayConnectionManager.RestartSshTunnelAsync | App and ConnectionPage invoke the manager and present the result | a restart succeeds only after a fresh generation-bound hello-ok and current registry, config, tunnel generation, and owned listener verification | AppRefactorContractTests.UserSshRestart_StaysDelegatedToConnectionManager | source-shape | when App no longer owns any SSH tunnel UI actions | | hub-page-registry | authoritative | src/OpenClaw.Tray.WinUI/Windows/HubWindow.xaml.cs and GatewayNavVisibilityDebouncePolicy | navigation aliases, page mapping, command metadata and search, and gateway-page classification | HubPageRegistry | HubWindow keeps Frame and NavigationView application, back-stack mutation, command cache lifetime, and semantic action execution; GatewayNavVisibilityDebouncePolicy keeps disconnect timing | every current direct, legacy, and agent-scoped tag resolves identically; command order, titles, actions, search caps, and gateway prune set remain stable | HubPageRegistryTests.BuildCommands_PreservesBaseOrderActionsIconsAndResourceKeys | behavioral | - | | hub-page-registry-closed | closed | src/OpenClaw.Tray.WinUI/Windows/HubWindow.xaml.cs and GatewayNavVisibilityDebouncePolicy | private tag/page switches, command catalogs or search predicates, and gateway-page tag lists | HubPageRegistry | view-only navigation application and debounce timing listed in hub-page-registry | HubWindow and the debounce policy do not regain catalog or page-classification copies | HubPresentationContractTests.HubPageRegistry_OwnsMappingsCommandsAndGatewayClassification | source-shape | when HubWindow is replaced by a different shell and GatewayNavVisibilityDebouncePolicy is retired | | app-notification-infobar-presentation | authoritative | src/OpenClaw.Tray.WinUI/Windows/HubWindow.xaml.cs | banner severity filtering, selected-banner fallback, notification action versus Show more, and action enabled state | AppNotificationInfoBarPresenter | HubWindow keeps notification subscription, bell reconciliation, WinUI control assignment, navigation, and dismissal side effects | Warning and Error banners retain priority and hiding semantics while action projection stays WinUI-free | AppNotificationInfoBarPresenterTests.Present_ActionableNotificationWinsOverShowMore | behavioral | - | diff --git a/src/OpenClaw.Connection/GatewayConnectionManager.cs b/src/OpenClaw.Connection/GatewayConnectionManager.cs index 9b5aa4cda..4a1caeba1 100644 --- a/src/OpenClaw.Connection/GatewayConnectionManager.cs +++ b/src/OpenClaw.Connection/GatewayConnectionManager.cs @@ -62,6 +62,9 @@ public sealed class GatewayConnectionManager : IGatewayConnectionManager private readonly Func>? _endpointProvenanceProbe; private readonly Func _validationTunnelFactory; + private readonly TimeSpan _credentialHandoffTimeout; + private readonly TimeSpan _manualSshRestartTimeout; + private readonly TimeSpan _manualSshRestartCleanupTimeout; private readonly SemaphoreSlim _transitionSemaphore = new(1, 1); private readonly SemaphoreSlim _nodeStartSemaphore = new(1, 1); private readonly object _nodeOperationLock = new(); @@ -107,6 +110,8 @@ private readonly Func? reconnectDelay = null, Func>? endpointProvenanceProbe = null, - Func? validationTunnelFactory = null) + Func? validationTunnelFactory = null, + TimeSpan? credentialHandoffTimeout = null, + TimeSpan? manualSshRestartTimeout = null, + TimeSpan? manualSshRestartCleanupTimeout = null) { _credentialResolver = credentialResolver ?? throw new ArgumentNullException(nameof(credentialResolver)); _clientFactory = clientFactory ?? throw new ArgumentNullException(nameof(clientFactory)); @@ -153,6 +161,16 @@ public GatewayConnectionManager( _reconnectDelay = reconnectDelay ?? Task.Delay; _endpointProvenanceProbe = endpointProvenanceProbe; _validationTunnelFactory = validationTunnelFactory ?? (() => new SshTunnelService(_logger)); + _credentialHandoffTimeout = credentialHandoffTimeout ?? TimeSpan.FromSeconds(5); + _manualSshRestartTimeout = manualSshRestartTimeout ?? TimeSpan.FromSeconds(35); + _manualSshRestartCleanupTimeout = + manualSshRestartCleanupTimeout ?? TimeSpan.FromSeconds(5); + if (_credentialHandoffTimeout <= TimeSpan.Zero) + throw new ArgumentOutOfRangeException(nameof(credentialHandoffTimeout)); + if (_manualSshRestartTimeout <= TimeSpan.Zero) + throw new ArgumentOutOfRangeException(nameof(manualSshRestartTimeout)); + if (_manualSshRestartCleanupTimeout <= TimeSpan.Zero) + throw new ArgumentOutOfRangeException(nameof(manualSshRestartCleanupTimeout)); _diagnostics = diagnostics ?? new ConnectionDiagnostics(clock: clock); _diagnostics.EventRecorded += (_, e) => DiagnosticEvent?.Invoke(this, e); @@ -257,7 +275,10 @@ public async Task ConnectNodeOnlyAsync(string? gatewayId = null) } /// Core connect logic. Caller must hold . - private async Task ConnectCoreAsync(string? gatewayId = null, string operation = "connect") + private async Task ConnectCoreAsync( + string? gatewayId = null, + string operation = "connect", + CancellationToken externalCancellationToken = default) { var id = gatewayId ?? _registry.ActiveGatewayId; if (id == null) @@ -281,7 +302,10 @@ private async Task ConnectCoreAsync(string? gatewayId = null, string operation = // Cancel any in-flight operation var gen = Interlocked.Increment(ref _generation); - var oldCts = Interlocked.Exchange(ref _operationCts, new CancellationTokenSource()); + var newOperationCts = externalCancellationToken.CanBeCanceled + ? CancellationTokenSource.CreateLinkedTokenSource(externalCancellationToken) + : new CancellationTokenSource(); + var oldCts = Interlocked.Exchange(ref _operationCts, newOperationCts); oldCts?.Cancel(); oldCts?.Dispose(); @@ -411,6 +435,7 @@ private async Task ConnectCoreAsync(string? gatewayId = null, string operation = // logs appear in the Connection Status window timeline. // When SSH tunnel is configured, start the tunnel and connect to the local URL. var connectUrl = record.Url; + SshTunnelStartResult? startedTunnel = null; if (record.SshTunnel != null && _tunnelManager != null) { var tunnel = record.SshTunnel; @@ -430,14 +455,33 @@ tunnel.SshPort is < 1 or > 65535 || } try { - connectUrl = await _tunnelManager.StartAsync(tunnel, _operationCts!.Token); + startedTunnel = await _tunnelManager + .StartOwnedAsync(tunnel, _operationCts!.Token) + .ConfigureAwait(false); + connectUrl = startedTunnel.Url; var tunnelAuthorization = await AuthorizeCredentialForEndpointAsync( record, credential, _operationCts.Token, requireSshTunnelOwnership: true).ConfigureAwait(false); if (!tunnelAuthorization.Allowed) - throw new InvalidOperationException(tunnelAuthorization.Detail); + { + _stateMachine.SetOperatorErrorKind(tunnelAuthorization.FailureKind); + _stateMachine.TryTransition( + tunnelAuthorization.FailureKind == GatewayErrorKind.Network + ? ConnectionTrigger.WebSocketError + : ConnectionTrigger.AuthenticationFailed, + tunnelAuthorization.Detail); + await StopOwnedTunnelAfterFailedConnectionAsync( + startedTunnel, + "SSH tunnel ownership authorization failure"); + CompleteOperatorTelemetryAttempt( + gen, + "failure", + ConnectionErrorCategory.SshTunnelFailure); + EmitStateChanged(); + return; + } expectedEndpointOwnership = tunnelAuthorization.OwnershipProof; _diagnostics.Record("tunnel", $"SSH tunnel started → {connectUrl}"); } @@ -446,7 +490,12 @@ tunnel.SshPort is < 1 or > 65535 || _logger.Error($"[ConnMgr] SSH tunnel start failed: {ex.Message}"); _diagnostics.Record("tunnel", "SSH tunnel start failed", ex.Message); _stateMachine.TryTransition(ConnectionTrigger.WebSocketError, $"SSH tunnel failed: {ex.Message}"); - await StopTunnelAfterFailedConnectionAsync("SSH tunnel startup or authorization failure"); + if (startedTunnel is not null) + { + await StopOwnedTunnelAfterFailedConnectionAsync( + startedTunnel, + "SSH tunnel startup or authorization failure"); + } CompleteOperatorTelemetryAttempt( gen, "failure", @@ -486,7 +535,12 @@ tunnel.SshPort is < 1 or > 65535 || _stateMachine.TryTransition( ConnectionTrigger.WebSocketError, DeviceIdentityLoadException.RecoveryMessage); - await StopTunnelAfterFailedConnectionAsync("operator identity load failure"); + if (startedTunnel is not null) + { + await StopOwnedTunnelAfterFailedConnectionAsync( + startedTunnel, + "operator identity load failure"); + } CompleteOperatorTelemetryAttempt( gen, "failure", @@ -495,35 +549,32 @@ tunnel.SshPort is < 1 or > 65535 || return; } + var operatorOperationCancellation = _operationCts!.Token; async Task AuthorizeLiveCredentialHandoffAsync( CancellationToken cancellationToken) { - if (!IsCurrentGatewayAttempt(gen, record.Id) || - !IsAutomaticReconnectAllowed(record.Id)) - { - return new ReconnectAuthorizationResult( - false, - GatewayErrorKind.Unknown, - "Connection attempt was superseded or explicitly disconnected."); - } - var authorization = await AuthorizeCredentialForEndpointAsync( - record, - credential, - cancellationToken, - requireSshTunnelOwnership: true).ConfigureAwait(false); - if (authorization.Allowed && - expectedEndpointOwnership is not null && - authorization.OwnershipProof != expectedEndpointOwnership) + var authorization = await AuthorizeCredentialHandoffAsync( + record, + credential, + expectedEndpointOwnership, + () => IsCurrentGatewayAttempt(gen, record.Id) && + IsAutomaticReconnectAllowed(record.Id), + operatorOperationCancellation, + cancellationToken, + "operator") + .ConfigureAwait(false); + if (!authorization.Allowed && + authorization.FailureKind != GatewayErrorKind.Unknown) { - return new ReconnectAuthorizationResult( - false, - GatewayErrorKind.LocalPortConflict, - "Endpoint ownership changed after preflight, so credentials were not sent."); + await RecordOperatorCredentialHandoffFailureAsync( + authorization.Detail ?? "Operator credential handoff was not authorized.", + authorization.FailureKind, + operatorOperationCancellation, + gen, + record.Id) + .ConfigureAwait(false); } - return new ReconnectAuthorizationResult( - authorization.Allowed, - authorization.FailureKind, - authorization.Detail); + return authorization; } lifecycle.DataClient.HandshakeAuthorizationAsync = @@ -727,13 +778,18 @@ expectedEndpointOwnership is not null && _diagnostics.Record("node", $"Starting node-only connection to {record.Url}", $"Credential source: {nodeCredential.Source}"); - if (!preservesOperatorConnection && !await TryStartTunnelForNodeOnlyAsync(record)) + SshTunnelStartResult? startedTunnel = null; + if (!preservesOperatorConnection && record.SshTunnel is not null) { - _stateMachine.SetNodeCredentialResolution(nodeCredentialResolution); - _stateMachine.BlockNodeStart(NodeTunnelStartFailedMessage, preserveCredentialResolution: true); - EmitStateChanged(); - RecordNodePreflightTelemetryFailure(ConnectionErrorCategory.SshTunnelFailure); - return null; + startedTunnel = await TryStartTunnelForNodeOnlyAsync(record); + if (startedTunnel is null) + { + _stateMachine.SetNodeCredentialResolution(nodeCredentialResolution); + _stateMachine.BlockNodeStart(NodeTunnelStartFailedMessage, preserveCredentialResolution: true); + EmitStateChanged(); + RecordNodePreflightTelemetryFailure(ConnectionErrorCategory.SshTunnelFailure); + return null; + } } if (record.SshTunnel is not null) @@ -746,7 +802,11 @@ expectedEndpointOwnership is not null && if (!tunnelAuthorization.Allowed) { if (!preservesOperatorConnection) - await StopTunnelAfterFailedConnectionAsync("node-only ownership proof failure"); + { + await StopOwnedTunnelAfterFailedConnectionAsync( + startedTunnel!, + "node-only ownership proof failure"); + } _stateMachine.SetNodeCredentialResolution(nodeCredentialResolution); _stateMachine.BlockNodeStart( @@ -761,15 +821,16 @@ expectedEndpointOwnership is not null && return Interlocked.Read(ref _generation) == gen ? gen : null; } - private async Task TryStartTunnelForNodeOnlyAsync(GatewayRecord record) + private async Task TryStartTunnelForNodeOnlyAsync( + GatewayRecord record) { if (record.SshTunnel == null) - return true; + return null; if (_tunnelManager == null) { _diagnostics.Record("tunnel", "No tunnel manager available for node-only SSH connection"); - return false; + return null; } var tunnel = record.SshTunnel; @@ -780,20 +841,52 @@ tunnel.RemotePort is < 1 or > 65535 || { _logger.Warn("[ConnMgr] SSH tunnel config is incomplete for node-only connect"); _diagnostics.Record("tunnel", "SSH tunnel config is incomplete for node-only connect"); - return false; + return null; } try { - var connectUrl = await _tunnelManager.StartAsync(tunnel, _operationCts!.Token); - _diagnostics.Record("tunnel", $"SSH tunnel started for node-only connect → {connectUrl}"); - return true; + var startedTunnel = await _tunnelManager + .StartOwnedAsync(tunnel, _operationCts!.Token) + .ConfigureAwait(false); + _diagnostics.Record( + "tunnel", + $"SSH tunnel started for node-only connect → {startedTunnel.Url}"); + return startedTunnel; } catch (Exception ex) when (ex is not OperationCanceledException) { _logger.Error($"[ConnMgr] SSH tunnel start failed for node-only connect: {ex.Message}"); _diagnostics.Record("tunnel", "SSH tunnel start failed for node-only connect", ex.Message); - return false; + return null; + } + } + + private async Task StopOwnedTunnelAfterFailedConnectionAsync( + SshTunnelStartResult startedTunnel, + string operation) + { + if (_tunnelManager is null) + return; + + using var stopCts = new CancellationTokenSource(TimeSpan.FromSeconds(5)); + try + { + if (await _tunnelManager.StopIfOwnedAsync( + startedTunnel.Config, + startedTunnel.OwnershipGeneration, + stopCts.Token).ConfigureAwait(false)) + { + _diagnostics.Record("tunnel", $"SSH tunnel stopped after {operation}"); + } + } + catch (OperationCanceledException) + { + _logger.Warn($"[ConnMgr] Tunnel stop timed out after {operation}"); + } + catch (Exception ex) + { + _logger.Warn($"[ConnMgr] Tunnel stop failed after {operation}: {ex.Message}"); } } @@ -980,6 +1073,271 @@ public async Task RecoverSshTunnelAsync(SshTunnelExit tunnelExit) } } + public async Task RestartSshTunnelAsync(CancellationToken cancellationToken = default) + { + ThrowIfDisposed(); + var restartGeneration = Interlocked.Increment(ref _manualSshRestartGeneration); + using var restartCts = CancellationTokenSource.CreateLinkedTokenSource(cancellationToken); + restartCts.CancelAfter(_manualSshRestartTimeout); + var previousRestartCts = Interlocked.Exchange(ref _manualSshRestartCts, restartCts); + try { previousRestartCts?.Cancel(); } + catch (ObjectDisposedException) { } + + EventHandler? stateHandler = null; + TaskCompletionSource? handshakeCompletion = null; + long connectionGeneration = 0; + long tunnelGeneration = 0; + IGatewayClientLifecycle? lifecycle = null; + GatewayRecord? expectedGatewayRecord = null; + string? gatewayId = null; + SshTunnelConfig? tunnelConfig = null; + SshTunnelConfig? ownedTunnelConfig = null; + var succeeded = false; + + try + { + await _transitionSemaphore.WaitAsync(restartCts.Token).ConfigureAwait(false); + try + { + if (restartGeneration != Interlocked.Read(ref _manualSshRestartGeneration)) + return false; + + var activeRecord = _registry.GetActive(); + expectedGatewayRecord = activeRecord; + gatewayId = activeRecord?.Id; + tunnelConfig = activeRecord?.SshTunnel; + if (activeRecord is null || + tunnelConfig is null || + _tunnelManager is null) + { + return false; + } + ownedTunnelConfig = tunnelConfig with + { + User = tunnelConfig.User.Trim(), + Host = tunnelConfig.Host.Trim(), + }; + + var previousTunnelGeneration = _tunnelManager.OwnershipGeneration; + var previousTunnelWasActive = _tunnelManager.IsActive; + var previousTunnelConfig = _tunnelManager.ActiveConfig; + if (previousTunnelWasActive && previousTunnelConfig is null) + return false; + + if (previousTunnelWasActive) + { + var stopped = await _tunnelManager.StopIfOwnedAsync( + previousTunnelConfig!, + previousTunnelGeneration, + restartCts.Token) + .ConfigureAwait(false); + if (!stopped) + return false; + } + else if (_tunnelManager.IsActive) + { + return false; + } + + var stoppedTunnelGeneration = _tunnelManager.OwnershipGeneration; + SetGatewayConnectionIntent(activeRecord.Id, shouldBeConnected: true); + await DisconnectCoreAsync().ConfigureAwait(false); + if (restartGeneration != Interlocked.Read(ref _manualSshRestartGeneration) || + !IsCurrentSshRegistryRecord(activeRecord) || + _tunnelManager.OwnershipGeneration != stoppedTunnelGeneration) + { + return false; + } + + try + { + await ConnectCoreAsync( + activeRecord.Id, + operation: "reconnect", + restartCts.Token) + .ConfigureAwait(false); + } + finally + { + connectionGeneration = Interlocked.Read(ref _generation); + lifecycle = _activeLifecycle; + tunnelGeneration = _tunnelManager.OwnershipGeneration; + } + if (lifecycle is null || + !IsCurrentGatewayAttempt(connectionGeneration, activeRecord.Id) || + !_tunnelManager.IsActive || + _tunnelManager.ActiveConfig != ownedTunnelConfig) + { + return false; + } + + handshakeCompletion = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + stateHandler = (_, snapshot) => + { + if (restartGeneration != Interlocked.Read(ref _manualSshRestartGeneration) || + connectionGeneration != Interlocked.Read(ref _generation) || + !ReferenceEquals(lifecycle, _activeLifecycle)) + { + handshakeCompletion.TrySetResult(false); + return; + } + + if (snapshot.OperatorState == RoleConnectionState.Connected) + handshakeCompletion.TrySetResult(true); + else if (snapshot.OperatorState is RoleConnectionState.Error + or RoleConnectionState.PairingRequired) + handshakeCompletion.TrySetResult(false); + }; + StateChanged += stateHandler; + if (_stateMachine.Current.OperatorState == RoleConnectionState.Connected) + handshakeCompletion.TrySetResult(true); + else if (_stateMachine.Current.OperatorState is RoleConnectionState.Error + or RoleConnectionState.PairingRequired) + handshakeCompletion.TrySetResult(false); + } + finally + { + _transitionSemaphore.Release(); + } + + if (handshakeCompletion is null || + !await handshakeCompletion.Task.WaitAsync(restartCts.Token).ConfigureAwait(false)) + { + return false; + } + + await _transitionSemaphore.WaitAsync(restartCts.Token).ConfigureAwait(false); + try + { + var activeRecordId = _registry.ActiveGatewayId; + if (restartGeneration != Interlocked.Read(ref _manualSshRestartGeneration) || + connectionGeneration != Interlocked.Read(ref _generation) || + !ReferenceEquals(lifecycle, _activeLifecycle) || + _stateMachine.Current.OperatorState != RoleConnectionState.Connected || + !lifecycle.DataClient.IsConnectedToGateway || + gatewayId is null || + activeRecordId is null || + !string.Equals(activeRecordId, gatewayId, StringComparison.Ordinal) || + expectedGatewayRecord is null || + tunnelConfig is null || + ownedTunnelConfig is null || + !IsCurrentSshRegistryRecord(expectedGatewayRecord) || + _tunnelManager is null || + !_tunnelManager.IsActive || + _tunnelManager.OwnershipGeneration != tunnelGeneration || + _tunnelManager.ActiveConfig != ownedTunnelConfig) + { + return false; + } + + succeeded = await _tunnelManager.IsOwnedListenerReadyAsync( + tunnelConfig, + tunnelConfig.LocalPort, + restartCts.Token) + .ConfigureAwait(false); + succeeded = succeeded && + restartGeneration == Interlocked.Read(ref _manualSshRestartGeneration) && + connectionGeneration == Interlocked.Read(ref _generation) && + ReferenceEquals(lifecycle, _activeLifecycle) && + _stateMachine.Current.OperatorState == RoleConnectionState.Connected && + lifecycle.DataClient.IsConnectedToGateway && + _tunnelManager.OwnershipGeneration == tunnelGeneration && + _tunnelManager.ActiveConfig == ownedTunnelConfig && + IsCurrentSshRegistryRecord(expectedGatewayRecord); + return succeeded; + } + finally + { + _transitionSemaphore.Release(); + } + } + catch (OperationCanceledException) + { + return false; + } + finally + { + if (stateHandler is not null) + StateChanged -= stateHandler; + + if (!succeeded) + { + await CleanupManualSshRestartAsync( + restartGeneration, + connectionGeneration, + lifecycle, + tunnelGeneration, + tunnelConfig) + .ConfigureAwait(false); + } + + Interlocked.CompareExchange(ref _manualSshRestartCts, null, restartCts); + } + } + + private bool IsCurrentSshRegistryRecord(GatewayRecord expectedRecord) + { + var currentRecord = _registry.GetById(expectedRecord.Id); + return currentRecord is not null && + string.Equals(_registry.ActiveGatewayId, expectedRecord.Id, StringComparison.Ordinal) && + IsSameCredentialHandoffRecord(currentRecord, expectedRecord); + } + + private async Task CleanupManualSshRestartAsync( + long restartGeneration, + long connectionGeneration, + IGatewayClientLifecycle? lifecycle, + long tunnelGeneration, + SshTunnelConfig? tunnelConfig) + { + if (restartGeneration != Interlocked.Read(ref _manualSshRestartGeneration) || + connectionGeneration == 0) + return; + + using var cleanupCts = new CancellationTokenSource(_manualSshRestartCleanupTimeout); + try + { + await _transitionSemaphore.WaitAsync(cleanupCts.Token).ConfigureAwait(false); + } + catch (OperationCanceledException) + { + _logger.Warn("[ConnMgr] Timed out waiting to clean up a failed manual SSH restart."); + return; + } + catch (ObjectDisposedException) when (_disposed) + { + return; + } + try + { + if (restartGeneration != Interlocked.Read(ref _manualSshRestartGeneration) || + connectionGeneration != Interlocked.Read(ref _generation) || + !ReferenceEquals(lifecycle, _activeLifecycle)) + { + return; + } + + await DisconnectCoreAsync().ConfigureAwait(false); + if (tunnelConfig is not null && _tunnelManager is not null) + { + await _tunnelManager.StopIfOwnedAsync( + tunnelConfig, + tunnelGeneration, + cleanupCts.Token) + .ConfigureAwait(false); + } + } + catch (OperationCanceledException) when (cleanupCts.IsCancellationRequested) + { + _logger.Warn("[ConnMgr] Timed out cleaning up a failed manual SSH restart."); + } + finally + { + _transitionSemaphore.Release(); + } + } + public void SetGatewayConnectionIntent(string gatewayId, bool shouldBeConnected) { if (string.IsNullOrWhiteSpace(gatewayId)) @@ -1670,8 +2028,10 @@ private async Task HandleOperatorStatusChangedAsync(ConnectionStatus status, lon _stateMachine.Current.OperatorErrorKind is null) { _stateMachine.SetOperatorErrorKind(ReadOperatorFailureKind(gen)); + _stateMachine.TryTransition( + ConnectionTrigger.WebSocketError, + "Transport error"); } - _stateMachine.TryTransition(ConnectionTrigger.WebSocketError, "Transport error"); } CompleteOperatorTelemetryAttempt( gen, @@ -1942,6 +2302,120 @@ private async Task AuthorizeCredentialForEndpoi "The managed gateway address is owned by an unverified process. OpenClaw did not send the shared or bootstrap token."); } + private async Task AuthorizeCredentialHandoffAsync( + GatewayRecord expectedRecord, + GatewayCredential credential, + EndpointOwnershipProof? expectedOwnership, + Func isCurrentAttempt, + CancellationToken operationCancellationToken, + CancellationToken handshakeCancellationToken, + string role) + { + using var timeoutCts = new CancellationTokenSource(_credentialHandoffTimeout); + using var linkedCts = CancellationTokenSource.CreateLinkedTokenSource( + operationCancellationToken, + handshakeCancellationToken, + timeoutCts.Token); + + try + { + if (!isCurrentAttempt()) + { + return new ReconnectAuthorizationResult( + false, + GatewayErrorKind.Unknown, + $"{role} connection attempt was superseded."); + } + + var currentRecord = _registry.GetById(expectedRecord.Id); + if (currentRecord is null || + !string.Equals(_registry.ActiveGatewayId, expectedRecord.Id, StringComparison.Ordinal) || + !IsSameCredentialHandoffRecord(currentRecord, expectedRecord)) + { + return new ReconnectAuthorizationResult( + false, + GatewayErrorKind.LocalPortConflict, + $"The active gateway endpoint, credentials, or SSH configuration changed before the {role} credential handoff."); + } + + var authorization = await AuthorizeCredentialForEndpointAsync( + currentRecord, + credential, + linkedCts.Token, + requireSshTunnelOwnership: true) + .ConfigureAwait(false); + if (!isCurrentAttempt()) + { + return new ReconnectAuthorizationResult( + false, + GatewayErrorKind.Unknown, + $"{role} connection attempt was superseded."); + } + + var verifiedRecord = _registry.GetById(expectedRecord.Id); + if (verifiedRecord is null || + !string.Equals(_registry.ActiveGatewayId, expectedRecord.Id, StringComparison.Ordinal) || + !IsSameCredentialHandoffRecord(verifiedRecord, expectedRecord)) + { + return new ReconnectAuthorizationResult( + false, + GatewayErrorKind.LocalPortConflict, + $"The active gateway endpoint, credentials, or SSH configuration changed during the {role} credential handoff."); + } + + if (authorization.Allowed && + expectedOwnership is not null && + authorization.OwnershipProof != expectedOwnership) + { + return new ReconnectAuthorizationResult( + false, + GatewayErrorKind.LocalPortConflict, + $"Endpoint ownership changed after preflight, so {role} credentials were not sent."); + } + + return new ReconnectAuthorizationResult( + authorization.Allowed, + authorization.FailureKind, + authorization.Detail); + } + catch (OperationCanceledException) + { + if (!isCurrentAttempt() || operationCancellationToken.IsCancellationRequested) + { + return new ReconnectAuthorizationResult( + false, + GatewayErrorKind.Unknown, + $"{role} connection attempt was superseded or canceled."); + } + + return new ReconnectAuthorizationResult( + false, + GatewayErrorKind.Network, + $"Timed out re-verifying the owned SSH listener before the {role} credential handoff."); + } + } + + private static bool IsSameCredentialHandoffRecord( + GatewayRecord current, + GatewayRecord expected) => + string.Equals(current.Id, expected.Id, StringComparison.Ordinal) && + string.Equals(current.Url, expected.Url, StringComparison.Ordinal) && + string.Equals( + current.SharedGatewayToken, + expected.SharedGatewayToken, + StringComparison.Ordinal) && + string.Equals( + current.BootstrapToken, + expected.BootstrapToken, + StringComparison.Ordinal) && + current.IsLocal == expected.IsLocal && + current.RequiresV2Signature == expected.RequiresV2Signature && + string.Equals( + current.SetupManagedDistroName, + expected.SetupManagedDistroName, + StringComparison.Ordinal) && + current.SshTunnel == expected.SshTunnel; + private readonly record struct EndpointCredentialAuthorization( bool Allowed, GatewayErrorKind FailureKind, @@ -2828,6 +3302,36 @@ private async Task BlockNodeStartAsync( } } + private async Task RecordOperatorCredentialHandoffFailureAsync( + string detail, + GatewayErrorKind failureKind, + CancellationToken cancellationToken, + long expectedLifecycleGeneration, + string expectedGatewayId) + { + if (!IsCurrentGatewayAttempt(expectedLifecycleGeneration, expectedGatewayId)) + return; + + await _transitionSemaphore.WaitAsync(cancellationToken); + try + { + if (!IsCurrentGatewayAttempt(expectedLifecycleGeneration, expectedGatewayId)) + return; + + _stateMachine.SetOperatorErrorKind(failureKind); + _stateMachine.TryTransition( + failureKind == GatewayErrorKind.Network + ? ConnectionTrigger.WebSocketError + : ConnectionTrigger.AuthenticationFailed, + detail); + EmitStateChanged(); + } + finally + { + _transitionSemaphore.Release(); + } + } + private async Task StartNodeConnectionCoreAsync( long expectedLifecycleGeneration, long nodeGeneration, @@ -2998,38 +3502,28 @@ await BlockNodeStartAsync( async Task AuthorizeNodeCredentialHandoffAsync( CancellationToken authorizationCancellationToken) { - if (!IsExpectedNodeStartCurrent(expectedLifecycleGeneration, nodeGeneration)) - return new ReconnectAuthorizationResult( - false, - GatewayErrorKind.Unknown, - "Node attempt was superseded."); - var authorization = await AuthorizeCredentialForEndpointAsync( - record, - nodeCredential, - authorizationCancellationToken, - requireSshTunnelOwnership: true).ConfigureAwait(false); - if (authorization.Allowed && - expectedNodeEndpointOwnership is not null && - authorization.OwnershipProof != expectedNodeEndpointOwnership) - { - authorization = new EndpointCredentialAuthorization( - false, - GatewayErrorKind.LocalPortConflict, - "Endpoint ownership changed after preflight, so node credentials were not sent."); - } - if (!authorization.Allowed) + var authorization = await AuthorizeCredentialHandoffAsync( + record, + nodeCredential, + expectedNodeEndpointOwnership, + () => IsExpectedNodeStartCurrent( + expectedLifecycleGeneration, + nodeGeneration), + cancellationToken, + authorizationCancellationToken, + "node") + .ConfigureAwait(false); + if (!authorization.Allowed && + authorization.FailureKind != GatewayErrorKind.Unknown) { await BlockNodeStartAsync( - authorization.Detail, + authorization.Detail ?? "Node credential handoff was not authorized.", authorizationCancellationToken, expectedLifecycleGeneration, nodeGeneration) .ConfigureAwait(false); } - return new ReconnectAuthorizationResult( - authorization.Allowed, - authorization.FailureKind, - authorization.Detail); + return authorization; } reconnectPolicy.HandshakeAuthorizationAsync = @@ -4149,6 +4643,8 @@ private async Task DisposeCoreAsync() _disposed = true; CancelOperatorTelemetryAttempt("disposed", ConnectionErrorCategory.Disposed); CancelNodeTelemetryAttempt("disposed", ConnectionErrorCategory.Disposed); + try { Volatile.Read(ref _manualSshRestartCts)?.Cancel(); } + catch (ObjectDisposedException) { } _operationCts?.Cancel(); // Unsubscribe from node events before disposing the semaphore diff --git a/src/OpenClaw.Connection/IGatewayConnectionManager.cs b/src/OpenClaw.Connection/IGatewayConnectionManager.cs index 702036125..a98958dfe 100644 --- a/src/OpenClaw.Connection/IGatewayConnectionManager.cs +++ b/src/OpenClaw.Connection/IGatewayConnectionManager.cs @@ -26,6 +26,8 @@ public interface IGatewayConnectionManager : IDisposable, IAsyncDisposable Task ReconnectAsync(); Task ReconnectIfCurrentAsync(string gatewayId, CancellationToken cancellationToken = default); Task RecoverSshTunnelAsync(SshTunnelExit tunnelExit); + Task RestartSshTunnelAsync(CancellationToken cancellationToken = default) => + Task.FromResult(false); Task SwitchGatewayAsync(string gatewayId); void SetGatewayConnectionIntent(string gatewayId, bool shouldBeConnected); bool IsAutomaticReconnectAllowed(string gatewayId); diff --git a/src/OpenClaw.Connection/ISshTunnelManager.cs b/src/OpenClaw.Connection/ISshTunnelManager.cs index ef105644f..d9b1afc67 100644 --- a/src/OpenClaw.Connection/ISshTunnelManager.cs +++ b/src/OpenClaw.Connection/ISshTunnelManager.cs @@ -1,5 +1,10 @@ namespace OpenClaw.Connection; +public sealed record SshTunnelStartResult( + string Url, + SshTunnelConfig Config, + long OwnershipGeneration); + /// /// Manages an SSH tunnel lifecycle for a gateway connection. /// Wraps the existing SshTunnelService behind a clean interface. @@ -7,7 +12,7 @@ namespace OpenClaw.Connection; public interface ISshTunnelManager : IDisposable { bool IsActive { get; } - long OwnershipGeneration => 0; + long OwnershipGeneration { get; } bool IsRestartPending(SshTunnelExit tunnelExit); SshTunnelConfig? ActiveConfig { get; } Task IsOwnedListenerReadyAsync( @@ -15,6 +20,13 @@ Task IsOwnedListenerReadyAsync( int destinationPort, CancellationToken ct); Task StartAsync(SshTunnelConfig config, CancellationToken ct); + Task StartOwnedAsync( + SshTunnelConfig config, + CancellationToken ct); Task StopAsync(); + Task StopIfOwnedAsync( + SshTunnelConfig config, + long ownershipGeneration, + CancellationToken ct); string? LocalTunnelUrl { get; } } diff --git a/src/OpenClaw.Connection/SshTunnelService.cs b/src/OpenClaw.Connection/SshTunnelService.cs index 79e3fdff9..14c054c14 100644 --- a/src/OpenClaw.Connection/SshTunnelService.cs +++ b/src/OpenClaw.Connection/SshTunnelService.cs @@ -11,7 +11,6 @@ namespace OpenClaw.Connection; public sealed class SshTunnelService : ISshTunnelManager { private readonly IOpenClawLogger _logger; - private readonly string? _sshConfigFile; private readonly object _operationLock = new(); private readonly object _stateLock = new(); private Process? _process; @@ -24,10 +23,9 @@ public sealed class SshTunnelService : ISshTunnelManager /// Raised when the SSH tunnel exits unexpectedly (not during shutdown). public event EventHandler? TunnelExited; - public SshTunnelService(IOpenClawLogger logger, string? sshConfigFile = null) + public SshTunnelService(IOpenClawLogger logger) { _logger = logger; - _sshConfigFile = sshConfigFile; } public bool IsRunning @@ -288,8 +286,7 @@ private void StartProcess(SshTunnelConfig tunnel, SshTunnelOwner owner, string s remotePort, localPort, includeBrowserProxyForward, - sshPort, - _sshConfigFile), + sshPort), UseShellExecute = false, RedirectStandardOutput = true, RedirectStandardError = true, @@ -558,7 +555,12 @@ _process is null || } } - public async Task StartAsync(SshTunnelConfig config, CancellationToken ct) + public async Task StartAsync(SshTunnelConfig config, CancellationToken ct) => + (await StartOwnedAsync(config, ct).ConfigureAwait(false)).Url; + + public async Task StartOwnedAsync( + SshTunnelConfig config, + CancellationToken ct) { Process? process = null; long generation = 0; @@ -613,7 +615,10 @@ await WaitForOwnedLocalListenerAsync( processStartTimeUtc, ct).ConfigureAwait(false); } - return $"ws://localhost:{config.LocalPort}"; + return new SshTunnelStartResult( + $"ws://localhost:{config.LocalPort}", + normalizedConfig, + generation); } catch { @@ -629,6 +634,42 @@ public Task StopAsync() return Task.CompletedTask; } + public Task StopIfOwnedAsync( + SshTunnelConfig config, + long ownershipGeneration, + CancellationToken ct) + { + ct.ThrowIfCancellationRequested(); + while (!Monitor.TryEnter(_operationLock, millisecondsTimeout: 50)) + ct.ThrowIfCancellationRequested(); + try + { + ct.ThrowIfCancellationRequested(); + lock (_stateLock) + { + var normalizedConfig = config with + { + User = config.User.Trim(), + Host = config.Host.Trim(), + }; + if (_lifecycleGeneration != ownershipGeneration || + !Equals(_currentConfig, normalizedConfig) || + _currentOwner != SshTunnelOwner.GatewayConnectionManager) + { + return Task.FromResult(false); + } + } + + ct.ThrowIfCancellationRequested(); + StopLocked(); + return Task.FromResult(true); + } + finally + { + Monitor.Exit(_operationLock); + } + } + private async Task WaitForOwnedLocalListenerAsync( int localPort, Process process, diff --git a/src/OpenClaw.Shared/OpenClawGatewayClient.cs b/src/OpenClaw.Shared/OpenClawGatewayClient.cs index 7587af48a..e929b0727 100644 --- a/src/OpenClaw.Shared/OpenClawGatewayClient.cs +++ b/src/OpenClaw.Shared/OpenClawGatewayClient.cs @@ -95,8 +95,7 @@ internal IConnectEnvelopeSigner ConnectEnvelopeSigner private bool _pairingRequiredAwaitingApproval; private string? _pairingRequiredRequestId; private bool _authFailed; - private int _handshakeAuthorizationBlocked; - private int _handshakeChallengeActive; + private readonly HandshakeChallengeGate _handshakeChallengeGate = new(); private string? _lastSkillsStatusAgentId; private readonly bool _tokenIsBootstrapToken; private readonly bool _bootstrapPairAsNode; @@ -152,14 +151,21 @@ public void SetUserRules(IReadOnlyList? rules) protected override Task ProcessMessageAsync(string json) { - ProcessMessage(json); + ProcessMessageForConnection(json, CurrentConnectionGeneration); + return Task.CompletedTask; + } + + protected override Task ProcessMessageForConnectionAsync( + string json, + long sourceConnectionGeneration) + { + ProcessMessageForConnection(json, sourceConnectionGeneration); return Task.CompletedTask; } protected override Task OnConnectedAsync() { - Volatile.Write(ref _handshakeAuthorizationBlocked, 0); - Volatile.Write(ref _handshakeChallengeActive, 0); + _handshakeChallengeGate.Reset(CurrentConnectionGeneration); _pendingRequests.OpenConnection(); ResetUnsupportedMethodFlags(); RaiseTransportConnected(); @@ -1722,7 +1728,10 @@ public async Task LogoutChannelAsync(string channelName, int timeoutMs = 1 } } - private async Task SendConnectMessageAsync(string? nonce = null) + private async Task SendConnectMessageAsync( + string? nonce, + long connectionGeneration, + CancellationToken cancellationToken) { var requestId = Guid.NewGuid().ToString(); var pending = IsConnected @@ -1771,7 +1780,14 @@ private async Task SendConnectMessageAsync(string? nonce = null) try { - await SendRawAsync(envelope.Serialize(signature)); + var sent = await SendRawAsync( + envelope.Serialize(signature), + connectionGeneration, + cancellationToken) + .ConfigureAwait(false); + if (!sent && pending is not null) + _pendingRequests.TryRemove(pending); + return sent; } catch { @@ -1893,8 +1909,16 @@ private static string SerializeRequest(string requestId, string method, object? // --- Message processing --- - private void ProcessMessage(string json) + private void ProcessMessage(string json) => + ProcessMessageForConnection(json, CurrentConnectionGeneration); + + private void ProcessMessageForConnection( + string json, + long sourceConnectionGeneration) { + if (!IsCurrentConnectionGeneration(sourceConnectionGeneration)) + return; + try { using var doc = JsonDocument.Parse(json); @@ -1906,10 +1930,13 @@ private void ProcessMessage(string json) switch (type) { case "res": - HandleResponse(root); + HandleResponseForConnection(root, sourceConnectionGeneration); break; case "event": - HandleEvent(root, json.Length); + HandleEventForConnection( + root, + json.Length, + sourceConnectionGeneration); break; } } @@ -1923,7 +1950,9 @@ private void ProcessMessage(string json) } } - private void HandleResponse(JsonElement root) + private void HandleResponseForConnection( + JsonElement root, + long sourceConnectionGeneration) { string? requestId = null; if (root.TryGetProperty("id", out var idProp)) @@ -2013,10 +2042,11 @@ private void HandleResponse(JsonElement root) // Handle handshake acknowledgement payload. if (payload.TryGetProperty("type", out var t) && t.GetString() == "hello-ok") { - if (HandshakeAuthorizationAsync is not null && + if (!IsCurrentConnectionGeneration(sourceConnectionGeneration) || + !_handshakeChallengeGate.IsAuthorized(sourceConnectionGeneration) || !string.Equals(requestMethod, "connect", StringComparison.Ordinal)) { - _logger.Warn("[HANDSHAKE] Ignoring uncorrelated hello-ok on guarded validation connection."); + _logger.Warn("[HANDSHAKE] Ignoring stale or uncorrelated hello-ok."); return; } @@ -2933,7 +2963,10 @@ private static bool IsSupportedDeviceTokenRole(string role) => return null; } - private void HandleEvent(JsonElement root, int rawMessageLength) + private void HandleEventForConnection( + JsonElement root, + int rawMessageLength, + long sourceConnectionGeneration) { if (!root.TryGetProperty("event", out var eventProp)) return; var eventType = eventProp.GetString(); @@ -2942,7 +2975,9 @@ private void HandleEvent(JsonElement root, int rawMessageLength) switch (eventType) { case "connect.challenge": - HandleConnectChallenge(root); + HandleConnectChallengeForConnection( + root, + sourceConnectionGeneration); break; case "agent": HandleAgentEvent(root, rawMessageLength); @@ -3175,31 +3210,46 @@ static string SafeStr(JsonElement obj, string name) } } - private void HandleConnectChallenge(JsonElement root) + private void HandleConnectChallenge(JsonElement root) => + HandleConnectChallengeForConnection(root, CurrentConnectionGeneration); + + private void HandleConnectChallengeForConnection( + JsonElement root, + long sourceConnectionGeneration) { - if (Volatile.Read(ref _handshakeAuthorizationBlocked) != 0 || - Interlocked.CompareExchange(ref _handshakeChallengeActive, 1, 0) != 0) + if (!IsCurrentConnectionGeneration(sourceConnectionGeneration)) + return; + + if (!root.TryGetProperty("payload", out var payload) || + payload.ValueKind != JsonValueKind.Object) { - _logger.Warn("[HANDSHAKE] Ignoring duplicate challenge on the current socket."); + _logger.Warn("[HANDSHAKE] Ignoring malformed challenge without an object payload."); return; } string? nonce = null; - long? ts = null; - if (root.TryGetProperty("payload", out var payload)) + if (payload.TryGetProperty("nonce", out var nonceProp)) { - if (payload.TryGetProperty("nonce", out var nonceProp)) + if (nonceProp.ValueKind != JsonValueKind.String) { - nonce = nonceProp.GetString(); + _logger.Warn("[HANDSHAKE] Ignoring malformed challenge with a non-string nonce."); + return; } - ts = ConnectAuthTimestamp.ReadChallengeTimestamp(payload); + nonce = nonceProp.GetString(); + } + var ts = ConnectAuthTimestamp.ReadChallengeTimestamp(payload); + + if (!_handshakeChallengeGate.TryBegin(sourceConnectionGeneration)) + { + _logger.Warn("[HANDSHAKE] Ignoring duplicate challenge on the current socket."); + return; } _challengeTimestampMs = ts; _currentChallengeNonce = nonce; _logger.Info($"[HANDSHAKE] Received connect.challenge: nonce={nonce}, ts={ts}"); - _ = SendConnectSafeAsync(nonce, CurrentConnectionGeneration); + _ = SendConnectSafeAsync(nonce, sourceConnectionGeneration); } private async Task SendConnectSafeAsync(string? nonce, long connectionGeneration) @@ -3208,36 +3258,58 @@ private async Task SendConnectSafeAsync(string? nonce, long connectionGeneration { if (HandshakeAuthorizationAsync is not null) { - var authorization = await HandshakeAuthorizationAsync(CancellationToken.None) + var authorization = await HandshakeAuthorizationAsync(CancellationToken) .ConfigureAwait(false); if (!IsCurrentConnectionGeneration(connectionGeneration)) return; if (!authorization.Allowed) { - Volatile.Write(ref _handshakeAuthorizationBlocked, 1); + if (!_handshakeChallengeGate.TryBlock(connectionGeneration)) + return; + _logger.Warn( + $"[HANDSHAKE] Operator credential handoff blocked: {authorization.Detail}"); RaiseConnectionFailure(authorization.FailureKind); - RaiseAuthenticationFailed( - authorization.Detail ?? "Connection authorization failed."); + if (authorization.FailureKind is GatewayErrorKind.Auth + or GatewayErrorKind.DeviceTokenMismatch + or GatewayErrorKind.TokenDrift) + { + RaiseAuthenticationFailed( + authorization.Detail ?? "Connection authorization failed."); + } AbortCurrentWebSocket(connectionGeneration); RaiseStatusChanged(ConnectionStatus.Error); return; } + if (!_handshakeChallengeGate.TryAuthorize(connectionGeneration)) + return; } - - if (!IsCurrentConnectionGeneration(connectionGeneration) || - Volatile.Read(ref _handshakeAuthorizationBlocked) != 0) + else if (!_handshakeChallengeGate.TryAuthorize(connectionGeneration)) + { return; + } - await SendConnectMessageAsync(nonce); + var sent = await SendConnectMessageAsync( + nonce, + connectionGeneration, + CancellationToken) + .ConfigureAwait(false); + if (!sent && IsCurrentConnectionGeneration(connectionGeneration)) + AbortCurrentWebSocket(connectionGeneration); } - catch (Exception ex) + catch (OperationCanceledException) when ( + CancellationToken.IsCancellationRequested || + !IsCurrentConnectionGeneration(connectionGeneration)) { - _logger.Error($"[HANDSHAKE] FATAL: SendConnectMessageAsync threw: {ex}"); } - finally + catch (Exception ex) { - if (IsCurrentConnectionGeneration(connectionGeneration)) - Volatile.Write(ref _handshakeChallengeActive, 0); + _logger.Error($"[HANDSHAKE] FATAL: SendConnectMessageAsync threw: {ex}"); + if (_handshakeChallengeGate.TryBlock(connectionGeneration)) + { + RaiseConnectionFailure(GatewayErrorKind.Network); + AbortCurrentWebSocket(connectionGeneration); + RaiseStatusChanged(ConnectionStatus.Error); + } } } diff --git a/src/OpenClaw.Shared/SshTunnelCommandLine.cs b/src/OpenClaw.Shared/SshTunnelCommandLine.cs index a8c15f652..48076d8b9 100644 --- a/src/OpenClaw.Shared/SshTunnelCommandLine.cs +++ b/src/OpenClaw.Shared/SshTunnelCommandLine.cs @@ -36,23 +36,6 @@ public static string BuildArguments( int localPort, bool includeBrowserProxyForward, int sshPort) - => BuildArguments( - user, - host, - remotePort, - localPort, - includeBrowserProxyForward, - sshPort, - sshConfigFile: null); - - public static string BuildArguments( - string user, - string host, - int remotePort, - int localPort, - bool includeBrowserProxyForward, - int sshPort, - string? sshConfigFile) { user = user.Trim(); host = host.Trim(); @@ -71,18 +54,6 @@ public static string BuildArguments( } var sb = new StringBuilder(); - if (sshConfigFile is not null) - { - if (string.IsNullOrWhiteSpace(sshConfigFile) || - sshConfigFile.IndexOfAny(['"', '\r', '\n', '\0']) >= 0) - { - throw new ArgumentException("SSH config path contains invalid characters.", nameof(sshConfigFile)); - } - - sb.Append("-F \""); - sb.Append(sshConfigFile); - sb.Append("\" "); - } sb.Append(BaseOptions); AppendLocalForward(sb, localPort, remotePort); if (includeBrowserProxyForward) diff --git a/src/OpenClaw.Shared/WebSocketClientBase.cs b/src/OpenClaw.Shared/WebSocketClientBase.cs index 131b9e849..534151000 100644 --- a/src/OpenClaw.Shared/WebSocketClientBase.cs +++ b/src/OpenClaw.Shared/WebSocketClientBase.cs @@ -15,6 +15,79 @@ public readonly record struct ReconnectAuthorizationResult( public static ReconnectAuthorizationResult AllowedResult { get; } = new(true); } +internal enum HandshakeChallengeState +{ + Idle, + Active, + Authorized, + Blocked, +} + +internal sealed class HandshakeChallengeGate +{ + private readonly object _lock = new(); + private long _generation; + private HandshakeChallengeState _state; + + public void Reset(long generation) + { + lock (_lock) + { + _generation = generation; + _state = HandshakeChallengeState.Idle; + } + } + + public bool TryBegin(long generation) + { + lock (_lock) + { + if (_generation != generation) + return false; + + if (_state != HandshakeChallengeState.Idle) + return false; + + _state = HandshakeChallengeState.Active; + return true; + } + } + + public bool TryAuthorize(long generation) + { + lock (_lock) + { + if (_generation != generation || _state != HandshakeChallengeState.Active) + return false; + + _state = HandshakeChallengeState.Authorized; + return true; + } + } + + public bool TryBlock(long generation) + { + lock (_lock) + { + if (_generation != generation || + _state is not (HandshakeChallengeState.Active or HandshakeChallengeState.Authorized)) + return false; + + _state = HandshakeChallengeState.Blocked; + return true; + } + } + + public bool IsAuthorized(long generation) + { + lock (_lock) + { + return _generation == generation && + _state == HandshakeChallengeState.Authorized; + } + } +} + /// /// Abstract base class for WebSocket-based gateway clients. /// Extracts shared connection lifecycle: connect, listen, reconnect, send, dispose. @@ -97,6 +170,15 @@ protected void RaiseAuthenticationFailed(string message) /// protected abstract Task ProcessMessageAsync(string json); + /// + /// Process a message attributed to the socket generation that received it. + /// Override when message side effects or responses must remain bound to that socket. + /// + protected virtual Task ProcessMessageForConnectionAsync( + string json, + long sourceConnectionGeneration) => + ProcessMessageAsync(json); + /// Receive buffer size in bytes. Gateway: 16384, Node: 65536. protected abstract int ReceiveBufferSize { get; } @@ -316,7 +398,9 @@ private async Task ListenForMessagesAsync(ClientWebSocket ws, long connectionGen if (result.EndOfMessage && sb.Length == 0) { // Fast path: single-frame message — decode directly, skip StringBuilder round-trip - await ProcessMessageAsync(Encoding.UTF8.GetString(buffer, 0, result.Count)); + await ProcessMessageForConnectionAsync( + Encoding.UTF8.GetString(buffer, 0, result.Count), + connectionGeneration); } else { @@ -346,7 +430,9 @@ private async Task ListenForMessagesAsync(ClientWebSocket ws, long connectionGen if (result.EndOfMessage) { - await ProcessMessageAsync(sb.ToString()); + await ProcessMessageForConnectionAsync( + sb.ToString(), + connectionGeneration); sb.Clear(); } } @@ -526,6 +612,7 @@ protected virtual async Task SendRawAsync(string message) { await _sendLock.WaitAsync(_cts.Token); } + catch (OperationCanceledException) { // Shutdown canceled the wait; drop the send silently. @@ -581,6 +668,77 @@ await ws.SendAsync(buffer.AsMemory(0, written), } } + /// + /// Sends only when the captured socket generation still owns the transport immediately before + /// the write. Used for credential-bearing handshake frames. + /// + protected virtual async Task SendRawAsync( + string message, + long expectedConnectionGeneration, + CancellationToken cancellationToken) + { + using var linkedCancellation = CancellationTokenSource.CreateLinkedTokenSource( + _cts.Token, + cancellationToken); + try + { + await _sendLock.WaitAsync(linkedCancellation.Token).ConfigureAwait(false); + } + catch (OperationCanceledException) + { + return false; + } + catch (ObjectDisposedException) + { + return false; + } + + try + { + var ws = _webSocket; + if (ws?.State != WebSocketState.Open || + !IsCurrentConnection(ws, expectedConnectionGeneration)) + { + return false; + } + + var byteCount = Encoding.UTF8.GetByteCount(message); + var buffer = ArrayPool.Shared.Rent(byteCount); + try + { + var written = Encoding.UTF8.GetBytes(message, buffer); + await ws.SendAsync( + buffer.AsMemory(0, written), + WebSocketMessageType.Text, + true, + linkedCancellation.Token) + .ConfigureAwait(false); + return IsCurrentConnection(ws, expectedConnectionGeneration); + } + catch (OperationCanceledException) + { + return false; + } + catch (ObjectDisposedException) + { + return false; + } + catch (WebSocketException ex) when (ex.WebSocketErrorCode == WebSocketError.InvalidState) + { + _logger.Warn($"WebSocket send failed (state changed): {ex.Message}"); + return false; + } + finally + { + ArrayPool.Shared.Return(buffer); + } + } + finally + { + _sendLock.Release(); + } + } + /// Gracefully close the WebSocket connection. protected async Task CloseWebSocketAsync() { diff --git a/src/OpenClaw.Shared/WindowsNodeClient.cs b/src/OpenClaw.Shared/WindowsNodeClient.cs index 497272ec2..171f49fe6 100644 --- a/src/OpenClaw.Shared/WindowsNodeClient.cs +++ b/src/OpenClaw.Shared/WindowsNodeClient.cs @@ -40,7 +40,7 @@ public class WindowsNodeClient : WebSocketClientBase private volatile bool _rateLimited; private bool _useV2Signature; // true after v3 signature rejected by gateway public bool UseV2Signature { get => _useV2Signature; set => _useV2Signature = value; } - private int _handshakeAuthorizationBlocked; + private readonly HandshakeChallengeGate _handshakeChallengeGate = new(); internal IConnectEnvelopeSigner ConnectEnvelopeSigner { get => _connectEnvelopeSigner; @@ -137,7 +137,7 @@ protected override void OnReconnectAuthorizationDenied( protected override Task OnConnectedAsync() { _isConnected = false; - Volatile.Write(ref _handshakeAuthorizationBlocked, 0); + _handshakeChallengeGate.Reset(CurrentConnectionGeneration); Volatile.Write(ref _pendingConnectRequestId, null); TransportConnected?.Invoke(this, EventArgs.Empty); return Task.CompletedTask; @@ -269,6 +269,19 @@ public Task DisconnectAsync() protected override async Task ProcessMessageAsync(string json) { + await ProcessMessageForConnectionAsync( + json, + CurrentConnectionGeneration) + .ConfigureAwait(false); + } + + protected override async Task ProcessMessageForConnectionAsync( + string json, + long sourceConnectionGeneration) + { + if (!IsCurrentConnectionGeneration(sourceConnectionGeneration)) + return; + try { // Log raw messages at debug level (visible in dbgview, not in log file noise) @@ -288,10 +301,12 @@ protected override async Task ProcessMessageAsync(string json) switch (type) { case "event": - await HandleEventAsync(root); + await HandleEventForConnectionAsync( + root, + sourceConnectionGeneration); break; case "res": - HandleResponse(root); + HandleResponseForConnection(root, sourceConnectionGeneration); break; case "req": await HandleRequestAsync(root); @@ -311,7 +326,12 @@ protected override async Task ProcessMessageAsync(string json) } } - private async Task HandleEventAsync(JsonElement root) + private Task HandleEventAsync(JsonElement root) => + HandleEventForConnectionAsync(root, CurrentConnectionGeneration); + + private async Task HandleEventForConnectionAsync( + JsonElement root, + long sourceConnectionGeneration) { if (!root.TryGetProperty("event", out var eventProp)) return; var eventType = eventProp.GetString(); @@ -336,7 +356,9 @@ and not "node.pair.resolved" switch (eventType) { case "connect.challenge": - await HandleConnectChallengeAsync(root); + await HandleConnectChallengeForConnectionAsync( + root, + sourceConnectionGeneration); break; case "node.pair.requested": case "device.pair.requested": @@ -633,58 +655,106 @@ private async Task SendNodeInvokeResultAsync(string requestId, bool success, obj await SendRawAsync(json); } - private async Task HandleConnectChallengeAsync(JsonElement root) + private Task HandleConnectChallengeAsync(JsonElement root) => + HandleConnectChallengeForConnectionAsync( + root, + CurrentConnectionGeneration); + + private async Task HandleConnectChallengeForConnectionAsync( + JsonElement root, + long sourceConnectionGeneration) { - var connectionGeneration = CurrentConnectionGeneration; - if (Volatile.Read(ref _handshakeAuthorizationBlocked) != 0) + if (!IsCurrentConnectionGeneration(sourceConnectionGeneration)) + return; + + if (!root.TryGetProperty("payload", out var payload) || + payload.ValueKind != JsonValueKind.Object) { - _logger.Warn("[HANDSHAKE] Ignoring duplicate node challenge on the current socket."); + _logger.Warn("[HANDSHAKE] Ignoring malformed node challenge without an object payload."); return; } string? nonce = null; - long? challengeTimestampMs = null; - - if (root.TryGetProperty("payload", out var payload)) + if (payload.TryGetProperty("nonce", out var nonceProp)) { - if (payload.TryGetProperty("nonce", out var nonceProp)) + if (nonceProp.ValueKind != JsonValueKind.String) { - nonce = nonceProp.GetString(); + _logger.Warn("[HANDSHAKE] Ignoring malformed node challenge with a non-string nonce."); + return; } - challengeTimestampMs = ConnectAuthTimestamp.ReadChallengeTimestamp(payload); + nonce = nonceProp.GetString(); + } + var challengeTimestampMs = ConnectAuthTimestamp.ReadChallengeTimestamp(payload); + + if (!_handshakeChallengeGate.TryBegin(sourceConnectionGeneration)) + { + _logger.Warn("[HANDSHAKE] Ignoring duplicate node challenge on the current socket."); + return; } _logger.Info($"[HANDSHAKE] Received connect.challenge: nonce={nonce}, ts={challengeTimestampMs?.ToString() ?? "missing"}"); _pendingNonce = nonce; - if (HandshakeAuthorizationAsync is not null) + try { - var authorization = await HandshakeAuthorizationAsync(CancellationToken.None) - .ConfigureAwait(false); - if (!IsCurrentConnectionGeneration(connectionGeneration)) + if (HandshakeAuthorizationAsync is not null) + { + var authorization = await HandshakeAuthorizationAsync(CancellationToken) + .ConfigureAwait(false); + if (!IsCurrentConnectionGeneration(sourceConnectionGeneration)) + return; + if (!authorization.Allowed) + { + if (!_handshakeChallengeGate.TryBlock(sourceConnectionGeneration)) + return; + _logger.Warn( + $"[HANDSHAKE] Node credential handoff blocked: {authorization.Detail}"); + ConnectionFailure?.Invoke(this, authorization.FailureKind); + AbortCurrentWebSocket(sourceConnectionGeneration); + RaiseStatusChanged(ConnectionStatus.Error); + return; + } + if (!_handshakeChallengeGate.TryAuthorize(sourceConnectionGeneration)) + return; + } + else if (!_handshakeChallengeGate.TryAuthorize(sourceConnectionGeneration)) + { return; - if (!authorization.Allowed) + } + + var sent = await SendNodeConnectAsync( + nonce, + challengeTimestampMs, + sourceConnectionGeneration, + CancellationToken) + .ConfigureAwait(false); + if (!sent && IsCurrentConnectionGeneration(sourceConnectionGeneration)) + AbortCurrentWebSocket(sourceConnectionGeneration); + } + catch (OperationCanceledException) when ( + CancellationToken.IsCancellationRequested || + !IsCurrentConnectionGeneration(sourceConnectionGeneration)) + { + } + catch (Exception ex) + { + _logger.Error($"[HANDSHAKE] Node connect handoff failed: {ex.Message}"); + if (_handshakeChallengeGate.TryBlock(sourceConnectionGeneration)) { - Volatile.Write(ref _handshakeAuthorizationBlocked, 1); - _logger.Warn( - $"[HANDSHAKE] Node credential handoff blocked: {authorization.Detail}"); - ConnectionFailure?.Invoke(this, authorization.FailureKind); - AbortCurrentWebSocket(connectionGeneration); + ConnectionFailure?.Invoke(this, GatewayErrorKind.Network); + AbortCurrentWebSocket(sourceConnectionGeneration); RaiseStatusChanged(ConnectionStatus.Error); - return; } } - - if (!IsCurrentConnectionGeneration(connectionGeneration) || - Volatile.Read(ref _handshakeAuthorizationBlocked) != 0) - return; - - await SendNodeConnectAsync(nonce, challengeTimestampMs); } private const string ClientId = "node-host"; // Must be "node-host" for nodes - private async Task SendNodeConnectAsync(string? nonce, long? challengeTimestampMs) + private async Task SendNodeConnectAsync( + string? nonce, + long? challengeTimestampMs, + long connectionGeneration, + CancellationToken cancellationToken) { var isPaired = !string.IsNullOrEmpty(_deviceIdentity.NodeDeviceToken); var usingBootstrap = !isPaired && !string.IsNullOrEmpty(_bootstrapToken); @@ -704,7 +774,14 @@ private async Task SendNodeConnectAsync(string? nonce, long? challengeTimestampM var requestId = Guid.NewGuid().ToString(); Volatile.Write(ref _pendingConnectRequestId, requestId); - await SendRawAsync(BuildNodeConnectMessage(nonce, challengeTimestampMs, requestId)); + var sent = await SendRawAsync( + BuildNodeConnectMessage(nonce, challengeTimestampMs, requestId), + connectionGeneration, + cancellationToken) + .ConfigureAwait(false); + if (!sent) + Interlocked.CompareExchange(ref _pendingConnectRequestId, null, requestId); + return sent; } private string BuildNodeConnectMessage( @@ -762,8 +839,16 @@ private ConnectCredential SelectConnectCredential() return new TokenConnectCredential(_gatewayToken); } - internal void HandleResponse(JsonElement root) + internal void HandleResponse(JsonElement root) => + HandleResponseForConnection(root, CurrentConnectionGeneration); + + private void HandleResponseForConnection( + JsonElement root, + long sourceConnectionGeneration) { + if (!IsCurrentConnectionGeneration(sourceConnectionGeneration)) + return; + var responseId = root.TryGetProperty("id", out var idProp) ? idProp.GetString() : null; @@ -775,6 +860,12 @@ internal void HandleResponse(JsonElement root) if (root.TryGetProperty("ok", out var okProp) && okProp.ValueKind == JsonValueKind.False) { + if (isConnectResponse && + !_handshakeChallengeGate.IsAuthorized(sourceConnectionGeneration)) + { + _logger.Warn("[HANDSHAKE] Ignoring stale node connect denial."); + return; + } if (isConnectResponse) Volatile.Write(ref _pendingConnectRequestId, null); HandleRequestError(root); @@ -790,9 +881,10 @@ internal void HandleResponse(JsonElement root) // Handle hello-ok (successful registration) if (payload.TryGetProperty("type", out var t) && t.GetString() == "hello-ok") { - if (!isConnectResponse) + if (!isConnectResponse || + !_handshakeChallengeGate.IsAuthorized(sourceConnectionGeneration)) { - _logger.Warn("[HANDSHAKE] Ignoring uncorrelated node hello-ok."); + _logger.Warn("[HANDSHAKE] Ignoring stale or uncorrelated node hello-ok."); return; } diff --git a/src/OpenClaw.Tray.WinUI/App.xaml.cs b/src/OpenClaw.Tray.WinUI/App.xaml.cs index 050c0f544..a57b2c8c1 100644 --- a/src/OpenClaw.Tray.WinUI/App.xaml.cs +++ b/src/OpenClaw.Tray.WinUI/App.xaml.cs @@ -278,19 +278,6 @@ public IntPtr GetHubWindowHandle() => private readonly AppCrashLogger _crashLogger = new(Path.Combine(DataPath, "crash.log")); private static readonly AppRunMarker s_runMarker = new(Path.Combine(DataPath, "run.marker")); - private static string? ResolveE2eSshConfigFile() - { - if (Environment.GetEnvironmentVariable("OPENCLAW_RUN_E2E") != "1") - return null; - - var path = Environment.GetEnvironmentVariable("OPENCLAW_E2E_SSH_CONFIG_FILE"); - if (string.IsNullOrWhiteSpace(path)) - return null; - if (!File.Exists(path)) - throw new FileNotFoundException("E2E SSH config file was not found.", path); - return Path.GetFullPath(path); - } - public App() { WaitForRestartSourceIfRequested(Environment.GetCommandLineArgs()); @@ -719,8 +706,7 @@ _dispatcherQueue is null // Register toast activation handler ToastNotificationManagerCompat.OnActivated += OnToastActivated; - var e2eSshConfigFile = ResolveE2eSshConfigFile(); - _sshTunnelService = new SshTunnelService(new AppLogger(), e2eSshConfigFile); + _sshTunnelService = new SshTunnelService(new AppLogger()); _sshTunnelService.TunnelExited += OnSshTunnelExited; // Initialize tray icon FIRST (window-less pattern from WinUIEx). @@ -812,7 +798,7 @@ _dispatcherQueue is null diagnostics: diagnostics, tunnelManager: _sshTunnelService, endpointProvenanceProbe: managedLocalPortProvenance.InspectAsync, - validationTunnelFactory: () => new SshTunnelService(appLogger, e2eSshConfigFile)); + validationTunnelFactory: () => new SshTunnelService(appLogger)); _connectionManager.OperatorClientChanged += OnOperatorClientChanged; _connectionManager.StateChanged += OnManagerStateChanged; _gatewayDirectConnectService = new GatewayDirectConnectService( @@ -3502,7 +3488,19 @@ private void ShowPairingApprovalDialog(bool bringToFront) _pairingApprovalDialog.ShowForeground(); } - private void RestartSshTunnel() + public async Task RestartSshTunnelAsync() + { + return _connectionManager is not null && + await _connectionManager.RestartSshTunnelAsync(); + } + + private void RestartSshTunnel() => + AsyncEventHandlerGuard.Run( + RestartSshTunnelCoreAsync, + new AppLogger(), + nameof(RestartSshTunnel)); + + private async Task RestartSshTunnelCoreAsync() { if (_settings?.UseSshTunnel != true) { @@ -3521,26 +3519,21 @@ private void RestartSshTunnel() remotePort = _settings.SshTunnelRemotePort }); - _sshTunnelService?.Stop(); - // Status is updated by OnManagerStateChanged when reconnect completes. - UpdateTrayIcon(); - - if (!EnsureSshTunnelConfigured()) + var restarted = await RestartSshTunnelAsync(); + UpdateStatusDetailWindow(); + if (restarted) + { + _sshTunnelRecoveryBudget.Reset(); + _toastService!.ShowToast(new ToastContentBuilder() + .AddText("SSH tunnel") + .AddText("Restarted and authenticated.")); + } + else { - UpdateStatusDetailWindow(); _toastService!.ShowToast(new ToastContentBuilder() .AddText("SSH tunnel restart failed") - .AddText(_sshTunnelService?.LastError ?? "Check SSH tunnel settings and logs.")); - return; + .AddText("The owned tunnel or authenticated gateway connection could not be verified.")); } - - _sshTunnelRecoveryBudget.Reset(); - ReconnectWithSyncedBrowserProxyForward(); - - UpdateStatusDetailWindow(); - _toastService!.ShowToast(new ToastContentBuilder() - .AddText("SSH tunnel") - .AddText("Restarted; reconnecting to gateway.")); } catch (Exception ex) { @@ -4098,9 +4091,7 @@ _settings.SshTunnelRemotePort is < 1 or > 65535 || try { - _sshTunnelService ??= new SshTunnelService( - new AppLogger(), - ResolveE2eSshConfigFile()); + _sshTunnelService ??= new SshTunnelService(new AppLogger()); var includeBrowserProxy = BrowserProxySshTunnelForwardPolicy.ShouldInclude( _settings.NodeBrowserProxyEnabled, _settings.SshTunnelRemotePort, diff --git a/src/OpenClaw.Tray.WinUI/Pages/ConnectionPage.xaml.cs b/src/OpenClaw.Tray.WinUI/Pages/ConnectionPage.xaml.cs index e3df58493..7800cd709 100644 --- a/src/OpenClaw.Tray.WinUI/Pages/ConnectionPage.xaml.cs +++ b/src/OpenClaw.Tray.WinUI/Pages/ConnectionPage.xaml.cs @@ -2379,13 +2379,23 @@ private void OnCopyApproveCommand(object sender, RoutedEventArgs e) ClipboardHelper.CopyText(RecoveryApproveCmdText.Text); } - private void OnRestartTunnel(object sender, RoutedEventArgs e) + private void OnRestartTunnel(object sender, RoutedEventArgs e) => + AsyncEventHandlerGuard.Run( + OnRestartTunnelAsync, + new OpenClawTray.AppLogger(), + nameof(OnRestartTunnel)); + + private async Task OnRestartTunnelAsync() { try { var app = (App)Microsoft.UI.Xaml.Application.Current; - app.EnsureSshTunnelStarted(); - AddResultText.Text = LocalizationHelper.GetString("ConnectionPage_TunnelRestartTriggered"); + var restarted = await app.RestartSshTunnelAsync(); + AddResultText.Text = restarted + ? LocalizationHelper.GetString("ConnectionPage_TunnelRestartTriggered") + : string.Format( + LocalizationHelper.GetString("ConnectionPage_TunnelRestartFailed"), + "The owned tunnel or authenticated gateway connection could not be verified."); } catch (Exception ex) { diff --git a/tests/OpenClaw.Connection.Tests/GatewayClientFactoryTests.cs b/tests/OpenClaw.Connection.Tests/GatewayClientFactoryTests.cs index 73996456f..d19ecd7c6 100644 --- a/tests/OpenClaw.Connection.Tests/GatewayClientFactoryTests.cs +++ b/tests/OpenClaw.Connection.Tests/GatewayClientFactoryTests.cs @@ -119,10 +119,11 @@ public async Task HandshakeAuthorization_BlocksCredentialBearingConnectMessage() try { + var logger = new TestLogger(); using var client = new OpenClawGatewayClient( "ws://127.0.0.1:18789", "replacement-token", - NullLogger.Instance, + logger, identityPath: tempDir, ignoreStoredDeviceToken: true, persistHandshakeDeviceTokens: false); @@ -138,13 +139,21 @@ public async Task HandshakeAuthorization_BlocksCredentialBearingConnectMessage() "validation listener ownership lost")); }; client.AuthenticationFailed += (_, message) => failure = message; + GatewayErrorKind? failureKind = null; + client.ConnectionFailure += (_, kind) => failureKind = kind; client.StatusChanged += (_, status) => lastStatus = status; await InvokeSendConnectSafeAsync(client); Assert.Equal(1, authorizationCalls); - Assert.Equal("validation listener ownership lost", failure); + Assert.Null(failure); + Assert.Equal(GatewayErrorKind.LocalPortConflict, failureKind); Assert.Equal(ConnectionStatus.Error, lastStatus); + Assert.Contains( + logger.Warnings, + message => message.Contains( + "validation listener ownership lost", + StringComparison.Ordinal)); } finally { @@ -152,6 +161,16 @@ public async Task HandshakeAuthorization_BlocksCredentialBearingConnectMessage() } } + private sealed class TestLogger : IOpenClawLogger + { + public List Warnings { get; } = []; + + public void Info(string message) { } + public void Debug(string message) { } + public void Warn(string message) => Warnings.Add(message); + public void Error(string message, Exception? ex = null) { } + } + private static string GetConnectRole(OpenClawGatewayClient client) { var method = typeof(OpenClawGatewayClient).GetMethod( @@ -187,6 +206,16 @@ private static bool TryStoreHandshakeDeviceToken( private static async Task InvokeSendConnectSafeAsync(OpenClawGatewayClient client) { + var gateField = typeof(OpenClawGatewayClient).GetField( + "_handshakeChallengeGate", + BindingFlags.Instance | BindingFlags.NonPublic); + Assert.NotNull(gateField); + var gate = gateField.GetValue(client); + Assert.NotNull(gate); + var tryBegin = gate.GetType().GetMethod("TryBegin"); + Assert.NotNull(tryBegin); + Assert.True(Assert.IsType(tryBegin.Invoke(gate, [0L]))); + var method = typeof(OpenClawGatewayClient).GetMethod( "SendConnectSafeAsync", BindingFlags.Instance | BindingFlags.NonPublic); diff --git a/tests/OpenClaw.Connection.Tests/GatewayConnectionManagerTests.cs b/tests/OpenClaw.Connection.Tests/GatewayConnectionManagerTests.cs index c49632453..53bf12a58 100644 --- a/tests/OpenClaw.Connection.Tests/GatewayConnectionManagerTests.cs +++ b/tests/OpenClaw.Connection.Tests/GatewayConnectionManagerTests.cs @@ -261,6 +261,461 @@ public async Task RecoverSshTunnelAsync_SettingsOwnedTunnel_DoesNotReconnectGate Assert.Equal(0, tunnel.StartCount); } + [Fact] + public async Task RestartSshTunnelAsync_ReturnsSuccessOnlyAfterFreshHandshake() + { + var (manager, tunnel, factory, tunnelConfig) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var listenerChecksBeforeRestart = tunnel.OwnedListenerCheckCount; + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + + Assert.False(restart.IsCompleted); + factory.CreatedClients[1].SimulateHandshake(); + + Assert.True(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.True(tunnel.IsActive); + Assert.Equal(tunnelConfig, tunnel.ActiveConfig); + Assert.Equal( + listenerChecksBeforeRestart + 2, + tunnel.OwnedListenerCheckCount); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_ConnectionLossDuringListenerCheckFails() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var listenerCheckStarted = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var releaseListenerCheck = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var checkCount = 0; + tunnel.OwnedListenerCheckAsync = cancellationToken => + { + if (Interlocked.Increment(ref checkCount) == 2) + { + listenerCheckStarted.TrySetResult(); + return releaseListenerCheck.Task.WaitAsync(cancellationToken); + } + + return Task.FromResult(true); + }; + + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + factory.CreatedClients[1].SimulateHandshake(); + await listenerCheckStarted.Task.WaitAsync(TimeSpan.FromSeconds(2)); + factory.CreatedClients[1].SimulateStatusChanged(ConnectionStatus.Error); + releaseListenerCheck.TrySetResult(true); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.False(tunnel.IsActive); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_ManagerDisposalCancelsInFlightRestart() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + + await manager.DisposeAsync(); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.False(tunnel.IsActive); + } + + [Fact] + public async Task RestartSshTunnelAsync_InactiveTunnelStartsFreshOwnedGeneration() + { + var (manager, tunnel, factory, tunnelConfig) = await CreateConnectedSshManagerAsync(); + using (manager) + { + tunnel.SimulateExit(); + + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + factory.CreatedClients[1].SimulateHandshake(); + + Assert.True(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.True(tunnel.IsActive); + Assert.Equal(tunnelConfig, tunnel.ActiveConfig); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_UnhealthyOwnedListenerStartsFreshGeneration() + { + var (manager, tunnel, factory, tunnelConfig) = await CreateConnectedSshManagerAsync(); + using (manager) + { + tunnel.OwnedListenerReady = false; + tunnel.BecomeReadyOnStart = true; + + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + factory.CreatedClients[1].SimulateHandshake(); + + Assert.True(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.True(tunnel.IsActive); + Assert.Equal(tunnelConfig, tunnel.ActiveConfig); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_NormalizedConfigMatchesOwnedTunnel() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync( + tunnelUser: " user ", + tunnelHost: " host.example "); + using (manager) + { + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + factory.CreatedClients[1].SimulateHandshake(); + + Assert.True(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.Equal("user", tunnel.ActiveConfig?.User); + Assert.Equal("host.example", tunnel.ActiveConfig?.Host); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_ChangedConfigReplacesPreviouslyOwnedTunnel() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var current = Assert.IsType(_registry.GetActive()); + var updatedTunnel = Assert.IsType(current.SshTunnel) with + { + Host = "replacement.example", + LocalPort = 45679, + }; + _registry.AddOrUpdate(current with { SshTunnel = updatedTunnel }); + + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + factory.CreatedClients[1].SimulateHandshake(); + + Assert.True(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.Equal(updatedTunnel, tunnel.ActiveConfig); + Assert.Equal(updatedTunnel, tunnel.StartedConfigs[^1]); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_TimeoutFailsAndCleansCapturedGeneration() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync( + restartTimeout: TimeSpan.FromMilliseconds(100)); + using (manager) + { + var restarted = await manager.RestartSshTunnelAsync() + .WaitAsync(TimeSpan.FromSeconds(2)); + + Assert.False(restarted); + Assert.Equal(2, factory.CreatedClients.Count); + Assert.False(tunnel.IsActive); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_CleanupStopReceivesBoundedCancellation() + { + var (manager, tunnel, _, _) = await CreateConnectedSshManagerAsync( + restartTimeout: TimeSpan.FromMilliseconds(50), + cleanupTimeout: TimeSpan.FromMilliseconds(50)); + using (manager) + { + CancellationToken cleanupToken = default; + var stopCalls = 0; + tunnel.StopIfOwnedAsyncOverride = token => + { + if (Interlocked.Increment(ref stopCalls) == 1) + { + tunnel.SimulateExit(); + return Task.FromResult(true); + } + + cleanupToken = token; + return Task.FromResult(false); + }; + + Assert.False( + await manager.RestartSshTunnelAsync() + .WaitAsync(TimeSpan.FromSeconds(2))); + Assert.True(cleanupToken.CanBeCanceled); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_CleanupCancellationDoesNotEscapeAsFault() + { + var (manager, tunnel, _, _) = await CreateConnectedSshManagerAsync( + restartTimeout: TimeSpan.FromMilliseconds(50), + cleanupTimeout: TimeSpan.FromMilliseconds(50)); + using (manager) + { + var stopCalls = 0; + tunnel.StopIfOwnedAsyncOverride = async token => + { + if (Interlocked.Increment(ref stopCalls) == 1) + { + tunnel.SimulateExit(); + return true; + } + + await Task.Delay(Timeout.InfiniteTimeSpan, token); + return false; + }; + + Assert.False( + await manager.RestartSshTunnelAsync() + .WaitAsync(TimeSpan.FromSeconds(2))); + Assert.Equal(2, stopCalls); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_CancellationFailsAndCleansCapturedGeneration() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + using (manager) + using (var cts = new CancellationTokenSource()) + { + var restart = manager.RestartSshTunnelAsync(cts.Token); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + cts.Cancel(); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.False(tunnel.IsActive); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_AuthenticationFailureDoesNotReportSuccess() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + factory.CreatedClients[1].SimulateAuthFailed("token mismatch"); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.False(tunnel.IsActive); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_SupersededConnectionDoesNotCleanReplacement() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + + await manager.ReconnectAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 3); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.True(tunnel.IsActive); + Assert.Equal(3, factory.CreatedClients.Count); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_OwnerReplacementFailsWithoutStoppingReplacement() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + tunnel.OwnershipGeneration++; + factory.CreatedClients[1].SimulateHandshake(); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.True(tunnel.IsActive); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_UnownedTunnelFailsWithoutDisconnecting() + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + using (manager) + { + tunnel.StopIfOwnedAsyncOverride = _ => Task.FromResult(false); + + Assert.False(await manager.RestartSshTunnelAsync()); + Assert.Single(factory.CreatedClients); + Assert.Equal(RoleConnectionState.Connected, manager.CurrentSnapshot.OperatorState); + Assert.True(tunnel.IsActive); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_CurrentRecordMutationFails() + { + var (manager, tunnel, factory, tunnelConfig) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + _registry.AddOrUpdate(new GatewayRecord + { + Id = "gw-restart-user", + Url = "wss://test", + SshTunnel = tunnelConfig with { RemotePort = tunnelConfig.RemotePort + 1 }, + SharedGatewayToken = "gateway-token", + }); + factory.CreatedClients[1].SimulateHandshake(); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.False(tunnel.IsActive); + } + } + + [Theory] + [InlineData(true)] + [InlineData(false)] + public async Task RestartSshTunnelAsync_RecordMutationDuringListenerCheckFails( + bool mutateEndpoint) + { + var (manager, tunnel, factory, _) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var original = Assert.IsType(_registry.GetActive()); + var listenerCheckStarted = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var releaseListenerCheck = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var checkCount = 0; + tunnel.OwnedListenerCheckAsync = cancellationToken => + { + if (Interlocked.Increment(ref checkCount) == 2) + { + listenerCheckStarted.TrySetResult(); + return releaseListenerCheck.Task.WaitAsync(cancellationToken); + } + + return Task.FromResult(true); + }; + + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + factory.CreatedClients[1].SimulateHandshake(); + await listenerCheckStarted.Task.WaitAsync(TimeSpan.FromSeconds(2)); + _registry.AddOrUpdate( + mutateEndpoint + ? original with { Url = "wss://replacement.example" } + : original with { SharedGatewayToken = "second" }); + releaseListenerCheck.TrySetResult(true); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.False(tunnel.IsActive); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_ActiveGatewayReplacementWithSameConfigFails() + { + var (manager, tunnel, factory, tunnelConfig) = await CreateConnectedSshManagerAsync(); + using (manager) + { + var restart = manager.RestartSshTunnelAsync(); + await WaitUntilAsync(() => factory.CreatedClients.Count == 2); + _registry.AddOrUpdate(new GatewayRecord + { + Id = "gw-restart-replacement", + Url = "wss://replacement.example", + SshTunnel = tunnelConfig, + SharedGatewayToken = "gateway-token", + }); + _registry.SetActive("gw-restart-replacement"); + factory.CreatedClients[1].SimulateHandshake(); + + Assert.False(await restart.WaitAsync(TimeSpan.FromSeconds(2))); + Assert.False(tunnel.IsActive); + } + } + + [Fact] + public async Task RestartSshTunnelAsync_TimeoutBeforeTransitionLockReturnsPromptly() + { + var (manager, _, _, _) = await CreateConnectedSshManagerAsync( + restartTimeout: TimeSpan.FromMilliseconds(100)); + using (manager) + { + var semaphoreField = typeof(GatewayConnectionManager).GetField( + "_transitionSemaphore", + System.Reflection.BindingFlags.Instance | + System.Reflection.BindingFlags.NonPublic); + var semaphore = Assert.IsType(semaphoreField?.GetValue(manager)); + await semaphore.WaitAsync(); + try + { + Assert.False( + await manager.RestartSshTunnelAsync() + .WaitAsync(TimeSpan.FromSeconds(2))); + } + finally + { + semaphore.Release(); + } + } + } + + private async Task<( + GatewayConnectionManager Manager, + CountingTunnelManager Tunnel, + MockClientFactory Factory, + SshTunnelConfig TunnelConfig)> CreateConnectedSshManagerAsync( + TimeSpan? restartTimeout = null, + TimeSpan? cleanupTimeout = null, + string tunnelUser = "user", + string tunnelHost = "host.example") + { + var tunnelConfig = new SshTunnelConfig( + tunnelUser, + tunnelHost, + 18789, + 45678); + _registry.AddOrUpdate(new GatewayRecord + { + Id = "gw-restart-user", + Url = "wss://test", + SharedGatewayToken = "gateway-token", + SshTunnel = tunnelConfig, + }); + _registry.SetActive("gw-restart-user"); + _resolver.OperatorCredential = new GatewayCredential("gateway-token", false, "test"); + var tunnel = new CountingTunnelManager(); + var factory = new MockClientFactory(); + var manager = new GatewayConnectionManager( + _resolver, + factory, + _registry, + NullLogger.Instance, + tunnelManager: tunnel, + manualSshRestartTimeout: restartTimeout, + manualSshRestartCleanupTimeout: cleanupTimeout); + + await manager.ConnectAsync("gw-restart-user"); + factory.CreatedClients[0].SimulateHandshake(); + await WaitUntilAsync( + () => manager.CurrentSnapshot.OperatorState == RoleConnectionState.Connected); + return (manager, tunnel, factory, tunnelConfig); + } + [Fact] public async Task PassiveGatewayRestart_ReusesLiveClientsAndPreservesDurableIdentity() { @@ -1183,6 +1638,96 @@ public async Task SshReconnectAuthorization_RequiresCurrentOwnedListener() Assert.False(authorization.Allowed); Assert.Equal(GatewayErrorKind.LocalPortConflict, authorization.FailureKind); Assert.Contains("credentials were not sent", authorization.Detail); + + lifecycle.SimulateStatusChanged(ConnectionStatus.Error); + await WaitUntilAsync(() => + manager.CurrentSnapshot.OperatorState == RoleConnectionState.Error); + Assert.Equal( + GatewayErrorKind.LocalPortConflict, + manager.CurrentSnapshot.OperatorErrorKind); + Assert.Contains( + "credentials were not sent", + manager.CurrentSnapshot.OperatorError); + } + + [Theory] + [InlineData(true)] + [InlineData(false)] + public async Task SshHandshakeAuthorization_RejectsEndpointOrCredentialMutation( + bool mutateEndpoint) + { + var ssh = new SshTunnelConfig("user", "host.example", 18789, 45678); + var original = new GatewayRecord + { + Id = "gw-ssh-mutation", + Url = "wss://remote.example", + SharedGatewayToken = "first", + SshTunnel = ssh, + }; + _registry.AddOrUpdate(original); + _registry.SetActive(original.Id); + _resolver.OperatorCredential = new GatewayCredential( + "first", + IsBootstrapToken: false, + CredentialResolver.SourceSharedGatewayToken); + var tunnel = new CountingTunnelManager(); + using var manager = new GatewayConnectionManager( + _resolver, + _factory, + _registry, + NullLogger.Instance, + tunnelManager: tunnel); + + await manager.ConnectAsync(original.Id); + var authorizeHandshake = Assert.IsType< + Func>>( + Assert.Single(_factory.CreatedClients).DataClient.HandshakeAuthorizationAsync); + _registry.AddOrUpdate( + mutateEndpoint + ? original with { Url = "wss://replacement.example" } + : original with { SharedGatewayToken = "second" }); + + var authorization = await authorizeHandshake(CancellationToken.None); + + Assert.False(authorization.Allowed); + Assert.Equal(GatewayErrorKind.LocalPortConflict, authorization.FailureKind); + Assert.Contains("changed before", authorization.Detail); + } + + [Fact] + public async Task SshOwnershipFailure_DoesNotStopReplacementTunnelGeneration() + { + var ssh = new SshTunnelConfig("user", "host.example", 18789, 45678); + _registry.AddOrUpdate(new GatewayRecord + { + Id = "gw-ssh-owner-replacement", + Url = "wss://remote.example", + SharedGatewayToken = "first", + SshTunnel = ssh, + }); + _registry.SetActive("gw-ssh-owner-replacement"); + _resolver.OperatorCredential = new GatewayCredential( + "first", + IsBootstrapToken: false, + CredentialResolver.SourceSharedGatewayToken); + var tunnel = new CountingTunnelManager(); + tunnel.OwnedListenerCheckAsync = _ => + { + tunnel.OwnershipGeneration++; + return Task.FromResult(false); + }; + using var manager = new GatewayConnectionManager( + _resolver, + _factory, + _registry, + NullLogger.Instance, + tunnelManager: tunnel); + + await manager.ConnectAsync("gw-ssh-owner-replacement"); + + Assert.True(tunnel.IsActive); + Assert.Equal(0, tunnel.StopCount); + Assert.Empty(_factory.CreatedClients); } [Fact] @@ -1279,12 +1824,12 @@ public async Task SshInitialHandshakeAuthorization_RechecksOwnershipAfterTranspo await manager.ConnectAsync("gw-ssh"); var lifecycle = Assert.Single(_factory.CreatedClients); Assert.NotNull(lifecycle.DataClient.HandshakeAuthorizationAsync); - var failure = new TaskCompletionSource( + var failure = new TaskCompletionSource( TaskCreationOptions.RunContinuationsAsynchronously); var errorStatus = new TaskCompletionSource( TaskCreationOptions.RunContinuationsAsynchronously); - lifecycle.DataClient.AuthenticationFailed += (_, message) => - failure.TrySetResult(message); + lifecycle.DataClient.ConnectionFailure += (_, kind) => + failure.TrySetResult(kind); lifecycle.DataClient.StatusChanged += (_, status) => { if (status == ConnectionStatus.Error) @@ -1296,11 +1841,11 @@ public async Task SshInitialHandshakeAuthorization_RechecksOwnershipAfterTranspo tunnel.OwnedListenerReady = false; lifecycle.SimulateConnectChallenge(); - var failureMessage = await failure.Task.WaitAsync(TimeSpan.FromSeconds(2)); + var failureKind = await failure.Task.WaitAsync(TimeSpan.FromSeconds(2)); await errorStatus.Task.WaitAsync(TimeSpan.FromSeconds(2)); Assert.Equal(checksBeforeChallenge + 1, tunnel.OwnedListenerCheckCount); - Assert.Contains("credentials were not sent", failureMessage); + Assert.Equal(GatewayErrorKind.LocalPortConflict, failureKind); } [Fact] @@ -1347,6 +1892,95 @@ public async Task SshNodeInitialHandshakeAuthorization_RequiresCurrentOwnedListe Assert.Contains("credentials were not sent", manager.CurrentSnapshot.NodeError); } + [Fact] + public async Task SshOperatorHandshakeAuthorization_TimesOutAsTunnelFailure() + { + var ssh = new SshTunnelConfig("user", "host.example", 18789, 45678); + _registry.AddOrUpdate(new GatewayRecord + { + Id = "gw-ssh-timeout", + Url = "wss://remote.example", + SharedGatewayToken = "gateway-token", + SshTunnel = ssh, + }); + _registry.SetActive("gw-ssh-timeout"); + _resolver.OperatorCredential = new GatewayCredential( + "gateway-token", + IsBootstrapToken: false, + CredentialResolver.SourceSharedGatewayToken); + var tunnel = new CountingTunnelManager(); + using var manager = new GatewayConnectionManager( + _resolver, + _factory, + _registry, + NullLogger.Instance, + tunnelManager: tunnel, + credentialHandoffTimeout: TimeSpan.FromMilliseconds(50)); + + await manager.ConnectAsync("gw-ssh-timeout"); + var authorizeHandshake = Assert.IsType< + Func>>( + Assert.Single(_factory.CreatedClients).DataClient.HandshakeAuthorizationAsync); + tunnel.OwnedListenerCheckAsync = async ct => + { + await Task.Delay(Timeout.InfiniteTimeSpan, ct); + return true; + }; + + var authorization = await authorizeHandshake(CancellationToken.None) + .WaitAsync(TimeSpan.FromSeconds(2)); + + Assert.False(authorization.Allowed); + Assert.Equal(GatewayErrorKind.Network, authorization.FailureKind); + Assert.Contains("Timed out", authorization.Detail); + } + + [Fact] + public async Task SshNodeHandshakeAuthorization_TimesOutAsTunnelFailure() + { + var ssh = new SshTunnelConfig("user", "host.example", 18789, 45678); + _registry.AddOrUpdate(new GatewayRecord + { + Id = "gw-node-ssh-timeout", + Url = "wss://remote.example", + SharedGatewayToken = "gateway-token", + SshTunnel = ssh, + }); + _registry.SetActive("gw-node-ssh-timeout"); + _resolver.NodeCredential = new GatewayCredential( + "gateway-token", + IsBootstrapToken: false, + CredentialResolver.SourceSharedGatewayToken); + var tunnel = new CountingTunnelManager(); + var node = new CountingNodeConnector(); + using var manager = new GatewayConnectionManager( + _resolver, + _factory, + _registry, + NullLogger.Instance, + nodeConnector: node, + isNodeEnabled: () => true, + tunnelManager: tunnel, + credentialHandoffTimeout: TimeSpan.FromMilliseconds(50)); + + await manager.ConnectNodeOnlyAsync("gw-node-ssh-timeout"); + var authorizeHandshake = Assert.IsType< + Func>>( + node.HandshakeAuthorizationAsync); + tunnel.OwnedListenerCheckAsync = async ct => + { + await Task.Delay(Timeout.InfiniteTimeSpan, ct); + return true; + }; + + var authorization = await authorizeHandshake(CancellationToken.None) + .WaitAsync(TimeSpan.FromSeconds(2)); + + Assert.False(authorization.Allowed); + Assert.Equal(GatewayErrorKind.Network, authorization.FailureKind); + Assert.Contains("Timed out", authorization.Detail); + } + [Fact] public async Task ConnectWithSharedTokenAsync_ExactActiveSshConfigUsesIsolatedValidationTunnel() { @@ -4063,8 +4697,11 @@ public MockLifecycle(string url, string identityPath) public Task ConnectAsync(CancellationToken ct) => Task.CompletedTask; - public void SimulateStatusChanged(ConnectionStatus status) => + public void SimulateStatusChanged(ConnectionStatus status) + { + _client.SetConnected(status == ConnectionStatus.Connected); StatusChanged?.Invoke(this, status); + } public void SimulateAuthFailed(string msg) => AuthenticationFailed?.Invoke(this, msg); @@ -4092,9 +4729,15 @@ public void SimulateDeviceTokenReceived(string token, string role, string[]? sco private sealed class MockGatewayClient : OpenClawGatewayClient { + private bool _isConnected = true; + public MockGatewayClient(string url, string identityPath) : base(url, "mock-token", NullLogger.Instance, identityPath: identityPath) { } + public override bool IsConnectedToGateway => _isConnected; + + public void SetConnected(bool connected) => _isConnected = connected; + public void SimulateTransportConnected() => RaiseTransportConnected(); @@ -4663,21 +5306,26 @@ private sealed class CountingTunnelManager : ISshTunnelManager public string? LocalTunnelUrl { get; private set; } public bool RestartPending { get; set; } public bool OwnedListenerReady { get; set; } = true; + public bool BecomeReadyOnStart { get; set; } + public Func>? OwnedListenerCheckAsync { get; set; } + public Func>? StopIfOwnedAsyncOverride { get; set; } public int OwnedListenerCheckCount { get; private set; } public bool IsRestartPending(SshTunnelExit tunnelExit) => RestartPending; - public Task IsOwnedListenerReadyAsync( + public async Task IsOwnedListenerReadyAsync( SshTunnelConfig config, int destinationPort, CancellationToken ct) { ct.ThrowIfCancellationRequested(); OwnedListenerCheckCount++; - return Task.FromResult( - OwnedListenerReady && + var ready = OwnedListenerCheckAsync is null + ? OwnedListenerReady + : await OwnedListenerCheckAsync(ct); + return ready && IsActive && - ActiveConfig == config && - IsConfiguredForward(config, destinationPort)); + ActiveConfig == Normalize(config) && + IsConfiguredForward(config, destinationPort); } public Task StartAsync(SshTunnelConfig config, CancellationToken ct) @@ -4685,6 +5333,11 @@ public Task StartAsync(SshTunnelConfig config, CancellationToken ct) ct.ThrowIfCancellationRequested(); StartCount++; OwnershipGeneration++; + var normalizedConfig = config with + { + User = config.User.Trim(), + Host = config.Host.Trim(), + }; StartedConfigs.Add(config); if (FailStart || config == FailForConfig) { @@ -4694,11 +5347,34 @@ public Task StartAsync(SshTunnelConfig config, CancellationToken ct) throw new InvalidOperationException("tunnel failed"); } IsActive = true; - LastConfig = config; + LastConfig = normalizedConfig; LocalTunnelUrl = $"ws://localhost:{config.LocalPort}"; + if (BecomeReadyOnStart) + OwnedListenerReady = true; return Task.FromResult(LocalTunnelUrl); } + public async Task StartOwnedAsync( + SshTunnelConfig config, + CancellationToken ct) + { + var url = await StartAsync(config, ct); + return new SshTunnelStartResult(url, Normalize(config), OwnershipGeneration); + } + + public void SimulateExit() + { + IsActive = false; + LocalTunnelUrl = null; + } + + private static SshTunnelConfig Normalize(SshTunnelConfig config) => + config with + { + User = config.User.Trim(), + Host = config.Host.Trim(), + }; + public Task StopAsync() { StopCount++; @@ -4708,16 +5384,40 @@ public Task StopAsync() return Task.CompletedTask; } + public Task StopIfOwnedAsync( + SshTunnelConfig config, + long ownershipGeneration, + CancellationToken ct) + { + ct.ThrowIfCancellationRequested(); + if (StopIfOwnedAsyncOverride is not null) + return StopIfOwnedAsyncOverride(ct); + if (!IsActive || + ActiveConfig != Normalize(config) || + OwnershipGeneration != ownershipGeneration) + { + return Task.FromResult(false); + } + + StopCount++; + OwnershipGeneration++; + IsActive = false; + LocalTunnelUrl = null; + return Task.FromResult(true); + } + public void Dispose() => IsDisposed = true; } private sealed class BlockingTunnelManager : ISshTunnelManager { private SshTunnelConfig? _activeConfig; + private long _ownershipGeneration; public TaskCompletionSource Started { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); public TaskCompletionSource AllowStart { get; } = new(TaskCreationOptions.RunContinuationsAsynchronously); public bool IsActive => _activeConfig is not null; + public long OwnershipGeneration => _ownershipGeneration; public SshTunnelConfig? ActiveConfig => _activeConfig; public string? LocalTunnelUrl => _activeConfig is null ? null : $"ws://localhost:{_activeConfig.LocalPort}"; public bool RestartPending { get; set; } @@ -4740,20 +5440,48 @@ public async Task StartAsync(SshTunnelConfig config, CancellationToken c Started.SetResult(true); await AllowStart.Task.WaitAsync(ct); _activeConfig = config; + _ownershipGeneration++; return $"ws://localhost:{config.LocalPort}"; } + public async Task StartOwnedAsync( + SshTunnelConfig config, + CancellationToken ct) + { + var url = await StartAsync(config, ct); + return new SshTunnelStartResult(url, config, OwnershipGeneration); + } + public Task StopAsync() { _activeConfig = null; return Task.CompletedTask; } + + public Task StopIfOwnedAsync( + SshTunnelConfig config, + long ownershipGeneration, + CancellationToken ct) + { + ct.ThrowIfCancellationRequested(); + if (_activeConfig != config || + _ownershipGeneration != ownershipGeneration) + { + return Task.FromResult(false); + } + + _activeConfig = null; + _ownershipGeneration++; + return Task.FromResult(true); + } + public void Dispose() { } } private sealed class FailingTunnelManager : ISshTunnelManager { public bool IsActive => false; + public long OwnershipGeneration => 0; public SshTunnelConfig? ActiveConfig => null; public string? LocalTunnelUrl => null; @@ -4770,8 +5498,18 @@ public Task IsOwnedListenerReadyAsync( public Task StartAsync(SshTunnelConfig config, CancellationToken ct) => throw new InvalidOperationException("tunnel failed"); + public Task StartOwnedAsync( + SshTunnelConfig config, + CancellationToken ct) => + throw new InvalidOperationException("tunnel failed"); + public Task StopAsync() => Task.CompletedTask; + public Task StopIfOwnedAsync( + SshTunnelConfig config, + long ownershipGeneration, + CancellationToken ct) => Task.FromResult(false); + public void Dispose() { } } diff --git a/tests/OpenClaw.Connection.Tests/SetupCodeFlowTests.cs b/tests/OpenClaw.Connection.Tests/SetupCodeFlowTests.cs index e289b9b16..0a590a278 100644 --- a/tests/OpenClaw.Connection.Tests/SetupCodeFlowTests.cs +++ b/tests/OpenClaw.Connection.Tests/SetupCodeFlowTests.cs @@ -485,6 +485,14 @@ public Task StartAsync(SshTunnelConfig config, CancellationToken ct) return Task.FromResult(LocalTunnelUrl); } + public async Task StartOwnedAsync( + SshTunnelConfig config, + CancellationToken ct) + { + var url = await StartAsync(config, ct); + return new SshTunnelStartResult(url, config, OwnershipGeneration); + } + public Task StopAsync() { OwnershipGeneration++; @@ -494,6 +502,26 @@ public Task StopAsync() return Task.CompletedTask; } + public Task StopIfOwnedAsync( + SshTunnelConfig config, + long ownershipGeneration, + CancellationToken ct) + { + ct.ThrowIfCancellationRequested(); + if (!IsActive || + ActiveConfig != config || + OwnershipGeneration != ownershipGeneration) + { + return Task.FromResult(false); + } + + OwnershipGeneration++; + IsActive = false; + ActiveConfig = null; + LocalTunnelUrl = null; + return Task.FromResult(true); + } + public void Dispose() { } diff --git a/tests/OpenClaw.Connection.Tests/SshTunnelServiceTests.cs b/tests/OpenClaw.Connection.Tests/SshTunnelServiceTests.cs index 704349544..0125d270e 100644 --- a/tests/OpenClaw.Connection.Tests/SshTunnelServiceTests.cs +++ b/tests/OpenClaw.Connection.Tests/SshTunnelServiceTests.cs @@ -193,6 +193,47 @@ public void Dispose_DoesNotThrow() Assert.Null(ex); } + [Fact] + public async Task StopIfOwnedAsync_CancellationInterruptsOperationLockWait() + { + using var service = new SshTunnelService(NullLogger.Instance); + var operationLockField = typeof(SshTunnelService).GetField( + "_operationLock", + System.Reflection.BindingFlags.Instance | + System.Reflection.BindingFlags.NonPublic); + var operationLock = Assert.IsType(operationLockField?.GetValue(service)); + using var cts = new CancellationTokenSource(); + using var releaseLock = new ManualResetEventSlim(); + var lockEntered = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var lockHolder = Task.Run(() => + { + lock (operationLock) + { + lockEntered.TrySetResult(); + releaseLock.Wait(); + } + }); + await lockEntered.Task.WaitAsync(TimeSpan.FromSeconds(2)); + try + { + var stop = Task.Run( + async () => await service.StopIfOwnedAsync( + new SshTunnelConfig("user", "host", 18789, 45678), + ownershipGeneration: 1, + cts.Token)); + cts.Cancel(); + + await Assert.ThrowsAnyAsync( + () => stop.WaitAsync(TimeSpan.FromSeconds(2))); + } + finally + { + releaseLock.Set(); + await lockHolder.WaitAsync(TimeSpan.FromSeconds(2)); + } + } + [Fact] public void EnsurePortIsUnoccupied_RejectsExistingListener() { diff --git a/tests/OpenClaw.E2ETests/Setup/SshOwnershipAdversarialProofTests.cs b/tests/OpenClaw.E2ETests/Setup/SshOwnershipAdversarialProofTests.cs index f3c459890..d22879418 100644 --- a/tests/OpenClaw.E2ETests/Setup/SshOwnershipAdversarialProofTests.cs +++ b/tests/OpenClaw.E2ETests/Setup/SshOwnershipAdversarialProofTests.cs @@ -1,11 +1,8 @@ using System.Diagnostics; -using System.Drawing; using System.Net; -using System.Net.Sockets; using System.Net.WebSockets; using System.Text; using System.Text.Json; -using System.Text.Json.Nodes; using OpenClaw.Connection; using OpenClaw.Shared; @@ -27,315 +24,88 @@ public SshOwnershipAdversarialProofTests(E2ESetupFixture fixture) public async Task UnownedListenerIsRejectedThenOwnedTunnelRecoversWithoutRepairing() { var proofDir = Path.Combine(_fixture.ArtifactDir, "pr1076-proof"); - var profileDir = Path.Combine(proofDir, "profile"); - var sshDir = Path.Combine(profileDir, ".ssh"); - var sshConfigPath = Path.Combine(profileDir, "ssh_config"); - Directory.CreateDirectory(sshDir); - var gatewayPath = Path.Combine(_fixture.DataDir, "gateways.json"); - var settingsPath = Path.Combine(_fixture.DataDir, "settings.json"); - var originalGatewayBytes = File.ReadAllBytes(gatewayPath); - var originalSettingsBytes = File.ReadAllBytes(settingsPath); - using var sshPortLease = MirroredWslPortLease.Acquire(); - var sshPort = sshPortLease.Port; - var tunnelPort = AllocateFreeForwardPortPair(); - var screenshotDegraded = Path.Combine(proofDir, "04-connection-degraded.png"); - var captureUiProof = string.Equals( - Environment.GetEnvironmentVariable("OPENCLAW_CAPTURE_UI_PROOF"), - "1", - StringComparison.Ordinal); - TcpListener? adversary = null; - (string? Operator, string? Node) beforeTokens = (null, null); + Directory.CreateDirectory(proofDir); + var registryDir = Path.Combine( + Path.GetTempPath(), + $"openclaw-owned-listener-recovery-proof-{Guid.NewGuid():N}"); + Directory.CreateDirectory(registryDir); + await using var server = new InitialHandshakeChallengeServer(); try { - beforeTokens = ReadRoleTokens(); - Assert.False(string.IsNullOrWhiteSpace(beforeTokens.Operator)); - Assert.False(string.IsNullOrWhiteSpace(beforeTokens.Node)); - - await _fixture.StopTrayAsync(); - await ConfigureProofSshAsync(profileDir, sshDir, sshPort); - var proofSshd = await StartProofSshdAsync(profileDir, sshDir, sshPort); - WriteObject("00-proof-sshd.json", new - { - unitName = proofSshd.UnitName, - processId = proofSshd.ProcessId, - executablePath = proofSshd.ExecutablePath, - commandLine = proofSshd.CommandLine, - }); - var identityFile = Path.Combine(sshDir, "id_ed25519").Replace('\\', '/'); - await File.WriteAllTextAsync( - sshConfigPath, - $""" - Host * - BatchMode yes - IdentitiesOnly yes - IdentityFile "{identityFile}" - UserKnownHostsFile NUL - StrictHostKeyChecking no - ProxyCommand wsl.exe -d {_fixture.DistroName} -- nc 127.0.0.1 {sshPort} - """); - PatchActiveGateway(tunnelPort, proofSshd.HostAddress, sshPort, browserControlPort: null); - _fixture.SetTrayEnvironmentVariable("HOME", profileDir); - _fixture.SetTrayEnvironmentVariable("USERPROFILE", profileDir); - _fixture.SetTrayEnvironmentVariable("OPENCLAW_E2E_SSH_CONFIG_FILE", sshConfigPath); - if (captureUiProof) + var registry = new GatewayRegistry(registryDir); + var tunnelConfig = new SshTunnelConfig( + "proof-user", + "proof-host", + RemotePort: server.Port, + LocalPort: server.Port); + var record = new GatewayRecord { - _fixture.SetTrayEnvironmentVariable("OPENCLAW_VISUAL_TEST", "1"); - _fixture.SetTrayEnvironmentVariable("OPENCLAW_VISUAL_TEST_DIR", proofDir); - } - await _fixture.StartTrayAsync(); - - var ownedSnapshot = await WaitForListenerSnapshotAsync( - listeners => listeners.Any(listener => - listener.Port == tunnelPort && - string.Equals(listener.ProcessName, "ssh", StringComparison.OrdinalIgnoreCase)), - TimeSpan.FromSeconds(30)); - var ownedListeners = ownedSnapshot.Listeners - .Where(listener => listener.Port == tunnelPort) - .ToArray(); - var sshListeners = ownedListeners - .Where(listener => - string.Equals(listener.ProcessName, "ssh", StringComparison.OrdinalIgnoreCase)) - .ToArray(); - Assert.NotEmpty(sshListeners); - var sshProcessId = Assert.Single(sshListeners.Select(listener => listener.ProcessId).Distinct()); - var sshListener = Assert.Single( - sshListeners, - listener => listener.Address.Equals(IPAddress.Loopback)); - using var ready = await WaitForReadyStatusAsync(TimeSpan.FromSeconds(30)); - AssertReady(ready.RootElement); - WriteJson("01-valid-ready.json", ready.RootElement); - WriteObject("02-owned-listener.json", new + Id = "owned-listener-recovery-proof", + Url = "wss://proof.invalid", + SharedGatewayToken = "synthetic-proof-credential", + SshTunnel = tunnelConfig, + }; + registry.AddOrUpdate(record); + registry.SetActive(record.Id); + var tunnel = new InitialHandshakeRaceTunnelManager(server.WebSocketUrl) { - tunnelPort, - owned = true, - listenerCount = ownedListeners.Length, - processName = sshListener.ProcessName, - processId = sshProcessId, - addresses = sshListeners.Select(listener => listener.Address.ToString()).ToArray(), - }); + OwnedListenerReady = false, + }; + using var manager = new GatewayConnectionManager( + new CredentialResolver(DeviceIdentityFileReader.Instance), + new GatewayClientFactory(), + registry, + NullLogger.Instance, + tunnelManager: tunnel); - adversary = new TcpListener(IPAddress.Parse("127.0.0.2"), tunnelPort); - adversary.Start(); - var competingSnapshot = await WaitForListenerSnapshotAsync( - listeners => - listeners.Any(listener => - listener.Port == tunnelPort && - listener.ProcessId == Environment.ProcessId) && - listeners.Any(listener => - listener.Port == tunnelPort && - listener.ProcessId == sshListener.ProcessId), - TimeSpan.FromSeconds(30)); - var competingListeners = competingSnapshot.Listeners - .Where(listener => listener.Port == tunnelPort) - .ToArray(); - Assert.Contains(competingListeners, listener => listener.ProcessId == Environment.ProcessId); - Assert.Contains(competingListeners, listener => listener.ProcessId == sshListener.ProcessId); - WriteObject("03-competing-listener.json", new - { - tunnelPort, - listenerCount = competingListeners.Length, - unrelatedListenerPresent = true, - ownedProcessId = sshListener.ProcessId, - unrelatedProcessId = Environment.ProcessId, - }); + await manager.ConnectAsync(record.Id); + var rejected = await WaitForOperatorErrorAsync(manager, TimeSpan.FromSeconds(5)); + var checksAfterRejection = tunnel.OwnedListenerCheckCount; - using (var reconnect = await _fixture.Client!.CallToolExpectSuccessAsync( - "app.connection.reconnectNode")) - { - Assert.True(reconnect.RootElement.GetProperty("reconnected").GetBoolean()); - } - using var degraded = await WaitForStatusAsync( - status => - status.GetProperty("overallState").GetString() == "Degraded" && - status.GetProperty("nodeState").GetString() == "Error" && - status.GetProperty("nodeError").GetString()?.Contains( - "credentials were not sent", - StringComparison.OrdinalIgnoreCase) == true, - TimeSpan.FromSeconds(45)); - adversary.Stop(); - adversary = null; + Assert.False(server.Accepted.IsCompleted); Assert.Contains( "credentials were not sent", - degraded.RootElement.GetProperty("nodeError").GetString(), + rejected.OperatorError, StringComparison.OrdinalIgnoreCase); - WriteJson("04-degraded-status.json", degraded.RootElement); - using (var connectionStatus = await _fixture.Client!.CallToolExpectSuccessAsync( - "app.connection.status")) - { - WriteJson("04-connection-diagnostics.json", connectionStatus.RootElement); - } - if (captureUiProof) - await NavigateAndCaptureAsync("connection", screenshotDegraded); - - using (var reconnect = await _fixture.Client!.CallToolExpectSuccessAsync( - "app.connection.reconnectNode")) - { - Assert.True(reconnect.RootElement.GetProperty("reconnected").GetBoolean()); - } - await _fixture.WaitForConnectionReady(TimeSpan.FromSeconds(90)); - await _fixture.WaitForNodeListReady(TimeSpan.FromSeconds(60)); - using var recovered = await WaitForReadyStatusAsync(TimeSpan.FromSeconds(30)); - AssertReady(recovered.RootElement); - WriteJson("05-recovered-ready.json", recovered.RootElement); - - var afterTokens = ReadRoleTokens(); - Assert.Equal(beforeTokens.Operator, afterTokens.Operator); - Assert.Equal(beforeTokens.Node, afterTokens.Node); - using (var approvals = await WaitForConnectedPendingApprovalsAsync( - TimeSpan.FromSeconds(30))) - { - Assert.True(approvals.RootElement.GetProperty("connected").GetBoolean()); - Assert.Equal(0, approvals.RootElement.GetProperty("totalPending").GetInt32()); - Assert.Empty(approvals.RootElement.GetProperty("devicePending").EnumerateArray()); - Assert.Empty(approvals.RootElement.GetProperty("nodePending").EnumerateArray()); - WriteJson("05-pending-approvals.json", approvals.RootElement); - } - - await _fixture.StopTrayAsync(); - var identityDir = _fixture.ReadActiveGatewayCredentialState().IdentityDir; - var clear = DeviceIdentityStore.BeginTransactionalTokenClear(identityDir); - Assert.True(clear.Success, clear.Error); - Assert.NotNull(clear.Transaction); - var newerOperatorToken = $"proof-operator-{Guid.NewGuid():N}"; - var newerNodeToken = $"proof-node-{Guid.NewGuid():N}"; - var lateWriter = new DeviceIdentity(identityDir); - lateWriter.Initialize(); - lateWriter.StoreDeviceTokenForRole("operator", newerOperatorToken); - lateWriter.StoreDeviceTokenForRole("node", newerNodeToken); - var restore = DeviceIdentityStore.RestoreTransactionalTokenClear(clear.Transaction!); - Assert.Equal(DeviceTokenRestoreOutcome.Superseded, restore.Outcome); - var lateWriterTokens = ReadRoleTokens(); - Assert.Equal(newerOperatorToken, lateWriterTokens.Operator); - Assert.Equal(newerNodeToken, lateWriterTokens.Node); - Assert.NotEqual(beforeTokens.Operator, lateWriterTokens.Operator); - Assert.NotEqual(beforeTokens.Node, lateWriterTokens.Node); - WriteObject("07-late-writer-rollback.json", new - { - restoreOutcome = restore.Outcome.ToString(), - newerOperatorCredentialPreserved = true, - newerNodeCredentialPreserved = true, - originalOperatorCredentialWasNotRestored = true, - originalNodeCredentialWasNotRestored = true, - }); + Assert.False(tunnel.IsActive); - WriteObject("proof-summary.json", new - { - head = ResolveHeadSha(), - distro = _fixture.DistroName, - gatewayPort = _fixture.GatewayPort, - sshPort, - tunnelPort, - ambiguousListenerOwnershipRejectedBeforeCredentialSend = true, - recoveredReady = true, - sameOperatorCredential = true, - sameNodeCredential = true, - lateWriterWonRollback = true, - degradedScreenshot = captureUiProof && File.Exists(screenshotDegraded), - }); - WriteRedactedTrayLog(); - } - finally - { - adversary?.Stop(); - await _fixture.StopTrayAsync(); - File.WriteAllBytes(gatewayPath, originalGatewayBytes); - File.WriteAllBytes(settingsPath, originalSettingsBytes); - if (!string.IsNullOrWhiteSpace(beforeTokens.Operator) && - !string.IsNullOrWhiteSpace(beforeTokens.Node)) - { - var identityDir = _fixture.ReadActiveGatewayCredentialState().IdentityDir; - var originalIdentity = new DeviceIdentity(identityDir); - originalIdentity.Initialize(); - originalIdentity.StoreDeviceTokenForRole("operator", beforeTokens.Operator); - originalIdentity.StoreDeviceTokenForRole("node", beforeTokens.Node); - } - _fixture.RemoveTrayEnvironmentVariable("HOME"); - _fixture.RemoveTrayEnvironmentVariable("USERPROFILE"); - _fixture.RemoveTrayEnvironmentVariable("OPENCLAW_E2E_SSH_CONFIG_FILE"); - _fixture.RemoveTrayEnvironmentVariable("OPENCLAW_VISUAL_TEST"); - _fixture.RemoveTrayEnvironmentVariable("OPENCLAW_VISUAL_TEST_DIR"); - var sshdUnit = $"openclaw-pr1076-sshd-{sshPort}.service"; - await _fixture.RunInWslAsync( - $"systemctl stop '{sshdUnit}' 2>/dev/null || true; " + - $"systemctl reset-failed '{sshdUnit}' 2>/dev/null || true", - TimeSpan.FromSeconds(15), - inputViaStdin: true, - user: "root"); - try { Directory.Delete(profileDir, recursive: true); } catch { } - await _fixture.StartTrayAsync(); - } - - return; - - (string? Operator, string? Node) ReadRoleTokens() - { - var identityDir = _fixture.ReadActiveGatewayCredentialState().IdentityDir; - using var document = JsonDocument.Parse( - File.ReadAllText(Path.Combine(identityDir, "device-key-ed25519.json"))); - return ( - ReadString(document.RootElement, "DeviceToken"), - ReadString(document.RootElement, "NodeDeviceToken")); - } - - void PatchActiveGateway( - int localTunnelPort, - string sshHost, - int localSshPort, - int? browserControlPort) - { - var root = JsonNode.Parse(File.ReadAllText(gatewayPath))!.AsObject(); - var activeId = root["activeId"]!.GetValue(); - var records = root["gateways"]!.AsArray(); - var active = records - .Select(node => node!.AsObject()) - .Single(record => record["id"]!.GetValue() == activeId); - active["sshTunnel"] = JsonSerializer.SerializeToNode( - new - { - user = "openclaw", - host = sshHost, - remotePort = _fixture.GatewayPort, - localPort = localTunnelPort, - includeBrowserProxyForward = true, - sshPort = localSshPort, - }); - if (browserControlPort.HasValue) - active["browserControlPort"] = browserControlPort.Value; - else - active.Remove("browserControlPort"); - File.WriteAllText( - gatewayPath, - root.ToJsonString(new JsonSerializerOptions { WriteIndented = true })); - } + tunnel.OwnedListenerReady = true; + await manager.ConnectAsync(record.Id); + await server.Accepted.WaitAsync(TimeSpan.FromSeconds(10)); + var connectFrames = await server.CompleteHandshakeAsync(TimeSpan.FromSeconds(5)); + var recovered = await WaitForOperatorConnectedAsync(manager, TimeSpan.FromSeconds(5)); + + Assert.Equal(1, connectFrames); + Assert.True(tunnel.OwnedListenerCheckCount > checksAfterRejection); + Assert.True(tunnel.IsActive); + Assert.Equal(RoleConnectionState.Connected, recovered.OperatorState); + var current = Assert.IsType(registry.GetById(record.Id)); + Assert.Equal(record.Id, current.Id); + Assert.Equal(record.Url, current.Url); + Assert.Equal(record.SharedGatewayToken, current.SharedGatewayToken); + Assert.Equal(record.SshTunnel, current.SshTunnel); - void WriteJson(string fileName, JsonElement element) => File.WriteAllText( - Path.Combine(proofDir, fileName), + Path.Combine(proofDir, "owned-listener-recovery.json"), JsonSerializer.Serialize( - JsonSerializer.Deserialize(element.GetRawText()), + new + { + head = ResolveHeadSha(), + unownedListenerRejectedBeforeWebSocketConnect = true, + rejection = rejected.OperatorError, + ownershipChecksAfterRejection = checksAfterRejection, + ownershipChecksAfterRecovery = tunnel.OwnedListenerCheckCount, + credentialBearingConnectFramesReceived = connectFrames, + recoveredOperatorState = recovered.OperatorState.ToString(), + gatewayCredentialUnchanged = true, + tunnelConfigurationUnchanged = true, + }, new JsonSerializerOptions { WriteIndented = true })); - - void WriteObject(string fileName, object value) => - File.WriteAllText( - Path.Combine(proofDir, fileName), - JsonSerializer.Serialize(value, new JsonSerializerOptions { WriteIndented = true })); - - void WriteRedactedTrayLog() + } + finally { - var logPath = Path.Combine(_fixture.DataDir, "openclaw-tray.log"); - if (!File.Exists(logPath)) - return; - var selected = File.ReadLines(logPath) - .Where(line => - line.Contains("listener", StringComparison.OrdinalIgnoreCase) || - line.Contains("credential", StringComparison.OrdinalIgnoreCase) || - line.Contains("tunnel", StringComparison.OrdinalIgnoreCase) || - line.Contains("Degraded", StringComparison.OrdinalIgnoreCase)) - .TakeLast(200); - File.WriteAllLines( - Path.Combine(proofDir, "selected-tray-log.redacted.txt"), - selected.Select(TokenSanitizer.SanitizeLogMessage)); + try { Directory.Delete(registryDir, recursive: true); } catch (IOException) { } } } @@ -491,168 +261,6 @@ public async Task InitialNodeHandshakeListenerReplacementWithholdsCredentialFram } } - private async Task ConfigureProofSshAsync( - string profileDir, - string sshDir, - int sshPort) - { - var keyPath = Path.Combine(sshDir, "id_ed25519"); - var keygen = await RunProcessAsync( - "ssh-keygen.exe", - ["-q", "-t", "ed25519", "-N", "", "-f", keyPath]); - Assert.Equal(0, keygen.ExitCode); - - var install = await _fixture.RunInWslAsync( - "set -e; export DEBIAN_FRONTEND=noninteractive; " + - "if ! command -v sshd >/dev/null || ! command -v nc >/dev/null; then " + - "apt-get update -qq; apt-get install -y -qq --no-install-recommends openssh-server netcat-openbsd; fi; " + - "ssh-keygen -A; install -d -m 0755 /run/sshd", - TimeSpan.FromMinutes(3), - user: "root"); - Assert.Equal(0, install.ExitCode); - - var publicKey = await File.ReadAllTextAsync(keyPath + ".pub"); - var publicKeyBase64 = Convert.ToBase64String(Encoding.UTF8.GetBytes(publicKey)); - var authorize = await _fixture.RunInWslAsync( - $"set -e; install -d -m 700 -o openclaw -g openclaw /home/openclaw/.ssh; echo '{publicKeyBase64}' | base64 -d > /home/openclaw/.ssh/authorized_keys; chown openclaw:openclaw /home/openclaw/.ssh/authorized_keys; chmod 600 /home/openclaw/.ssh/authorized_keys", - TimeSpan.FromSeconds(30), - user: "root"); - Assert.Equal(0, authorize.ExitCode); - - _fixture.SetTrayEnvironmentVariable("HOME", profileDir); - _fixture.SetTrayEnvironmentVariable("USERPROFILE", profileDir); - } - - private async Task StartProofSshdAsync( - string profileDir, - string sshDir, - int sshPort) - { - var unitName = $"openclaw-pr1076-sshd-{sshPort}.service"; - var start = await _fixture.RunInWslAsync( - $"set -e; systemctl stop '{unitName}' 2>/dev/null || true; " + - $"systemctl reset-failed '{unitName}' 2>/dev/null || true; " + - $"systemd-run --quiet --unit='{unitName}' --collect --property=Type=exec " + - $"/usr/sbin/sshd -D -e -p {sshPort} -o KexAlgorithms=curve25519-sha256", - TimeSpan.FromSeconds(15), - inputViaStdin: true, - user: "root"); - Assert.Equal(0, start.ExitCode); - - var inspect = await _fixture.RunInWslAsync( - "for i in $(seq 1 50); do " + - $"if systemctl is-active --quiet '{unitName}'; then " + - $"pid=$(systemctl show '{unitName}' -p MainPID --value); " + - "if [ \"$pid\" != '0' ] && [ -r \"/proc/$pid/cmdline\" ]; then " + - "exe=$(readlink -f \"/proc/$pid/exe\" 2>/dev/null || true); " + - "cmd=$(tr '\\0' ' ' < \"/proc/$pid/cmdline\" 2>/dev/null || true); " + - $"if [ \"$exe\" = '/usr/sbin/sshd' ] && [[ \"$cmd\" == *'-D'* ]] && [[ \"$cmd\" == *'-e'* ]] && [[ \"$cmd\" == *'-p {sshPort}'* ]]; then " + - "printf '%s\\n%s\\n%s\\n' \"$pid\" \"$exe\" \"$cmd\"; exit 0; fi; fi; " + - "fi; sleep 0.1; done; " + - $"systemctl status '{unitName}' --no-pager >&2 || true; " + - $"journalctl -u '{unitName}' -n 50 --no-pager >&2 || true; exit 1", - TimeSpan.FromSeconds(15), - inputViaStdin: true, - user: "root"); - Assert.Equal(0, inspect.ExitCode); - var inspectionLines = inspect.Stdout - .Split(['\r', '\n'], StringSplitOptions.RemoveEmptyEntries | StringSplitOptions.TrimEntries); - Assert.True(inspectionLines.Length >= 3, $"Missing sshd PID/command proof: {inspect.Stdout}"); - Assert.True(int.TryParse(inspectionLines[0], out var pid), $"Invalid sshd PID: {inspectionLines[0]}"); - Assert.Equal("/usr/sbin/sshd", inspectionLines[1]); - Assert.Contains("-D", inspectionLines[2], StringComparison.Ordinal); - Assert.Contains("-e", inspectionLines[2], StringComparison.Ordinal); - Assert.Contains($"-p {sshPort}", inspectionLines[2], StringComparison.Ordinal); - - const string hostAddress = "127.0.0.1"; - - ProcessResult? preflight = null; - var preflightTimeout = TimeSpan.FromSeconds(30); - var preflightStopwatch = Stopwatch.StartNew(); - while (true) - { - var remaining = preflightTimeout - preflightStopwatch.Elapsed; - if (remaining <= TimeSpan.Zero) - break; - var attemptTimeout = TimeSpan.FromMilliseconds( - Math.Min(TimeSpan.FromSeconds(5).TotalMilliseconds, remaining.TotalMilliseconds)); - preflight = await RunProcessAsync( - "ssh.exe", - [ - "-o", "BatchMode=yes", - "-o", "IdentitiesOnly=yes", - "-o", "StrictHostKeyChecking=no", - "-o", "UserKnownHostsFile=NUL", - "-o", $"ProxyCommand=wsl.exe -d {_fixture.DistroName} -- nc 127.0.0.1 {sshPort}", - "-i", Path.Combine(sshDir, "id_ed25519"), - "-p", sshPort.ToString(), - $"openclaw@{hostAddress}", - "true" - ], - new Dictionary - { - ["HOME"] = profileDir, - ["USERPROFILE"] = profileDir, - }, - attemptTimeout); - if (preflight.ExitCode == 0) - break; - remaining = preflightTimeout - preflightStopwatch.Elapsed; - if (remaining <= TimeSpan.Zero) - break; - await Task.Delay(TimeSpan.FromMilliseconds( - Math.Min(TimeSpan.FromMilliseconds(250).TotalMilliseconds, remaining.TotalMilliseconds))); - } - Assert.NotNull(preflight); - Assert.True( - preflight.ExitCode == 0, - $"SSH preflight failed ({preflight.ExitCode}): " + - TokenSanitizer.SanitizeLogMessage(preflight.Stderr)); - return new ProofSshdProcess( - unitName, - pid, - inspectionLines[1], - inspectionLines[2], - hostAddress); - } - - private async Task ReadStatusAsync() => - await _fixture.Client!.CallToolExpectSuccessAsync("app.status"); - - private async Task WaitForConnectedPendingApprovalsAsync(TimeSpan timeout) - { - var deadline = DateTime.UtcNow.Add(timeout); - string last = ""; - while (DateTime.UtcNow < deadline) - { - using var document = await _fixture.Client!.CallToolExpectSuccessAsync( - "app.connection.pendingApprovals"); - last = document.RootElement.GetRawText(); - if (document.RootElement.GetProperty("connected").GetBoolean()) - return JsonDocument.Parse(last); - await Task.Delay(500); - } - - throw new TimeoutException($"Pending approvals never reached connected state. Last: {last}"); - } - - private async Task WaitForStatusAsync( - Func predicate, - TimeSpan timeout) - { - var deadline = DateTime.UtcNow.Add(timeout); - string last = ""; - while (DateTime.UtcNow < deadline) - { - using var document = await ReadStatusAsync(); - last = document.RootElement.GetRawText(); - if (predicate(document.RootElement)) - return JsonDocument.Parse(last); - await Task.Delay(500); - } - throw new TimeoutException($"Status predicate was not satisfied. Last: {last}"); - } - private static async Task WaitForOperatorErrorAsync( GatewayConnectionManager manager, TimeSpan timeout) @@ -670,7 +278,7 @@ private static async Task WaitForOperatorErrorAsync( $"Operator did not enter Error state. Last: {manager.CurrentSnapshot.OperatorState}"); } - private static async Task WaitForNodeErrorAsync( + private static async Task WaitForOperatorConnectedAsync( GatewayConnectionManager manager, TimeSpan timeout) { @@ -678,115 +286,30 @@ private static async Task WaitForNodeErrorAsync( while (DateTime.UtcNow < deadline) { var snapshot = manager.CurrentSnapshot; - if (snapshot.NodeState == RoleConnectionState.Error) + if (snapshot.OperatorState == RoleConnectionState.Connected) return snapshot; await Task.Delay(25); } throw new TimeoutException( - $"Node did not enter Error state. Last: {manager.CurrentSnapshot.NodeState}"); + $"Operator did not enter Connected state. Last: {manager.CurrentSnapshot.OperatorState}"); } - private static void AssertReady(JsonElement status) - { - Assert.Equal("Ready", status.GetProperty("overallState").GetString()); - Assert.Equal("Connected", status.GetProperty("operatorState").GetString()); - Assert.Equal("Connected", status.GetProperty("nodeState").GetString()); - Assert.True(status.GetProperty("nodePaired").GetBoolean()); - } - - private Task WaitForReadyStatusAsync(TimeSpan timeout) => - WaitForStatusAsync( - status => - status.GetProperty("overallState").GetString() == "Ready" && - status.GetProperty("operatorState").GetString() == "Connected" && - status.GetProperty("nodeState").GetString() == "Connected" && - status.GetProperty("nodePaired").GetBoolean(), - timeout); - - private static async Task WaitForListenerSnapshotAsync( - Func, bool> predicate, + private static async Task WaitForNodeErrorAsync( + GatewayConnectionManager manager, TimeSpan timeout) { var deadline = DateTime.UtcNow.Add(timeout); - WindowsTcpListenerSnapshotResult? last = null; while (DateTime.UtcNow < deadline) { - last = WindowsTcpListenerSnapshot.Capture(); - if (last.Ipv4Complete && last.Ipv6Complete && predicate(last.Listeners)) - return last; - await Task.Delay(100); + var snapshot = manager.CurrentSnapshot; + if (snapshot.NodeState == RoleConnectionState.Error) + return snapshot; + await Task.Delay(25); } throw new TimeoutException( - $"Listener predicate was not satisfied. IPv4 complete: {last?.Ipv4Complete}; " + - $"IPv6 complete: {last?.Ipv6Complete}; listener count: {last?.Listeners.Count ?? 0}."); - } - - private async Task NavigateAndCaptureAsync(string page, string outputPath) - { - var captureStartedAt = DateTime.UtcNow; - using var navigate = await _fixture.Client!.CallToolExpectSuccessAsync( - "app.navigate", - new { page }); - Assert.True(navigate.RootElement.GetProperty("navigated").GetBoolean()); - var captureDirectory = Path.Combine( - Path.GetDirectoryName(outputPath)!, - "Connection"); - var deadline = DateTime.UtcNow.AddSeconds(30); - while (DateTime.UtcNow < deadline) - { - var capture = Directory.Exists(captureDirectory) - ? Directory.EnumerateFiles(captureDirectory, "capture-*.png") - .Select(path => new FileInfo(path)) - .Where(file => file.LastWriteTimeUtc >= captureStartedAt.AddSeconds(-1)) - .OrderByDescending(file => file.LastWriteTimeUtc) - .FirstOrDefault() - : null; - if (capture is not null) - { - try - { - var isComposedFrame = false; - using (var stream = new FileStream( - capture.FullName, - FileMode.Open, - FileAccess.Read, - FileShare.ReadWrite)) - using (var bitmap = new Bitmap(stream)) - { - var sampledColors = new HashSet(); - var xStep = Math.Max(1, bitmap.Width / 40); - var yStep = Math.Max(1, bitmap.Height / 40); - for (var y = 0; y < bitmap.Height; y += yStep) - { - for (var x = 0; x < bitmap.Width; x += xStep) - sampledColors.Add(bitmap.GetPixel(x, y).ToArgb()); - } - isComposedFrame = sampledColors.Count >= 8; - } - - if (isComposedFrame) - { - File.Copy(capture.FullName, outputPath, overwrite: true); - return; - } - } - catch (ArgumentException) - { - // Capture is still being encoded; retry. - } - catch (IOException) - { - // Capture is still being encoded; retry. - } - } - - await Task.Delay(250); - } - - throw new InvalidOperationException( - "Connection page did not produce a composed XAML frame for proof capture."); + $"Node did not enter Error state. Last: {manager.CurrentSnapshot.NodeState}"); } private static async Task RunProcessAsync( @@ -846,19 +369,6 @@ private static async Task RunProcessAsync( } } - private static string? ReadString(JsonElement root, string propertyName) - { - foreach (var property in root.EnumerateObject()) - { - if (string.Equals(property.Name, propertyName, StringComparison.OrdinalIgnoreCase) && - property.Value.ValueKind == JsonValueKind.String) - { - return property.Value.GetString(); - } - } - return null; - } - private static string ResolveHeadSha() { var result = RunProcessAsync("git.exe", ["rev-parse", "HEAD"]) @@ -867,79 +377,15 @@ private static string ResolveHeadSha() return result.ExitCode == 0 ? result.Stdout.Trim() : "unknown"; } - private static int AllocateFreeForwardPortPair() - { - for (var attempt = 0; attempt < 100; attempt++) - { - var candidate = Random.Shared.Next(20_000, 40_000); - IReadOnlyList? gatewayForward = null; - IReadOnlyList? browserForward = null; - try - { - gatewayForward = StartExclusiveLoopbackListeners(candidate); - browserForward = StartExclusiveLoopbackListeners(candidate + 2); - return candidate; - } - catch (SocketException) - { - // Try another pair. - } - finally - { - StopListeners(browserForward); - StopListeners(gatewayForward); - } - } - - throw new InvalidOperationException("Unable to allocate an SSH forward port pair."); - } - - private static IReadOnlyList StartExclusiveLoopbackListeners(int port) - { - var listeners = new List(MirroredWslPortLease.BindProbeAddresses.Count); - try - { - foreach (var address in MirroredWslPortLease.BindProbeAddresses) - { - var listener = new TcpListener(address, port); - listeners.Add(listener); - listener.Server.ExclusiveAddressUse = true; - listener.Start(); - } - - return listeners; - } - catch - { - StopListeners(listeners); - throw; - } - } - - private static void StopListeners(IReadOnlyList? listeners) - { - if (listeners is null) - return; - - foreach (var listener in listeners) - listener.Stop(); - } - private sealed record ProcessResult(int ExitCode, string Stdout, string Stderr); - private sealed record ProofSshdProcess( - string UnitName, - int ProcessId, - string ExecutablePath, - string CommandLine, - string HostAddress); - private sealed class InitialHandshakeRaceTunnelManager(string webSocketUrl) : ISshTunnelManager { public bool OwnedListenerReady { get; set; } = true; public int OwnedListenerCheckCount { get; private set; } public bool IsActive { get; private set; } + public long OwnershipGeneration { get; private set; } public SshTunnelConfig? ActiveConfig { get; private set; } public string? LocalTunnelUrl => IsActive ? webSocketUrl : null; @@ -964,9 +410,18 @@ public Task StartAsync(SshTunnelConfig config, CancellationToken ct) ct.ThrowIfCancellationRequested(); IsActive = true; ActiveConfig = config; + OwnershipGeneration++; return Task.FromResult(webSocketUrl); } + public async Task StartOwnedAsync( + SshTunnelConfig config, + CancellationToken ct) + { + var url = await StartAsync(config, ct); + return new SshTunnelStartResult(url, config, OwnershipGeneration); + } + public Task StopAsync() { IsActive = false; @@ -974,6 +429,24 @@ public Task StopAsync() return Task.CompletedTask; } + public Task StopIfOwnedAsync( + SshTunnelConfig config, + long ownershipGeneration, + CancellationToken ct) + { + ct.ThrowIfCancellationRequested(); + if (!IsActive || + ActiveConfig != config || + OwnershipGeneration != ownershipGeneration) + { + return Task.FromResult(false); + } + + IsActive = false; + ActiveConfig = null; + return Task.FromResult(true); + } + public void Dispose() { } @@ -1076,6 +549,70 @@ await socket.SendAsync( return connectFrames; } + public async Task CompleteHandshakeAsync(TimeSpan observation) + { + var socket = await _socket.Task.WaitAsync(TimeSpan.FromSeconds(10)); + var challenge = JsonSerializer.Serialize(new + { + type = "event", + @event = "connect.challenge", + payload = new + { + nonce = "owned-listener-recovery-proof", + ts = DateTimeOffset.UtcNow.ToUnixTimeMilliseconds(), + }, + }); + await socket.SendAsync( + Encoding.UTF8.GetBytes(challenge), + WebSocketMessageType.Text, + endOfMessage: true, + CancellationToken.None); + + var connectFrames = 0; + var deadline = DateTime.UtcNow.Add(observation); + while (DateTime.UtcNow < deadline) + { + using var receiveCts = new CancellationTokenSource( + deadline - DateTime.UtcNow); + var frame = await ReceiveTextAsync(socket, receiveCts.Token); + if (frame is null) + break; + + using var document = JsonDocument.Parse(frame); + var root = document.RootElement; + if (!root.TryGetProperty("type", out var type) || + type.GetString() != "req" || + !root.TryGetProperty("method", out var method) || + method.GetString() != "connect" || + !root.TryGetProperty("id", out var id)) + { + continue; + } + + connectFrames++; + var response = JsonSerializer.Serialize(new + { + type = "res", + id = id.GetString(), + ok = true, + payload = new + { + type = "hello-ok", + protocol = 4, + server = new { version = "proof" }, + }, + }); + await socket.SendAsync( + Encoding.UTF8.GetBytes(response), + WebSocketMessageType.Text, + endOfMessage: true, + CancellationToken.None); + break; + } + + return connectFrames; + } + private async Task AcceptAsync() { try diff --git a/tests/OpenClaw.Shared.Tests/ConnectEnvelopeBuilderTests.cs b/tests/OpenClaw.Shared.Tests/ConnectEnvelopeBuilderTests.cs index 54ae001a4..b3d721dce 100644 --- a/tests/OpenClaw.Shared.Tests/ConnectEnvelopeBuilderTests.cs +++ b/tests/OpenClaw.Shared.Tests/ConnectEnvelopeBuilderTests.cs @@ -231,9 +231,15 @@ public async Task OperatorClient_SignerThrow_ReachesExistingSafeHandler() identityPath: identityPath); client.ConnectEnvelopeSigner = new ThrowingSigner(); - var unsafeTask = InvokePrivateTask(client, "SendConnectMessageAsync", Nonce); + var unsafeTask = InvokePrivateTask( + client, + "SendConnectMessageAsync", + Nonce, + 0L, + CancellationToken.None); await Assert.ThrowsAsync(() => unsafeTask); + BeginCurrentHandshake(client, 0L); await InvokePrivateTask(client, "SendConnectSafeAsync", Nonce, 0L); Assert.Contains( logger.Errors, @@ -476,6 +482,22 @@ private static Task InvokePrivateTask(object instance, string methodName, params return (Task)method!.Invoke(instance, arguments)!; } + private static void BeginCurrentHandshake( + OpenClawGatewayClient client, + long connectionGeneration) + { + var gateField = typeof(OpenClawGatewayClient).GetField( + "_handshakeChallengeGate", + BindingFlags.Instance | BindingFlags.NonPublic); + Assert.NotNull(gateField); + var gate = gateField!.GetValue(client); + Assert.NotNull(gate); + var gateType = gate!.GetType(); + gateType.GetMethod("Reset")!.Invoke(gate, [connectionGeneration]); + Assert.True( + (bool)gateType.GetMethod("TryBegin")!.Invoke(gate, [connectionGeneration])!); + } + private static string CreateDataPath() { var path = Path.Combine(Path.GetTempPath(), $"connect-envelope-{Guid.NewGuid():N}"); diff --git a/tests/OpenClaw.Shared.Tests/ModelsTests.cs b/tests/OpenClaw.Shared.Tests/ModelsTests.cs index 7826d0af7..da50c1e94 100644 --- a/tests/OpenClaw.Shared.Tests/ModelsTests.cs +++ b/tests/OpenClaw.Shared.Tests/ModelsTests.cs @@ -212,42 +212,6 @@ public void BuildArguments_OmitsDefaultSshPort() Assert.EndsWith("scott@mac-mini.local", args); } - [Fact] - public void BuildArguments_CanUseExplicitSshConfigFile() - { - var args = SshTunnelCommandLine.BuildArguments( - "scott", - "mac-mini.local", - 18789, - 28789, - includeBrowserProxyForward: false, - sshPort: 2222, - sshConfigFile: @"C:\proof data\ssh_config"); - - Assert.StartsWith( - "-F \"C:\\proof data\\ssh_config\" -o BatchMode=yes ", - args); - Assert.EndsWith("-p 2222 scott@mac-mini.local", args); - } - - [Theory] - [InlineData("")] - [InlineData(" ")] - [InlineData("C:\\proof\"bad\\ssh_config")] - [InlineData("C:\\proof\r\nbad\\ssh_config")] - public void BuildArguments_RejectsUnsafeSshConfigFile(string sshConfigFile) - { - Assert.Throws(() => - SshTunnelCommandLine.BuildArguments( - "scott", - "mac-mini.local", - 18789, - 28789, - includeBrowserProxyForward: false, - sshPort: 22, - sshConfigFile)); - } - [Fact] public void BuildArguments_RejectsInvalidSshPort() { diff --git a/tests/OpenClaw.Shared.Tests/OpenClawGatewayClientApprovalTranslationTests.cs b/tests/OpenClaw.Shared.Tests/OpenClawGatewayClientApprovalTranslationTests.cs index 00bee8184..95ba8f5d0 100644 --- a/tests/OpenClaw.Shared.Tests/OpenClawGatewayClientApprovalTranslationTests.cs +++ b/tests/OpenClaw.Shared.Tests/OpenClawGatewayClientApprovalTranslationTests.cs @@ -23,10 +23,10 @@ private static void InvokeHandleEvent(OpenClawGatewayClient client, string json) { using var doc = JsonDocument.Parse(json); var method = typeof(OpenClawGatewayClient).GetMethod( - "HandleEvent", + "HandleEventForConnection", BindingFlags.NonPublic | BindingFlags.Instance); Assert.NotNull(method); - method!.Invoke(client, new object[] { doc.RootElement, json.Length }); + method!.Invoke(client, new object[] { doc.RootElement, json.Length, 0L }); } [Fact] @@ -303,10 +303,10 @@ private static void InvokeHandleResponse(OpenClawGatewayClient client, string js { using var doc = JsonDocument.Parse(json); var method = typeof(OpenClawGatewayClient).GetMethod( - "HandleResponse", + "HandleResponseForConnection", BindingFlags.NonPublic | BindingFlags.Instance); Assert.NotNull(method); - method!.Invoke(client, new object[] { doc.RootElement }); + method!.Invoke(client, new object[] { doc.RootElement, 0L }); } [Fact] diff --git a/tests/OpenClaw.Shared.Tests/OpenClawGatewayClientTests.cs b/tests/OpenClaw.Shared.Tests/OpenClawGatewayClientTests.cs index 88c1a9202..e42f3591e 100644 --- a/tests/OpenClaw.Shared.Tests/OpenClawGatewayClientTests.cs +++ b/tests/OpenClaw.Shared.Tests/OpenClawGatewayClientTests.cs @@ -441,6 +441,26 @@ public void TrackPendingRequest(string requestId, string method) { EnsurePendingRegistryOpen(); GetPendingRegistry().RegisterTracked(requestId, method); + if (string.Equals(method, "connect", StringComparison.Ordinal)) + AuthorizeCurrentHandshake(); + } + + private void AuthorizeCurrentHandshake() + { + var generationProperty = typeof(WebSocketClientBase).GetProperty( + "CurrentConnectionGeneration", + System.Reflection.BindingFlags.NonPublic | + System.Reflection.BindingFlags.Instance); + var generation = (long)generationProperty!.GetValue(_client)!; + var gateField = typeof(OpenClawGatewayClient).GetField( + "_handshakeChallengeGate", + System.Reflection.BindingFlags.NonPublic | + System.Reflection.BindingFlags.Instance); + var gate = gateField!.GetValue(_client)!; + var gateType = gate.GetType(); + gateType.GetMethod("Reset")!.Invoke(gate, [generation]); + Assert.True((bool)gateType.GetMethod("TryBegin")!.Invoke(gate, [generation])!); + Assert.True((bool)gateType.GetMethod("TryAuthorize")!.Invoke(gate, [generation])!); } public bool GetPairingRequiredFlag() => @@ -725,6 +745,7 @@ public async Task SendWizardRequestAsync_ServiceRestartClose_PreservesCloseStatu identityPath: identity.Path); using var client = helper.Client; await client.ConnectAsync(); + helper.TrackPendingRequest("req-hello-restart", "connect"); helper.ProcessRawMessage(""" { "type": "res", @@ -922,6 +943,7 @@ public void BootstrapNodeHandoff_HelloOkWithNodeRole_DoesNotStorePrimaryNodeToke bootstrapPairAsNode: true, identityPath: CreateTempIdentityPath()); helper.SetDeviceTokenForTest(null); + helper.TrackPendingRequest("req-hello-node", "connect"); helper.ProcessRawMessage(""" { @@ -950,6 +972,7 @@ public void BootstrapNodeHandoff_HelloOkWithOperatorHandoffToken_StoresOperatorT bootstrapPairAsNode: true, identityPath: CreateTempIdentityPath()); helper.SetDeviceTokenForTest(null); + helper.TrackPendingRequest("req-hello-node", "connect"); helper.ProcessRawMessage(""" { @@ -1044,7 +1067,7 @@ public async Task HandshakeAuthorizationDenial_BlocksLaterChallengesWithoutDisab "listener ownership lost") : ReconnectAuthorizationResult.AllowedResult); }; - helper.Client.AuthenticationFailed += (_, _) => denied.TrySetResult(); + helper.Client.ConnectionFailure += (_, _) => denied.TrySetResult(); const string challenge = """ { @@ -1066,6 +1089,81 @@ public async Task HandshakeAuthorizationDenial_BlocksLaterChallengesWithoutDisab Assert.False(helper.GetAuthFailedFlag()); } + [Fact] + public async Task DuplicateChallengeWhileAuthorizationActive_IsSuppressed() + { + var helper = new GatewayClientTestHelper(); + var authorizationCalls = 0; + var authorizationStarted = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var releaseAuthorization = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + helper.Client.HandshakeAuthorizationAsync = async _ => + { + Interlocked.Increment(ref authorizationCalls); + authorizationStarted.TrySetResult(); + await releaseAuthorization.Task; + return ReconnectAuthorizationResult.AllowedResult; + }; + + const string challenge = """ + { + "type": "event", + "event": "connect.challenge", + "payload": { + "nonce": "duplicate", + "ts": 1785824000000 + } + } + """; + + helper.ProcessRawMessage(challenge); + await authorizationStarted.Task.WaitAsync(TimeSpan.FromSeconds(2)); + helper.ProcessRawMessage(challenge); + releaseAuthorization.TrySetResult(); + await Task.Delay(50); + + Assert.Equal(1, authorizationCalls); + } + + [Fact] + public async Task MalformedChallenge_DoesNotConsumeCurrentSocketGate() + { + var helper = new GatewayClientTestHelper(); + var authorizationCalls = 0; + var authorizationObserved = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + helper.Client.HandshakeAuthorizationAsync = _ => + { + authorizationCalls++; + authorizationObserved.TrySetResult(); + return Task.FromResult(ReconnectAuthorizationResult.AllowedResult); + }; + + helper.ProcessRawMessage( + """ + { + "type": "event", + "event": "connect.challenge", + "payload": { "nonce": 42 } + } + """); + helper.ProcessRawMessage( + """ + { + "type": "event", + "event": "connect.challenge", + "payload": { + "nonce": "valid-after-malformed", + "ts": 1785824000000 + } + } + """); + await authorizationObserved.Task.WaitAsync(TimeSpan.FromSeconds(2)); + + Assert.Equal(1, authorizationCalls); + } + [Fact] public async Task StaleHandshakeAuthorizationDenial_DoesNotAbortNewerSocket() { @@ -1128,6 +1226,7 @@ public void OperatorBootstrap_HelloOkWithNodeHandoffToken_StoresNodeToken() bootstrapPairAsNode: false, identityPath: CreateTempIdentityPath()); helper.SetDeviceTokenForTest(null); + helper.TrackPendingRequest("req-hello-operator", "connect"); helper.ProcessRawMessage(""" { @@ -1166,6 +1265,7 @@ public void HelloOkWhenTokenWriteFails_CompletesHandshakeAndPublishesToken() DeviceTokenReceivedEventArgs? receivedToken = null; helper.Client.HandshakeSucceeded += (_, _) => handshakeSucceeded = true; helper.Client.DeviceTokenReceived += (_, e) => receivedToken = e; + helper.TrackPendingRequest("req-hello-operator", "connect"); using (new FileStream( Path.Combine(identityPath, "device-key-ed25519.json"), @@ -4463,6 +4563,7 @@ public void HandleHelloOk_AfterAuthFailed_ClearsAuthFailedFlag() Assert.True(helper.GetAuthFailedFlag()); // Now receive hello-ok — flag must be cleared + helper.TrackPendingRequest("req-hello-1", "connect"); helper.ProcessRawMessage(""" { "type": "res", diff --git a/tests/OpenClaw.Shared.Tests/WebSocketClientBaseTests.cs b/tests/OpenClaw.Shared.Tests/WebSocketClientBaseTests.cs index 86c19d53b..79a1a4ee7 100644 --- a/tests/OpenClaw.Shared.Tests/WebSocketClientBaseTests.cs +++ b/tests/OpenClaw.Shared.Tests/WebSocketClientBaseTests.cs @@ -80,6 +80,18 @@ public class WebSocketClientBaseTests { private readonly TestLogger _logger = new(); + [Fact] + public void HandshakeChallengeGate_StaleGenerationCannotReplaceCurrentState() + { + var gate = new HandshakeChallengeGate(); + gate.Reset(2); + + Assert.True(gate.TryBegin(2)); + Assert.False(gate.TryBegin(1)); + Assert.True(gate.TryAuthorize(2)); + Assert.True(gate.IsAuthorized(2)); + } + [Theory] [InlineData("http://localhost:18789", "ws://localhost:18789")] [InlineData("https://gateway.example.com", "wss://gateway.example.com")] diff --git a/tests/OpenClaw.Shared.Tests/WindowsClientMetadataTests.cs b/tests/OpenClaw.Shared.Tests/WindowsClientMetadataTests.cs index b8972adf9..65c4440cc 100644 --- a/tests/OpenClaw.Shared.Tests/WindowsClientMetadataTests.cs +++ b/tests/OpenClaw.Shared.Tests/WindowsClientMetadataTests.cs @@ -229,14 +229,19 @@ public async Task BuildConnectMessageAsync(string nonce) "SendConnectMessageAsync", BindingFlags.NonPublic | BindingFlags.Instance); Assert.NotNull(method); - await (Task)method!.Invoke(this, [nonce])!; + await (Task)method!.Invoke( + this, + [nonce, 0L, CancellationToken.None])!; return Assert.Single(_messages); } - protected override Task SendRawAsync(string message) + protected override Task SendRawAsync( + string message, + long expectedConnectionGeneration, + CancellationToken cancellationToken) { _messages.Enqueue(message); - return Task.CompletedTask; + return Task.FromResult(true); } } } diff --git a/tests/OpenClaw.Shared.Tests/WindowsNodeClientTests.cs b/tests/OpenClaw.Shared.Tests/WindowsNodeClientTests.cs index 647a0855f..4c8dbbbe1 100644 --- a/tests/OpenClaw.Shared.Tests/WindowsNodeClientTests.cs +++ b/tests/OpenClaw.Shared.Tests/WindowsNodeClientTests.cs @@ -29,6 +29,15 @@ protected override Task SendRawAsync(string message) SentMessages.Enqueue(message); return Task.CompletedTask; } + + protected override Task SendRawAsync( + string message, + long expectedConnectionGeneration, + CancellationToken cancellationToken) + { + SentMessages.Enqueue(message); + return Task.FromResult(true); + } } private sealed class ThrowingWindowsNodeClient( @@ -38,6 +47,12 @@ private sealed class ThrowingWindowsNodeClient( { protected override Task SendRawAsync(string message) => Task.FromException(new IOException("simulated gateway send failure")); + + protected override Task SendRawAsync( + string message, + long expectedConnectionGeneration, + CancellationToken cancellationToken) => + Task.FromException(new IOException("simulated gateway send failure")); } [Theory] @@ -1598,6 +1613,7 @@ await InvokeHandleEventAsync( "ts": 1785824000000 } } + """); Assert.Equal(1, authorizationCalls); @@ -1627,6 +1643,252 @@ await InvokeHandleEventAsync( } } + [Fact] + public async Task HandleConnectChallenge_DuplicateWhileActive_IsSuppressed() + { + var dataPath = Path.Combine(Path.GetTempPath(), $"openclaw-node-test-{Guid.NewGuid():N}"); + Directory.CreateDirectory(dataPath); + + try + { + using var client = new CapturingWindowsNodeClient( + "ws://localhost:18789", + "gateway-token", + dataPath); + var authorizationCalls = 0; + var authorizationStarted = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var releaseAuthorization = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + client.HandshakeAuthorizationAsync = async _ => + { + Interlocked.Increment(ref authorizationCalls); + authorizationStarted.TrySetResult(); + await releaseAuthorization.Task; + return ReconnectAuthorizationResult.AllowedResult; + }; + const string challenge = """ + { + "type": "event", + "event": "connect.challenge", + "payload": { + "nonce": "duplicate-node", + "ts": 1785824000000 + } + } + """; + + var first = InvokeHandleEventAsync(client, challenge); + await authorizationStarted.Task.WaitAsync(TimeSpan.FromSeconds(2)); + await InvokeHandleEventAsync(client, challenge); + releaseAuthorization.TrySetResult(); + await first; + + Assert.Equal(1, authorizationCalls); + Assert.Single(client.SentMessages); + } + finally + { + Directory.Delete(dataPath, true); + } + } + + [Fact] + public async Task HandleConnectChallenge_MalformedFrameDoesNotConsumeSocketGate() + { + var dataPath = Path.Combine(Path.GetTempPath(), $"openclaw-node-test-{Guid.NewGuid():N}"); + Directory.CreateDirectory(dataPath); + + try + { + using var client = new CapturingWindowsNodeClient( + "ws://localhost:18789", + "fake-gateway-token", + dataPath); + var authorizationCalls = 0; + client.HandshakeAuthorizationAsync = _ => + { + authorizationCalls++; + return Task.FromResult(ReconnectAuthorizationResult.AllowedResult); + }; + + await InvokeHandleEventAsync( + client, + """ + { + "type": "event", + "event": "connect.challenge", + "payload": { "nonce": 42 } + } + """); + await InvokeHandleEventAsync( + client, + """ + { + "type": "event", + "event": "connect.challenge", + "payload": { + "nonce": "valid-after-malformed", + "ts": 1785824000000 + } + } + """); + + Assert.Equal(1, authorizationCalls); + Assert.Single(client.SentMessages); + } + finally + { + Directory.Delete(dataPath, true); + } + } + + [Fact] + public async Task HandleConnectChallenge_AuthorizationExceptionBlocksSocketGeneration() + { + var dataPath = Path.Combine(Path.GetTempPath(), $"openclaw-node-test-{Guid.NewGuid():N}"); + Directory.CreateDirectory(dataPath); + + try + { + using var client = new CapturingWindowsNodeClient( + "ws://localhost:18789", + "fake-gateway-token", + dataPath); + var authorizationCalls = 0; + GatewayErrorKind? failureKind = null; + ConnectionStatus? lastStatus = null; + client.HandshakeAuthorizationAsync = _ => + { + authorizationCalls++; + throw new IOException("listener verification failed"); + }; + client.ConnectionFailure += (_, kind) => failureKind = kind; + client.StatusChanged += (_, status) => lastStatus = status; + const string challenge = """ + { + "type": "event", + "event": "connect.challenge", + "payload": { + "nonce": "authorization-error", + "ts": 1785824000000 + } + } + """; + + await InvokeHandleEventAsync(client, challenge); + await InvokeHandleEventAsync(client, challenge); + + Assert.Equal(1, authorizationCalls); + Assert.Equal(GatewayErrorKind.Network, failureKind); + Assert.Equal(ConnectionStatus.Error, lastStatus); + Assert.Empty(client.SentMessages); + } + finally + { + Directory.Delete(dataPath, true); + } + } + + [Fact] + public async Task HandleConnectChallenge_SendExceptionBlocksSocketGeneration() + { + var dataPath = Path.Combine(Path.GetTempPath(), $"openclaw-node-test-{Guid.NewGuid():N}"); + Directory.CreateDirectory(dataPath); + + try + { + using var client = new ThrowingWindowsNodeClient( + "ws://localhost:18789", + "fake-gateway-token", + dataPath); + GatewayErrorKind? failureKind = null; + ConnectionStatus? lastStatus = null; + client.ConnectionFailure += (_, kind) => failureKind = kind; + client.StatusChanged += (_, status) => lastStatus = status; + const string challenge = """ + { + "type": "event", + "event": "connect.challenge", + "payload": { + "nonce": "send-error", + "ts": 1785824000000 + } + } + """; + + await InvokeHandleEventAsync(client, challenge); + await InvokeHandleEventAsync(client, challenge); + + Assert.Equal(GatewayErrorKind.Network, failureKind); + Assert.Equal(ConnectionStatus.Error, lastStatus); + } + finally + { + Directory.Delete(dataPath, true); + } + } + + [Fact] + public async Task HandleConnectChallenge_StaleDenialCannotAffectReplacementGeneration() + { + var dataPath = Path.Combine(Path.GetTempPath(), $"openclaw-node-test-{Guid.NewGuid():N}"); + Directory.CreateDirectory(dataPath); + + try + { + using var client = new CapturingWindowsNodeClient( + "ws://localhost:18789", + "gateway-token", + dataPath); + var authorizationStarted = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var releaseAuthorization = new TaskCompletionSource( + TaskCreationOptions.RunContinuationsAsynchronously); + var failures = new List(); + var statuses = new List(); + client.ConnectionFailure += (_, kind) => failures.Add(kind); + client.StatusChanged += (_, status) => statuses.Add(status); + client.HandshakeAuthorizationAsync = async _ => + { + authorizationStarted.TrySetResult(); + await releaseAuthorization.Task; + return new ReconnectAuthorizationResult( + false, + GatewayErrorKind.LocalPortConflict, + "stale owner"); + }; + + var challenge = InvokeHandleEventAsync( + client, + """ + { + "type": "event", + "event": "connect.challenge", + "payload": { + "nonce": "old-generation", + "ts": 1785824000000 + } + } + """); + await authorizationStarted.Task.WaitAsync(TimeSpan.FromSeconds(2)); + var generationField = typeof(WebSocketClientBase).GetField( + "_connectionGeneration", + BindingFlags.NonPublic | BindingFlags.Instance); + generationField!.SetValue(client, 1L); + releaseAuthorization.TrySetResult(); + await challenge; + + Assert.Empty(failures); + Assert.DoesNotContain(ConnectionStatus.Error, statuses); + Assert.Empty(client.SentMessages); + } + finally + { + Directory.Delete(dataPath, true); + } + } + [Theory] [InlineData("event")] [InlineData("req")] @@ -2055,6 +2317,7 @@ private static void HandleCorrelatedHelloOk( BindingFlags.NonPublic | BindingFlags.Instance); Assert.NotNull(pendingRequestField); pendingRequestField.SetValue(client, requestId); + AuthorizeCurrentHandshake(client); using var correlated = JsonDocument.Parse( JsonSerializer.Serialize(new @@ -2067,6 +2330,22 @@ private static void HandleCorrelatedHelloOk( client.HandleResponse(correlated.RootElement); } + private static void AuthorizeCurrentHandshake(WindowsNodeClient client) + { + var generationProperty = typeof(WebSocketClientBase).GetProperty( + "CurrentConnectionGeneration", + BindingFlags.NonPublic | BindingFlags.Instance); + var generation = (long)generationProperty!.GetValue(client)!; + var gateField = typeof(WindowsNodeClient).GetField( + "_handshakeChallengeGate", + BindingFlags.NonPublic | BindingFlags.Instance); + var gate = gateField!.GetValue(client)!; + var gateType = gate.GetType(); + gateType.GetMethod("Reset")!.Invoke(gate, [generation]); + Assert.True((bool)gateType.GetMethod("TryBegin")!.Invoke(gate, [generation])!); + Assert.True((bool)gateType.GetMethod("TryAuthorize")!.Invoke(gate, [generation])!); + } + private static string InvokeBuildNodeConnectMessage( WindowsNodeClient client, long? challengeTimestampMs = null, diff --git a/tests/OpenClaw.Tray.Tests/AppRefactorContractTests.cs b/tests/OpenClaw.Tray.Tests/AppRefactorContractTests.cs index 9812e343c..b442dc582 100644 --- a/tests/OpenClaw.Tray.Tests/AppRefactorContractTests.cs +++ b/tests/OpenClaw.Tray.Tests/AppRefactorContractTests.cs @@ -469,6 +469,20 @@ public void SshTunnelExit_RecoversActiveRegistryGatewayThroughConnectionManager( Assert.DoesNotContain("_sshTunnelService.EnsureStarted", method); } + [Fact] + public void UserSshRestart_StaysDelegatedToConnectionManager() + { + var source = ReadAppSources(); + var wrapper = ExtractMethod(source, "RestartSshTunnelAsync"); + var action = ExtractMethod(source, "RestartSshTunnelCoreAsync"); + + Assert.Contains("_connectionManager.RestartSshTunnelAsync()", wrapper); + Assert.Contains("RestartSshTunnelAsync()", action); + Assert.DoesNotContain("_sshTunnelService", action); + Assert.DoesNotContain("EnsureSshTunnelConfigured", action); + Assert.DoesNotContain("ReconnectWithSyncedBrowserProxyForward", action); + } + [Fact] public void ConnectionIssueNotification_PrefersNodeOwnedFailuresBeforeGenericGatewayError() {