diff --git a/go.mod b/go.mod index e73807eb..bb05e235 100644 --- a/go.mod +++ b/go.mod @@ -52,9 +52,9 @@ require ( golang.org/x/net v0.58.0 golang.org/x/sync v0.22.0 golang.org/x/term v0.45.0 - google.golang.org/genproto/googleapis/rpc v0.0.0-20260729162451-8efbd57d26e0 - google.golang.org/grpc v1.83.0 - google.golang.org/protobuf v1.36.11 + google.golang.org/genproto/googleapis/rpc v0.0.0-20260831171406-18b4a7587f8a + google.golang.org/grpc v1.83.2 + google.golang.org/protobuf v1.36.12 gopkg.in/yaml.v3 v3.0.1 ) diff --git a/go.sum b/go.sum index 3640bd48..ab29f9d9 100644 --- a/go.sum +++ b/go.sum @@ -1121,16 +1121,16 @@ google.golang.org/genproto v0.0.0-20260729162451-8efbd57d26e0 h1:xJf8e9ReUqiexuI google.golang.org/genproto v0.0.0-20260729162451-8efbd57d26e0/go.mod h1:0MNk3ibJAyOwDZVlp18knQe3jRyFxpcU06xjxhgjx0M= google.golang.org/genproto/googleapis/api v0.0.0-20260729162451-8efbd57d26e0 h1:ybvH/ZpOcpCrjtkb7oW/fdlzbEmRVeumw19SRQmNFKU= google.golang.org/genproto/googleapis/api v0.0.0-20260729162451-8efbd57d26e0/go.mod h1:HJ9MpJLeDSstBkx1LILTpd5f41ADSMZcTPypw02qEGw= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260729162451-8efbd57d26e0 h1:mJiOtnGp0k/BcSgdu03G2NwnscCfCH+h2QKUBZr18KI= -google.golang.org/genproto/googleapis/rpc v0.0.0-20260729162451-8efbd57d26e0/go.mod h1:4Hqkh8ycfw05ld/3BWL7rJOSfebL2Q+DVDeRgYgxUU8= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260831171406-18b4a7587f8a h1:3Dnd1cDaZlB68lziofO+bJXpjOy8UfRv8Unt+yH8tQ4= +google.golang.org/genproto/googleapis/rpc v0.0.0-20260831171406-18b4a7587f8a/go.mod h1:DjtHYE8FKJLivXcBEjGwndXfIC23G0VpXiXKqG179uA= google.golang.org/grpc v1.19.0/go.mod h1:mqu4LbDTu4XGKhr4mRzUsmM4RtVoemTSY81AxZiDr8c= google.golang.org/grpc v1.23.0/go.mod h1:Y5yQAOtifL1yxbo5wqy6BxZv8vAUGQwXBOALyacEbxg= google.golang.org/grpc v1.25.1/go.mod h1:c3i+UQWmh7LiEpx4sFZnkU36qjEYZ0imhYfXVyQciAY= google.golang.org/grpc v1.27.0/go.mod h1:qbnxyOmOxrQa7FizSgH+ReBfzJrCY1pSN7KXBS8abTk= google.golang.org/grpc v1.29.1/go.mod h1:itym6AZVZYACWQqET3MqgPpjcuV5QH3BxFS3IjizoKk= google.golang.org/grpc v1.33.2/go.mod h1:JMHMWHQWaTccqQQlmk3MJZS+GWXOdAesneDmEnv2fbc= -google.golang.org/grpc v1.83.0 h1:JeNZEKJFbQxArAMl+hiytHauacDNqJUllNfmIMmpqnQ= -google.golang.org/grpc v1.83.0/go.mod h1:kDyl6SKsiHKt0uylY5gtn5cEjkrIOhQOGDgIc4JGwzQ= +google.golang.org/grpc v1.83.2 h1:EManeRomTObA0BU7I8vXgg/78uE5MJ9M8B39EX2WscU= +google.golang.org/grpc v1.83.2/go.mod h1:YPI1hK3kDked6iHvgX3tR0y+nX/qpMFKhPgFsokw1S8= google.golang.org/protobuf v0.0.0-20200109180630-ec00e32a8dfd/go.mod h1:DFci5gLYBciE7Vtevhsrf46CRTquxDuWsQurQQe4oz8= google.golang.org/protobuf v0.0.0-20200221191635-4d8936d0db64/go.mod h1:kwYJMbMJ01Woi6D6+Kah6886xMZcty6N08ah7+eCXa0= google.golang.org/protobuf v0.0.0-20200228230310-ab0ca4ff8a60/go.mod h1:cfTl7dwQJ+fmap5saPgwCLgHXTUD7jkjRqWcaiX5VyM= @@ -1140,8 +1140,8 @@ google.golang.org/protobuf v1.22.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2 google.golang.org/protobuf v1.23.0/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU= google.golang.org/protobuf v1.23.1-0.20200526195155-81db48ad09cc/go.mod h1:EGpADcykh3NcUnDUJcl1+ZksZNG86OlYog2l/sGQquU= google.golang.org/protobuf v1.25.0/go.mod h1:9JNX74DMeImyA3h4bdi1ymwjUzf21/xIlbajtzgsN7c= -google.golang.org/protobuf v1.36.11 h1:fV6ZwhNocDyBLK0dj+fg8ektcVegBBuEolpbTQyBNVE= -google.golang.org/protobuf v1.36.11/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= +google.golang.org/protobuf v1.36.12 h1:pJOKDDOyeXErUroCihFAd5LQuwXBSpVnKGrj5o/fwxc= +google.golang.org/protobuf v1.36.12/go.mod h1:HTf+CrKn2C3g5S8VImy6tdcUvCska2kB7j23XfzDpco= gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20180628173108-788fd7840127/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= gopkg.in/check.v1 v1.0.0-20190902080502-41f04d3bba15/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0= diff --git a/internal/cmd/cmd_test.go b/internal/cmd/cmd_test.go index 48413e10..54bc4ed1 100644 --- a/internal/cmd/cmd_test.go +++ b/internal/cmd/cmd_test.go @@ -50,6 +50,34 @@ func TestCommandOutput(t *testing.T) { flagErrorContains: "accepts 3 arg(s), received 0", expectUsageContains: "zed permission check ", }, + { + name: "prints usage when a subject is passed where a type is expected in an LR", + command: []string{"zed", "perm", "lookup-resources", "test/user:jimmy", "view", "test/resource"}, + expectFlagErrorCalled: true, + flagErrorContains: `invalid resource type "test/user:jimmy": expected format ` + "``", + expectUsageContains: "zed permission lookup-resources ", + }, + { + name: "prints usage on a subject missing an object ID in an LR", + command: []string{"zed", "perm", "lookup-resources", "test/resource", "view", "test/user"}, + expectFlagErrorCalled: true, + flagErrorContains: `invalid subject "test/user": expected format ` + "`:` or `:#`", + expectUsageContains: "zed permission lookup-resources ", + }, + { + name: "prints usage on a malformed resource in a check call", + command: []string{"zed", "perm", "check", "test/resource", "view", "test/user:jimmy"}, + expectFlagErrorCalled: true, + flagErrorContains: `invalid resource "test/resource": expected format ` + "`:`", + expectUsageContains: "zed permission check ", + }, + { + name: "prints usage when a subject is passed where a subject type is expected in LS", + command: []string{"zed", "perm", "lookup-subjects", "test/resource:someresource", "view", "test/user:jimmy"}, + expectFlagErrorCalled: true, + flagErrorContains: `invalid subject type "test/user:jimmy": expected format ` + "`` or `#`", + expectUsageContains: "zed permission lookup-subjects ", + }, { name: "does not print usage on command error", command: []string{"zed", "validate", uuid.NewString()}, diff --git a/internal/commands/permission.go b/internal/commands/permission.go index 2b284974..3f828667 100644 --- a/internal/commands/permission.go +++ b/internal/commands/permission.go @@ -8,7 +8,6 @@ import ( "strings" "github.com/jzelinskie/cobrautil/v2" - "github.com/jzelinskie/stringz" "github.com/rs/zerolog/log" "github.com/spf13/cobra" "github.com/spf13/pflag" @@ -150,8 +149,7 @@ func RegisterPermissionCmd(rootCmd *cobra.Command) *cobra.Command { } func checkCmdFunc(cmd *cobra.Command, args []string) error { - var objectNS, objectID string - err := stringz.SplitExact(args[0], ":", &objectNS, &objectID) + objectNS, objectID, err := ParseResource(args[0]) if err != nil { return err } @@ -362,8 +360,7 @@ func checkBulkCmdFunc(cmd *cobra.Command, args []string) error { func expandCmdFunc(cmd *cobra.Command, args []string) error { relation := args[0] - var objectNS, objectID string - err := stringz.SplitExact(args[1], ":", &objectNS, &objectID) + objectNS, objectID, err := ParseResource(args[1]) if err != nil { return err } @@ -413,8 +410,13 @@ func expandCmdFunc(cmd *cobra.Command, args []string) error { var newLookupResourcesPageCallbackForTests func(readByPage uint) func lookupResourcesCmdFunc(cmd *cobra.Command, args []string) error { - objectNS := args[0] + objectNS, err := ParseResourceType(args[0]) + if err != nil { + return err + } + relation := args[1] + subjectNS, subjectID, subjectRel, err := ParseSubject(args[2]) if err != nil { return err @@ -546,15 +548,17 @@ func handleLookupResourcesErr(err error) error { } func lookupSubjectsCmdFunc(cmd *cobra.Command, args []string) error { - var objectNS, objectID string - err := stringz.SplitExact(args[0], ":", &objectNS, &objectID) + objectNS, objectID, err := ParseResource(args[0]) if err != nil { return err } permission := args[1] - subjectType, subjectRelation := ParseType(args[2]) + subjectType, subjectRelation, err := ParseType(args[2]) + if err != nil { + return err + } caveatContext, err := GetCaveatContext(cmd) if err != nil { diff --git a/internal/commands/relationship.go b/internal/commands/relationship.go index bd6bc113..e8866a23 100644 --- a/internal/commands/relationship.go +++ b/internal/commands/relationship.go @@ -12,7 +12,6 @@ import ( "unicode" "github.com/jzelinskie/cobrautil/v2" - "github.com/jzelinskie/stringz" "github.com/rs/zerolog/log" "github.com/spf13/cobra" "google.golang.org/genproto/googleapis/rpc/errdetails" @@ -212,11 +211,11 @@ func buildRelationshipsFilter(cmd *cobra.Command, args []string) (*v1.Relationsh filter := &v1.RelationshipFilter{ResourceType: args[0]} if strings.Contains(args[0], ":") { - var resourceID string - err := stringz.SplitExact(args[0], ":", &filter.ResourceType, &resourceID) + resourceType, resourceID, err := ParseResource(args[0]) if err != nil { return nil, err } + filter.ResourceType = resourceType if strings.HasSuffix(resourceID, "%") { filter.OptionalResourceIdPrefix = strings.TrimSuffix(resourceID, "%") diff --git a/internal/commands/util.go b/internal/commands/util.go index eafbd03a..c4d08b6a 100644 --- a/internal/commands/util.go +++ b/internal/commands/util.go @@ -17,25 +17,67 @@ import ( "github.com/authzed/authzed-go/pkg/requestmeta" ) +const ( + resourceFormat = "`:`" + subjectFormat = "`:` or `:#`" + resourceTypeFormat = "``" + subjectTypeFormat = "`` or `#`" +) + +// invalidArgError returns a ValidationError describing an argument that is not +// in the expected format, so that the command's usage is printed alongside it. +func invalidArgError(kind, value, format string) error { + return ValidationError{error: fmt.Errorf("invalid %s %q: expected format %s", kind, value, format)} +} + +// ParseResource parses the given resource string into its object type and +// object ID, if valid. +func ParseResource(s string) (objectType, objectID string, err error) { + if err := stringz.SplitExact(s, ":", &objectType, &objectID); err != nil { + return "", "", invalidArgError("resource", s, resourceFormat) + } + if objectType == "" || objectID == "" { + return "", "", invalidArgError("resource", s, resourceFormat) + } + return objectType, objectID, nil +} + // ParseSubject parses the given subject string into its namespace, object ID // and relation, if valid. func ParseSubject(s string) (namespace, id, relation string, err error) { - err = stringz.SplitExact(s, ":", &namespace, &id) - if err != nil { - return namespace, id, relation, err + if err := stringz.SplitExact(s, ":", &namespace, &id); err != nil { + return "", "", "", invalidArgError("subject", s, subjectFormat) } - err = stringz.SplitExact(id, "#", &id, &relation) - if err != nil { - relation = "" - err = nil + if strings.Contains(id, "#") { + if err := stringz.SplitExact(id, "#", &id, &relation); err != nil || relation == "" { + return "", "", "", invalidArgError("subject", s, subjectFormat) + } + } + if namespace == "" || id == "" { + return "", "", "", invalidArgError("subject", s, subjectFormat) + } + return namespace, id, relation, nil +} + +// ParseResourceType parses a bare type reference of the form `namespace`. +func ParseResourceType(s string) (namespace string, err error) { + if s == "" || strings.ContainsAny(s, ":#") { + return "", invalidArgError("resource type", s, resourceTypeFormat) } - return namespace, id, relation, err + return s, nil } -// ParseType parses a type reference of the form `namespace#relaion`. -func ParseType(s string) (namespace, relation string) { +// ParseType parses a type reference of the form `namespace#relation`, where the +// relation is optional. +func ParseType(s string) (namespace, relation string, err error) { + if strings.Contains(s, ":") { + return "", "", invalidArgError("subject type", s, subjectTypeFormat) + } namespace, relation, _ = strings.Cut(s, "#") - return namespace, relation + if namespace == "" { + return "", "", invalidArgError("subject type", s, subjectTypeFormat) + } + return namespace, relation, nil } // GetCaveatContext returns the entered caveat caveat, if any. diff --git a/internal/commands/util_test.go b/internal/commands/util_test.go index d2e2964c..813b0d2a 100644 --- a/internal/commands/util_test.go +++ b/internal/commands/util_test.go @@ -7,6 +7,133 @@ import ( "github.com/stretchr/testify/require" ) +func TestParseResource(t *testing.T) { + tests := []struct { + input string + wantObjectType string + wantObjectID string + wantErr bool + }{ + {input: "test/resource:foo", wantObjectType: "test/resource", wantObjectID: "foo"}, + {input: "resource:foo", wantObjectType: "resource", wantObjectID: "foo"}, + {input: "test/resource", wantErr: true}, + {input: "test/resource:foo:bar", wantErr: true}, + {input: "test/resource:", wantErr: true}, + {input: ":foo", wantErr: true}, + {input: "", wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + objectType, objectID, err := ParseResource(tt.input) + if tt.wantErr { + requireValidationError(t, err, "invalid resource") + return + } + require.NoError(t, err) + require.Equal(t, tt.wantObjectType, objectType) + require.Equal(t, tt.wantObjectID, objectID) + }) + } +} + +func TestParseSubject(t *testing.T) { + tests := []struct { + input string + wantNamespace string + wantID string + wantRelation string + wantErr bool + }{ + {input: "test/user:jimmy", wantNamespace: "test/user", wantID: "jimmy"}, + {input: "test/group:eng#member", wantNamespace: "test/group", wantID: "eng", wantRelation: "member"}, + {input: "test/user", wantErr: true}, + {input: "test/user:jimmy:extra", wantErr: true}, + {input: "test/user:", wantErr: true}, + {input: ":jimmy", wantErr: true}, + {input: "test/group:eng#", wantErr: true}, + {input: "test/group:eng#member#extra", wantErr: true}, + {input: "", wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + namespace, id, relation, err := ParseSubject(tt.input) + if tt.wantErr { + requireValidationError(t, err, "invalid subject") + return + } + require.NoError(t, err) + require.Equal(t, tt.wantNamespace, namespace) + require.Equal(t, tt.wantID, id) + require.Equal(t, tt.wantRelation, relation) + }) + } +} + +func TestParseResourceType(t *testing.T) { + tests := []struct { + input string + wantErr bool + }{ + {input: "test/resource"}, + {input: "resource"}, + {input: "test/resource:foo", wantErr: true}, + {input: "test/resource#viewer", wantErr: true}, + {input: "", wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + namespace, err := ParseResourceType(tt.input) + if tt.wantErr { + requireValidationError(t, err, "invalid resource type") + return + } + require.NoError(t, err) + require.Equal(t, tt.input, namespace) + }) + } +} + +func TestParseType(t *testing.T) { + tests := []struct { + input string + wantNamespace string + wantRelation string + wantErr bool + }{ + {input: "test/user", wantNamespace: "test/user"}, + {input: "test/group#member", wantNamespace: "test/group", wantRelation: "member"}, + {input: "test/user:jimmy", wantErr: true}, + {input: "#member", wantErr: true}, + {input: "", wantErr: true}, + } + + for _, tt := range tests { + t.Run(tt.input, func(t *testing.T) { + namespace, relation, err := ParseType(tt.input) + if tt.wantErr { + requireValidationError(t, err, "invalid subject type") + return + } + require.NoError(t, err) + require.Equal(t, tt.wantNamespace, namespace) + require.Equal(t, tt.wantRelation, relation) + }) + } +} + +// requireValidationError asserts that the given error is a ValidationError, so +// that the command's usage is printed, and that it names the offending argument. +func requireValidationError(t *testing.T, err error, contains string) { + t.Helper() + + var validationError ValidationError + require.ErrorAs(t, err, &validationError) + require.ErrorContains(t, err, contains) +} + func TestValidationWrapper(t *testing.T) { tests := []struct { name string