diff --git a/server/datastore/mysql/certificate_templates_test.go b/server/datastore/mysql/certificate_templates_test.go index 2692b586d7..1d49dbf218 100644 --- a/server/datastore/mysql/certificate_templates_test.go +++ b/server/datastore/mysql/certificate_templates_test.go @@ -28,6 +28,7 @@ func TestCertificates(t *testing.T) { {"BatchDeleteCertificateTemplates", testBatchDeleteCertificateTemplates}, {"GetHostCertificateTemplates", testGetHostCertificateTemplates}, {"GetCertificateTemplateForHost", testGetCertificateTemplateForHost}, + {"GetHostCertificateTemplateRecord", testGetHostCertificateTemplateRecord}, } for _, c := range cases { @@ -1101,3 +1102,116 @@ func testGetCertificateTemplateForHost(t *testing.T, ds *Datastore) { }) } } + +func testGetHostCertificateTemplateRecord(t *testing.T, ds *Datastore) { + defer TruncateTables(t, ds) + + ctx := context.Background() + + // Create team + team1, err := ds.NewTeam(ctx, &fleet.Team{Name: "Team 1"}) + require.NoError(t, err) + + // Create host + h1 := test.NewHost(t, ds, "host_1", "127.0.0.1", "1", "1", time.Now()) + h1.TeamID = &team1.ID + err = ds.UpdateHost(ctx, h1) + require.NoError(t, err) + + // Create certificate authority + ca, err := ds.NewCertificateAuthority(ctx, &fleet.CertificateAuthority{ + Type: string(fleet.CATypeCustomSCEPProxy), + Name: ptr.String("Test SCEP CA"), + URL: ptr.String("http://localhost:8080/scep"), + Challenge: ptr.String("test-challenge"), + }) + require.NoError(t, err) + + // Create certificate template + ct1, err := ds.CreateCertificateTemplate(ctx, &fleet.CertificateTemplate{ + Name: "Template1", + TeamID: team1.ID, + CertificateAuthorityID: ca.ID, + SubjectName: "CN=Test Subject 1", + }) + require.NoError(t, err) + + // Create host_certificate_template record + err = ds.BulkInsertHostCertificateTemplates(ctx, []fleet.HostCertificateTemplate{ + { + HostUUID: h1.UUID, + CertificateTemplateID: ct1.ID, + FleetChallenge: ptr.String("challenge-123"), + Status: fleet.CertificateTemplateDelivered, + OperationType: fleet.MDMOperationTypeInstall, + Name: "Template1", + }, + }) + require.NoError(t, err) + + t.Run("Returns record when it exists", func(t *testing.T) { + result, err := ds.GetHostCertificateTemplateRecord(ctx, h1.UUID, ct1.ID) + require.NoError(t, err) + require.NotNil(t, result) + + require.Equal(t, h1.UUID, result.HostUUID) + require.Equal(t, ct1.ID, result.CertificateTemplateID) + require.NotNil(t, result.FleetChallenge) + require.Equal(t, "challenge-123", *result.FleetChallenge) + require.Equal(t, fleet.CertificateTemplateDelivered, result.Status) + require.Equal(t, fleet.MDMOperationTypeInstall, result.OperationType) + }) + + t.Run("Returns NotFound for non-existent host", func(t *testing.T) { + _, err := ds.GetHostCertificateTemplateRecord(ctx, "non-existent-uuid", ct1.ID) + require.Error(t, err) + require.True(t, fleet.IsNotFound(err)) + }) + + t.Run("Returns NotFound for non-existent template", func(t *testing.T) { + _, err := ds.GetHostCertificateTemplateRecord(ctx, h1.UUID, 99999) + require.Error(t, err) + require.True(t, fleet.IsNotFound(err)) + }) + + t.Run("Returns record even after parent certificate_template is deleted", func(t *testing.T) { + // Create a new template and record + ct2, err := ds.CreateCertificateTemplate(ctx, &fleet.CertificateTemplate{ + Name: "Template2", + TeamID: team1.ID, + CertificateAuthorityID: ca.ID, + SubjectName: "CN=Test Subject 2", + }) + require.NoError(t, err) + + err = ds.BulkInsertHostCertificateTemplates(ctx, []fleet.HostCertificateTemplate{ + { + HostUUID: h1.UUID, + CertificateTemplateID: ct2.ID, + FleetChallenge: ptr.String("challenge-456"), + Status: fleet.CertificateTemplateDelivered, + OperationType: fleet.MDMOperationTypeInstall, + Name: "Template2", + }, + }) + require.NoError(t, err) + + // Verify record exists + result, err := ds.GetHostCertificateTemplateRecord(ctx, h1.UUID, ct2.ID) + require.NoError(t, err) + require.Equal(t, "challenge-456", *result.FleetChallenge) + + // Delete the parent certificate_template + err = ds.DeleteCertificateTemplate(ctx, ct2.ID) + require.NoError(t, err) + + // Record should still be accessible + result, err = ds.GetHostCertificateTemplateRecord(ctx, h1.UUID, ct2.ID) + require.NoError(t, err) + require.NotNil(t, result) + require.Equal(t, h1.UUID, result.HostUUID) + require.Equal(t, ct2.ID, result.CertificateTemplateID) + require.Equal(t, "challenge-456", *result.FleetChallenge) + require.Equal(t, fleet.CertificateTemplateDelivered, result.Status) + }) +} diff --git a/server/datastore/mysql/host_certificate_templates.go b/server/datastore/mysql/host_certificate_templates.go index 4f5ba7abd4..7e794e266e 100644 --- a/server/datastore/mysql/host_certificate_templates.go +++ b/server/datastore/mysql/host_certificate_templates.go @@ -109,6 +109,36 @@ func (ds *Datastore) GetCertificateTemplateForHost(ctx context.Context, hostUUID return &result, nil } +// GetHostCertificateTemplateRecord returns the host_certificate_templates record directly without +// requiring the parent certificate_template to exist. Used for status updates on orphaned records. +func (ds *Datastore) GetHostCertificateTemplateRecord(ctx context.Context, hostUUID string, certificateTemplateID uint) (*fleet.HostCertificateTemplate, error) { + const stmt = ` + SELECT + id, + name, + host_uuid, + certificate_template_id, + fleet_challenge, + status, + operation_type, + detail, + created_at, + updated_at + FROM host_certificate_templates + WHERE host_uuid = ? AND certificate_template_id = ? + ` + + var result fleet.HostCertificateTemplate + if err := sqlx.GetContext(ctx, ds.reader(ctx), &result, stmt, hostUUID, certificateTemplateID); err != nil { + if errors.Is(err, sql.ErrNoRows) { + return nil, ctxerr.Wrap(ctx, notFound("HostCertificateTemplate")) + } + return nil, ctxerr.Wrap(ctx, err, "get host certificate template record") + } + + return &result, nil +} + // BulkInsertHostCertificateTemplates inserts multiple host_certificate_templates records func (ds *Datastore) BulkInsertHostCertificateTemplates(ctx context.Context, hostCertTemplates []fleet.HostCertificateTemplate) error { if len(hostCertTemplates) == 0 { diff --git a/server/fleet/datastore.go b/server/fleet/datastore.go index 71e53d9f59..9795cebc36 100644 --- a/server/fleet/datastore.go +++ b/server/fleet/datastore.go @@ -2568,6 +2568,9 @@ type Datastore interface { ListCertificateTemplatesForHosts(ctx context.Context, hostUUIDs []string) ([]CertificateTemplateForHost, error) // GetCertificateTemplateForHost returns a certificate template for the given host UUID and certificate template ID. GetCertificateTemplateForHost(ctx context.Context, hostUUID string, certificateTemplateID uint) (*CertificateTemplateForHost, error) + // GetHostCertificateTemplateRecord returns the host_certificate_templates record directly without + // requiring the parent certificate_template to exist. Used for status updates on orphaned records. + GetHostCertificateTemplateRecord(ctx context.Context, hostUUID string, certificateTemplateID uint) (*HostCertificateTemplate, error) // BulkInsertHostCertificateTemplates inserts multiple host_certificate_templates records. BulkInsertHostCertificateTemplates(ctx context.Context, hostCertTemplates []HostCertificateTemplate) error // DeleteHostCertificateTemplates deletes specific host_certificate_templates records diff --git a/server/mock/datastore_mock.go b/server/mock/datastore_mock.go index a8a2343ecd..e49c68f6dd 100644 --- a/server/mock/datastore_mock.go +++ b/server/mock/datastore_mock.go @@ -1681,6 +1681,8 @@ type ListCertificateTemplatesForHostsFunc func(ctx context.Context, hostUUIDs [] type GetCertificateTemplateForHostFunc func(ctx context.Context, hostUUID string, certificateTemplateID uint) (*fleet.CertificateTemplateForHost, error) +type GetHostCertificateTemplateRecordFunc func(ctx context.Context, hostUUID string, certificateTemplateID uint) (*fleet.HostCertificateTemplate, error) + type BulkInsertHostCertificateTemplatesFunc func(ctx context.Context, hostCertTemplates []fleet.HostCertificateTemplate) error type DeleteHostCertificateTemplatesFunc func(ctx context.Context, hostCertTemplates []fleet.HostCertificateTemplate) error @@ -4189,6 +4191,9 @@ type DataStore struct { GetCertificateTemplateForHostFunc GetCertificateTemplateForHostFunc GetCertificateTemplateForHostFuncInvoked bool + GetHostCertificateTemplateRecordFunc GetHostCertificateTemplateRecordFunc + GetHostCertificateTemplateRecordFuncInvoked bool + BulkInsertHostCertificateTemplatesFunc BulkInsertHostCertificateTemplatesFunc BulkInsertHostCertificateTemplatesFuncInvoked bool @@ -10025,6 +10030,13 @@ func (s *DataStore) GetCertificateTemplateForHost(ctx context.Context, hostUUID return s.GetCertificateTemplateForHostFunc(ctx, hostUUID, certificateTemplateID) } +func (s *DataStore) GetHostCertificateTemplateRecord(ctx context.Context, hostUUID string, certificateTemplateID uint) (*fleet.HostCertificateTemplate, error) { + s.mu.Lock() + s.GetHostCertificateTemplateRecordFuncInvoked = true + s.mu.Unlock() + return s.GetHostCertificateTemplateRecordFunc(ctx, hostUUID, certificateTemplateID) +} + func (s *DataStore) BulkInsertHostCertificateTemplates(ctx context.Context, hostCertTemplates []fleet.HostCertificateTemplate) error { s.mu.Lock() s.BulkInsertHostCertificateTemplatesFuncInvoked = true diff --git a/server/service/certificates.go b/server/service/certificates.go index 8e6051c8b0..e31221dea8 100644 --- a/server/service/certificates.go +++ b/server/service/certificates.go @@ -571,18 +571,20 @@ func (svc *Service) UpdateCertificateStatus( return fleet.NewInvalidArgumentError("operation_type", string(opType)) } - certificate, err := svc.ds.GetCertificateTemplateForHost(ctx, host.UUID, certificateTemplateID) + // Use GetHostCertificateTemplateRecord to query the host_certificate_templates table directly, + // allowing status updates even when the parent certificate_template has been deleted. + record, err := svc.ds.GetHostCertificateTemplateRecord(ctx, host.UUID, certificateTemplateID) if err != nil { return err } - if certificate.Status != nil && *certificate.Status != fleet.CertificateTemplateDelivered { - level.Info(svc.logger).Log("msg", "ignoring certificate status update for non-delivered certificate", "host_uuid", host.UUID, "certificate_template_id", certificateTemplateID, "current_status", certificate.Status, "new_status", status) + if record.Status != fleet.CertificateTemplateDelivered { + level.Info(svc.logger).Log("msg", "ignoring certificate status update for non-delivered certificate", "host_uuid", host.UUID, "certificate_template_id", certificateTemplateID, "current_status", record.Status, "new_status", status) return nil } - if certificate.OperationType != nil && *certificate.OperationType != opType { - level.Info(svc.logger).Log("msg", "ignoring certificate status update for different operation type", "host_uuid", host.UUID, "certificate_template_id", certificateTemplateID, "current_operation_type", certificate.OperationType, "new_operation_type", opType) + if record.OperationType != opType { + level.Info(svc.logger).Log("msg", "ignoring certificate status update for different operation type", "host_uuid", host.UUID, "certificate_template_id", certificateTemplateID, "current_operation_type", record.OperationType, "new_operation_type", opType) return nil }