Fix: Update android cert status for deleted template (#37537)
This commit is contained in:
@@ -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)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user