diff --git a/pkg/probo/policy_service.go b/pkg/probo/policy_service.go index 15d257103..7355e19e9 100644 --- a/pkg/probo/policy_service.go +++ b/pkg/probo/policy_service.go @@ -127,13 +127,14 @@ func (s *PolicyService) Create( policyID := gid.New(s.svc.scope.GetTenantID(), coredata.PolicyEntityType) policyVersionID := gid.New(s.svc.scope.GetTenantID(), coredata.PolicyVersionEntityType) + organization := &coredata.Organization{} + people := &coredata.People{} + policy := &coredata.Policy{ - ID: policyID, - OrganizationID: req.OrganizationID, - OwnerID: req.OwnerID, - Title: req.Title, - CreatedAt: now, - UpdatedAt: now, + ID: policyID, + Title: req.Title, + CreatedAt: now, + UpdatedAt: now, } policyVersion := &coredata.PolicyVersion{ @@ -149,6 +150,17 @@ func (s *PolicyService) Create( err := s.svc.pg.WithTx( ctx, func(conn pg.Conn) error { + if err := organization.LoadByID(ctx, conn, s.svc.scope, req.OrganizationID); err != nil { + return fmt.Errorf("cannot load organization: %w", err) + } + + if err := people.LoadByID(ctx, conn, s.svc.scope, req.OwnerID); err != nil { + return fmt.Errorf("cannot load people: %w", err) + } + + policy.OrganizationID = organization.ID + policy.OwnerID = people.ID + if err := policy.Insert(ctx, conn, s.svc.scope); err != nil { return fmt.Errorf("cannot insert policy: %w", err) } diff --git a/pkg/probo/vendor_service.go b/pkg/probo/vendor_service.go index e2ed754cb..d4e3c1fe3 100644 --- a/pkg/probo/vendor_service.go +++ b/pkg/probo/vendor_service.go @@ -86,15 +86,20 @@ func (s VendorService) ListForOrganizationID( cursor *page.Cursor[coredata.VendorOrderField], ) (*page.Page[*coredata.Vendor, coredata.VendorOrderField], error) { var vendors coredata.Vendors + organization := &coredata.Organization{} err := s.svc.pg.WithConn( ctx, func(conn pg.Conn) error { + if err := organization.LoadByID(ctx, conn, s.svc.scope, organizationID); err != nil { + return fmt.Errorf("cannot load organization: %w", err) + } + return vendors.LoadByOrganizationID( ctx, conn, s.svc.scope, - organizationID, + organization.ID, cursor, ) }, @@ -192,12 +197,22 @@ func (s VendorService) Update( vendor.TrustPageURL = req.TrustPageURL } + businessOwner := &coredata.People{} + if err := businessOwner.LoadByID(ctx, conn, s.svc.scope, *req.BusinessOwnerID); err != nil { + return fmt.Errorf("cannot load business owner: %w", err) + } + if req.BusinessOwnerID != nil { - vendor.BusinessOwnerID = req.BusinessOwnerID + vendor.BusinessOwnerID = &businessOwner.ID + } + + securityOwner := &coredata.People{} + if err := securityOwner.LoadByID(ctx, conn, s.svc.scope, *req.SecurityOwnerID); err != nil { + return fmt.Errorf("cannot load security owner: %w", err) } if req.SecurityOwnerID != nil { - vendor.SecurityOwnerID = req.SecurityOwnerID + vendor.SecurityOwnerID = &securityOwner.ID } vendor.UpdatedAt = time.Now() @@ -255,42 +270,59 @@ func (s VendorService) Create( req CreateVendorRequest, ) (*coredata.Vendor, error) { now := time.Now() - vendorID := gid.New(s.svc.scope.GetTenantID(), coredata.VendorEntityType) - - organization := &coredata.Organization{} - vendor := &coredata.Vendor{ - ID: vendorID, - OrganizationID: req.OrganizationID, - Name: req.Name, - CreatedAt: now, - UpdatedAt: now, - Description: req.Description, - HeadquarterAddress: req.HeadquarterAddress, - LegalName: req.LegalName, - WebsiteURL: req.WebsiteURL, - PrivacyPolicyURL: req.PrivacyPolicyURL, - ServiceLevelAgreementURL: req.ServiceLevelAgreementURL, - DataProcessingAgreementURL: req.DataProcessingAgreementURL, - Certifications: req.Certifications, - SecurityPageURL: req.SecurityPageURL, - TrustPageURL: req.TrustPageURL, - StatusPageURL: req.StatusPageURL, - TermsOfServiceURL: req.TermsOfServiceURL, - } - - if req.Category != nil { - vendor.Category = *req.Category - } else { - vendor.Category = "Other" - } + var vendor *coredata.Vendor err := s.svc.pg.WithTx( ctx, func(conn pg.Conn) error { + vendor := &coredata.Vendor{ + ID: gid.New(s.svc.scope.GetTenantID(), coredata.VendorEntityType), + Name: req.Name, + CreatedAt: now, + UpdatedAt: now, + Description: req.Description, + HeadquarterAddress: req.HeadquarterAddress, + LegalName: req.LegalName, + WebsiteURL: req.WebsiteURL, + PrivacyPolicyURL: req.PrivacyPolicyURL, + ServiceLevelAgreementURL: req.ServiceLevelAgreementURL, + DataProcessingAgreementURL: req.DataProcessingAgreementURL, + Certifications: req.Certifications, + SecurityPageURL: req.SecurityPageURL, + TrustPageURL: req.TrustPageURL, + StatusPageURL: req.StatusPageURL, + TermsOfServiceURL: req.TermsOfServiceURL, + } + + organization := &coredata.Organization{} if err := organization.LoadByID(ctx, conn, s.svc.scope, req.OrganizationID); err != nil { return fmt.Errorf("cannot load organization %q: %w", req.OrganizationID, err) } + vendor.OrganizationID = organization.ID + + if req.BusinessOwnerID != nil { + businessOwner := &coredata.People{} + if err := businessOwner.LoadByID(ctx, conn, s.svc.scope, *req.BusinessOwnerID); err != nil { + return fmt.Errorf("cannot load business owner: %w", err) + } + vendor.BusinessOwnerID = &businessOwner.ID + } + + if req.SecurityOwnerID != nil { + securityOwner := &coredata.People{} + if err := securityOwner.LoadByID(ctx, conn, s.svc.scope, *req.SecurityOwnerID); err != nil { + return fmt.Errorf("cannot load security owner: %w", err) + } + vendor.SecurityOwnerID = &securityOwner.ID + } + + if req.Category != nil { + vendor.Category = *req.Category + } else { + vendor.Category = "Other" + } + if err := vendor.Insert(ctx, conn, s.svc.scope); err != nil { return fmt.Errorf("cannot insert vendor: %w", err) }