mirror of
https://github.com/BishopFox/sliver
synced 2026-06-08 10:29:05 +00:00
Improve operator cert revocation
This commit is contained in:
+17
-3
@@ -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
|
||||
|
||||
@@ -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))
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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),
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user