Improve operator cert revocation

This commit is contained in:
moloch
2026-03-13 10:59:00 -07:00
parent 01dd3cac72
commit 95764945b8
9 changed files with 745 additions and 19 deletions
+17 -3
View File
@@ -20,10 +20,15 @@ set -u
set -o pipefail
SKIP_GENERATE=0
UNIT_ONLY=0
for arg in "$@"; do
if [ "$arg" = "--skip-generate" ]; then
SKIP_GENERATE=1
fi
if [ "$arg" = "--unit-only" ]; then
UNIT_ONLY=1
SKIP_GENERATE=1
fi
done
echo "----------------------------------------------------------------"
@@ -219,8 +224,13 @@ done
## Server c2 (unit + e2e)
run_test_cmd "./server/c2" go test -tags="server,$TAGS" ./server/c2 || exit 1
run_test_cmd "./server/c2 (e2e yamux)" go test -tags="server,$TAGS,sliver_e2e" ./server/c2 -run 'Test(MTLS|WG)Yamux_' -count=1 || exit 1
run_test_cmd "./server/c2 (e2e dns)" go test -tags="server,$TAGS,sliver_e2e" ./server/c2 -run 'TestDNS_' -count=1 || exit 1
if [ "$UNIT_ONLY" -eq 1 ]; then
echo
echo "Skipping ./server/c2 e2e tests (--unit-only)"
else
run_test_cmd "./server/c2 (e2e yamux)" go test -tags="server,$TAGS,sliver_e2e" ./server/c2 -run 'Test(MTLS|WG)Yamux_' -count=1 || exit 1
run_test_cmd "./server/c2 (e2e dns)" go test -tags="server,$TAGS,sliver_e2e" ./server/c2 -run 'TestDNS_' -count=1 || exit 1
fi
## Server generate
if [ "$SKIP_GENERATE" -eq 0 ]; then
@@ -235,5 +245,9 @@ if [ "$SKIP_GENERATE" -eq 0 ]; then
go test -timeout 6h -p "$GENERATE_GO_P" -parallel "$GENERATE_TEST_PARALLEL" -tags="server,$TAGS" ./server/generate || exit 1
else
echo
echo "Skipping ./server/generate tests (--skip-generate)"
if [ "$UNIT_ONLY" -eq 1 ]; then
echo "Skipping ./server/generate tests (--unit-only)"
else
echo "Skipping ./server/generate tests (--skip-generate)"
fi
fi
+58
View File
@@ -21,6 +21,7 @@ package certs
import (
"crypto/x509"
"encoding/pem"
"errors"
"fmt"
"github.com/bishopfox/sliver/server/db"
@@ -35,6 +36,17 @@ const (
serverNamespace = "server" // Operator servers
)
var (
// ErrOperatorClientCertificateNotFound indicates that the presented operator
// client certificate is no longer trusted because it is not present in the
// certificate store.
ErrOperatorClientCertificateNotFound = errors.New("operator client certificate not found in database")
// ErrInvalidOperatorClientCertificate indicates that the presented
// certificate is not shaped like an operator client leaf certificate.
ErrInvalidOperatorClientCertificate = errors.New("invalid operator client certificate")
)
// OperatorClientGenerateCertificate - Generate a certificate signed with a given CA
func OperatorClientGenerateCertificate(operator string) ([]byte, []byte, error) {
cert, key := GenerateECCCertificate(OperatorCA, operator, false, true, true)
@@ -52,6 +64,52 @@ func OperatorClientRemoveCertificate(operator string) error {
return RemoveCertificate(OperatorCA, ECCKey, fmt.Sprintf("%s.%s", clientNamespace, operator))
}
// ValidateOperatorClientCertificate ensures that the presented operator client
// certificate is still present in the database. A valid chain alone is not
// enough; the exact leaf certificate must still exist in storage.
func ValidateOperatorClientCertificate(peerCertificates []*x509.Certificate) error {
if len(peerCertificates) == 0 || peerCertificates[0] == nil {
return ErrInvalidOperatorClientCertificate
}
leaf := peerCertificates[0]
if leaf.IsCA || leaf.Subject.CommonName == "" || !hasExtKeyUsage(leaf, x509.ExtKeyUsageClientAuth) {
return ErrInvalidOperatorClientCertificate
}
pemBytes := pem.EncodeToMemory(&pem.Block{
Type: "CERTIFICATE",
Bytes: leaf.Raw,
})
if len(pemBytes) == 0 {
return ErrInvalidOperatorClientCertificate
}
record := &models.Certificate{}
result := db.Session().Select("id").Where(&models.Certificate{
CommonName: fmt.Sprintf("%s.%s", clientNamespace, leaf.Subject.CommonName),
CAType: OperatorCA,
KeyType: ECCKey,
CertificatePEM: string(pemBytes),
}).First(record)
if result.Error == nil {
return nil
}
if errors.Is(result.Error, db.ErrRecordNotFound) {
return ErrOperatorClientCertificateNotFound
}
return result.Error
}
func hasExtKeyUsage(cert *x509.Certificate, usage x509.ExtKeyUsage) bool {
for _, extKeyUsage := range cert.ExtKeyUsage {
if extKeyUsage == usage {
return true
}
}
return false
}
// OperatorServerGetCertificate - Helper function to fetch a server cert
func OperatorServerGetCertificate(hostname string) ([]byte, []byte, error) {
return GetECCCertificate(OperatorCA, fmt.Sprintf("%s.%s", serverNamespace, hostname))
+117
View File
@@ -0,0 +1,117 @@
package certs
import (
"crypto/x509"
"encoding/pem"
"errors"
"fmt"
"testing"
"time"
"github.com/bishopfox/sliver/server/db"
"github.com/bishopfox/sliver/server/db/models"
)
func TestValidateOperatorClientCertificateAdversarial(t *testing.T) {
SetupCAs()
t.Run("accepts stored operator client leaf", func(t *testing.T) {
operatorName := uniqueOperatorCertificateName(t, "stored")
certPEM, _, err := OperatorClientGenerateCertificate(operatorName)
if err != nil {
t.Fatalf("generate operator client certificate: %v", err)
}
err = ValidateOperatorClientCertificate([]*x509.Certificate{mustParseCertificate(t, certPEM)})
if err != nil {
t.Fatalf("expected stored operator client certificate to be accepted, got %v", err)
}
})
t.Run("rejects valid stored certificate hidden behind an unstored first certificate", func(t *testing.T) {
operatorName := uniqueOperatorCertificateName(t, "hidden")
storedCertPEM, _, err := OperatorClientGenerateCertificate(operatorName)
if err != nil {
t.Fatalf("generate stored operator certificate: %v", err)
}
rogueCertPEM, _ := GenerateECCCertificate(OperatorCA, operatorName, false, true, true)
err = ValidateOperatorClientCertificate([]*x509.Certificate{
mustParseCertificate(t, rogueCertPEM),
mustParseCertificate(t, storedCertPEM),
})
if !errors.Is(err, ErrOperatorClientCertificateNotFound) {
t.Fatalf("expected first unstored certificate to be rejected, got %v", err)
}
})
t.Run("rejects same-common-name different-leaf certificate", func(t *testing.T) {
operatorName := uniqueOperatorCertificateName(t, "same-cn")
_, _, err := OperatorClientGenerateCertificate(operatorName)
if err != nil {
t.Fatalf("generate stored operator certificate: %v", err)
}
rogueCertPEM, _ := GenerateECCCertificate(OperatorCA, operatorName, false, true, true)
err = ValidateOperatorClientCertificate([]*x509.Certificate{mustParseCertificate(t, rogueCertPEM)})
if !errors.Is(err, ErrOperatorClientCertificateNotFound) {
t.Fatalf("expected rotated or forged same-CN leaf to be rejected, got %v", err)
}
})
t.Run("rejects stored operator server certificate", func(t *testing.T) {
hostName := uniqueOperatorCertificateName(t, "server")
serverCertPEM, _, err := OperatorServerGenerateCertificate(hostName)
if err != nil {
t.Fatalf("generate operator server certificate: %v", err)
}
err = ValidateOperatorClientCertificate([]*x509.Certificate{mustParseCertificate(t, serverCertPEM)})
if !errors.Is(err, ErrInvalidOperatorClientCertificate) {
t.Fatalf("expected stored operator server certificate to be rejected, got %v", err)
}
})
t.Run("rejects operator CA certificate even if inserted into the certificates table", func(t *testing.T) {
caCertPEM, caKeyPEM, err := GetCertificateAuthorityPEM(OperatorCA)
if err != nil {
t.Fatalf("get operator CA: %v", err)
}
caCert := mustParseCertificate(t, caCertPEM)
record := &models.Certificate{
CommonName: fmt.Sprintf("%s.%s", clientNamespace, caCert.Subject.CommonName),
CAType: OperatorCA,
KeyType: ECCKey,
CertificatePEM: string(caCertPEM),
PrivateKeyPEM: string(caKeyPEM),
}
if err := db.Session().Create(record).Error; err != nil {
t.Fatalf("insert CA cert into certificates table: %v", err)
}
err = ValidateOperatorClientCertificate([]*x509.Certificate{caCert})
if !errors.Is(err, ErrInvalidOperatorClientCertificate) {
t.Fatalf("expected CA certificate to be rejected, got %v", err)
}
})
}
func mustParseCertificate(t *testing.T, certPEM []byte) *x509.Certificate {
t.Helper()
block, _ := pem.Decode(certPEM)
if block == nil {
t.Fatal("failed to decode certificate PEM")
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
t.Fatalf("parse x509 certificate: %v", err)
}
return cert
}
func uniqueOperatorCertificateName(t *testing.T, prefix string) string {
t.Helper()
return fmt.Sprintf("%s-%d", prefix, time.Now().UnixNano())
}
+29 -11
View File
@@ -218,18 +218,9 @@ func kickOperatorCmd(cmd *cobra.Command, args []string) {
}
fmt.Printf(Info+"Removing auth token(s) for %s, please wait ... \n", operator)
err = db.Session().Where(&models.Operator{
Name: operator,
}).Delete(&models.Operator{}).Error
err = kickOperator(operator)
if err != nil {
fmt.Printf(Warn+"Failed to remove operator %s: %v\n", operator, err)
return
}
transport.ClearTokenCache()
fmt.Printf(Info+"Removing client certificate(s) for %s, please wait ... \n", operator)
err = certs.OperatorClientRemoveCertificate(operator)
if err != nil {
fmt.Printf(Warn+"Failed to remove the operator certificate: %v \n", err)
fmt.Printf(Warn+"Failed to kick operator %s: %v\n", operator, err)
return
}
fmt.Printf(Info+"Operator %s has been kicked out.\n", operator)
@@ -242,6 +233,33 @@ func shouldPromptKickOperator(cmd *cobra.Command, args []string) bool {
return cmd.Flags().NFlag() == 0
}
func removeOperator(operator string) error {
err := db.Session().Where(&models.Operator{
Name: operator,
}).Delete(&models.Operator{}).Error
if err != nil {
return err
}
transport.ClearTokenCache()
return nil
}
func revokeOperatorClientCertificate(operator string) error {
return certs.OperatorClientRemoveCertificate(operator)
}
func closeOperatorStreams(operator string) {
transport.CloseOperatorStreams(operator)
}
func kickOperator(operator string) error {
if err := removeOperator(operator); err != nil {
return err
}
defer closeOperatorStreams(operator)
return revokeOperatorClientCertificate(operator)
}
func operatorNames() ([]string, error) {
operators, err := db.OperatorAll()
if err != nil {
+211
View File
@@ -0,0 +1,211 @@
package console
import (
"context"
"crypto/tls"
"crypto/x509"
"encoding/json"
"errors"
"fmt"
"io"
"net"
"testing"
"time"
clientassets "github.com/bishopfox/sliver/client/assets"
clienttransport "github.com/bishopfox/sliver/client/transport"
"github.com/bishopfox/sliver/protobuf/commonpb"
"github.com/bishopfox/sliver/protobuf/rpcpb"
"github.com/bishopfox/sliver/server/certs"
"github.com/bishopfox/sliver/server/core"
servertransport "github.com/bishopfox/sliver/server/transport"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/status"
"google.golang.org/grpc/test/bufconn"
)
func TestKickOperatorClosesActiveEventStreams(t *testing.T) {
certs.SetupCAs()
listener := bufconn.Listen(2 * 1024 * 1024)
grpcServer, err := servertransport.StartMtlsClientServer(listener)
if err != nil {
t.Fatalf("start mTLS client server: %v", err)
}
defer grpcServer.Stop()
defer listener.Close()
operatorName := uniqueKickOperatorName(t)
t.Cleanup(func() {
_ = removeOperator(operatorName)
_ = revokeOperatorClientCertificate(operatorName)
closeOperatorStreams(operatorName)
})
config := mustNewOperatorAssetsConfig(t, operatorName)
rpcClient, conn, err := mustMTLSBufconnClient(t, listener, config)
if err != nil {
t.Fatalf("connect operator client: %v", err)
}
defer conn.Close()
stream, err := rpcClient.Events(context.Background(), &commonpb.Empty{})
if err != nil {
t.Fatalf("start events stream: %v", err)
}
waitForCondition(t, 3*time.Second, func() bool {
return operatorActive(operatorName)
}, "operator to appear in the active client registry")
recvErr := make(chan error, 1)
go func() {
_, err := stream.Recv()
recvErr <- err
}()
if err := kickOperator(operatorName); err != nil {
t.Fatalf("kick operator: %v", err)
}
select {
case err := <-recvErr:
if err == nil {
t.Fatal("expected kicked operator event stream to close")
}
if errors.Is(err, io.EOF) {
break
}
code := status.Code(err)
if code != codes.Canceled && code != codes.Unavailable && code != codes.Unknown {
t.Fatalf("expected stream cancellation after kick, got %v", err)
}
case <-time.After(3 * time.Second):
t.Fatal("timed out waiting for kicked operator event stream to close")
}
waitForCondition(t, 3*time.Second, func() bool {
return !operatorActive(operatorName)
}, "operator to leave the active client registry")
_, err = rpcClient.GetVersion(context.Background(), &commonpb.Empty{})
if err == nil {
t.Fatal("expected kicked operator to be denied on subsequent RPCs")
}
code := status.Code(err)
if code != codes.Unauthenticated && code != codes.Unavailable {
t.Fatalf("expected unauthenticated or unavailable after kick, got %v", err)
}
exists, err := operatorExists(operatorName)
if err != nil {
t.Fatalf("lookup operator after kick: %v", err)
}
if exists {
t.Fatal("expected operator record to be removed after kick")
}
_, _, err = certs.OperatorClientGetCertificate(operatorName)
if !errors.Is(err, certs.ErrCertDoesNotExist) {
t.Fatalf("expected operator certificate to be removed after kick, got %v", err)
}
}
func mustNewOperatorAssetsConfig(t *testing.T, operatorName string) *clientassets.ClientConfig {
t.Helper()
configJSON, err := NewOperatorConfig(operatorName, "bufnet", 31337, []string{"all"})
if err != nil {
t.Fatalf("generate operator config: %v", err)
}
config := &clientassets.ClientConfig{}
if err := json.Unmarshal(configJSON, config); err != nil {
t.Fatalf("parse operator config: %v", err)
}
return config
}
func mustMTLSBufconnClient(t *testing.T, listener *bufconn.Listener, config *clientassets.ClientConfig) (rpcpb.SliverRPCClient, *grpc.ClientConn, error) {
t.Helper()
dialCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
conn, err := grpc.DialContext(dialCtx, "bufnet",
grpc.WithContextDialer(func(context.Context, string) (net.Conn, error) {
return listener.Dial()
}),
grpc.WithTransportCredentials(credentials.NewTLS(mustClientTLSConfig(t, config))),
grpc.WithPerRPCCredentials(staticTokenAuth(config.Token)),
grpc.WithDefaultCallOptions(grpc.MaxCallRecvMsgSize(clienttransport.ClientMaxReceiveMessageSize)),
grpc.WithBlock(),
)
if err != nil {
return nil, nil, err
}
return rpcpb.NewSliverRPCClient(conn), conn, nil
}
func mustClientTLSConfig(t *testing.T, config *clientassets.ClientConfig) *tls.Config {
t.Helper()
clientCert, err := tls.X509KeyPair([]byte(config.Certificate), []byte(config.PrivateKey))
if err != nil {
t.Fatalf("parse client certificate: %v", err)
}
caCertPool := x509.NewCertPool()
caCertPool.AppendCertsFromPEM([]byte(config.CACertificate))
return &tls.Config{
Certificates: []tls.Certificate{clientCert},
RootCAs: caCertPool,
InsecureSkipVerify: true,
MinVersion: tls.VersionTLS13,
VerifyPeerCertificate: func(rawCerts [][]byte, _ [][]*x509.Certificate) error {
return clienttransport.RootOnlyVerifyCertificate(config.CACertificate, rawCerts)
},
}
}
func operatorActive(name string) bool {
for _, active := range core.Clients.ActiveOperators() {
if active == name {
return true
}
}
return false
}
func waitForCondition(t *testing.T, timeout time.Duration, predicate func() bool, description string) {
t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
if predicate() {
return
}
time.Sleep(25 * time.Millisecond)
}
t.Fatalf("timed out waiting for %s", description)
}
func uniqueKickOperatorName(t *testing.T) string {
t.Helper()
return fmt.Sprintf("kick-operator-%d", time.Now().UnixNano())
}
type staticTokenAuth string
func (t staticTokenAuth) GetRequestMetadata(context.Context, ...string) (map[string]string, error) {
return map[string]string{
"Authorization": "Bearer " + string(t),
}, nil
}
func (staticTokenAuth) RequireTransportSecurity() bool {
return true
}
+2
View File
@@ -74,6 +74,7 @@ func initMiddleware(enableAuth bool) []grpc.ServerOption {
grpc.ChainStreamInterceptor(
grpc_auth.StreamServerInterceptor(tokenAuthFunc),
permissionsStreamServerInterceptor(),
trackOperatorStreamInterceptor(),
grpc_tags.StreamServerInterceptor(grpc_tags.WithFieldExtractor(grpc_tags.CodeGenRequestFieldExtractor)),
grpc_logrus.StreamServerInterceptor(logrusEntry, logrusOpts...),
grpc_logrus.PayloadStreamServerInterceptor(logrusEntry, deciderStream),
@@ -90,6 +91,7 @@ func initMiddleware(enableAuth bool) []grpc.ServerOption {
),
grpc.ChainStreamInterceptor(
grpc_auth.StreamServerInterceptor(serverAuthFunc),
trackOperatorStreamInterceptor(),
grpc_tags.StreamServerInterceptor(grpc_tags.WithFieldExtractor(grpc_tags.CodeGenRequestFieldExtractor)),
grpc_logrus.StreamServerInterceptor(logrusEntry, logrusOpts...),
grpc_logrus.PayloadStreamServerInterceptor(logrusEntry, deciderStream),
+28 -5
View File
@@ -21,6 +21,7 @@ package transport
import (
"crypto/tls"
"crypto/x509"
"errors"
"fmt"
"net"
"runtime/debug"
@@ -50,15 +51,34 @@ var (
// StartMtlsClientListener - Start a mutual TLS listener
func StartMtlsClientListener(host string, port uint16) (*grpc.Server, net.Listener, error) {
mtlsLog.Infof("Starting gRPC/mtls listener on %s:%d", host, port)
tlsConfig := getOperatorServerTLSConfig("multiplayer")
creds := credentials.NewTLS(tlsConfig)
ln, err := net.Listen("tcp", fmt.Sprintf("%s:%d", host, port))
if err != nil {
mtlsLog.Error(err)
return nil, nil, err
}
grpcServer, err := StartMtlsClientServer(ln)
if err != nil {
ln.Close()
return nil, nil, err
}
return grpcServer, ln, nil
}
// StartMtlsClientServer serves the authenticated multiplayer gRPC server on an
// existing listener. This is primarily useful for tests that need the full mTLS
// + auth stack without opening a real TCP socket.
func StartMtlsClientServer(ln net.Listener) (*grpc.Server, error) {
if ln == nil {
return nil, errors.New("listener is required")
}
tlsConfig := getOperatorServerTLSConfig("multiplayer")
if tlsConfig == nil {
return nil, errors.New("failed to create operator TLS config")
}
creds := credentials.NewTLS(tlsConfig)
options := []grpc.ServerOption{
grpc.Creds(creds),
grpc.MaxRecvMsgSize(ServerMaxMessageSize),
@@ -81,7 +101,7 @@ func StartMtlsClientListener(host string, port uint16) (*grpc.Server, net.Listen
panicked = false
}
}()
return grpcServer, ln, nil
return grpcServer, nil
}
// getOperatorServerTLSConfig - Generate the TLS configuration, we do now allow the end user
@@ -115,6 +135,9 @@ func getOperatorServerTLSConfig(host string) *tls.Config {
ClientCAs: caCertPool,
Certificates: []tls.Certificate{cert},
MinVersion: tls.VersionTLS13,
VerifyConnection: func(state tls.ConnectionState) error {
return certs.ValidateOperatorClientCertificate(state.PeerCertificates)
},
}
return tlsConfig
+153
View File
@@ -0,0 +1,153 @@
package transport
import (
"crypto/tls"
"errors"
"fmt"
"net"
"strings"
"testing"
"time"
"github.com/bishopfox/sliver/server/certs"
)
func TestOperatorClientCertificateValidation(t *testing.T) {
certs.SetupCAs()
t.Run("accepts stored operator certificate", func(t *testing.T) {
operatorName := uniqueOperatorName(t, "stored")
certPEM, keyPEM, err := certs.OperatorClientGenerateCertificate(operatorName)
if err != nil {
t.Fatalf("generate operator certificate: %v", err)
}
err = performMutualTLSHandshake(getOperatorServerTLSConfig("multiplayer"), newClientTLSConfig(t, certPEM, keyPEM))
if err != nil {
t.Fatalf("expected stored operator certificate to be accepted, got %v", err)
}
})
t.Run("rejects deleted operator certificate", func(t *testing.T) {
operatorName := uniqueOperatorName(t, "deleted")
certPEM, keyPEM, err := certs.OperatorClientGenerateCertificate(operatorName)
if err != nil {
t.Fatalf("generate operator certificate: %v", err)
}
if err := certs.OperatorClientRemoveCertificate(operatorName); err != nil {
t.Fatalf("remove operator certificate: %v", err)
}
err = performMutualTLSHandshake(getOperatorServerTLSConfig("multiplayer"), newClientTLSConfig(t, certPEM, keyPEM))
if err == nil {
t.Fatal("expected deleted operator certificate to be rejected")
}
if !strings.Contains(err.Error(), certs.ErrOperatorClientCertificateNotFound.Error()) {
t.Fatalf("expected database rejection error, got %v", err)
}
})
t.Run("rejects rotated-out certificate but accepts current one", func(t *testing.T) {
operatorName := uniqueOperatorName(t, "rotated")
oldCertPEM, oldKeyPEM, err := certs.OperatorClientGenerateCertificate(operatorName)
if err != nil {
t.Fatalf("generate original operator certificate: %v", err)
}
if err := certs.OperatorClientRemoveCertificate(operatorName); err != nil {
t.Fatalf("remove original operator certificate: %v", err)
}
newCertPEM, newKeyPEM, err := certs.OperatorClientGenerateCertificate(operatorName)
if err != nil {
t.Fatalf("generate replacement operator certificate: %v", err)
}
err = performMutualTLSHandshake(getOperatorServerTLSConfig("multiplayer"), newClientTLSConfig(t, oldCertPEM, oldKeyPEM))
if err == nil {
t.Fatal("expected rotated-out operator certificate to be rejected")
}
if !strings.Contains(err.Error(), certs.ErrOperatorClientCertificateNotFound.Error()) {
t.Fatalf("expected rotated certificate to fail database validation, got %v", err)
}
err = performMutualTLSHandshake(getOperatorServerTLSConfig("multiplayer"), newClientTLSConfig(t, newCertPEM, newKeyPEM))
if err != nil {
t.Fatalf("expected replacement operator certificate to be accepted, got %v", err)
}
})
t.Run("rejects certificate from another authority", func(t *testing.T) {
certPEM, keyPEM := certs.GenerateECCCertificate(certs.HTTPSCA, uniqueOperatorName(t, "wrong-ca"), false, true, false)
err := performMutualTLSHandshake(getOperatorServerTLSConfig("multiplayer"), newClientTLSConfig(t, certPEM, keyPEM))
if err == nil {
t.Fatal("expected certificate from another authority to be rejected")
}
if strings.Contains(err.Error(), certs.ErrOperatorClientCertificateNotFound.Error()) {
t.Fatalf("expected trust-chain rejection before database lookup, got %v", err)
}
})
}
func newClientTLSConfig(t *testing.T, certPEM []byte, keyPEM []byte) *tls.Config {
t.Helper()
clientCert, err := tls.X509KeyPair(certPEM, keyPEM)
if err != nil {
t.Fatalf("parse client certificate: %v", err)
}
return &tls.Config{
Certificates: []tls.Certificate{clientCert},
InsecureSkipVerify: true,
MinVersion: tls.VersionTLS13,
}
}
func performMutualTLSHandshake(serverConfig *tls.Config, clientConfig *tls.Config) error {
serverConn, clientConn := net.Pipe()
defer serverConn.Close()
defer clientConn.Close()
deadline := time.Now().Add(2 * time.Second)
_ = serverConn.SetDeadline(deadline)
_ = clientConn.SetDeadline(deadline)
serverTLS := tls.Server(serverConn, serverConfig)
clientTLS := tls.Client(clientConn, clientConfig)
defer serverTLS.Close()
defer clientTLS.Close()
type handshakeResult struct {
side string
err error
}
results := make(chan handshakeResult, 2)
go func() {
results <- handshakeResult{side: "server", err: serverTLS.Handshake()}
}()
go func() {
results <- handshakeResult{side: "client", err: clientTLS.Handshake()}
}()
var errs []string
for i := 0; i < 2; i++ {
result := <-results
if result.err == nil {
continue
}
if errors.Is(result.err, net.ErrClosed) {
continue
}
errs = append(errs, fmt.Sprintf("%s: %v", result.side, result.err))
}
if len(errs) == 0 {
return nil
}
return errors.New(strings.Join(errs, "; "))
}
func uniqueOperatorName(t *testing.T, prefix string) string {
t.Helper()
return fmt.Sprintf("%s-%d", prefix, time.Now().UnixNano())
}
+130
View File
@@ -0,0 +1,130 @@
package transport
import (
"context"
"strings"
"sync"
"sync/atomic"
"github.com/bishopfox/sliver/server/db/models"
"google.golang.org/grpc"
)
type trackedServerStream struct {
grpc.ServerStream
ctx context.Context
}
func (s *trackedServerStream) Context() context.Context {
return s.ctx
}
type operatorStreamRegistry struct {
mu sync.Mutex
nextID atomic.Uint64
operators map[string]map[uint64]context.CancelFunc
}
var activeOperatorStreams = &operatorStreamRegistry{
operators: map[string]map[uint64]context.CancelFunc{},
}
func (r *operatorStreamRegistry) register(operator string, cancel context.CancelFunc) uint64 {
if cancel == nil {
return 0
}
operator = strings.TrimSpace(operator)
if operator == "" {
return 0
}
streamID := r.nextID.Add(1)
r.mu.Lock()
defer r.mu.Unlock()
streams := r.operators[operator]
if streams == nil {
streams = map[uint64]context.CancelFunc{}
r.operators[operator] = streams
}
streams[streamID] = cancel
return streamID
}
func (r *operatorStreamRegistry) unregister(operator string, streamID uint64) {
if streamID == 0 {
return
}
operator = strings.TrimSpace(operator)
if operator == "" {
return
}
r.mu.Lock()
defer r.mu.Unlock()
streams := r.operators[operator]
if streams == nil {
return
}
delete(streams, streamID)
if len(streams) == 0 {
delete(r.operators, operator)
}
}
func (r *operatorStreamRegistry) close(operator string) int {
operator = strings.TrimSpace(operator)
if operator == "" {
return 0
}
r.mu.Lock()
streams := r.operators[operator]
if len(streams) == 0 {
r.mu.Unlock()
return 0
}
cancels := make([]context.CancelFunc, 0, len(streams))
for streamID, cancel := range streams {
cancels = append(cancels, cancel)
delete(streams, streamID)
}
delete(r.operators, operator)
r.mu.Unlock()
for _, cancel := range cancels {
cancel()
}
return len(cancels)
}
// CloseOperatorStreams cancels all tracked operator-owned gRPC streams.
func CloseOperatorStreams(operator string) int {
return activeOperatorStreams.close(operator)
}
func trackOperatorStreamInterceptor() grpc.StreamServerInterceptor {
return func(srv interface{}, ss grpc.ServerStream, _ *grpc.StreamServerInfo, handler grpc.StreamHandler) error {
operator, ok := ss.Context().Value(Operator).(*models.Operator)
if !ok || operator == nil || strings.TrimSpace(operator.Name) == "" {
return handler(srv, ss)
}
ctx, cancel := context.WithCancel(ss.Context())
streamID := activeOperatorStreams.register(operator.Name, cancel)
defer func() {
activeOperatorStreams.unregister(operator.Name, streamID)
cancel()
}()
return handler(srv, &trackedServerStream{
ServerStream: ss,
ctx: ctx,
})
}
}