Migrate audit reports to the files table

Signed-off-by: Ludovic Vielle <ludovic@probo.com>
This commit is contained in:
Ludovic Vielle
2026-06-03 22:19:17 +02:00
parent b0a0f0efc9
commit 0e4d73bb0f
31 changed files with 517 additions and 947 deletions

View File

@@ -53,18 +53,17 @@ func (s AuditService) Get(
return audit, nil
}
func (s AuditService) GetByReportID(
func (s AuditService) GetByReportFileID(
ctx context.Context,
scope coredata.Scoper,
reportID gid.GID,
fileID gid.GID,
) (*coredata.Audit, error) {
audit := &coredata.Audit{}
err := s.svc.pg.WithConn(
ctx,
func(ctx context.Context, conn pg.Querier) error {
err := audit.LoadByReportID(ctx, conn, scope, reportID)
if err != nil {
if err := audit.LoadByReportFileID(ctx, conn, scope, fileID); err != nil {
return fmt.Errorf("cannot load audit: %w", err)
}

View File

@@ -16,6 +16,7 @@ package trust
import (
"context"
"errors"
"fmt"
"io"
"time"
@@ -36,33 +37,43 @@ func (s ReportService) Get(
ctx context.Context,
scope coredata.Scoper,
organizationID gid.GID,
reportID gid.GID,
) (*coredata.Report, error) {
report, err := s.loadByID(ctx, scope, reportID)
fileID gid.GID,
) (*coredata.File, error) {
file, err := s.loadByID(ctx, scope, fileID)
if err != nil {
return nil, err
}
if report.OrganizationID != organizationID {
if file.OrganizationID != organizationID {
return nil, ErrReportNotFound
}
return report, nil
// check the given report file ID is linked to an audit in order to avoid
// being able to get any file from the report request.
_, err = s.svc.Audits.GetByReportFileID(ctx, scope, fileID)
if err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return nil, ErrReportNotFound
}
return nil, fmt.Errorf("cannot verify report file: %w", err)
}
return file, nil
}
func (s ReportService) loadByID(
ctx context.Context,
scope coredata.Scoper,
reportID gid.GID,
) (*coredata.Report, error) {
report := &coredata.Report{}
fileID gid.GID,
) (*coredata.File, error) {
file := &coredata.File{}
err := s.svc.pg.WithConn(
ctx,
func(ctx context.Context, conn pg.Querier) error {
err := report.LoadByID(ctx, conn, scope, reportID)
if err != nil {
return fmt.Errorf("cannot load report: %w", err)
if err := file.LoadActiveByID(ctx, conn, scope, fileID); err != nil {
return fmt.Errorf("cannot load file: %w", err)
}
return nil
@@ -72,28 +83,28 @@ func (s ReportService) loadByID(
return nil, err
}
return report, nil
return file, nil
}
func (s ReportService) GenerateDownloadURL(
ctx context.Context,
scope coredata.Scoper,
reportID gid.GID,
fileID gid.GID,
expiresIn time.Duration,
) (*string, error) {
report, err := s.loadByID(ctx, scope, reportID)
file, err := s.loadByID(ctx, scope, fileID)
if err != nil {
return nil, fmt.Errorf("cannot get report: %w", err)
return nil, fmt.Errorf("cannot get file: %w", err)
}
presignClient := s3.NewPresignClient(s.svc.s3)
presignedReq, err := presignClient.PresignGetObject(ctx, &s3.GetObjectInput{
Bucket: new(s.svc.bucket),
Key: new(report.ObjectKey),
Key: new(file.FileKey),
ResponseCacheControl: new("max-age=3600, public"),
ResponseContentType: new(report.MimeType),
ResponseContentDisposition: new(fmt.Sprintf("attachment; filename=\"%s\"", report.Filename)),
ResponseContentType: new(file.MimeType),
ResponseContentDisposition: new(fmt.Sprintf("attachment; filename=\"%s\"", file.FileName)),
}, func(opts *s3.PresignOptions) {
opts.Expires = expiresIn
})
@@ -134,16 +145,16 @@ func (s ReportService) ExportPDFWithoutWatermark(
func (s ReportService) exportPDFData(
ctx context.Context,
scope coredata.Scoper,
reportID gid.GID,
fileID gid.GID,
) ([]byte, error) {
report, err := s.loadByID(ctx, scope, reportID)
file, err := s.loadByID(ctx, scope, fileID)
if err != nil {
return nil, fmt.Errorf("cannot get report: %w", err)
return nil, fmt.Errorf("cannot get file: %w", err)
}
result, err := s.svc.s3.GetObject(ctx, &s3.GetObjectInput{
Bucket: new(s.svc.bucket),
Key: new(report.ObjectKey),
Key: new(file.FileKey),
})
if err != nil {
return nil, fmt.Errorf("cannot download PDF from S3: %w", err)

View File

@@ -100,8 +100,8 @@ func (s TrustCenterAccessService) Request(
}
for _, audit := range allAudits {
if audit.ReportID != nil {
reportIDs = append(reportIDs, *audit.ReportID)
if audit.ReportFileID != nil {
reportIDs = append(reportIDs, *audit.ReportFileID)
}
}
}
@@ -148,7 +148,7 @@ func (s TrustCenterAccessService) Request(
return fmt.Errorf("cannot bulk insert trust center document accesses: %w", err)
}
if err := accesses.BulkInsertReportAccesses(
if err := accesses.BulkInsertReportFileAccesses(
ctx,
tx,
scope,
@@ -255,12 +255,12 @@ func (s TrustCenterAccessService) GetDocumentAccess(
return documentAccess, nil
}
func (s TrustCenterAccessService) GetReportAccess(
func (s TrustCenterAccessService) GetReportFileAccess(
ctx context.Context,
scope coredata.Scoper,
trustCenterID gid.GID,
identityID gid.GID,
reportID gid.GID,
reportFileID gid.GID,
) (*coredata.TrustCenterDocumentAccess, error) {
var reportAccess *coredata.TrustCenterDocumentAccess
@@ -289,7 +289,7 @@ func (s TrustCenterAccessService) GetReportAccess(
reportAccess = &coredata.TrustCenterDocumentAccess{}
err = reportAccess.LoadByTrustCenterAccessIDAndReportID(ctx, conn, scope, access.ID, reportID)
err = reportAccess.LoadByTrustCenterAccessIDAndReportFileID(ctx, conn, scope, access.ID, reportFileID)
if err != nil {
if errors.Is(err, coredata.ErrResourceNotFound) {
return ErrDocumentAccessNotFound
@@ -405,7 +405,7 @@ func (s *TrustCenterAccessService) GrantByIDs(
}
if len(reportIDs) > 0 {
if err := coredata.GrantByReportIDs(ctx, tx, scope, access.ID, reportIDs, now); err != nil {
if err := coredata.GrantByReportFileIDs(ctx, tx, scope, access.ID, reportIDs, now); err != nil {
return fmt.Errorf("cannot grant report accesses: %w", err)
}
}
@@ -526,7 +526,7 @@ func (s *TrustCenterAccessService) RejectOrRevokeByIDs(
if len(reportIDs) > 0 {
shouldSendEmail = true
if err := coredata.RejectOrRevokeByReportIDs(ctx, tx, scope, access.ID, reportIDs, now); err != nil {
if err := coredata.RejectOrRevokeByReportFileIDs(ctx, tx, scope, access.ID, reportIDs, now); err != nil {
return fmt.Errorf("cannot reject/revoke report accesses: %w", err)
}
}
@@ -645,8 +645,8 @@ func extractExistingIDs(accesses coredata.TrustCenterDocumentAccesses) ([]gid.GI
documentIDs = append(documentIDs, *access.DocumentID)
}
if access.ReportID != nil {
reportIDs = append(reportIDs, *access.ReportID)
if access.ReportFileID != nil {
reportIDs = append(reportIDs, *access.ReportFileID)
}
if access.TrustCenterFileID != nil {
@@ -678,36 +678,36 @@ func reportAccessLabels(
ctx context.Context,
conn pg.Querier,
scope coredata.Scoper,
reportIDs []gid.GID,
reportFileIDs []gid.GID,
) ([]string, error) {
var reports coredata.Reports
if err := reports.LoadByIDs(ctx, conn, scope, reportIDs); err != nil {
return nil, fmt.Errorf("cannot load reports by IDs: %w", err)
var reportFiles coredata.Files
if err := reportFiles.LoadByIDs(ctx, conn, scope, reportFileIDs); err != nil {
return nil, fmt.Errorf("cannot load report files by IDs: %w", err)
}
reportByID := make(map[gid.GID]*coredata.Report, len(reports))
for _, report := range reports {
reportByID[report.ID] = report
fileByID := make(map[gid.GID]*coredata.File, len(reportFiles))
for _, f := range reportFiles {
fileByID[f.ID] = f
}
var audits coredata.Audits
if err := audits.LoadByReportIDs(ctx, conn, scope, reportIDs); err != nil {
return nil, fmt.Errorf("cannot load audits by report IDs: %w", err)
if err := audits.LoadByReportFileIDs(ctx, conn, scope, reportFileIDs); err != nil {
return nil, fmt.Errorf("cannot load audits by report file IDs: %w", err)
}
auditByReportID := make(map[gid.GID]*coredata.Audit, len(audits))
auditByFileID := make(map[gid.GID]*coredata.Audit, len(audits))
frameworkIDSet := make(map[gid.GID]struct{})
for _, audit := range audits {
if audit.ReportID == nil {
if audit.ReportFileID == nil {
continue
}
if _, exists := auditByReportID[*audit.ReportID]; exists {
if _, exists := auditByFileID[*audit.ReportFileID]; exists {
continue
}
auditByReportID[*audit.ReportID] = audit
auditByFileID[*audit.ReportFileID] = audit
frameworkIDSet[audit.FrameworkID] = struct{}{}
}
@@ -729,23 +729,23 @@ func reportAccessLabels(
}
}
labels := make([]string, 0, len(reportIDs))
labels := make([]string, 0, len(reportFileIDs))
for _, reportID := range reportIDs {
report, ok := reportByID[reportID]
for _, fileID := range reportFileIDs {
file, ok := fileByID[fileID]
if !ok {
return nil, fmt.Errorf("cannot load report %q: %w", reportID, coredata.ErrResourceNotFound)
return nil, fmt.Errorf("cannot load report file %q: %w", fileID, coredata.ErrResourceNotFound)
}
audit, ok := auditByReportID[reportID]
audit, ok := auditByFileID[fileID]
if !ok {
labels = append(labels, report.Filename)
labels = append(labels, file.FileName)
continue
}
framework, ok := frameworkByID[audit.FrameworkID]
if !ok {
labels = append(labels, report.Filename)
labels = append(labels, file.FileName)
continue
}