mirror of
https://github.com/gravitational/teleport.git
synced 2026-09-24 16:17:11 +08:00
Fix ReportSecrets hanging on upstream error (#66062)
* Add regression test * Add regression test for waiting for correctly waiting for upstream * Fix ReportSecrets hanging on upstream error * Replace `context.AfterFunc` with `defer`
This commit is contained in:
@@ -73,6 +73,7 @@ type Service struct {
|
||||
func (s *Service) ReportSecrets(client accessgraphsecretsv1pb.SecretsScannerService_ReportSecretsServer) error {
|
||||
ctx, cancel := context.WithCancel(client.Context())
|
||||
defer cancel()
|
||||
|
||||
upstream, err := s.authClient.AccessGraphSecretsScannerClient().ReportSecrets(ctx)
|
||||
if err != nil {
|
||||
return trace.Wrap(err)
|
||||
@@ -80,38 +81,45 @@ func (s *Service) ReportSecrets(client accessgraphsecretsv1pb.SecretsScannerServ
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() {
|
||||
err := trace.Wrap(s.forwardClientToServer(ctx, client, upstream))
|
||||
if err != nil {
|
||||
cancel()
|
||||
}
|
||||
errCh <- err
|
||||
errCh <- trace.Wrap(s.forwardClientToServer(ctx, cancel, client, upstream))
|
||||
}()
|
||||
|
||||
err = s.forwardServerToClient(ctx, client, upstream)
|
||||
return trace.NewAggregate(err, <-errCh)
|
||||
if err != nil {
|
||||
// Return immediately so gRPC closes the stream, which unblocks client.Recv()
|
||||
// in the forwardClientToServer goroutine. The buffered errCh prevents a leak.
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
return trace.Wrap(<-errCh)
|
||||
}
|
||||
|
||||
func (s *Service) forwardClientToServer(ctx context.Context,
|
||||
func (s *Service) forwardClientToServer(ctx context.Context, cancel context.CancelFunc,
|
||||
client accessgraphsecretsv1pb.SecretsScannerService_ReportSecretsServer,
|
||||
server accessgraphsecretsv1pb.SecretsScannerService_ReportSecretsClient) (err error) {
|
||||
defer func() {
|
||||
// CloseSend always returns nil error.
|
||||
_ = server.CloseSend()
|
||||
}()
|
||||
for {
|
||||
req, err := client.Recv()
|
||||
if errors.Is(err, io.EOF) {
|
||||
if err := server.CloseSend(); err != nil {
|
||||
s.log.WarnContext(ctx, "Failed to close upstream stream", "error", err)
|
||||
}
|
||||
break
|
||||
// The client closed the send direction and won't send more messages.
|
||||
// Close the send direction of the server stream by returning and _do not_
|
||||
// cancel the context so that the client can receive any messages that the
|
||||
// server sends after getting io.EOF from the client.
|
||||
return nil
|
||||
}
|
||||
if err != nil {
|
||||
s.log.WarnContext(ctx, "Failed to receive from client stream", "error", err)
|
||||
cancel()
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
if err := server.Send(req); err != nil {
|
||||
s.log.WarnContext(ctx, "Failed to send to upstream stream", "error", err)
|
||||
cancel()
|
||||
return trace.Wrap(err)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *Service) forwardServerToClient(ctx context.Context,
|
||||
|
||||
@@ -25,6 +25,7 @@ import (
|
||||
"io"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/gravitational/trace"
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -44,13 +45,13 @@ func TestProxy(t *testing.T) {
|
||||
// Disable the TLS routing connection upgrade
|
||||
t.Setenv(defaults.TLSRoutingConnUpgradeEnvVar, "false")
|
||||
|
||||
authClient := newFakefakeSecretsScannerSvc(t)
|
||||
_, authClient := newFakefakeSecretsScannerSvc(t)
|
||||
|
||||
lis, err := net.Listen("tcp", "localhost:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
newProxyService(t, lis, authClient)
|
||||
ctx := context.Background()
|
||||
ctx := t.Context()
|
||||
|
||||
client, err := secretscannerclient.NewSecretsScannerServiceClient(ctx, secretscannerclient.ClientConfig{
|
||||
ProxyServer: lis.Addr().String(),
|
||||
@@ -102,15 +103,114 @@ func TestProxy(t *testing.T) {
|
||||
// Receive the termination message
|
||||
_, err = stream.Recv()
|
||||
require.ErrorIs(t, err, io.EOF)
|
||||
|
||||
}
|
||||
|
||||
func newFakefakeSecretsScannerSvc(t *testing.T) *fakeSecretsClient {
|
||||
func TestProxy_HandlesServerReturningErr(t *testing.T) {
|
||||
// Disable the TLS routing connection upgrade
|
||||
t.Setenv(defaults.TLSRoutingConnUpgradeEnvVar, "false")
|
||||
|
||||
_, authClient := newFakefakeSecretsScannerSvc(t)
|
||||
|
||||
lis, err := net.Listen("tcp", "localhost:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
newProxyService(t, lis, authClient)
|
||||
// Add a short timeout so if the proxy hangs (as it did before introducing this regression test),
|
||||
// the test doesn't wait for a whole minute to fail.
|
||||
ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second)
|
||||
defer cancel()
|
||||
|
||||
client, err := secretscannerclient.NewSecretsScannerServiceClient(ctx, secretscannerclient.ClientConfig{
|
||||
ProxyServer: lis.Addr().String(),
|
||||
Insecure: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stream, err := client.ReportSecrets(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Send incomplete message which should cause the server to return an error.
|
||||
err = stream.Send(&accessgraphsecretsv1pb.ReportSecretsRequest{})
|
||||
require.NoError(t, err)
|
||||
_, err = stream.Recv()
|
||||
require.ErrorContains(t, err, "missing device init")
|
||||
}
|
||||
|
||||
// TestProxy_PropagatesUpstreamErrorAfterClientEOF asserts that a terminal
|
||||
// error produced by the upstream SecretsScannerService *after* the client has
|
||||
// half-closed (CloseSend) is still propagated through the proxy to the client.
|
||||
//
|
||||
// This exercises the handler path where forwardClientToServer returns first
|
||||
// (normal CloseSend) and forwardServerToClient is the one that ends up carrying
|
||||
// Auth's terminal status. A handler that treats forwardClientToServer as
|
||||
// authoritative will finish and the client will see io.EOF instead of the real
|
||||
// error, masking real upstream failures.
|
||||
func TestProxy_PropagatesUpstreamErrorAfterClientEOF(t *testing.T) {
|
||||
t.Setenv(defaults.TLSRoutingConnUpgradeEnvVar, "false")
|
||||
|
||||
service, authClient := newFakefakeSecretsScannerSvc(t)
|
||||
service.postClientEOFErr = trace.AccessDenied("post-EOF validation failed")
|
||||
|
||||
lis, err := net.Listen("tcp", "localhost:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
newProxyService(t, lis, authClient)
|
||||
ctx := t.Context()
|
||||
|
||||
client, err := secretscannerclient.NewSecretsScannerServiceClient(ctx, secretscannerclient.ClientConfig{
|
||||
ProxyServer: lis.Addr().String(),
|
||||
Insecure: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
stream, err := client.ReportSecrets(ctx)
|
||||
require.NoError(t, err)
|
||||
|
||||
// Full handshake so Auth reaches the final in.Recv() that ends in EOF.
|
||||
err = stream.Send(&accessgraphsecretsv1pb.ReportSecretsRequest{
|
||||
Payload: &accessgraphsecretsv1pb.ReportSecretsRequest_DeviceAssertion{
|
||||
DeviceAssertion: &devicepb.AssertDeviceRequest{
|
||||
Payload: &devicepb.AssertDeviceRequest_Init{
|
||||
Init: &devicepb.AssertDeviceInit{},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = stream.Recv()
|
||||
require.NoError(t, err)
|
||||
|
||||
err = stream.Send(&accessgraphsecretsv1pb.ReportSecretsRequest{
|
||||
Payload: &accessgraphsecretsv1pb.ReportSecretsRequest_DeviceAssertion{
|
||||
DeviceAssertion: &devicepb.AssertDeviceRequest{
|
||||
Payload: &devicepb.AssertDeviceRequest_ChallengeResponse{
|
||||
ChallengeResponse: &devicepb.AuthenticateDeviceChallengeResponse{Signature: []byte("response")},
|
||||
},
|
||||
},
|
||||
},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
_, err = stream.Recv()
|
||||
require.NoError(t, err)
|
||||
|
||||
err = stream.CloseSend()
|
||||
require.NoError(t, err)
|
||||
|
||||
// The client must see the upstream error, not a clean io.EOF.
|
||||
_, recvErr := stream.Recv()
|
||||
require.NotErrorIs(t, recvErr, io.EOF, "client saw clean EOF; upstream error was swallowed")
|
||||
require.ErrorContains(t, recvErr, "post-EOF validation failed")
|
||||
}
|
||||
|
||||
func newFakefakeSecretsScannerSvc(t *testing.T) (*fakeSecretsScannerSvc, *fakeSecretsClient) {
|
||||
lis, err := net.Listen("tcp", "localhost:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
server := grpc.NewServer()
|
||||
accessgraphsecretsv1pb.RegisterSecretsScannerServiceServer(server, &fakeSecretsScannerSvc{})
|
||||
service := &fakeSecretsScannerSvc{}
|
||||
accessgraphsecretsv1pb.RegisterSecretsScannerServiceServer(server, service)
|
||||
go func() {
|
||||
err := server.Serve(lis)
|
||||
assert.NoError(t, err)
|
||||
@@ -120,7 +220,7 @@ func newFakefakeSecretsScannerSvc(t *testing.T) *fakeSecretsClient {
|
||||
client, err := grpc.NewClient(lis.Addr().String(), grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
require.NoError(t, err)
|
||||
|
||||
return &fakeSecretsClient{
|
||||
return service, &fakeSecretsClient{
|
||||
SecretsScannerServiceClient: accessgraphsecretsv1pb.NewSecretsScannerServiceClient(client),
|
||||
}
|
||||
|
||||
@@ -136,6 +236,11 @@ func (s *fakeSecretsClient) AccessGraphSecretsScannerClient() accessgraphsecrets
|
||||
|
||||
type fakeSecretsScannerSvc struct {
|
||||
accessgraphsecretsv1pb.UnimplementedSecretsScannerServiceServer
|
||||
|
||||
// postClientEOFErr, if non-nil, is returned by ReportSecrets after it
|
||||
// receives EOF from the client, modeling Auth producing a terminal error
|
||||
// during post-upload processing (after the client has already half-closed).
|
||||
postClientEOFErr error
|
||||
}
|
||||
|
||||
func (f *fakeSecretsScannerSvc) ReportSecrets(in accessgraphsecretsv1pb.SecretsScannerService_ReportSecretsServer) error {
|
||||
@@ -184,6 +289,9 @@ func (f *fakeSecretsScannerSvc) ReportSecrets(in accessgraphsecretsv1pb.SecretsS
|
||||
|
||||
_, err = in.Recv()
|
||||
if errors.Is(err, io.EOF) {
|
||||
if f.postClientEOFErr != nil {
|
||||
return f.postClientEOFErr
|
||||
}
|
||||
return nil
|
||||
}
|
||||
return trace.BadParameter("unexpected message")
|
||||
|
||||
Reference in New Issue
Block a user