diff --git a/internal/servers/catalog_item_reference_checker.go b/internal/servers/catalog_item_reference_checker.go new file mode 100644 index 000000000..8a8fad25c --- /dev/null +++ b/internal/servers/catalog_item_reference_checker.go @@ -0,0 +1,46 @@ +/* +Copyright (c) 2025 Red Hat Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with the +License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software distributed under the License is distributed on an +"AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific +language governing permissions and limitations under the License. +*/ + +package servers + +import ( + "context" + "fmt" + + grpccodes "google.golang.org/grpc/codes" + grpcstatus "google.golang.org/grpc/status" + + "github.com/osac-project/fulfillment-service/internal/database/dao" +) + +// catalogItemReferenceChecker checks whether the caller has any resources that reference a catalog item. +type catalogItemReferenceChecker interface { + hasReference(ctx context.Context, catalogItemID string) (bool, error) +} + +// daoReferenceChecker implements catalogItemReferenceChecker using a GenericDAO. +type daoReferenceChecker[Resource dao.Object] struct { + resourceDao *dao.GenericDAO[Resource] +} + +func (c *daoReferenceChecker[Resource]) hasReference(ctx context.Context, catalogItemID string) (bool, error) { + filter := fmt.Sprintf("this.spec.catalog_item == %q", catalogItemID) + response, err := c.resourceDao.List(). + SetFilter(filter). + SetLimit(1). + Do(ctx) + if err != nil { + return false, grpcstatus.Errorf(grpccodes.Internal, "failed to check resource references") + } + return response.GetTotal() > 0, nil +} diff --git a/internal/servers/catalog_item_reference_checker_mock.go b/internal/servers/catalog_item_reference_checker_mock.go new file mode 100644 index 000000000..2b9b6e95f --- /dev/null +++ b/internal/servers/catalog_item_reference_checker_mock.go @@ -0,0 +1,58 @@ +/* +Copyright (c) 2025 Red Hat Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with the +License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software distributed under the License is distributed on an +"AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific +language governing permissions and limitations under the License. +*/ + +package servers + +import ( + "context" + "reflect" + + "go.uber.org/mock/gomock" +) + +type MockCatalogItemReferenceChecker struct { + ctrl *gomock.Controller + recorder *MockCatalogItemReferenceCheckerMockRecorder + isgomock struct{} +} + +type MockCatalogItemReferenceCheckerMockRecorder struct { + mock *MockCatalogItemReferenceChecker +} + +func NewMockCatalogItemReferenceChecker(ctrl *gomock.Controller) *MockCatalogItemReferenceChecker { + mock := &MockCatalogItemReferenceChecker{ctrl: ctrl} + mock.recorder = &MockCatalogItemReferenceCheckerMockRecorder{mock} + return mock +} + +func (m *MockCatalogItemReferenceChecker) EXPECT() *MockCatalogItemReferenceCheckerMockRecorder { + return m.recorder +} + +func (m *MockCatalogItemReferenceChecker) hasReference(ctx context.Context, catalogItemID string) (bool, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "hasReference", ctx, catalogItemID) + ret0, _ := ret[0].(bool) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +func (mr *MockCatalogItemReferenceCheckerMockRecorder) hasReference(ctx, catalogItemID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType( + mr.mock, "hasReference", + reflect.TypeOf((*MockCatalogItemReferenceChecker)(nil).hasReference), + ctx, catalogItemID, + ) +} diff --git a/internal/servers/catalog_item_validation.go b/internal/servers/catalog_item_validation.go index 866fd2ae5..92a98d185 100644 --- a/internal/servers/catalog_item_validation.go +++ b/internal/servers/catalog_item_validation.go @@ -15,8 +15,11 @@ package servers import ( "encoding/json" + "fmt" "strings" + "sync" + "github.com/google/cel-go/cel" "github.com/santhosh-tekuri/jsonschema/v6" grpccodes "google.golang.org/grpc/codes" grpcstatus "google.golang.org/grpc/status" @@ -27,6 +30,29 @@ import ( privatev1 "github.com/osac-project/fulfillment-service/internal/api/osac/private/v1" ) +var ( + celSyntaxEnv *cel.Env + celSyntaxEnvOnce sync.Once + celSyntaxEnvErr error +) + +// validateCELSyntax checks that a filter string is a syntactically valid, complete CEL expression. +// This prevents filter bypass attacks where a malicious filter like "true) || (true" could +// break out of parenthesized composition and change operator precedence. +func validateCELSyntax(filter string) error { + celSyntaxEnvOnce.Do(func() { + celSyntaxEnv, celSyntaxEnvErr = cel.NewEnv() + }) + if celSyntaxEnvErr != nil { + return fmt.Errorf("failed to create CEL environment: %w", celSyntaxEnvErr) + } + _, issues := celSyntaxEnv.Parse(filter) + if issues != nil && issues.Err() != nil { + return fmt.Errorf("syntax error: %w", issues.Err()) + } + return nil +} + // catalogItem is implemented by both ClusterCatalogItem and ComputeInstanceCatalogItem. type catalogItem interface { proto.Message diff --git a/internal/servers/catalog_item_validation_test.go b/internal/servers/catalog_item_validation_test.go new file mode 100644 index 000000000..6064c8417 --- /dev/null +++ b/internal/servers/catalog_item_validation_test.go @@ -0,0 +1,73 @@ +/* +Copyright (c) 2025 Red Hat Inc. + +Licensed under the Apache License, Version 2.0 (the "License"); you may not use this file except in compliance with the +License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + +Unless required by applicable law or agreed to in writing, software distributed under the License is distributed on an +"AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the License for the specific +language governing permissions and limitations under the License. +*/ + +package servers + +import ( + . "github.com/onsi/ginkgo/v2" + . "github.com/onsi/gomega" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" +) + +var _ = Describe("addPublishedFilter", func() { + var server *ClusterCatalogItemsServer + + BeforeEach(func() { + server = &ClusterCatalogItemsServer{} + }) + + DescribeTable("composes filter correctly", + func(input string, expected string) { + result, err := server.addPublishedFilter(input) + Expect(err).ToNot(HaveOccurred()) + Expect(result).To(Equal(expected)) + }, + Entry("empty filter", "", "this.published"), + Entry("simple filter", "this.id == '123'", "(this.id == '123') && this.published"), + Entry("compound filter", "this.title == 'a' && this.template == 'b'", + "(this.title == 'a' && this.template == 'b') && this.published"), + Entry("valid filter with OR is safely composed", "true || true", + "(true || true) && this.published"), + ) + + DescribeTable("rejects malformed filters", + func(input string) { + _, err := server.addPublishedFilter(input) + Expect(err).To(HaveOccurred()) + Expect(status.Code(err)).To(Equal(codes.InvalidArgument)) + }, + Entry("unbalanced parens to bypass published", `true) || (true`), + Entry("unbalanced closing paren", `true)`), + Entry("unbalanced opening paren", `(true`), + ) + + DescribeTable("validateCELSyntax", + func(input string, shouldPass bool) { + err := validateCELSyntax(input) + if shouldPass { + Expect(err).ToNot(HaveOccurred()) + } else { + Expect(err).To(HaveOccurred()) + } + }, + Entry("valid simple expression", "true", true), + Entry("valid field reference", "this.published", true), + Entry("valid comparison", "this.id == '123'", true), + Entry("valid compound", "this.a && this.b || this.c", true), + Entry("unbalanced closing paren", "true)", false), + Entry("unbalanced opening paren", "(true", false), + Entry("injection attempt", `true) || (true`, false), + Entry("empty string is not valid CEL", "", false), + ) +}) diff --git a/internal/servers/cluster_catalog_items_server.go b/internal/servers/cluster_catalog_items_server.go index c9e74c6c4..c7b188334 100644 --- a/internal/servers/cluster_catalog_items_server.go +++ b/internal/servers/cluster_catalog_items_server.go @@ -25,6 +25,7 @@ import ( privatev1 "github.com/osac-project/fulfillment-service/internal/api/osac/private/v1" publicv1 "github.com/osac-project/fulfillment-service/internal/api/osac/public/v1" "github.com/osac-project/fulfillment-service/internal/auth" + "github.com/osac-project/fulfillment-service/internal/database/dao" "github.com/osac-project/fulfillment-service/internal/events" ) @@ -41,10 +42,11 @@ var _ publicv1.ClusterCatalogItemsServer = (*ClusterCatalogItemsServer)(nil) type ClusterCatalogItemsServer struct { publicv1.UnimplementedClusterCatalogItemsServer - logger *slog.Logger - delegate privatev1.ClusterCatalogItemsServer - inMapper *GenericMapper[*publicv1.ClusterCatalogItem, *privatev1.ClusterCatalogItem] - outMapper *GenericMapper[*privatev1.ClusterCatalogItem, *publicv1.ClusterCatalogItem] + logger *slog.Logger + referenceChecker catalogItemReferenceChecker + delegate privatev1.ClusterCatalogItemsServer + inMapper *GenericMapper[*publicv1.ClusterCatalogItem, *privatev1.ClusterCatalogItem] + outMapper *GenericMapper[*privatev1.ClusterCatalogItem, *publicv1.ClusterCatalogItem] } func NewClusterCatalogItemsServer() *ClusterCatalogItemsServerBuilder { @@ -101,6 +103,16 @@ func (b *ClusterCatalogItemsServerBuilder) Build() (result *ClusterCatalogItemsS return } + clustersDao, err := dao.NewGenericDAO[*privatev1.Cluster](). + SetLogger(b.logger). + SetTenancyLogic(b.tenancyLogic). + SetMetricsRegisterer(b.metricsRegisterer). + Build() + if err != nil { + return + } + referenceChecker := &daoReferenceChecker[*privatev1.Cluster]{resourceDao: clustersDao} + delegate, err := NewPrivateClusterCatalogItemsServer(). SetLogger(b.logger). SetNotifier(b.notifier). @@ -113,10 +125,11 @@ func (b *ClusterCatalogItemsServerBuilder) Build() (result *ClusterCatalogItemsS } result = &ClusterCatalogItemsServer{ - logger: b.logger, - delegate: delegate, - inMapper: inMapper, - outMapper: outMapper, + logger: b.logger, + referenceChecker: referenceChecker, + delegate: delegate, + inMapper: inMapper, + outMapper: outMapper, } return } @@ -126,7 +139,11 @@ func (s *ClusterCatalogItemsServer) List(ctx context.Context, privateRequest := &privatev1.ClusterCatalogItemsListRequest{} privateRequest.SetOffset(request.GetOffset()) privateRequest.SetLimit(request.GetLimit()) - privateRequest.SetFilter(request.GetFilter()) + composedFilter, err := s.addPublishedFilter(request.GetFilter()) + if err != nil { + return nil, err + } + privateRequest.SetFilter(composedFilter) privateRequest.SetOrder(request.GetOrder()) privateResponse, err := s.delegate.List(ctx, privateRequest) @@ -163,6 +180,16 @@ func (s *ClusterCatalogItemsServer) Get(ctx context.Context, return nil, err } + if !privateResponse.GetObject().GetPublished() { + hasRef, refErr := s.referenceChecker.hasReference(ctx, request.GetId()) + if refErr != nil { + return nil, refErr + } + if !hasRef { + return nil, grpcstatus.Errorf(grpccodes.NotFound, "catalog item not found") + } + } + publicCatalogItem := &publicv1.ClusterCatalogItem{} err = s.outMapper.Copy(ctx, privateResponse.GetObject(), publicCatalogItem) if err != nil { @@ -259,6 +286,16 @@ func (s *ClusterCatalogItemsServer) Update(ctx context.Context, return } +func (s *ClusterCatalogItemsServer) addPublishedFilter(filter string) (string, error) { + if filter == "" { + return "this.published", nil + } + if err := validateCELSyntax(filter); err != nil { + return "", grpcstatus.Errorf(grpccodes.InvalidArgument, "invalid filter: %v", err) + } + return "(" + filter + ") && this.published", nil +} + func (s *ClusterCatalogItemsServer) Delete(ctx context.Context, request *publicv1.ClusterCatalogItemsDeleteRequest) (response *publicv1.ClusterCatalogItemsDeleteResponse, err error) { privateRequest := &privatev1.ClusterCatalogItemsDeleteRequest{} diff --git a/internal/servers/cluster_catalog_items_server_test.go b/internal/servers/cluster_catalog_items_server_test.go index d4ff648e1..e76d44cea 100644 --- a/internal/servers/cluster_catalog_items_server_test.go +++ b/internal/servers/cluster_catalog_items_server_test.go @@ -19,6 +19,9 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" + "go.uber.org/mock/gomock" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "google.golang.org/protobuf/proto" publicv1 "github.com/osac-project/fulfillment-service/internal/api/osac/public/v1" @@ -126,8 +129,9 @@ var _ = Describe("Cluster catalog items server", func() { for i := range count { _, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ Object: publicv1.ClusterCatalogItem_builder{ - Title: fmt.Sprintf("Catalog item %d", i), - Template: "my-template-id", + Title: fmt.Sprintf("Catalog item %d", i), + Template: "my-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -144,8 +148,9 @@ var _ = Describe("Cluster catalog items server", func() { for i := range count { _, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ Object: publicv1.ClusterCatalogItem_builder{ - Title: fmt.Sprintf("Catalog item %d", i), - Template: "my-template-id", + Title: fmt.Sprintf("Catalog item %d", i), + Template: "my-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -164,8 +169,9 @@ var _ = Describe("Cluster catalog items server", func() { for i := range count { response, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ Object: publicv1.ClusterCatalogItem_builder{ - Title: fmt.Sprintf("Catalog item %d", i), - Template: "my-template-id", + Title: fmt.Sprintf("Catalog item %d", i), + Template: "my-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -182,12 +188,119 @@ var _ = Describe("Cluster catalog items server", func() { } }) + It("List excludes unpublished objects", func() { + _, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ + Object: publicv1.ClusterCatalogItem_builder{ + Title: "Published item", + Template: "my-template-id", + Published: true, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + _, err = server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ + Object: publicv1.ClusterCatalogItem_builder{ + Title: "Unpublished item", + Template: "my-template-id", + Published: false, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + response, err := server.List(ctx, publicv1.ClusterCatalogItemsListRequest_builder{}.Build()) + Expect(err).ToNot(HaveOccurred()) + Expect(response.GetItems()).To(HaveLen(1)) + Expect(response.GetItems()[0].GetTitle()).To(Equal("Published item")) + }) + + It("List with user filter excludes unpublished objects", func() { + publishedResponse, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ + Object: publicv1.ClusterCatalogItem_builder{ + Title: "Target published", + Template: "my-template-id", + Published: true, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + _, err = server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ + Object: publicv1.ClusterCatalogItem_builder{ + Title: "Other published", + Template: "my-template-id", + Published: true, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + unpublishedResponse, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ + Object: publicv1.ClusterCatalogItem_builder{ + Title: "Target unpublished", + Template: "my-template-id", + Published: false, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + targetID := publishedResponse.GetObject().GetId() + unpublishedID := unpublishedResponse.GetObject().GetId() + filter := fmt.Sprintf("this.id == '%s' || this.id == '%s'", targetID, unpublishedID) + response, err := server.List(ctx, publicv1.ClusterCatalogItemsListRequest_builder{ + Filter: proto.String(filter), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + Expect(response.GetItems()).To(HaveLen(1)) + Expect(response.GetItems()[0].GetId()).To(Equal(targetID)) + }) + + It("Get returns unpublished item when caller has a referencing cluster", func() { + createResponse, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ + Object: publicv1.ClusterCatalogItem_builder{ + Title: "Unpublished item", + Template: "my-template-id", + Published: false, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + catalogItemID := createResponse.GetObject().GetId() + + mockCtrl := gomock.NewController(GinkgoT()) + mockChecker := NewMockCatalogItemReferenceChecker(mockCtrl) + mockChecker.EXPECT().hasReference(gomock.Any(), catalogItemID).Return(true, nil) + originalChecker := server.referenceChecker + server.referenceChecker = mockChecker + DeferCleanup(func() { server.referenceChecker = originalChecker }) + + getResponse, err := server.Get(ctx, publicv1.ClusterCatalogItemsGetRequest_builder{ + Id: catalogItemID, + }.Build()) + Expect(err).ToNot(HaveOccurred()) + Expect(getResponse.GetObject().GetTitle()).To(Equal("Unpublished item")) + }) + + It("Get returns not found for unpublished object", func() { + createResponse, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ + Object: publicv1.ClusterCatalogItem_builder{ + Title: "Unpublished item", + Template: "my-template-id", + Published: false, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + _, err = server.Get(ctx, publicv1.ClusterCatalogItemsGetRequest_builder{ + Id: createResponse.GetObject().GetId(), + }.Build()) + Expect(err).To(HaveOccurred()) + Expect(status.Code(err)).To(Equal(codes.NotFound)) + }) + It("Get object", func() { createResponse, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ Object: publicv1.ClusterCatalogItem_builder{ Title: "My catalog item", Description: "My description.", Template: "my-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -205,6 +318,7 @@ var _ = Describe("Cluster catalog items server", func() { Title: "Original title", Description: "Original description.", Template: "my-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -216,6 +330,7 @@ var _ = Describe("Cluster catalog items server", func() { Title: "Updated title", Description: "Updated description.", Template: "my-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -233,8 +348,9 @@ var _ = Describe("Cluster catalog items server", func() { It("Delete object", func() { createResponse, err := server.Create(ctx, publicv1.ClusterCatalogItemsCreateRequest_builder{ Object: publicv1.ClusterCatalogItem_builder{ - Title: "My catalog item", - Template: "my-template-id", + Title: "My catalog item", + Template: "my-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) diff --git a/internal/servers/compute_instance_catalog_items_server.go b/internal/servers/compute_instance_catalog_items_server.go index 793fb3004..54bdef5f5 100644 --- a/internal/servers/compute_instance_catalog_items_server.go +++ b/internal/servers/compute_instance_catalog_items_server.go @@ -25,6 +25,7 @@ import ( privatev1 "github.com/osac-project/fulfillment-service/internal/api/osac/private/v1" publicv1 "github.com/osac-project/fulfillment-service/internal/api/osac/public/v1" "github.com/osac-project/fulfillment-service/internal/auth" + "github.com/osac-project/fulfillment-service/internal/database/dao" "github.com/osac-project/fulfillment-service/internal/events" ) @@ -41,10 +42,11 @@ var _ publicv1.ComputeInstanceCatalogItemsServer = (*ComputeInstanceCatalogItems type ComputeInstanceCatalogItemsServer struct { publicv1.UnimplementedComputeInstanceCatalogItemsServer - logger *slog.Logger - delegate privatev1.ComputeInstanceCatalogItemsServer - inMapper *GenericMapper[*publicv1.ComputeInstanceCatalogItem, *privatev1.ComputeInstanceCatalogItem] - outMapper *GenericMapper[*privatev1.ComputeInstanceCatalogItem, *publicv1.ComputeInstanceCatalogItem] + logger *slog.Logger + referenceChecker catalogItemReferenceChecker + delegate privatev1.ComputeInstanceCatalogItemsServer + inMapper *GenericMapper[*publicv1.ComputeInstanceCatalogItem, *privatev1.ComputeInstanceCatalogItem] + outMapper *GenericMapper[*privatev1.ComputeInstanceCatalogItem, *publicv1.ComputeInstanceCatalogItem] } func NewComputeInstanceCatalogItemsServer() *ComputeInstanceCatalogItemsServerBuilder { @@ -101,6 +103,16 @@ func (b *ComputeInstanceCatalogItemsServerBuilder) Build() (result *ComputeInsta return } + computeInstancesDao, err := dao.NewGenericDAO[*privatev1.ComputeInstance](). + SetLogger(b.logger). + SetTenancyLogic(b.tenancyLogic). + SetMetricsRegisterer(b.metricsRegisterer). + Build() + if err != nil { + return + } + referenceChecker := &daoReferenceChecker[*privatev1.ComputeInstance]{resourceDao: computeInstancesDao} + delegate, err := NewPrivateComputeInstanceCatalogItemsServer(). SetLogger(b.logger). SetNotifier(b.notifier). @@ -113,10 +125,11 @@ func (b *ComputeInstanceCatalogItemsServerBuilder) Build() (result *ComputeInsta } result = &ComputeInstanceCatalogItemsServer{ - logger: b.logger, - delegate: delegate, - inMapper: inMapper, - outMapper: outMapper, + logger: b.logger, + referenceChecker: referenceChecker, + delegate: delegate, + inMapper: inMapper, + outMapper: outMapper, } return } @@ -126,7 +139,11 @@ func (s *ComputeInstanceCatalogItemsServer) List(ctx context.Context, privateRequest := &privatev1.ComputeInstanceCatalogItemsListRequest{} privateRequest.SetOffset(request.GetOffset()) privateRequest.SetLimit(request.GetLimit()) - privateRequest.SetFilter(request.GetFilter()) + composedFilter, err := s.addPublishedFilter(request.GetFilter()) + if err != nil { + return nil, err + } + privateRequest.SetFilter(composedFilter) privateRequest.SetOrder(request.GetOrder()) privateResponse, err := s.delegate.List(ctx, privateRequest) @@ -163,6 +180,16 @@ func (s *ComputeInstanceCatalogItemsServer) Get(ctx context.Context, return nil, err } + if !privateResponse.GetObject().GetPublished() { + hasRef, refErr := s.referenceChecker.hasReference(ctx, request.GetId()) + if refErr != nil { + return nil, refErr + } + if !hasRef { + return nil, grpcstatus.Errorf(grpccodes.NotFound, "catalog item not found") + } + } + publicCatalogItem := &publicv1.ComputeInstanceCatalogItem{} err = s.outMapper.Copy(ctx, privateResponse.GetObject(), publicCatalogItem) if err != nil { @@ -259,6 +286,16 @@ func (s *ComputeInstanceCatalogItemsServer) Update(ctx context.Context, return } +func (s *ComputeInstanceCatalogItemsServer) addPublishedFilter(filter string) (string, error) { + if filter == "" { + return "this.published", nil + } + if err := validateCELSyntax(filter); err != nil { + return "", grpcstatus.Errorf(grpccodes.InvalidArgument, "invalid filter: %v", err) + } + return "(" + filter + ") && this.published", nil +} + func (s *ComputeInstanceCatalogItemsServer) Delete(ctx context.Context, request *publicv1.ComputeInstanceCatalogItemsDeleteRequest) (response *publicv1.ComputeInstanceCatalogItemsDeleteResponse, err error) { privateRequest := &privatev1.ComputeInstanceCatalogItemsDeleteRequest{} diff --git a/internal/servers/compute_instance_catalog_items_server_test.go b/internal/servers/compute_instance_catalog_items_server_test.go index 018787d87..f357628a2 100644 --- a/internal/servers/compute_instance_catalog_items_server_test.go +++ b/internal/servers/compute_instance_catalog_items_server_test.go @@ -19,6 +19,9 @@ import ( . "github.com/onsi/ginkgo/v2" . "github.com/onsi/gomega" + "go.uber.org/mock/gomock" + "google.golang.org/grpc/codes" + "google.golang.org/grpc/status" "google.golang.org/protobuf/proto" publicv1 "github.com/osac-project/fulfillment-service/internal/api/osac/public/v1" @@ -126,8 +129,9 @@ var _ = Describe("Compute instance catalog items server", func() { for i := range count { _, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ Object: publicv1.ComputeInstanceCatalogItem_builder{ - Title: fmt.Sprintf("CI catalog item %d", i), - Template: "my-ci-template-id", + Title: fmt.Sprintf("CI catalog item %d", i), + Template: "my-ci-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -144,8 +148,9 @@ var _ = Describe("Compute instance catalog items server", func() { for i := range count { _, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ Object: publicv1.ComputeInstanceCatalogItem_builder{ - Title: fmt.Sprintf("CI catalog item %d", i), - Template: "my-ci-template-id", + Title: fmt.Sprintf("CI catalog item %d", i), + Template: "my-ci-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -164,8 +169,9 @@ var _ = Describe("Compute instance catalog items server", func() { for i := range count { response, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ Object: publicv1.ComputeInstanceCatalogItem_builder{ - Title: fmt.Sprintf("CI catalog item %d", i), - Template: "my-ci-template-id", + Title: fmt.Sprintf("CI catalog item %d", i), + Template: "my-ci-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -182,12 +188,119 @@ var _ = Describe("Compute instance catalog items server", func() { } }) + It("List excludes unpublished objects", func() { + _, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ + Object: publicv1.ComputeInstanceCatalogItem_builder{ + Title: "Published item", + Template: "my-ci-template-id", + Published: true, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + _, err = server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ + Object: publicv1.ComputeInstanceCatalogItem_builder{ + Title: "Unpublished item", + Template: "my-ci-template-id", + Published: false, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + response, err := server.List(ctx, publicv1.ComputeInstanceCatalogItemsListRequest_builder{}.Build()) + Expect(err).ToNot(HaveOccurred()) + Expect(response.GetItems()).To(HaveLen(1)) + Expect(response.GetItems()[0].GetTitle()).To(Equal("Published item")) + }) + + It("List with user filter excludes unpublished objects", func() { + publishedResponse, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ + Object: publicv1.ComputeInstanceCatalogItem_builder{ + Title: "Target published", + Template: "my-ci-template-id", + Published: true, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + _, err = server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ + Object: publicv1.ComputeInstanceCatalogItem_builder{ + Title: "Other published", + Template: "my-ci-template-id", + Published: true, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + unpublishedResponse, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ + Object: publicv1.ComputeInstanceCatalogItem_builder{ + Title: "Target unpublished", + Template: "my-ci-template-id", + Published: false, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + targetID := publishedResponse.GetObject().GetId() + unpublishedID := unpublishedResponse.GetObject().GetId() + filter := fmt.Sprintf("this.id == '%s' || this.id == '%s'", targetID, unpublishedID) + response, err := server.List(ctx, publicv1.ComputeInstanceCatalogItemsListRequest_builder{ + Filter: proto.String(filter), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + Expect(response.GetItems()).To(HaveLen(1)) + Expect(response.GetItems()[0].GetId()).To(Equal(targetID)) + }) + + It("Get returns unpublished item when caller has a referencing compute instance", func() { + createResponse, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ + Object: publicv1.ComputeInstanceCatalogItem_builder{ + Title: "Unpublished item", + Template: "my-ci-template-id", + Published: false, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + catalogItemID := createResponse.GetObject().GetId() + + mockCtrl := gomock.NewController(GinkgoT()) + mockChecker := NewMockCatalogItemReferenceChecker(mockCtrl) + mockChecker.EXPECT().hasReference(gomock.Any(), catalogItemID).Return(true, nil) + originalChecker := server.referenceChecker + server.referenceChecker = mockChecker + DeferCleanup(func() { server.referenceChecker = originalChecker }) + + getResponse, err := server.Get(ctx, publicv1.ComputeInstanceCatalogItemsGetRequest_builder{ + Id: catalogItemID, + }.Build()) + Expect(err).ToNot(HaveOccurred()) + Expect(getResponse.GetObject().GetTitle()).To(Equal("Unpublished item")) + }) + + It("Get returns not found for unpublished object", func() { + createResponse, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ + Object: publicv1.ComputeInstanceCatalogItem_builder{ + Title: "Unpublished item", + Template: "my-ci-template-id", + Published: false, + }.Build(), + }.Build()) + Expect(err).ToNot(HaveOccurred()) + + _, err = server.Get(ctx, publicv1.ComputeInstanceCatalogItemsGetRequest_builder{ + Id: createResponse.GetObject().GetId(), + }.Build()) + Expect(err).To(HaveOccurred()) + Expect(status.Code(err)).To(Equal(codes.NotFound)) + }) + It("Get object", func() { createResponse, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ Object: publicv1.ComputeInstanceCatalogItem_builder{ Title: "My CI catalog item", Description: "My description.", Template: "my-ci-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -205,6 +318,7 @@ var _ = Describe("Compute instance catalog items server", func() { Title: "Original title", Description: "Original description.", Template: "my-ci-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -216,6 +330,7 @@ var _ = Describe("Compute instance catalog items server", func() { Title: "Updated title", Description: "Updated description.", Template: "my-ci-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred()) @@ -233,8 +348,9 @@ var _ = Describe("Compute instance catalog items server", func() { It("Delete object", func() { createResponse, err := server.Create(ctx, publicv1.ComputeInstanceCatalogItemsCreateRequest_builder{ Object: publicv1.ComputeInstanceCatalogItem_builder{ - Title: "My CI catalog item", - Template: "my-ci-template-id", + Title: "My CI catalog item", + Template: "my-ci-template-id", + Published: true, }.Build(), }.Build()) Expect(err).ToNot(HaveOccurred())