diff --git a/pkg/client/columns.go b/pkg/client/columns.go index 249d5d9f..2d73162b 100644 --- a/pkg/client/columns.go +++ b/pkg/client/columns.go @@ -85,11 +85,19 @@ func (c *Client) ListColumns(ctx context.Context, parentResourceID *v2.ResourceI // If the privilege is "grant", it grants SELECT, INSERT, UPDATE, and REFERENCES privileges. func (c *Client) GrantColumnPrivilege(ctx context.Context, table string, column string, user string, privilege string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format: %s", user) + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format: %s: %w", user, err) + } + userEsc, err := escapeMySQLUserHost(userName) + if err != nil { + return err + } + hostEsc, err := escapeMySQLUserHost(host) + if err != nil { + return err } - userGrant := fmt.Sprintf("%s'@'%s", userSplit[0], userSplit[1]) + userGrant := fmt.Sprintf("%s'@'%s", userEsc, hostEsc) var privileges []string if strings.ToLower(privilege) == "grant" { @@ -120,11 +128,19 @@ func (c *Client) GrantColumnPrivilege(ctx context.Context, table string, column } func (c *Client) RevokeColumnPrivilege(ctx context.Context, table string, column string, user string, privilege string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format: %s", user) + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format: %s: %w", user, err) + } + userEsc, err := escapeMySQLUserHost(userName) + if err != nil { + return err + } + hostEsc, err := escapeMySQLUserHost(host) + if err != nil { + return err } - userRevoke := fmt.Sprintf("%s'@'%s", userSplit[0], userSplit[1]) + userRevoke := fmt.Sprintf("%s'@'%s", userEsc, hostEsc) var privileges []string if strings.ToLower(privilege) == "grant" { diff --git a/pkg/client/databases.go b/pkg/client/databases.go index e58e396f..5dd211cc 100644 --- a/pkg/client/databases.go +++ b/pkg/client/databases.go @@ -74,15 +74,15 @@ func (c *Client) ListDatabases(ctx context.Context, pager *Pager) ([]*DbModel, s } func (c *Client) GrantDatabasePrivilege(ctx context.Context, database string, user string, privilege string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } @@ -99,15 +99,15 @@ func (c *Client) GrantDatabasePrivilege(ctx context.Context, database string, us } func (c *Client) RevokeDatabasePrivilege(ctx context.Context, database string, user string, privilege string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } diff --git a/pkg/client/helper.go b/pkg/client/helper.go index 0b29f19c..39f49ea9 100644 --- a/pkg/client/helper.go +++ b/pkg/client/helper.go @@ -20,8 +20,14 @@ func escapeMySQLIdent(ident string) (string, error) { return strings.Join(parts, "."), nil } -// Helper for user/host. -var validUserHost = regexp.MustCompile(`^[a-zA-Z0-9_%\\.\\-]+$`) +// Helper for user/host. Empty is allowed: every caller derives ident via +// SplitUserHost, which guarantees the host half is always non-empty, so an +// empty string here can only be the username of MySQL's anonymous account +// (''@'host'). ":" and "/" are allowed because MySQL host specs include IPv6 +// literals (the stock root@::1) and netmask forms (198.51.100.0/255.255.255.0); +// both are inert inside the single-quoted '%s'@'%s' the callers build. "'" and +// "\" stay excluded, as those are what could break out of that quoting. +var validUserHost = regexp.MustCompile(`^[a-zA-Z0-9_%.@:/\-]*$`) func escapeMySQLUserHost(ident string) (string, error) { if !validUserHost.MatchString(ident) { @@ -29,3 +35,16 @@ func escapeMySQLUserHost(ident string) (string, error) { } return ident, nil } + +// SplitUserHost splits a "name@host" identifier into its name and host parts. +// Names (MySQL usernames or role names) may themselves legally contain "@", +// but MySQL host specifications (hostnames, IPs, netmasks, or "%" wildcards) +// never do, so splitting on the last "@" unambiguously recovers both parts. +func SplitUserHost(s string) (string, string, error) { + idx := strings.LastIndex(s, "@") + // An empty name is valid: MySQL's anonymous account is ''@'host'. + if idx < 0 || idx == len(s)-1 { + return "", "", fmt.Errorf("invalid user@host format: %s", s) + } + return s[:idx], s[idx+1:], nil +} diff --git a/pkg/client/helper_test.go b/pkg/client/helper_test.go new file mode 100644 index 00000000..5e118c02 --- /dev/null +++ b/pkg/client/helper_test.go @@ -0,0 +1,118 @@ +package client + +import ( + "testing" +) + +func Test_SplitUserHost(t *testing.T) { + type want struct { + user string + host string + } + tests := []struct { + name string + in string + want want + wantErr bool + }{ + { + name: "simple user and host", + in: "someone@%", + want: want{user: "someone", host: "%"}, + }, + { + name: "username containing @", + in: "someone@orion.com@%", + want: want{user: "someone@orion.com", host: "%"}, + }, + { + name: "username containing multiple @", + in: "a@b@c@10.0.0.1", + want: want{user: "a@b@c", host: "10.0.0.1"}, + }, + { + name: "collapsed comma-separated hosts", + in: "someone@orion.com@localhost,%", + want: want{user: "someone@orion.com", host: "localhost,%"}, + }, + { + name: "ipv6 loopback host", + in: "root@::1", + want: want{user: "root", host: "::1"}, + }, + { + name: "username with @ and ipv6 host", + in: "someone@orion.com@::1", + want: want{user: "someone@orion.com", host: "::1"}, + }, + { + name: "netmask host", + in: "someone@198.51.100.0/255.255.255.0", + want: want{user: "someone", host: "198.51.100.0/255.255.255.0"}, + }, + { + name: "no @", + in: "someone", + wantErr: true, + }, + { + name: "empty user (MySQL anonymous account)", + in: "@%", + want: want{user: "", host: "%"}, + }, + { + name: "empty host", + in: "someone@", + wantErr: true, + }, + { + name: "empty string", + in: "", + wantErr: true, + }, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + user, host, err := SplitUserHost(tt.in) + if (err != nil) != tt.wantErr { + t.Fatalf("SplitUserHost(%q) error = %v, wantErr %v", tt.in, err, tt.wantErr) + } + if tt.wantErr { + return + } + if user != tt.want.user || host != tt.want.host { + t.Errorf("SplitUserHost(%q) = (%q, %q), want (%q, %q)", tt.in, user, host, tt.want.user, tt.want.host) + } + }) + } +} + +func Test_escapeMySQLUserHost(t *testing.T) { + tests := []struct { + name string + in string + wantErr bool + }{ + {name: "empty (MySQL anonymous account username)", in: ""}, + {name: "plain username", in: "someone"}, + {name: "username with @", in: "someone@orion.com"}, + {name: "wildcard host", in: "%"}, + {name: "hostname", in: "%.example.com"}, + {name: "ipv4 host", in: "127.0.0.1"}, + {name: "ipv6 loopback host", in: "::1"}, + {name: "ipv6 full host", in: "2001:db8::8a2e:370:7334"}, + {name: "netmask host", in: "198.51.100.0/255.255.255.0"}, + {name: "wildcard octet host", in: "198.51.100.%"}, + {name: "quote injection attempt", in: "someone' OR '1'='1", wantErr: true}, + {name: "space", in: "some one", wantErr: true}, + {name: "trailing backslash", in: `someone\`, wantErr: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + _, err := escapeMySQLUserHost(tt.in) + if (err != nil) != tt.wantErr { + t.Errorf("escapeMySQLUserHost(%q) error = %v, wantErr %v", tt.in, err, tt.wantErr) + } + }) + } +} diff --git a/pkg/client/roles.go b/pkg/client/roles.go index 4bd53e99..23246192 100644 --- a/pkg/client/roles.go +++ b/pkg/client/roles.go @@ -3,33 +3,32 @@ package client import ( "context" "fmt" - "strings" ) func (c *Client) GrantRolePrivilege(ctx context.Context, role, user, privilege string) error { - roleParts := strings.Split(role, "@") - if len(roleParts) != 2 { - return fmt.Errorf("invalid role format: %s", role) + roleName, roleHostRaw, err := SplitUserHost(role) + if err != nil { + return fmt.Errorf("invalid role format: %s: %w", role, err) } - userParts := strings.Split(user, "@") - if len(userParts) != 2 { - return fmt.Errorf("invalid user format: %s", user) + userName, userHostRaw, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format: %s: %w", user, err) } - roleUser, err := escapeMySQLUserHost(roleParts[0]) + roleUser, err := escapeMySQLUserHost(roleName) if err != nil { return err } - roleHost, err := escapeMySQLUserHost(roleParts[1]) + roleHost, err := escapeMySQLUserHost(roleHostRaw) if err != nil { return err } - targetUser, err := escapeMySQLUserHost(userParts[0]) + targetUser, err := escapeMySQLUserHost(userName) if err != nil { return err } - targetHost, err := escapeMySQLUserHost(userParts[1]) + targetHost, err := escapeMySQLUserHost(userHostRaw) if err != nil { return err } @@ -53,29 +52,29 @@ func (c *Client) GrantRolePrivilege(ctx context.Context, role, user, privilege s } func (c *Client) RevokeRolePrivilege(ctx context.Context, role, user, privilege string) error { - roleParts := strings.Split(role, "@") - if len(roleParts) != 2 { - return fmt.Errorf("invalid role format: %s", role) + roleName, roleHostRaw, err := SplitUserHost(role) + if err != nil { + return fmt.Errorf("invalid role format: %s: %w", role, err) } - userParts := strings.Split(user, "@") - if len(userParts) != 2 { - return fmt.Errorf("invalid user format: %s", user) + userName, userHostRaw, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format: %s: %w", user, err) } - roleUser, err := escapeMySQLUserHost(roleParts[0]) + roleUser, err := escapeMySQLUserHost(roleName) if err != nil { return err } - roleHost, err := escapeMySQLUserHost(roleParts[1]) + roleHost, err := escapeMySQLUserHost(roleHostRaw) if err != nil { return err } - targetUser, err := escapeMySQLUserHost(userParts[0]) + targetUser, err := escapeMySQLUserHost(userName) if err != nil { return err } - targetHost, err := escapeMySQLUserHost(userParts[1]) + targetHost, err := escapeMySQLUserHost(userHostRaw) if err != nil { return err } diff --git a/pkg/client/routines.go b/pkg/client/routines.go index c0b72875..9998955c 100644 --- a/pkg/client/routines.go +++ b/pkg/client/routines.go @@ -98,15 +98,15 @@ func (c *Client) GrantRoutinePrivilege(ctx context.Context, privilege string, sc return err } - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } @@ -134,15 +134,15 @@ func (c *Client) RevokeRoutinePrivilege(ctx context.Context, privilege string, s return err } - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } diff --git a/pkg/client/servers.go b/pkg/client/servers.go index 42101df1..a7a80c66 100644 --- a/pkg/client/servers.go +++ b/pkg/client/servers.go @@ -40,15 +40,15 @@ func (c *Client) ExecContext(ctx context.Context, query string) (sql.Result, err } func (c *Client) GrantServerPrivilege(ctx context.Context, user string, privilege string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } @@ -60,15 +60,15 @@ func (c *Client) GrantServerPrivilege(ctx context.Context, user string, privileg } func (c *Client) RevokeServerPrivilege(ctx context.Context, user string, privilege string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } diff --git a/pkg/client/tables.go b/pkg/client/tables.go index 0c2f0dd5..21267f1d 100644 --- a/pkg/client/tables.go +++ b/pkg/client/tables.go @@ -83,15 +83,15 @@ func (c *Client) ListTables(ctx context.Context, parentResourceID *v2.ResourceId } func (c *Client) GrantTablePrivilege(ctx context.Context, table string, user string, privilege string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } @@ -108,15 +108,15 @@ func (c *Client) GrantTablePrivilege(ctx context.Context, table string, user str } func (c *Client) RevokeTablePrivilege(ctx context.Context, table string, user string, privilege string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } diff --git a/pkg/client/users.go b/pkg/client/users.go index ee3b177a..10ad4cfa 100644 --- a/pkg/client/users.go +++ b/pkg/client/users.go @@ -228,15 +228,15 @@ func (c *Client) GetHost(ctx context.Context) (string, error) { } func (c *Client) CreateUser(ctx context.Context, user string, password string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } @@ -248,15 +248,15 @@ func (c *Client) CreateUser(ctx context.Context, user string, password string) e } func (c *Client) DropUser(ctx context.Context, user string) error { - userSplit := strings.Split(user, "@") - if len(userSplit) != 2 { - return fmt.Errorf("invalid user format, expected user@host") + userName, host, err := SplitUserHost(user) + if err != nil { + return fmt.Errorf("invalid user format, expected user@host: %w", err) } - userEsc, err := escapeMySQLUserHost(userSplit[0]) + userEsc, err := escapeMySQLUserHost(userName) if err != nil { return err } - hostEsc, err := escapeMySQLUserHost(userSplit[1]) + hostEsc, err := escapeMySQLUserHost(host) if err != nil { return err } diff --git a/pkg/connector/grants.go b/pkg/connector/grants.go index 71369500..7721b66d 100644 --- a/pkg/connector/grants.go +++ b/pkg/connector/grants.go @@ -22,19 +22,18 @@ func grantsForUserOrRole( var ret []*v2.Grant grantMap := make(map[string]struct{}) - parts := strings.Split(strings.TrimPrefix(resource.Id.Resource, fmt.Sprintf("%s:", resource.Id.ResourceType)), "@") - if len(parts) != 2 { - return nil, fmt.Errorf("malformed principal ID") + idStr := strings.TrimPrefix(resource.Id.Resource, fmt.Sprintf("%s:", resource.Id.ResourceType)) + user, hostPart, err := client.SplitUserHost(idStr) + if err != nil { + return nil, fmt.Errorf("malformed principal ID: %w", err) } - user := parts[0] - hosts := []string{parts[1]} + hosts := []string{hostPart} // If we are collapsing users, we will want to split the host portion of the ID to inspect each real user's grants if collapseUsers { - hosts = strings.Split(parts[1], ",") + hosts = strings.Split(hostPart, ",") } - var err error for _, host := range hosts { err = listGlobalGrants(ctx, resource.ParentResourceId, user, host, grantMap, c) if err != nil { diff --git a/pkg/connector/user.go b/pkg/connector/user.go index cedef0dc..e53e25b2 100644 --- a/pkg/connector/user.go +++ b/pkg/connector/user.go @@ -41,32 +41,11 @@ func (s *userSyncer) List( var ret []*v2.Resource for _, u := range users { - var annos annotations.Annotations - - ut, err := rs.NewUserTrait( - rs.WithUserProfile(map[string]interface{}{ - "user": u.User, - "host": u.Host, - "first_name": fmt.Sprintf("%s@%s", u.User, u.Host), - "user_id": fmt.Sprintf("%s@%s", u.User, u.Host), - }), - rs.WithUserLogin(u.User), - rs.WithStatus(v2.UserTrait_Status_STATUS_ENABLED), - ) + resource, err := parseIntoUserResource(u, parentResourceID) if err != nil { return nil, "", nil, err } - annos.Update(ut) - - ret = append(ret, &v2.Resource{ - DisplayName: fmt.Sprintf("%s@%s", u.User, u.Host), - Id: &v2.ResourceId{ - ResourceType: s.resourceType.Id, - Resource: u.GetID(), - }, - Annotations: annos, - ParentResourceId: parentResourceID, - }) + ret = append(ret, resource) } return ret, nextPageToken, nil, nil @@ -128,6 +107,12 @@ func (o *userSyncer) CreateAccount( if !ok { return nil, nil, nil, fmt.Errorf("missing or invalid 'username' in profile") } + // An empty username would create MySQL's anonymous account (''@'host'). + // That is a legitimate entity to read and delete, but never something we + // should provision on request. + if username == "" { + return nil, nil, nil, fmt.Errorf("baton-mysql: 'username' in profile must not be empty") + } host, err := o.client.GetHost(ctx) if err != nil { @@ -145,10 +130,13 @@ func (o *userSyncer) CreateAccount( return nil, nil, nil, fmt.Errorf("create user failed: %w", err) } - // Build resource + // Build resource. UserType must be set: GetID() renders it as the + // ":@" prefix, and omitting it yields ":user@host", + // which every consumer of the composite ID then fails to parse. user := &client.User{ - User: username, - Host: host, + UserType: client.UserType, + User: username, + Host: host, } userResource, err := parseIntoUserResource(user, nil) if err != nil { @@ -168,47 +156,35 @@ func (o *userSyncer) CreateAccount( } func parseIntoUserResource(user *client.User, parent *v2.ResourceId) (*v2.Resource, error) { - ut, err := rs.NewUserTrait( - rs.WithUserProfile(map[string]interface{}{ + return rs.NewUserResource( + fmt.Sprintf("%s@%s", user.User, user.Host), + resourceTypeUser, + user.GetID(), + []rs.UserTraitOption{rs.WithUserLogin(user.User)}, + rs.WithParentResourceID(parent), + rs.WithResourceProfile(map[string]interface{}{ "user": user.User, "host": user.Host, "first_name": fmt.Sprintf("%s@%s", user.User, user.Host), "user_id": fmt.Sprintf("%s@%s", user.User, user.Host), }), - rs.WithUserLogin(user.User), - rs.WithStatus(v2.UserTrait_Status_STATUS_ENABLED), + rs.WithResourceStatus(v2.Status_RESOURCE_STATUS_ENABLED, ""), ) - if err != nil { - return nil, err - } - - annos := annotations.Annotations{} - annos.Update(ut) - - return &v2.Resource{ - DisplayName: fmt.Sprintf("%s@%s", user.User, user.Host), - Id: &v2.ResourceId{ - ResourceType: resourceTypeUser.Id, - Resource: user.GetID(), - }, - Annotations: annos, - ParentResourceId: parent, - }, nil } func (s *userSyncer) Delete(ctx context.Context, resourceId *v2.ResourceId) (annotations.Annotations, error) { if resourceId.ResourceType != resourceTypeUser.Id { return nil, fmt.Errorf("baton-mysql: non-user resource passed to user delete") } - userID := strings.TrimSpace(strings.Split(resourceId.Resource, ":")[1]) - parts := strings.Split(userID, "@") - if len(parts) != 2 { - return nil, fmt.Errorf("baton-mysql: invalid user ID format, expected 'user@host'") + userID := strings.TrimSpace(strings.TrimPrefix(resourceId.Resource, fmt.Sprintf("%s:", resourceId.ResourceType))) + userPart, hostPart, err := client.SplitUserHost(userID) + if err != nil { + return nil, fmt.Errorf("baton-mysql: invalid user ID format, expected 'user@host': %w", err) } - user, host := strings.TrimSpace(parts[0]), strings.TrimSpace(parts[1]) + user, host := strings.TrimSpace(userPart), strings.TrimSpace(hostPart) userStr := fmt.Sprintf("%s@%s", user, host) - err := s.client.DropUser(ctx, userStr) + err = s.client.DropUser(ctx, userStr) if err != nil { return nil, fmt.Errorf("drop user failed: %w", err) }