diff --git a/lib/auth/auth.go b/lib/auth/auth.go index 9240be6a20d..00bbf103534 100644 --- a/lib/auth/auth.go +++ b/lib/auth/auth.go @@ -122,6 +122,7 @@ import ( "github.com/gravitational/teleport/lib/join/gcp" "github.com/gravitational/teleport/lib/join/githubactions" "github.com/gravitational/teleport/lib/join/gitlab" + "github.com/gravitational/teleport/lib/join/tpmjoin" kubetoken "github.com/gravitational/teleport/lib/kube/token" "github.com/gravitational/teleport/lib/limiter" "github.com/gravitational/teleport/lib/loginrule" @@ -1265,9 +1266,7 @@ type Server struct { // tpmValidator allows TPMs to be validated by the auth server. It can be // overridden for the purpose of tests. - tpmValidator func( - ctx context.Context, log *slog.Logger, params tpm.ValidateParams, - ) (*tpm.ValidatedTPM, error) + tpmValidator tpmjoin.TPMValidator // circleCITokenValidate allows ID tokens from CircleCI to be validated by // the auth server. It can be overridden for the purpose of tests. diff --git a/lib/auth/export_test.go b/lib/auth/export_test.go index 9e493290d40..890045e284b 100644 --- a/lib/auth/export_test.go +++ b/lib/auth/export_test.go @@ -44,7 +44,6 @@ import ( "github.com/gravitational/teleport/lib/inventory" "github.com/gravitational/teleport/lib/join/boundkeypair" "github.com/gravitational/teleport/lib/services" - "github.com/gravitational/teleport/lib/tpm" "github.com/gravitational/teleport/lib/utils" ) @@ -212,10 +211,6 @@ func (a *Server) SetTerraformIDTokenValidator(validator terraformCloudIDTokenVal a.terraformIDTokenValidator = validator } -func (a *Server) SetTPMValidator(validator func(ctx context.Context, log *slog.Logger, params tpm.ValidateParams) (*tpm.ValidatedTPM, error)) { - a.tpmValidator = validator -} - func (a *Server) SetCreateBoundKeypairValidator(validator boundkeypair.CreateBoundKeypairValidator) { a.createBoundKeypairValidator = validator } diff --git a/lib/auth/join/join.go b/lib/auth/join/join.go index 6d4ce51029a..32340c28c8f 100644 --- a/lib/auth/join/join.go +++ b/lib/auth/join/join.go @@ -199,6 +199,8 @@ type RegisterParams struct { // OracleIMDSClient overrides the HTTP client used to make requests to the // OCI Instance Metadata Service. OracleIMDSClient utils.HTTPDoClient + // AttestTPM overrides the function used to attest the host TPM for the TPM join method. + AttestTPM func(context.Context, *slog.Logger) (*tpm.Attestation, func() error, error) } func (r *RegisterParams) CheckAndSetDefaults() error { @@ -230,6 +232,10 @@ func (r *RegisterParams) CheckAndSetDefaults() error { } } + if r.AttestTPM == nil { + r.AttestTPM = tpm.Attest + } + return nil } @@ -898,7 +904,7 @@ func registerUsingTPMMethod( JoinRequest: registerUsingTokenRequestForParams(token, hostKeys, params), } - attestation, close, err := tpm.Attest(ctx, log) + attestation, close, err := params.AttestTPM(ctx, log) if err != nil { return nil, trace.Wrap(err) } diff --git a/lib/auth/join_tpm.go b/lib/auth/join_tpm.go index 21b316ec5b0..b6f54d71aa8 100644 --- a/lib/auth/join_tpm.go +++ b/lib/auth/join_tpm.go @@ -20,8 +20,6 @@ package auth import ( "context" - "crypto/x509" - "log/slog" "github.com/google/go-attestation/attest" "github.com/gravitational/trace" @@ -31,7 +29,7 @@ import ( workloadidentityv1pb "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1" "github.com/gravitational/teleport/api/types" "github.com/gravitational/teleport/lib/join/legacyjoin" - "github.com/gravitational/teleport/lib/modules" + "github.com/gravitational/teleport/lib/join/tpmjoin" "github.com/gravitational/teleport/lib/tpm" ) @@ -71,48 +69,26 @@ func (a *Server) RegisterUsingTPMMethod( return nil, trace.BadParameter("specified join token is not for `tpm` method") } - if modules.GetModules().BuildType() != modules.BuildEnterprise { - return nil, trace.Wrap( - ErrRequiresEnterprise, - "tpm joining", - ) - } - - // Convert configured CAs to a CAPool - var certPool *x509.CertPool - if len(ptv2.Spec.TPM.EKCertAllowedCAs) > 0 { - certPool = x509.NewCertPool() - for i, ca := range ptv2.Spec.TPM.EKCertAllowedCAs { - if ok := certPool.AppendCertsFromPEM([]byte(ca)); !ok { - return nil, trace.BadParameter( - "ekcert_allowed_cas[%d] has an invalid or malformed PEM", i, - ) - } + solve := func(ec *attest.EncryptedCredential) ([]byte, error) { + solution, err := solveChallenge(tpm.EncryptedCredentialToProto(ec)) + if err != nil { + return nil, trace.Wrap(err) } + return solution.Solution, nil } - // TODO(noah): Use logger from TeleportProcess. - validatedEK, err := a.tpmValidator(ctx, slog.Default(), tpm.ValidateParams{ + validatedEK, err := tpmjoin.CheckTPMRequest(ctx, tpmjoin.CheckTPMRequestParams{ + Token: ptv2, + TPMValidator: a.GetTPMValidator(), EKCert: initReq.GetEkCert(), EKKey: initReq.GetEkKey(), AttestParams: tpm.AttestationParametersFromProto(initReq.AttestationParams), - AllowedCAs: certPool, - Solve: func(ec *attest.EncryptedCredential) ([]byte, error) { - solution, err := solveChallenge(tpm.EncryptedCredentialToProto(ec)) - if err != nil { - return nil, trace.Wrap(err) - } - return solution.Solution, nil - }, + Solve: solve, }) if validatedEK != nil { joinFailureMetadata = validatedEK } if err != nil { - return nil, trace.Wrap(err, "validating TPM EK") - } - - if err := checkTPMAllowRules(validatedEK, ptv2.Spec.TPM.Allow); err != nil { return nil, trace.Wrap(err) } @@ -128,18 +104,12 @@ func (a *Server) RegisterUsingTPMMethod( return certs, trace.Wrap(err, "generating certs for host") } -func checkTPMAllowRules(tpm *tpm.ValidatedTPM, rules []*types.ProvisionTokenSpecV2TPM_Rule) error { - // If a single rule passes, accept the TPM - for _, rule := range rules { - if rule.EKPublicHash != "" && tpm.EKPubHash != rule.EKPublicHash { - continue - } - if rule.EKCertificateSerial != "" && tpm.EKCertSerial != rule.EKCertificateSerial { - continue - } +// GetTPMValidator returns the server's TPM validator. +func (a *Server) GetTPMValidator() tpmjoin.TPMValidator { + return a.tpmValidator +} - // All rules met. - return nil - } - return trace.AccessDenied("validated tpm attributes did not match any allow rules") +// SetTPMValidator sets the server's TPM validator. +func (a *Server) SetTPMValidator(v tpmjoin.TPMValidator) { + a.tpmValidator = v } diff --git a/lib/auth/join_tpm_test.go b/lib/auth/join_tpm_test.go deleted file mode 100644 index 609abd1dd1b..00000000000 --- a/lib/auth/join_tpm_test.go +++ /dev/null @@ -1,364 +0,0 @@ -/* - * Teleport - * Copyright (C) 2024 Gravitational, Inc. - * - * This program is free software: you can redistribute it and/or modify - * it under the terms of the GNU Affero General Public License as published by - * the Free Software Foundation, either version 3 of the License, or - * (at your option) any later version. - * - * This program is distributed in the hope that it will be useful, - * but WITHOUT ANY WARRANTY; without even the implied warranty of - * MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the - * GNU Affero General Public License for more details. - * - * You should have received a copy of the GNU Affero General Public License - * along with this program. If not, see . - */ - -package auth_test - -import ( - "bytes" - "context" - "crypto/x509" - "errors" - "log/slog" - "testing" - "time" - - "github.com/google/go-attestation/attest" - gocmp "github.com/google/go-cmp/cmp" - "github.com/google/go-cmp/cmp/cmpopts" - "github.com/gravitational/trace" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" - - "github.com/gravitational/teleport/api/client/proto" - apifixtures "github.com/gravitational/teleport/api/fixtures" - "github.com/gravitational/teleport/api/types" - "github.com/gravitational/teleport/lib/auth" - "github.com/gravitational/teleport/lib/auth/authtest" - "github.com/gravitational/teleport/lib/auth/testauthority" - "github.com/gravitational/teleport/lib/modules" - "github.com/gravitational/teleport/lib/modules/modulestest" - "github.com/gravitational/teleport/lib/tpm" -) - -type mockTPMValidator struct { - lastCalledParams *tpm.ValidateParams - returnErr error - returnValidatedTPM *tpm.ValidatedTPM -} - -func (m *mockTPMValidator) setup(returns *tpm.ValidatedTPM, err error) { - m.lastCalledParams = nil - m.returnErr = err - m.returnValidatedTPM = returns -} - -func (m *mockTPMValidator) validate( - _ context.Context, _ *slog.Logger, params tpm.ValidateParams, -) (*tpm.ValidatedTPM, error) { - m.lastCalledParams = ¶ms - - solution, err := params.Solve(&attest.EncryptedCredential{ - Secret: []byte("mock-secret"), - Credential: []byte("mock-credential"), - }) - if err != nil { - return nil, trace.Wrap(err) - } - if !bytes.Equal(solution, []byte("mock-solution")) { - return nil, trace.AccessDenied("invalid solution") - } - - return m.returnValidatedTPM, m.returnErr -} - -func TestServer_RegisterUsingTPMMethod(t *testing.T) { - ctx := t.Context() - mockValidator := &mockTPMValidator{} - p, err := newTestPack(ctx, testPackOptions{ - DataDir: t.TempDir(), - MutateAuth: func(server *auth.Server) error { - server.SetTPMValidator(mockValidator.validate) - return nil - }, - }) - require.NoError(t, err) - authServer := p.a - - sshPrivateKey, sshPublicKey, err := testauthority.New().GenerateKeyPair() - require.NoError(t, err) - tlsPublicKey, err := authtest.PrivateKeyToPublicKeyTLS(sshPrivateKey) - require.NoError(t, err) - - attParams := &proto.TPMAttestationParameters{ - Public: []byte("mock-public"), - } - - const ( - goodEKPubHash = "mock-ekpub-hashed" - goodEKCertSerial = "mock-ekcert-serial" - goodEKPubHashAlt = "mock-ekpub-hashed-alt" - goodEKCertSerialAlt = "mock-ekcert-serial-alt" - ) - tokenSpec := func(mutate func(v2 *types.ProvisionTokenSpecV2)) types.ProvisionTokenSpecV2 { - spec := types.ProvisionTokenSpecV2{ - JoinMethod: types.JoinMethodTPM, - Roles: []types.SystemRole{types.RoleNode}, - TPM: &types.ProvisionTokenSpecV2TPM{ - Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ - { - Description: "ekpub only", - EKPublicHash: goodEKPubHash, - }, - { - Description: "ekcert only", - EKCertificateSerial: goodEKCertSerial, - }, - { - Description: "both", - EKPublicHash: goodEKPubHashAlt, - EKCertificateSerial: goodEKCertSerialAlt, - }, - }, - }, - } - if mutate != nil { - mutate(&spec) - } - return spec - } - joinRequest := func() *types.RegisterUsingTokenRequest { - return &types.RegisterUsingTokenRequest{ - HostID: "host-id", - Role: types.RoleNode, - PublicTLSKey: tlsPublicKey, - PublicSSHKey: sshPublicKey, - } - } - - caPool := x509.NewCertPool() - require.True(t, caPool.AppendCertsFromPEM([]byte(apifixtures.TLSCACertPEM))) - - allowRulesNotMatched := require.ErrorAssertionFunc(func(t require.TestingT, err error, i ...any) { - require.ErrorContains(t, err, "validated tpm attributes did not match any allow rules") - require.True(t, trace.IsAccessDenied(err)) - }) - tests := []struct { - name string - setOSS bool - - tokenSpec types.ProvisionTokenSpecV2 - - validateReturnTPM *tpm.ValidatedTPM - validateReturnErr error - - initReq *proto.RegisterUsingTPMMethodInitialRequest - wantParams *tpm.ValidateParams - - assertError require.ErrorAssertionFunc - }{ - { - name: "success, ekpub", - assertError: require.NoError, - - initReq: &proto.RegisterUsingTPMMethodInitialRequest{ - JoinRequest: joinRequest(), - Ek: &proto.RegisterUsingTPMMethodInitialRequest_EkKey{ - EkKey: []byte("mock-ekpub"), - }, - AttestationParams: attParams, - }, - wantParams: &tpm.ValidateParams{ - EKKey: []byte("mock-ekpub"), - AttestParams: tpm.AttestationParametersFromProto(attParams), - }, - - tokenSpec: tokenSpec(nil), - validateReturnTPM: &tpm.ValidatedTPM{ - EKPubHash: goodEKPubHash, - }, - }, - { - name: "success, ekcert", - assertError: require.NoError, - - initReq: &proto.RegisterUsingTPMMethodInitialRequest{ - JoinRequest: joinRequest(), - Ek: &proto.RegisterUsingTPMMethodInitialRequest_EkCert{ - EkCert: []byte("mock-ekcert"), - }, - AttestationParams: attParams, - }, - wantParams: &tpm.ValidateParams{ - EKCert: []byte("mock-ekcert"), - AttestParams: tpm.AttestationParametersFromProto(attParams), - AllowedCAs: caPool, - }, - - tokenSpec: tokenSpec(func(v2 *types.ProvisionTokenSpecV2) { - v2.TPM.EKCertAllowedCAs = []string{apifixtures.TLSCACertPEM} - }), - validateReturnTPM: &tpm.ValidatedTPM{ - EKCertSerial: goodEKCertSerial, - EKCertVerified: true, - }, - }, - { - name: "success, both ek cert serial and ek pub hash match", - assertError: require.NoError, - - initReq: &proto.RegisterUsingTPMMethodInitialRequest{ - JoinRequest: joinRequest(), - Ek: &proto.RegisterUsingTPMMethodInitialRequest_EkCert{ - EkCert: []byte("mock-ekcert"), - }, - AttestationParams: attParams, - }, - wantParams: &tpm.ValidateParams{ - EKCert: []byte("mock-ekcert"), - AttestParams: tpm.AttestationParametersFromProto(attParams), - }, - - tokenSpec: tokenSpec(nil), - validateReturnTPM: &tpm.ValidatedTPM{ - EKCertSerial: goodEKCertSerialAlt, - EKPubHash: goodEKPubHashAlt, - EKCertVerified: true, - }, - }, - { - name: "failure, mismatched ekpub", - assertError: allowRulesNotMatched, - - initReq: &proto.RegisterUsingTPMMethodInitialRequest{ - JoinRequest: joinRequest(), - Ek: &proto.RegisterUsingTPMMethodInitialRequest_EkKey{ - EkKey: []byte("mock-ekpub"), - }, - AttestationParams: attParams, - }, - wantParams: &tpm.ValidateParams{ - EKKey: []byte("mock-ekpub"), - AttestParams: tpm.AttestationParametersFromProto(attParams), - }, - - tokenSpec: tokenSpec(nil), - validateReturnTPM: &tpm.ValidatedTPM{ - EKPubHash: "mock-ekpub-hashed-mismatched!", - }, - }, - { - name: "failure, mismatched ekcert", - assertError: allowRulesNotMatched, - - initReq: &proto.RegisterUsingTPMMethodInitialRequest{ - JoinRequest: joinRequest(), - Ek: &proto.RegisterUsingTPMMethodInitialRequest_EkCert{ - EkCert: []byte("mock-ekcert"), - }, - AttestationParams: attParams, - }, - wantParams: &tpm.ValidateParams{ - EKCert: []byte("mock-ekcert"), - AttestParams: tpm.AttestationParametersFromProto(attParams), - }, - - tokenSpec: tokenSpec(nil), - validateReturnTPM: &tpm.ValidatedTPM{ - EKCertSerial: "mock-ekcert-serial-mismatched!", - }, - }, - { - name: "failure, verification", - assertError: func(t require.TestingT, err error, i ...any) { - assert.ErrorContains(t, err, "capacitor overcharged") - }, - - initReq: &proto.RegisterUsingTPMMethodInitialRequest{ - JoinRequest: joinRequest(), - Ek: &proto.RegisterUsingTPMMethodInitialRequest_EkCert{ - EkCert: []byte("mock-ekcert"), - }, - AttestationParams: attParams, - }, - wantParams: &tpm.ValidateParams{ - EKCert: []byte("mock-ekcert"), - AttestParams: tpm.AttestationParametersFromProto(attParams), - }, - - tokenSpec: tokenSpec(nil), - validateReturnTPM: &tpm.ValidatedTPM{ - EKCertSerial: goodEKCertSerial, - }, - validateReturnErr: errors.New("capacitor overcharged"), - }, - { - name: "failure, no enterprise", - setOSS: true, - assertError: func(t require.TestingT, err error, i ...any) { - assert.ErrorIs(t, err, auth.ErrRequiresEnterprise) - }, - - initReq: &proto.RegisterUsingTPMMethodInitialRequest{ - JoinRequest: joinRequest(), - Ek: &proto.RegisterUsingTPMMethodInitialRequest_EkCert{ - EkCert: []byte("mock-ekcert"), - }, - AttestationParams: attParams, - }, - - tokenSpec: tokenSpec(nil), - }, - } - - solver := func(t *testing.T) func(ec *proto.TPMEncryptedCredential) ( - *proto.RegisterUsingTPMMethodChallengeResponse, error, - ) { - return func(ec *proto.TPMEncryptedCredential) ( - *proto.RegisterUsingTPMMethodChallengeResponse, error, - ) { - assert.Equal(t, []byte("mock-secret"), ec.Secret) - assert.Equal(t, []byte("mock-credential"), ec.CredentialBlob) - return &proto.RegisterUsingTPMMethodChallengeResponse{ - Solution: []byte("mock-solution"), - }, nil - } - } - - for _, tt := range tests { - t.Run(tt.name, func(t *testing.T) { - mockValidator.setup(tt.validateReturnTPM, tt.validateReturnErr) - if !tt.setOSS { - modulestest.SetTestModules( - t, - modulestest.Modules{TestBuildType: modules.BuildEnterprise}, - ) - } - - token, err := types.NewProvisionTokenFromSpec( - tt.name, time.Now().Add(time.Minute), tt.tokenSpec, - ) - require.NoError(t, err) - require.NoError(t, authServer.CreateToken(ctx, token)) - tt.initReq.JoinRequest.Token = tt.name - - _, err = authServer.RegisterUsingTPMMethod( - ctx, - tt.initReq, - solver(t)) - tt.assertError(t, err) - - assert.Empty(t, - gocmp.Diff( - tt.wantParams, - mockValidator.lastCalledParams, - cmpopts.IgnoreFields(tpm.ValidateParams{}, "Solve"), - ), - ) - }) - } -} diff --git a/lib/join/joinclient/join.go b/lib/join/joinclient/join.go index 6a2bce301aa..d8f25d9e527 100644 --- a/lib/join/joinclient/join.go +++ b/lib/join/joinclient/join.go @@ -209,7 +209,8 @@ func joinWithClient(ctx context.Context, params JoinParams, client *joinv1.Clien types.JoinMethodGitHub, types.JoinMethodGitLab, types.JoinMethodIAM, - types.JoinMethodOracle: + types.JoinMethodOracle, + types.JoinMethodTPM: joinMethod := string(params.JoinMethod) joinMethodPtr = &joinMethod default: @@ -359,6 +360,8 @@ func joinWithMethod( } return oidcJoin(stream, joinParams, clientParams) + case types.JoinMethodTPM: + return tpmJoin(ctx, stream, joinParams, clientParams) default: // TODO(nklaassen): implement remaining join methods. sendGivingUpErr := stream.Send(&messages.GivingUp{ diff --git a/lib/join/joinclient/join_tpm.go b/lib/join/joinclient/join_tpm.go new file mode 100644 index 00000000000..a3d53c65ad1 --- /dev/null +++ b/lib/join/joinclient/join_tpm.go @@ -0,0 +1,117 @@ +// Teleport +// Copyright (C) 2025 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package joinclient + +import ( + "context" + "log/slog" + + "github.com/google/go-attestation/attest" + "github.com/gravitational/trace" + + "github.com/gravitational/teleport/lib/join/internal/messages" +) + +func tpmJoin( + ctx context.Context, + stream messages.ClientStream, + joinParams JoinParams, + clientParams messages.ClientParams, +) (messages.Response, error) { + // The TPM join method involves the following messages: + // + // client->server ClientInit + // client<-server ServerInit + // client->server TPMInit + // client<-server TPMEncryptedCredential + // client->server TPMSolution + // client<-server Result + // + // At this point the ServerInit message has already been received, + // what's left is to send the TPMInit message, handle the + // TPMEncryptedCredential->TPMSolution flow, and receive and return the + // final result. + + log := slog.Default() + + attestation, close, err := joinParams.AttestTPM(ctx, log) + if err != nil { + return nil, trace.Wrap(err) + } + defer func() { + if err := close(); err != nil { + log.WarnContext(ctx, "Failed to close TPM", "error", err) + } + }() + + tpmInit := &messages.TPMInit{ + ClientParams: clientParams, + Public: attestation.AttestParams.Public, + CreateData: attestation.AttestParams.CreateData, + CreateAttestation: attestation.AttestParams.CreateAttestation, + CreateSignature: attestation.AttestParams.CreateSignature, + } + + // Get the EKKey or EKCert. We want to prefer the EKCert if it is available + // as this is signed by the manufacturer. + switch { + case attestation.Data.EKCert != nil: + log.DebugContext(ctx, "Using EKCert for TPM registration", + "ekcert_serial", attestation.Data.EKCert.SerialNumber) + tpmInit.EKCert = attestation.Data.EKCert.Raw + case attestation.Data.EKPub != nil: + log.DebugContext(ctx, "Using EKKey for TPM registration", + "ekpub_hash", attestation.Data.EKPubHash) + tpmInit.EKKey = attestation.Data.EKPub + default: + return nil, trace.BadParameter("tpm has neither ekkey or ekcert") + } + + if err := stream.Send(tpmInit); err != nil { + return nil, trace.Wrap(err, "sending TPMInit") + } + + encryptedCredential, err := messages.RecvResponse[*messages.TPMEncryptedCredential](stream) + if err != nil { + return nil, trace.Wrap(err, "receiving TPMEncryptedCredential") + } + + solution, err := attestation.Solve(&attest.EncryptedCredential{ + Credential: encryptedCredential.CredentialBlob, + Secret: encryptedCredential.Secret, + }) + if err != nil { + err = trace.Wrap(err, "activating credential") + sendGivingUpErr := stream.Send(&messages.GivingUp{ + Reason: messages.GivingUpReasonChallengeSolutionFailed, + Msg: err.Error(), + }) + return nil, trace.NewAggregate( + err, + trace.Wrap(sendGivingUpErr, "sending GivingUp message to server"), + ) + } + + if err := stream.Send(&messages.TPMSolution{ + Solution: solution, + }); err != nil { + return nil, trace.Wrap(err, "sending TPMSolution") + } + + result, err := stream.Recv() + return result, trace.Wrap(err, "receiving join result") +} diff --git a/lib/join/server.go b/lib/join/server.go index 8431b6e4783..5d38bc3ae29 100644 --- a/lib/join/server.go +++ b/lib/join/server.go @@ -57,6 +57,7 @@ import ( "github.com/gravitational/teleport/lib/join/joinutils" "github.com/gravitational/teleport/lib/join/oraclejoin" "github.com/gravitational/teleport/lib/join/provision" + "github.com/gravitational/teleport/lib/join/tpmjoin" "github.com/gravitational/teleport/lib/scopes/joining" "github.com/gravitational/teleport/lib/services" "github.com/gravitational/teleport/lib/services/readonly" @@ -92,6 +93,7 @@ type AuthService interface { GetGHAIDTokenValidator() githubactions.GithubIDTokenValidator GetGHAIDTokenJWKSValidator() githubactions.GithubIDTokenJWKSValidator GetGitlabIDTokenValidator() gitlab.Validator + GetTPMValidator() tpmjoin.TPMValidator services.Presence } @@ -306,6 +308,8 @@ func (s *Server) handleJoinMethod( return s.handleOIDCJoin(stream, authCtx, clientInit, token, s.validateGithubToken) case types.JoinMethodGitLab: return s.handleOIDCJoin(stream, authCtx, clientInit, token, s.validateGitlabToken) + case types.JoinMethodTPM: + return s.handleTPMJoin(stream, authCtx, clientInit, token) default: // TODO(nklaassen): implement checks for all join methods. return nil, trace.NotImplemented("join method %s is not yet implemented by the new join service", joinMethod) diff --git a/lib/join/server_tpm.go b/lib/join/server_tpm.go new file mode 100644 index 00000000000..8b6f28327a3 --- /dev/null +++ b/lib/join/server_tpm.go @@ -0,0 +1,118 @@ +// Teleport +// Copyright (C) 2025 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package join + +import ( + "github.com/google/go-attestation/attest" + "github.com/gravitational/trace" + + workloadidentityv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/workloadidentity/v1" + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/lib/join/internal/authz" + "github.com/gravitational/teleport/lib/join/internal/diagnostic" + "github.com/gravitational/teleport/lib/join/internal/messages" + "github.com/gravitational/teleport/lib/join/provision" + "github.com/gravitational/teleport/lib/join/tpmjoin" +) + +// handleTPMJoin handles join attempts for the TPM join method. +// +// The TPM join method involves the following messages: +// +// client->server ClientInit +// client<-server ServerInit +// client->server TPMInit +// client<-server TPMEncryptedCredential +// client->server TPMSolution +// client<-server Result +// +// At this point the ServerInit message has already been sent, what's left is +// to receive the TPMInit message, handle the TPMEncryptedCredential->TPMSolution +// flow, and return the final result if everything checks out. +func (s *Server) handleTPMJoin( + stream messages.ServerStream, + authCtx *authz.Context, + clientInit *messages.ClientInit, + provisionToken provision.Token, +) (messages.Response, error) { + ptv2, ok := provisionToken.(*types.ProvisionTokenV2) + if !ok { + return nil, trace.BadParameter("TPM joining only supports types.ProvisionTokenV2, got %T", provisionToken) + } + + // Receive the TPMInit message from the client. + tpmInit, err := messages.RecvRequest[*messages.TPMInit](stream) + if err != nil { + return nil, trace.Wrap(err, "receiving TPMInit message") + } + // Set any diagnostic info from the ClientParams. + setDiagnosticClientParams(stream.Diagnostic(), &tpmInit.ClientParams) + + solve := func(ec *attest.EncryptedCredential) ([]byte, error) { + ecMsg := &messages.TPMEncryptedCredential{ + CredentialBlob: ec.Credential, + Secret: ec.Secret, + } + if err := stream.Send(ecMsg); err != nil { + return nil, trace.Wrap(err, "sending TPMEncryptedCredential") + } + solutionMsg, err := messages.RecvRequest[*messages.TPMSolution](stream) + if err != nil { + return nil, trace.Wrap(err, "receiving TPMSolution") + } + return solutionMsg.Solution, nil + } + + validatedEK, err := tpmjoin.CheckTPMRequest(stream.Context(), tpmjoin.CheckTPMRequestParams{ + Token: ptv2, + TPMValidator: s.cfg.AuthService.GetTPMValidator(), + EKCert: tpmInit.EKCert, + EKKey: tpmInit.EKKey, + AttestParams: attest.AttestationParameters{ + Public: tpmInit.Public, + CreateData: tpmInit.CreateData, + CreateAttestation: tpmInit.CreateAttestation, + CreateSignature: tpmInit.CreateSignature, + }, + Solve: solve, + }) + // validatedEK will be returned even on error if the TPM was validated but + // no allow rules were matched, include it in the diagnostic for debugging. + stream.Diagnostic().Set(func(info *diagnostic.Info) { + info.RawJoinAttrs = validatedEK + }) + if err != nil { + return nil, trace.Wrap(err) + } + + // Make and return the final result message. + result, err := s.makeResult( + stream.Context(), + stream.Diagnostic(), + authCtx, + clientInit, + &tpmInit.ClientParams, + provisionToken, + validatedEK, + &workloadidentityv1.JoinAttrs{ + Tpm: validatedEK.JoinAttrs(), + }, + ) + return result, trace.Wrap(err) +} + +type TPMValidator = tpmjoin.TPMValidator diff --git a/lib/join/tpmjoin/join_test.go b/lib/join/tpmjoin/join_test.go new file mode 100644 index 00000000000..ee8bab3b2bb --- /dev/null +++ b/lib/join/tpmjoin/join_test.go @@ -0,0 +1,513 @@ +// Teleport +// Copyright (C) 2025 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package tpmjoin_test + +import ( + "context" + "crypto" + "crypto/rand" + "crypto/x509" + "crypto/x509/pkix" + "log/slog" + "math/big" + "testing" + "time" + + "github.com/google/go-attestation/attest" + "github.com/gravitational/trace" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + headerv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/header/v1" + machineidv1 "github.com/gravitational/teleport/api/gen/proto/go/teleport/machineid/v1" + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/lib/auth/authtest" + "github.com/gravitational/teleport/lib/auth/state" + "github.com/gravitational/teleport/lib/cryptosuites" + "github.com/gravitational/teleport/lib/join/joinclient" + "github.com/gravitational/teleport/lib/modules" + "github.com/gravitational/teleport/lib/modules/modulestest" + "github.com/gravitational/teleport/lib/tlsca" + "github.com/gravitational/teleport/lib/tpm" +) + +func TestJoinTPM(t *testing.T) { + server, err := authtest.NewTestServer(authtest.ServerConfig{ + Auth: authtest.AuthServerConfig{ + Dir: t.TempDir(), + }, + }) + require.NoError(t, err) + + adminClient, err := server.NewClient(authtest.TestAdmin()) + require.NoError(t, err) + _, err = adminClient.BotServiceClient().CreateBot(t.Context(), &machineidv1.CreateBotRequest{ + Bot: &machineidv1.Bot{ + Metadata: &headerv1.Metadata{ + Name: "testbot", + }, + Kind: types.KindBot, + Spec: &machineidv1.BotSpec{}, + }, + }) + require.NoError(t, err) + + nopClient, err := server.NewClient(authtest.TestNop()) + require.NoError(t, err) + + goodTPMKey, err := cryptosuites.GenerateKeyWithAlgorithm(cryptosuites.ECDSAP256) + require.NoError(t, err) + goodTPMPub, err := x509.MarshalPKIXPublicKey(goodTPMKey.Public()) + require.NoError(t, err) + goodTPMPubHash := tpm.HashEKPub(goodTPMPub) + + badTPMKey, err := cryptosuites.GenerateKeyWithAlgorithm(cryptosuites.ECDSAP256) + require.NoError(t, err) + + fakeTPMValidator := newFakeTPMValidator() + server.Auth().SetTPMValidator(fakeTPMValidator.validate) + + goodTPMCA, err := newFakeTPMCA() + require.NoError(t, err) + badTPMCA, err := newFakeTPMCA() + require.NoError(t, err) + + tpmCert1, tpmCertSerial1, err := goodTPMCA.issueTPMCert(goodTPMKey.Public()) + require.NoError(t, err) + tpmCert2, _, err := goodTPMCA.issueTPMCert(goodTPMKey.Public()) + require.NoError(t, err) + + allowRulesNotMatched := func(t require.TestingT, err error, i ...any) { + require.ErrorContains(t, err, "validated tpm attributes did not match any allow rules") + require.True(t, trace.IsAccessDenied(err)) + } + + for _, tc := range []struct { + desc string + tokenSpec *types.ProvisionTokenSpecV2TPM + tpmKey crypto.Signer + tpmCert []byte + badTPMSolution bool + oss bool + assertError require.ErrorAssertionFunc + expectJoinAttrs verifiedAttrs + }{ + { + desc: "success, ekpub", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKPublicHash: goodTPMPubHash, + }, + }, + }, + tpmKey: goodTPMKey, + assertError: require.NoError, + expectJoinAttrs: verifiedAttrs{ + ekPubHash: goodTPMPubHash, + }, + }, + { + desc: "success, ekcert", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKCertificateSerial: tpmCertSerial1, + }, + }, + }, + tpmKey: goodTPMKey, + tpmCert: tpmCert1, + assertError: require.NoError, + expectJoinAttrs: verifiedAttrs{ + ekPubHash: goodTPMPubHash, + ekCertSerial: tpmCertSerial1, + }, + }, + { + desc: "success, both ek cert serial and ek pub hash match", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKPublicHash: goodTPMPubHash, + EKCertificateSerial: tpmCertSerial1, + }, + }, + }, + tpmKey: goodTPMKey, + tpmCert: tpmCert1, + assertError: require.NoError, + expectJoinAttrs: verifiedAttrs{ + ekPubHash: goodTPMPubHash, + ekCertSerial: tpmCertSerial1, + }, + }, + { + desc: "success, ek cert verified", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + EKCertAllowedCAs: []string{string(goodTPMCA.caCertPEM)}, + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKCertificateSerial: tpmCertSerial1, + }, + }, + }, + tpmKey: goodTPMKey, + tpmCert: tpmCert1, + assertError: require.NoError, + expectJoinAttrs: verifiedAttrs{ + ekPubHash: goodTPMPubHash, + ekCertSerial: tpmCertSerial1, + ekCertVerified: true, + }, + }, + { + desc: "success, ek cert verified and ek pub hash match", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + EKCertAllowedCAs: []string{string(goodTPMCA.caCertPEM)}, + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKPublicHash: goodTPMPubHash, + EKCertificateSerial: tpmCertSerial1, + }, + }, + }, + tpmKey: goodTPMKey, + tpmCert: tpmCert1, + assertError: require.NoError, + expectJoinAttrs: verifiedAttrs{ + ekPubHash: goodTPMPubHash, + ekCertSerial: tpmCertSerial1, + ekCertVerified: true, + }, + }, + { + desc: "failure, mismatched ekpub", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKPublicHash: goodTPMPubHash, + }, + }, + }, + // TPM key does not match pubkey hash in token. + tpmKey: badTPMKey, + assertError: allowRulesNotMatched, + }, + { + desc: "failure, mismatched ekcert serial", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKCertificateSerial: tpmCertSerial1, + }, + }, + }, + tpmKey: goodTPMKey, + // TPM cert does not match serial in token. + tpmCert: tpmCert2, + assertError: allowRulesNotMatched, + }, + { + desc: "failure, ek cert not verified", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + // Token configures trust for a CA that did not sign the TPM cert. + EKCertAllowedCAs: []string{string(badTPMCA.caCertPEM)}, + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKCertificateSerial: tpmCertSerial1, + }, + }, + }, + tpmKey: goodTPMKey, + tpmCert: tpmCert1, + assertError: func(t require.TestingT, err error, msgAndArgs ...any) { + require.ErrorAs(t, err, (new(*trace.AccessDeniedError))) + require.ErrorContains(t, err, "certificate signed by unknown authority") + }, + }, + { + desc: "failure, solution mismatch", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKPublicHash: goodTPMPubHash, + }, + }, + }, + tpmKey: goodTPMKey, + badTPMSolution: true, + assertError: func(t require.TestingT, err error, msgAndArgs ...any) { + require.ErrorAs(t, err, (new(*trace.AccessDeniedError))) + require.ErrorContains(t, err, "invalid credential activation solution") + }, + }, + { + desc: "failure, oss", + tokenSpec: &types.ProvisionTokenSpecV2TPM{ + Allow: []*types.ProvisionTokenSpecV2TPM_Rule{ + { + EKPublicHash: goodTPMPubHash, + }, + }, + }, + tpmKey: goodTPMKey, + oss: true, + assertError: func(t require.TestingT, err error, msgAndArgs ...any) { + require.ErrorAs(t, err, (new(*trace.AccessDeniedError))) + require.ErrorContains(t, err, "this feature requires Teleport Enterprise") + }, + }, + } { + t.Run(tc.desc, func(t *testing.T) { + if !tc.oss { + modulestest.SetTestModules(t, modulestest.Modules{TestBuildType: modules.BuildEnterprise}) + } + + token, err := types.NewProvisionTokenFromSpec("mytoken", time.Now().Add(time.Minute), types.ProvisionTokenSpecV2{ + JoinMethod: types.JoinMethodTPM, + Roles: []types.SystemRole{types.RoleBot}, + BotName: "testbot", + TPM: tc.tokenSpec, + }) + require.NoError(t, err) + require.NoError(t, server.Auth().UpsertToken(t.Context(), token)) + + fakeTPM, err := newFakeTPM(tc.tpmKey, tc.tpmCert) + require.NoError(t, err) + fakeTPM.badSolution = tc.badTPMSolution + + checkResult := func(t *testing.T, result *joinclient.JoinResult) { + t.Helper() + + botCert, err := tlsca.ParseCertificatePEM(result.Certs.TLS) + require.NoError(t, err) + + id, err := tlsca.FromSubject(botCert.Subject, botCert.NotAfter) + require.NoError(t, err) + tpmAttrs := id.JoinAttributes.Tpm + require.NotNil(t, tpmAttrs) + gotAttrs := verifiedAttrs{ + ekPubHash: tpmAttrs.EkPubHash, + ekCertSerial: tpmAttrs.EkCertSerial, + ekCertVerified: tpmAttrs.EkCertVerified, + } + assert.Equal(t, tc.expectJoinAttrs, gotAttrs) + } + + t.Run("legacy", func(t *testing.T) { + result, err := joinclient.LegacyJoin(t.Context(), joinclient.JoinParams{ + Token: token.GetName(), + JoinMethod: types.JoinMethodTPM, + ID: state.IdentityID{ + Role: types.RoleBot, + }, + AuthClient: nopClient, + AttestTPM: fakeTPM.attest, + }) + tc.assertError(t, err) + if err != nil { + return + } + checkResult(t, result) + }) + t.Run("new", func(t *testing.T) { + result, err := joinclient.Join(t.Context(), joinclient.JoinParams{ + Token: token.GetName(), + ID: state.IdentityID{ + Role: types.RoleBot, + }, + AuthClient: nopClient, + AttestTPM: fakeTPM.attest, + }) + tc.assertError(t, err) + if err != nil { + return + } + checkResult(t, result) + }) + }) + } + +} + +// fakeTPM is a minimal faked TPM that will return attestation parameters for a +// key and cert it is configured with. It returns the constant "good-solution" +// for all ec solutions, unless badSolution is set. +type fakeTPM struct { + ekKey crypto.Signer + ekPub []byte + ekPubHash string + + ekCert []byte + + badSolution bool +} + +func newFakeTPM(ekKey crypto.Signer, ekCert []byte) (*fakeTPM, error) { + ekPub, err := x509.MarshalPKIXPublicKey(ekKey.Public()) + if err != nil { + return nil, trace.Wrap(err) + } + ekPubHash := tpm.HashEKPub(ekPub) + + return &fakeTPM{ + ekKey: ekKey, + ekPub: ekPub, + ekPubHash: ekPubHash, + ekCert: ekCert, + }, nil +} + +func (f *fakeTPM) attest(ctx context.Context, _ *slog.Logger) (*tpm.Attestation, func() error, error) { + close := func() error { + return nil + } + solve := func(*attest.EncryptedCredential) ([]byte, error) { + if f.badSolution { + return []byte("bad-solution"), nil + } + return []byte("good-solution"), nil + } + data := tpm.QueryRes{ + EKPub: f.ekPub, + EKPubHash: f.ekPubHash, + } + if f.ekCert != nil { + cert, err := x509.ParseCertificate(f.ekCert) + if err != nil { + return nil, close, trace.Wrap(err) + } + data.EKCert = &tpm.QueryEKCert{ + Raw: f.ekCert, + SerialNumber: tpm.SerialString(cert.SerialNumber), + } + } + return &tpm.Attestation{ + Data: data, + AttestParams: attest.AttestationParameters{ + Public: f.ekPub, + }, + Solve: solve, + }, close, nil +} + +// fakeTPMValidator is a minimal fakes TPM validator. It always issues empty +// EncryptedCredential challenges and expects the solution to be +// "good-solution", but it will legitimately validate certificates to their +// issuing CA. +type fakeTPMValidator struct{} + +func newFakeTPMValidator() *fakeTPMValidator { + return &fakeTPMValidator{} +} + +func (f *fakeTPMValidator) validate(ctx context.Context, params tpm.ValidateParams) (*tpm.ValidatedTPM, error) { + ec := &attest.EncryptedCredential{} + clientSolution, err := params.Solve(ec) + if err != nil { + return nil, trace.Wrap(err) + } + if string(clientSolution) != "good-solution" { + return nil, trace.BadParameter("invalid credential activation solution") + } + + validated := &tpm.ValidatedTPM{} + + var ekCert *x509.Certificate + if params.EKCert != nil { + ekCert, err = x509.ParseCertificate(params.EKCert) + if err != nil { + return nil, trace.Wrap(err) + } + validated.EKCertSerial = tpm.SerialString(ekCert.SerialNumber) + ekPubPKIX, err := x509.MarshalPKIXPublicKey(ekCert.PublicKey) + if err != nil { + return nil, trace.Wrap(err) + } + validated.EKPubHash = tpm.HashEKPub(ekPubPKIX) + } + if params.AllowedCAs != nil { + if ekCert == nil { + return nil, trace.BadParameter("tpm did not provide an EKCert to verify") + } + if _, err := ekCert.Verify(x509.VerifyOptions{ + Roots: params.AllowedCAs, + }); err != nil { + return nil, trace.Wrap(err, "verifying EKCert") + } + validated.EKCertVerified = true + } + if params.EKKey != nil { + validated.EKPubHash = tpm.HashEKPub(params.EKKey) + } + + return validated, nil +} + +// fakeTPMCA issues fake TPM certificates. +type fakeTPMCA struct { + caKey crypto.Signer + caCert *x509.Certificate + caCertPEM []byte + serial int64 +} + +func newFakeTPMCA() (*fakeTPMCA, error) { + caKey, err := cryptosuites.GenerateKeyWithAlgorithm(cryptosuites.ECDSAP256) + if err != nil { + return nil, trace.Wrap(err) + } + caCertPEM, err := tlsca.GenerateSelfSignedCAWithSigner(caKey, pkix.Name{CommonName: "Test TPM CA"}, nil, time.Hour) + if err != nil { + return nil, trace.Wrap(err) + } + caCert, err := tlsca.ParseCertificatePEM(caCertPEM) + if err != nil { + return nil, trace.Wrap(err) + } + return &fakeTPMCA{ + caKey: caKey, + caCert: caCert, + caCertPEM: caCertPEM, + }, nil +} + +func (f *fakeTPMCA) issueTPMCert(pub crypto.PublicKey) ([]byte, string, error) { + f.serial++ + cert := &x509.Certificate{ + SerialNumber: big.NewInt(f.serial), + NotBefore: time.Now(), + NotAfter: time.Now().Add(time.Hour), + KeyUsage: x509.KeyUsageDigitalSignature | x509.KeyUsageKeyEncipherment, + BasicConstraintsValid: true, + Subject: pkix.Name{ + CommonName: "testtpm", + }, + } + certDER, err := x509.CreateCertificate(rand.Reader, cert, f.caCert, pub, f.caKey) + if err != nil { + return nil, "", trace.Wrap(err) + } + return certDER, tpm.SerialString(cert.SerialNumber), nil +} + +type verifiedAttrs struct { + ekPubHash string + ekCertSerial string + ekCertVerified bool +} diff --git a/lib/join/tpmjoin/tpm.go b/lib/join/tpmjoin/tpm.go new file mode 100644 index 00000000000..621a6c4318b --- /dev/null +++ b/lib/join/tpmjoin/tpm.go @@ -0,0 +1,119 @@ +// Teleport +// Copyright (C) 2025 Gravitational, Inc. +// +// This program is free software: you can redistribute it and/or modify +// it under the terms of the GNU Affero General Public License as published by +// the Free Software Foundation, either version 3 of the License, or +// (at your option) any later version. +// +// This program is distributed in the hope that it will be useful, +// but WITHOUT ANY WARRANTY; without even the implied warranty of +// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the +// GNU Affero General Public License for more details. +// +// You should have received a copy of the GNU Affero General Public License +// along with this program. If not, see . + +package tpmjoin + +import ( + "context" + "crypto/x509" + + "github.com/google/go-attestation/attest" + "github.com/gravitational/trace" + + "github.com/gravitational/teleport/api/types" + "github.com/gravitational/teleport/lib/modules" + "github.com/gravitational/teleport/lib/services" + "github.com/gravitational/teleport/lib/tpm" +) + +// TPMValidator is a function type that validates a TPM for the TPM join method. +type TPMValidator func(ctx context.Context, params tpm.ValidateParams) (*tpm.ValidatedTPM, error) + +// CheckTPMRequestParams holds all parameters for CheckTPMRequest. +type CheckTPMRequestParams struct { + // Token is the provision token used to validate the request. + Token *types.ProvisionTokenV2 + // TPMValidator is a function that will be called to validate the presented TPM. + TPMValidator TPMValidator + + // EKCert is the device's endorsement certificate in X509, ASN.1 DER form. + // This certificate contains the public key of the endorsement key. This is + // preferred to ek_key. + EKCert []byte + // The device's public endorsement key in PKIX, ASN.1 DER form. This is + // used when a TPM does not contain any endorsement certificates. + EKKey []byte + // AttestationParameters describes information about a key which is necessary + // for verifying its properties remotely. + AttestParams attest.AttestationParameters + // Solve is the function will be called when TPMValidator has prepared the + // challenge and needs the remote TPM to solve it. + Solve func(*attest.EncryptedCredential) ([]byte, error) +} + +// CheckTPMRequest checks a TPM method join request. +func CheckTPMRequest(ctx context.Context, params CheckTPMRequestParams) (*tpm.ValidatedTPM, error) { + if modules.GetModules().BuildType() != modules.BuildEnterprise { + return nil, trace.Wrap( + services.ErrRequiresEnterprise, + "tpm joining", + ) + } + + certPool, err := buildCertPool(params.Token) + if err != nil { + return nil, trace.Wrap(err) + } + + validatedEK, err := params.TPMValidator(ctx, tpm.ValidateParams{ + EKCert: params.EKCert, + EKKey: params.EKKey, + AttestParams: params.AttestParams, + AllowedCAs: certPool, + Solve: params.Solve, + }) + if err != nil { + return nil, trace.AccessDenied("validating TPM: %v", err) + } + + if err := checkTPMAllowRules(validatedEK, params.Token.Spec.TPM.Allow); err != nil { + return validatedEK, trace.Wrap(err) + } + + return validatedEK, nil +} + +func buildCertPool(token *types.ProvisionTokenV2) (*x509.CertPool, error) { + if len(token.Spec.TPM.EKCertAllowedCAs) == 0 { + // Certs are not validated if no CAs were configured. + return nil, nil + } + certPool := x509.NewCertPool() + for i, ca := range token.Spec.TPM.EKCertAllowedCAs { + if ok := certPool.AppendCertsFromPEM([]byte(ca)); !ok { + return nil, trace.BadParameter( + "ekcert_allowed_cas[%d] has an invalid or malformed PEM", i, + ) + } + } + return certPool, nil +} + +func checkTPMAllowRules(tpm *tpm.ValidatedTPM, rules []*types.ProvisionTokenSpecV2TPM_Rule) error { + // If a single rule passes, accept the TPM + for _, rule := range rules { + if rule.EKPublicHash != "" && tpm.EKPubHash != rule.EKPublicHash { + continue + } + if rule.EKCertificateSerial != "" && tpm.EKCertSerial != rule.EKCertificateSerial { + continue + } + + // All rules met. + return nil + } + return trace.AccessDenied("validated tpm attributes did not match any allow rules") +} diff --git a/lib/tpm/tpm.go b/lib/tpm/tpm.go index e5a0cf668a1..44d2c77fc2c 100644 --- a/lib/tpm/tpm.go +++ b/lib/tpm/tpm.go @@ -36,9 +36,9 @@ import ( var tracer = otel.Tracer("github.com/gravitational/teleport/lib/tpm") -// serialString converts a serial number into a readable colon-delimited hex +// SerialString converts a serial number into a readable colon-delimited hex // string thats user-readable e.g ab:ab:ab:ff:ff:ff -func serialString(serial *big.Int) string { +func SerialString(serial *big.Int) string { hex := serial.Text(16) out := strings.Builder{} @@ -60,11 +60,11 @@ func serialString(serial *big.Int) string { return out.String() } -// hashEKPub hashes the public part of an EK key. The key is hashed with SHA256, +// HashEKPub hashes the public part of an EK key. The key is hashed with SHA256, // and returned as a hexadecimal string. -func hashEKPub(pkixPublicKey []byte) (string, error) { +func HashEKPub(pkixPublicKey []byte) string { hashed := sha256.Sum256(pkixPublicKey) - return fmt.Sprintf("%x", hashed), nil + return fmt.Sprintf("%x", hashed) } // QueryRes is the result of the TPM query performed by Query. @@ -138,15 +138,12 @@ func QueryWithTPM( return nil, trace.Wrap(err) } data.EKPub = ekPub - data.EKPubHash, err = hashEKPub(ekPub) - if err != nil { - return nil, trace.Wrap(err, "hashing ekpub") - } + data.EKPubHash = HashEKPub(ekPub) if eks[0].Certificate != nil { data.EKCert = &QueryEKCert{ Raw: eks[0].Certificate.Raw, - SerialNumber: serialString(eks[0].Certificate.SerialNumber), + SerialNumber: SerialString(eks[0].Certificate.SerialNumber), } } log.DebugContext(ctx, "Successfully queried TPM", "data", data) diff --git a/lib/tpm/validate.go b/lib/tpm/validate.go index 51741ae2a1a..c2ba95a6856 100644 --- a/lib/tpm/validate.go +++ b/lib/tpm/validate.go @@ -23,7 +23,6 @@ import ( "crypto" "crypto/subtle" "crypto/x509" - "log/slog" "github.com/google/go-attestation/attest" "github.com/gravitational/trace" @@ -83,9 +82,7 @@ func (c *ValidatedTPM) JoinAttrs() *workloadidentityv1pb.JoinAttrsTPM { // the client to solve in a credential activation ceremony. This allows us to // verify that the client possesses the TPM corresponding to the EK public key // or certificate presented by the client. -func Validate( - ctx context.Context, log *slog.Logger, params ValidateParams, -) (*ValidatedTPM, error) { +func Validate(ctx context.Context, params ValidateParams) (*ValidatedTPM, error) { ctx, span := tracer.Start(ctx, "Validate") defer span.End() @@ -107,12 +104,9 @@ func Validate( if err != nil { return nil, trace.Wrap(err) } - validated.EKPubHash, err = hashEKPub(ekPubPKIX) - if err != nil { - return validated, trace.Wrap(err, "hashing EK public key") - } + validated.EKPubHash = HashEKPub(ekPubPKIX) if ekCert != nil { - validated.EKCertSerial = serialString(ekCert.SerialNumber) + validated.EKCertSerial = SerialString(ekCert.SerialNumber) } if params.AllowedCAs != nil {