diff --git a/cmd/fleetctl/api.go b/cmd/fleetctl/api.go index e8a3f3cb88..5b3498339f 100644 --- a/cmd/fleetctl/api.go +++ b/cmd/fleetctl/api.go @@ -1,7 +1,9 @@ package main import ( + "flag" "fmt" + "os" "runtime" "github.com/fleetdm/fleet/v4/server/service" @@ -10,6 +12,10 @@ import ( ) func unauthenticatedClientFromCLI(c *cli.Context) (*service.Client, error) { + if flag.Lookup("test.v") != nil { + return service.NewClient(os.Getenv("FLEET_SERVER_ADDRESS"), true, "", "") + } + if err := makeConfigIfNotExists(c.String("config")); err != nil { return nil, errors.Wrapf(err, "error verifying that config exists at %s", c.String("config")) } @@ -59,6 +65,11 @@ func clientFromCLI(c *cli.Context) (*service.Client, error) { configPath, context := c.String("config"), c.String("context") + if flag.Lookup("test.v") != nil { + fleet.SetToken("AAAA") + return fleet, nil + } + // Add authentication token t, err := getConfigValue(configPath, context, "token") if err != nil { diff --git a/cmd/fleetctl/apply.go b/cmd/fleetctl/apply.go index 273751fea1..5ad1301d45 100644 --- a/cmd/fleetctl/apply.go +++ b/cmd/fleetctl/apply.go @@ -2,7 +2,6 @@ package main import ( "encoding/json" - "fmt" "io/ioutil" "regexp" "strings" @@ -29,6 +28,7 @@ type specGroup struct { Labels []*fleet.LabelSpec AppConfig *fleet.AppConfigPayload EnrollSecret *fleet.EnrollSecretSpec + UsersRoles *fleet.UsersRoleSpec } func specGroupFromBytes(b []byte) (*specGroup, error) { @@ -94,6 +94,13 @@ func specGroupFromBytes(b []byte) (*specGroup, error) { } specs.EnrollSecret = enrollSecretSpec + case fleet.UserRolesKind: + var userRoleSpec *fleet.UsersRoleSpec + if err := yaml.Unmarshal(s.Spec, &userRoleSpec); err != nil { + return nil, errors.Wrap(err, "unmarshaling "+kind+" spec") + } + specs.UsersRoles = userRoleSpec + default: return nil, errors.Errorf("unknown kind %q", s.Kind) } @@ -132,7 +139,7 @@ func applyCommand() *cli.Command { return err } - fleet, err := clientFromCLI(c) + fleetClient, err := clientFromCLI(c) if err != nil { return err } @@ -143,40 +150,46 @@ func applyCommand() *cli.Command { } if len(specs.Queries) > 0 { - if err := fleet.ApplyQueries(specs.Queries); err != nil { + if err := fleetClient.ApplyQueries(specs.Queries); err != nil { return errors.Wrap(err, "applying queries") } - fmt.Printf("[+] applied %d queries\n", len(specs.Queries)) + logf(c, "[+] applied %d queries\n", len(specs.Queries)) } if len(specs.Labels) > 0 { - if err := fleet.ApplyLabels(specs.Labels); err != nil { + if err := fleetClient.ApplyLabels(specs.Labels); err != nil { return errors.Wrap(err, "applying labels") } - fmt.Printf("[+] applied %d labels\n", len(specs.Labels)) + logf(c, "[+] applied %d labels\n", len(specs.Labels)) } if len(specs.Packs) > 0 { - if err := fleet.ApplyPacks(specs.Packs); err != nil { + if err := fleetClient.ApplyPacks(specs.Packs); err != nil { return errors.Wrap(err, "applying packs") } - fmt.Printf("[+] applied %d packs\n", len(specs.Packs)) + logf(c, "[+] applied %d packs\n", len(specs.Packs)) } if specs.AppConfig != nil { - if err := fleet.ApplyAppConfig(specs.AppConfig); err != nil { + if err := fleetClient.ApplyAppConfig(specs.AppConfig); err != nil { return errors.Wrap(err, "applying fleet config") } - fmt.Printf("[+] applied fleet config\n") + log(c, "[+] applied fleet config\n") } if specs.EnrollSecret != nil { - if err := fleet.ApplyEnrollSecretSpec(specs.EnrollSecret); err != nil { + if err := fleetClient.ApplyEnrollSecretSpec(specs.EnrollSecret); err != nil { return errors.Wrap(err, "applying enroll secrets") } - fmt.Printf("[+] applied enroll secrets\n") + log(c, "[+] applied enroll secrets\n") + } + if specs.UsersRoles != nil { + if err := fleetClient.ApplyUsersRoleSecretSpec(specs.UsersRoles); err != nil { + return errors.Wrap(err, "applying user roles") + } + log(c, "[+] applied user roles\n") } return nil diff --git a/cmd/fleetctl/apply_test.go b/cmd/fleetctl/apply_test.go new file mode 100644 index 0000000000..311075e712 --- /dev/null +++ b/cmd/fleetctl/apply_test.go @@ -0,0 +1,97 @@ +package main + +import ( + "io/ioutil" + "os" + "testing" + "time" + + "github.com/fleetdm/fleet/v4/server/fleet" + "github.com/fleetdm/fleet/v4/server/ptr" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +var userRoleSpecList = []*fleet.User{ + &fleet.User{ + UpdateCreateTimestamps: fleet.UpdateCreateTimestamps{ + CreateTimestamp: fleet.CreateTimestamp{CreatedAt: time.Now()}, + UpdateTimestamp: fleet.UpdateTimestamp{UpdatedAt: time.Now()}, + }, + ID: 42, + Name: "Test Name admin1@example.com", + Email: "admin1@example.com", + GlobalRole: ptr.String(fleet.RoleAdmin), + }, + &fleet.User{ + UpdateCreateTimestamps: fleet.UpdateCreateTimestamps{ + CreateTimestamp: fleet.CreateTimestamp{CreatedAt: time.Now()}, + UpdateTimestamp: fleet.UpdateTimestamp{UpdatedAt: time.Now()}, + }, + ID: 23, + Name: "Test Name2 admin2@example.com", + Email: "admin2@example.com", + GlobalRole: nil, + Teams: []fleet.UserTeam{}, + }, +} + +func TestApplyUserRoles(t *testing.T) { + server, ds := runServerWithMockedDS(t) + defer server.Close() + + ds.ListUsersFunc = func(opt fleet.UserListOptions) ([]*fleet.User, error) { + return userRoleSpecList, nil + } + + ds.UserByEmailFunc = func(email string) (*fleet.User, error) { + if email == "admin1@example.com" { + return userRoleSpecList[0], nil + } + return userRoleSpecList[1], nil + } + + ds.TeamByNameFunc = func(name string) (*fleet.Team, error) { + return &fleet.Team{ + ID: 1, + CreatedAt: time.Now(), + Name: "team1", + }, nil + } + + ds.SaveUsersFunc = func(users []*fleet.User) error { + for _, u := range users { + switch u.Email { + case "admin1@example.com": + userRoleList[0] = u + case "admin2@example.com": + userRoleList[1] = u + } + } + return nil + } + + tmpFile, err := ioutil.TempFile(os.TempDir(), "*.yml") + require.NoError(t, err) + defer os.Remove(tmpFile.Name()) + + tmpFile.WriteString(` +--- +apiVersion: v1 +kind: user_roles +spec: + roles: + admin1@example.com: + global_role: admin + teams: null + admin2@example.com: + global_role: null + teams: + - role: maintainer + team: team1 +`) + + assert.Equal(t, "[+] applied user roles\n", runAppForTest(t, []string{"apply", "-f", tmpFile.Name()})) + require.Len(t, userRoleSpecList[1].Teams, 1) + assert.Equal(t, fleet.RoleMaintainer, userRoleSpecList[1].Teams[0].Role) +} diff --git a/cmd/fleetctl/fleetctl.go b/cmd/fleetctl/fleetctl.go index 60d3976696..cc7e7d5c32 100644 --- a/cmd/fleetctl/fleetctl.go +++ b/cmd/fleetctl/fleetctl.go @@ -1,7 +1,9 @@ package main import ( + "io" "math/rand" + "os" "time" eefleetctl "github.com/fleetdm/fleet/v4/ee/fleetctl" @@ -18,13 +20,23 @@ func init() { } func main() { + app := createApp(os.Stdin, os.Stdout, nil) + + app.RunAndExitOnError() +} + +func createApp(reader io.Reader, writer io.Writer, exitErrHandler cli.ExitErrHandlerFunc) *cli.App { app := cli.NewApp() app.Name = "fleetctl" app.Usage = "CLI for operating Fleet" app.Version = version.Version().Version + app.ExitErrHandler = exitErrHandler cli.VersionPrinter = func(c *cli.Context) { version.PrintFull() } + app.Reader = reader + app.Writer = writer + app.ErrWriter = writer app.Commands = []*cli.Command{ applyCommand(), @@ -49,6 +61,5 @@ func main() { previewCommand(), eefleetctl.UpdatesCommand(), } - - app.RunAndExitOnError() + return app } diff --git a/cmd/fleetctl/get.go b/cmd/fleetctl/get.go index b6f0db789b..48f0ebfee5 100644 --- a/cmd/fleetctl/get.go +++ b/cmd/fleetctl/get.go @@ -7,6 +7,8 @@ import ( "os" "strconv" + "gopkg.in/guregu/null.v3" + "github.com/fleetdm/fleet/v4/server/fleet" "github.com/ghodss/yaml" "github.com/olekukonko/tablewriter" @@ -28,12 +30,22 @@ type specGeneric struct { Spec interface{} `json:"spec"` } -func defaultTable() *tablewriter.Table { - table := tablewriter.NewWriter(os.Stdout) +func defaultTable(writer io.Writer) *tablewriter.Table { + w := writerOrStdout(writer) + table := tablewriter.NewWriter(w) table.SetRowLine(true) return table } +func writerOrStdout(writer io.Writer) io.Writer { + var w io.Writer + w = os.Stdout + if writer != nil { + w = writer + } + return w +} + func yamlFlag() cli.Flag { return &cli.BoolFlag{ Name: yamlFlagName, @@ -48,21 +60,23 @@ func jsonFlag() cli.Flag { } } -func printJSON(spec interface{}) error { +func printJSON(spec interface{}, writer io.Writer) error { + w := writerOrStdout(writer) b, err := json.Marshal(spec) if err != nil { return err } - fmt.Printf("%s\n", b) + fmt.Fprintf(w, "%s\n", b) return nil } -func printYaml(spec interface{}) error { +func printYaml(spec interface{}, writer io.Writer) error { + w := writerOrStdout(writer) b, err := yaml.Marshal(spec) if err != nil { return err } - fmt.Printf("---\n%s", string(b)) + fmt.Fprintf(w, "---\n%s", string(b)) return nil } @@ -73,15 +87,7 @@ func printLabel(c *cli.Context, label *fleet.LabelSpec) error { Spec: label, } - var err error - - if c.Bool(jsonFlagName) { - err = printJSON(spec) - } else { - err = printYaml(spec) - } - - return err + return printSpec(c, spec) } func printQuery(c *cli.Context, query *fleet.QuerySpec) error { @@ -91,15 +97,7 @@ func printQuery(c *cli.Context, query *fleet.QuerySpec) error { Spec: query, } - var err error - - if c.Bool(jsonFlagName) { - err = printJSON(spec) - } else { - err = printYaml(spec) - } - - return err + return printSpec(c, spec) } func printPack(c *cli.Context, pack *fleet.PackSpec) error { @@ -109,15 +107,7 @@ func printPack(c *cli.Context, pack *fleet.PackSpec) error { Spec: pack, } - var err error - - if c.Bool(jsonFlagName) { - err = printJSON(spec) - } else { - err = printYaml(spec) - } - - return err + return printSpec(c, spec) } func printSecret(c *cli.Context, secret *fleet.EnrollSecretSpec) error { @@ -127,15 +117,7 @@ func printSecret(c *cli.Context, secret *fleet.EnrollSecretSpec) error { Spec: secret, } - var err error - - if c.Bool(jsonFlagName) { - err = printJSON(spec) - } else { - err = printYaml(spec) - } - - return err + return printSpec(c, spec) } func printHost(c *cli.Context, host *fleet.Host) error { @@ -145,15 +127,7 @@ func printHost(c *cli.Context, host *fleet.Host) error { Spec: host, } - var err error - - if c.Bool(jsonFlagName) { - err = printJSON(spec) - } else { - err = printYaml(spec) - } - - return err + return printSpec(c, spec) } func printConfig(c *cli.Context, config *fleet.AppConfigPayload) error { @@ -162,14 +136,60 @@ func printConfig(c *cli.Context, config *fleet.AppConfigPayload) error { Version: fleet.ApiVersion, Spec: config, } + + return printSpec(c, spec) +} + +type UserRoles struct { + Roles map[string]UserRole `json:"roles"` +} + +type TeamRole struct { + Team string `json:"team"` + Role string `json:"role"` +} + +type UserRole struct { + GlobalRole *string `json:"global_role"` + Teams []TeamRole `json:"teams"` +} + +func usersToUserRoles(users []fleet.User) UserRoles { + roles := make(map[string]UserRole) + for _, u := range users { + var teams []TeamRole + for _, t := range u.Teams { + teams = append(teams, TeamRole{ + Team: t.Name, + Role: t.Role, + }) + } + roles[u.Email] = UserRole{ + GlobalRole: u.GlobalRole, + Teams: teams, + } + } + return UserRoles{Roles: roles} +} + +func printUserRoles(c *cli.Context, users []fleet.User) error { + spec := specGeneric{ + Kind: fleet.UserRolesKind, + Version: fleet.ApiVersion, + Spec: usersToUserRoles(users), + } + + return printSpec(c, spec) +} + +func printSpec(c *cli.Context, spec specGeneric) error { var err error if c.Bool(jsonFlagName) { - err = printJSON(spec) + err = printJSON(spec, c.App.Writer) } else { - err = printYaml(spec) + err = printYaml(spec, c.App.Writer) } - return err } @@ -186,6 +206,7 @@ func getCommand() *cli.Command { getAppConfigCommand(), getCarveCommand(), getCarvesCommand(), + getUserRolesCommand(), }, } } @@ -240,10 +261,8 @@ func getQueriesCommand() *cli.Command { }) } - table := defaultTable() - table.SetHeader([]string{"name", "description", "query"}) - table.AppendBulk(data) - table.Render() + columns := []string{"name", "description", "query"} + printTable(c, columns, data) } return nil } @@ -357,10 +376,8 @@ func getPacksCommand() *cli.Command { }) } - table := defaultTable() - table.SetHeader([]string{"name", "platform", "description", "disabled"}) - table.AppendBulk(data) - table.Render() + columns := []string{"name", "platform", "description", "disabled"} + printTable(c, columns, data) return nil } @@ -434,10 +451,8 @@ func getLabelsCommand() *cli.Command { }) } - table := defaultTable() - table.SetHeader([]string{"name", "platform", "description", "query"}) - table.AppendBulk(data) - table.Render() + columns := []string{"name", "platform", "description", "query"} + printTable(c, columns, data) return nil } @@ -574,10 +589,8 @@ func getHostsCommand() *cli.Command { }) } - table := defaultTable() - table.SetHeader([]string{"uuid", "hostname", "platform", "osquery_version", "status"}) - table.AppendBulk(data) - table.Render() + columns := []string{"uuid", "hostname", "platform", "osquery_version", "status"} + printTable(c, columns, data) } else { host, err := client.HostByIdentifier(identifier) if err != nil { @@ -645,10 +658,8 @@ func getCarvesCommand() *cli.Command { }) } - table := defaultTable() - table.SetHeader([]string{"id", "created_at", "request_id", "carve_size", "completion"}) - table.AppendBulk(data) - table.Render() + columns := []string{"id", "created_at", "request_id", "carve_size", "completion"} + printTable(c, columns, data) return nil }, @@ -721,7 +732,7 @@ func getCarveCommand() *cli.Command { return err } - if err := printYaml(carve); err != nil { + if err := printYaml(carve, c.App.Writer); err != nil { return errors.Wrap(err, "print carve yaml") } @@ -729,3 +740,71 @@ func getCarveCommand() *cli.Command { }, } } + +func log(c *cli.Context, msg ...interface{}) { + fmt.Fprint(c.App.Writer, msg...) +} + +func logf(c *cli.Context, format string, a ...interface{}) { + fmt.Fprintf(c.App.Writer, format, a...) +} + +func getUserRolesCommand() *cli.Command { + return &cli.Command{ + Name: "user_roles", + Aliases: []string{"user_role", "ur"}, + Usage: "List global and team roles for users", + Flags: []cli.Flag{ + jsonFlag(), + yamlFlag(), + configFlag(), + contextFlag(), + debugFlag(), + }, + Action: func(c *cli.Context) error { + client, err := clientFromCLI(c) + if err != nil { + return err + } + + users, err := client.ListUsers() + if err != nil { + return errors.Wrap(err, "could not list users") + } + + if len(users) == 0 { + log(c, "No users found") + return nil + } + + if c.Bool(jsonFlagName) || c.Bool(yamlFlagName) { + err = printUserRoles(c, users) + if err != nil { + return err + } + return nil + } + + // Default to printing as table + data := [][]string{} + + for _, u := range users { + data = append(data, []string{ + u.Name, + null.StringFromPtr(u.GlobalRole).ValueOrZero(), + }) + } + columns := []string{"User", "Global Role"} + printTable(c, columns, data) + + return nil + }, + } +} + +func printTable(c *cli.Context, columns []string, data [][]string) { + table := defaultTable(c.App.Writer) + table.SetHeader(columns) + table.AppendBulk(data) + table.Render() +} diff --git a/cmd/fleetctl/get_test.go b/cmd/fleetctl/get_test.go new file mode 100644 index 0000000000..6623bcc13a --- /dev/null +++ b/cmd/fleetctl/get_test.go @@ -0,0 +1,83 @@ +package main + +import ( + "testing" + "time" + + "github.com/fleetdm/fleet/v4/server/fleet" + "github.com/fleetdm/fleet/v4/server/ptr" + "github.com/stretchr/testify/assert" +) + +var userRoleList = []*fleet.User{ + &fleet.User{ + UpdateCreateTimestamps: fleet.UpdateCreateTimestamps{ + CreateTimestamp: fleet.CreateTimestamp{CreatedAt: time.Now()}, + UpdateTimestamp: fleet.UpdateTimestamp{UpdatedAt: time.Now()}, + }, + ID: 42, + Name: "Test Name admin1@example.com", + Email: "admin1@example.com", + GlobalRole: ptr.String(fleet.RoleAdmin), + }, + &fleet.User{ + UpdateCreateTimestamps: fleet.UpdateCreateTimestamps{ + CreateTimestamp: fleet.CreateTimestamp{CreatedAt: time.Now()}, + UpdateTimestamp: fleet.UpdateTimestamp{UpdatedAt: time.Now()}, + }, + ID: 23, + Name: "Test Name2 admin2@example.com", + Email: "admin2@example.com", + GlobalRole: nil, + Teams: []fleet.UserTeam{ + fleet.UserTeam{ + Team: fleet.Team{ + ID: 1, + CreatedAt: time.Now(), + Name: "team1", + UserCount: 1, + HostCount: 1, + }, + Role: fleet.RoleMaintainer, + }, + }, + }, +} + +func TestGetUserRoles(t *testing.T) { + server, ds := runServerWithMockedDS(t) + defer server.Close() + + ds.ListUsersFunc = func(opt fleet.UserListOptions) ([]*fleet.User, error) { + return userRoleList, nil + } + + expectedText := `+-------------------------------+-------------+ +| USER | GLOBAL ROLE | ++-------------------------------+-------------+ +| Test Name admin1@example.com | admin | ++-------------------------------+-------------+ +| Test Name2 admin2@example.com | | ++-------------------------------+-------------+ +` + expectedYaml := `--- +apiVersion: v1 +kind: user_roles +spec: + roles: + admin1@example.com: + global_role: admin + teams: null + admin2@example.com: + global_role: null + teams: + - role: maintainer + team: team1 +` + expectedJson := `{"kind":"user_roles","apiVersion":"v1","spec":{"roles":{"admin1@example.com":{"global_role":"admin","teams":null},"admin2@example.com":{"global_role":null,"teams":[{"team":"team1","role":"maintainer"}]}}}} +` + + assert.Equal(t, expectedText, runAppForTest(t, []string{"get", "user_roles"})) + assert.Equal(t, expectedYaml, runAppForTest(t, []string{"get", "user_roles", "--yaml"})) + assert.Equal(t, expectedJson, runAppForTest(t, []string{"get", "user_roles", "--json"})) +} diff --git a/cmd/fleetctl/testing_utils.go b/cmd/fleetctl/testing_utils.go new file mode 100644 index 0000000000..13b1ced04d --- /dev/null +++ b/cmd/fleetctl/testing_utils.go @@ -0,0 +1,59 @@ +package main + +import ( + "bytes" + "net/http/httptest" + "os" + "testing" + "time" + + "github.com/fleetdm/fleet/v4/server/fleet" + "github.com/fleetdm/fleet/v4/server/mock" + "github.com/fleetdm/fleet/v4/server/service" + "github.com/stretchr/testify/require" + "github.com/urfave/cli/v2" +) + +func runServerWithMockedDS(t *testing.T) (*httptest.Server, *mock.Store) { + ds := new(mock.Store) + var users []*fleet.User + ds.NewUserFunc = func(user *fleet.User) (*fleet.User, error) { + users = append(users, user) + return user, nil + } + ds.SessionByKeyFunc = func(key string) (*fleet.Session, error) { + return &fleet.Session{ + CreateTimestamp: fleet.CreateTimestamp{CreatedAt: time.Now()}, + ID: 1, + AccessedAt: time.Now(), + UserID: 1, + Key: key, + }, nil + } + ds.MarkSessionAccessedFunc = func(session *fleet.Session) error { + return nil + } + ds.UserByIDFunc = func(id uint) (*fleet.User, error) { + return users[0], nil + } + ds.ListUsersFunc = func(opt fleet.UserListOptions) ([]*fleet.User, error) { + return users, nil + } + _, server := service.RunServerForTestsWithDS(t, ds) + os.Setenv("FLEET_SERVER_ADDRESS", server.URL) + + return server, ds +} + +func runAppForTest(t *testing.T, args []string) string { + w := new(bytes.Buffer) + r, _, _ := os.Pipe() + var exitErr error + app := createApp(r, w, func(context *cli.Context, err error) { + exitErr = err + }) + err := app.Run(append([]string{""}, args...)) + require.Nil(t, err) + require.Nil(t, exitErr) + return w.String() +} diff --git a/server/datastore/datastore_users.go b/server/datastore/datastore_users.go index ad4fb1fbd5..1d2268614e 100644 --- a/server/datastore/datastore_users.go +++ b/server/datastore/datastore_users.go @@ -124,6 +124,16 @@ func testUserGlobalRole(t *testing.T, ds fleet.Datastore, users []*fleet.User) { assert.Nil(t, err) assert.Equal(t, user.GlobalRole, verify.GlobalRole) } + err := ds.SaveUser(&fleet.User{ + Name: "some@email.asd", + Password: []byte("asdasd"), + Email: "some@email.asd", + GlobalRole: ptr.String(fleet.RoleObserver), + Teams: []fleet.UserTeam{{Role: fleet.RoleMaintainer}}, + }) + require.IsType(t, &fleet.Error{}, err) + flErr := err.(*fleet.Error) + assert.Equal(t, "Cannot specify both Global Role and Team Roles", flErr.Message) } func testListUsers(t *testing.T, ds fleet.Datastore) { @@ -169,9 +179,10 @@ func testUserTeams(t *testing.T, ds fleet.Datastore) { users[0].Teams = []fleet.UserTeam{ { Team: fleet.Team{ID: 3}, - Role: "foobar", + Role: fleet.RoleObserver, }, } + users[0].GlobalRole = nil err = ds.SaveUser(users[0]) require.NoError(t, err) @@ -188,17 +199,18 @@ func testUserTeams(t *testing.T, ds fleet.Datastore) { users[1].Teams = []fleet.UserTeam{ { Team: fleet.Team{ID: 1}, - Role: "foobar", + Role: fleet.RoleObserver, }, { Team: fleet.Team{ID: 2}, - Role: "foobar", + Role: fleet.RoleObserver, }, { Team: fleet.Team{ID: 3}, - Role: "foobar", + Role: fleet.RoleObserver, }, } + users[1].GlobalRole = nil err = ds.SaveUser(users[1]) require.NoError(t, err) @@ -214,6 +226,7 @@ func testUserTeams(t *testing.T, ds fleet.Datastore) { // Clear teams users[1].Teams = []fleet.UserTeam{} + users[1].GlobalRole = ptr.String(fleet.RoleObserver) err = ds.SaveUser(users[1]) require.NoError(t, err) diff --git a/server/datastore/mysql/users.go b/server/datastore/mysql/users.go index fc4f9e6e39..e9879a41e4 100644 --- a/server/datastore/mysql/users.go +++ b/server/datastore/mysql/users.go @@ -18,7 +18,8 @@ func (d *Datastore) NewUser(user *fleet.User) (*fleet.User, error) { return nil, err } - sqlStatement := ` + err := d.withTx(func(tx *sqlx.Tx) error { + sqlStatement := ` INSERT INTO users ( password, salt, @@ -32,25 +33,30 @@ func (d *Datastore) NewUser(user *fleet.User) (*fleet.User, error) { global_role ) VALUES (?,?,?,?,?,?,?,?,?,?) ` - result, err := d.db.Exec(sqlStatement, - user.Password, - user.Salt, - user.Name, - user.Email, - user.AdminForcedPasswordReset, - user.GravatarURL, - user.Position, - user.SSOEnabled, - user.APIOnly, - user.GlobalRole) + result, err := d.db.Exec(sqlStatement, + user.Password, + user.Salt, + user.Name, + user.Email, + user.AdminForcedPasswordReset, + user.GravatarURL, + user.Position, + user.SSOEnabled, + user.APIOnly, + user.GlobalRole) + if err != nil { + return errors.Wrap(err, "create new user") + } + + id, _ := result.LastInsertId() + user.ID = uint(id) + + if err := d.saveTeamsForUser(tx, user); err != nil { + return err + } + return nil + }) if err != nil { - return nil, errors.Wrap(err, "create new user") - } - - id, _ := result.LastInsertId() - user.ID = uint(id) - - if err := d.saveTeamsForUser(user); err != nil { return nil, err } @@ -119,10 +125,27 @@ func (d *Datastore) UserByID(id uint) (*fleet.User, error) { } func (d *Datastore) SaveUser(user *fleet.User) error { + return d.withTx(func(tx *sqlx.Tx) error { + return d.saveUser(tx, user) + }) +} + +func (d *Datastore) SaveUsers(users []*fleet.User) error { + return d.withTx(func(tx *sqlx.Tx) error { + for _, user := range users { + err := d.saveUser(tx, user) + if err != nil { + return err + } + } + return nil + }) +} + +func (d *Datastore) saveUser(tx *sqlx.Tx, user *fleet.User) error { if err := fleet.ValidateRole(user.GlobalRole, user.Teams); err != nil { return err } - sqlStatement := ` UPDATE users SET password = ?, @@ -161,7 +184,7 @@ func (d *Datastore) SaveUser(user *fleet.User) error { } // REVIEW: Check if teams have been set? - if err := d.saveTeamsForUser(user); err != nil { + if err := d.saveTeamsForUser(tx, user); err != nil { return err } @@ -211,37 +234,33 @@ func (d *Datastore) loadTeamsForUsers(users []*fleet.User) error { return nil } -func (d *Datastore) saveTeamsForUser(user *fleet.User) error { +func (d *Datastore) saveTeamsForUser(tx *sqlx.Tx, user *fleet.User) error { // Do a full teams update by deleting existing teams and then inserting all // the current teams in a single transaction. - if err := d.withRetryTxx(func(tx *sqlx.Tx) error { - // Delete before insert - sql := `DELETE FROM user_teams WHERE user_id = ?` - if _, err := tx.Exec(sql, user.ID); err != nil { - return errors.Wrap(err, "delete existing teams") - } - if len(user.Teams) == 0 { - return nil - } - - // Bulk insert - const valueStr = "(?,?,?)," - var args []interface{} - for _, userTeam := range user.Teams { - args = append(args, user.ID, userTeam.Team.ID, userTeam.Role) - } - sql = "INSERT INTO user_teams (user_id, team_id, role) VALUES " + - strings.Repeat(valueStr, len(user.Teams)) - sql = strings.TrimSuffix(sql, ",") - if _, err := tx.Exec(sql, args...); err != nil { - return errors.Wrap(err, "insert teams") - } - - return nil - }); err != nil { - return errors.Wrap(err, "save teams for user") + // Delete before insert + sql := `DELETE FROM user_teams WHERE user_id = ?` + if _, err := tx.Exec(sql, user.ID); err != nil { + return errors.Wrap(err, "delete existing teams") } + + if len(user.Teams) == 0 { + return nil + } + + // Bulk insert + const valueStr = "(?,?,?)," + var args []interface{} + for _, userTeam := range user.Teams { + args = append(args, user.ID, userTeam.Team.ID, userTeam.Role) + } + sql = "INSERT INTO user_teams (user_id, team_id, role) VALUES " + + strings.Repeat(valueStr, len(user.Teams)) + sql = strings.TrimSuffix(sql, ",") + if _, err := tx.Exec(sql, args...); err != nil { + return errors.Wrap(err, "insert teams") + } + return nil } diff --git a/server/fleet/errors.go b/server/fleet/errors.go index a6e473767a..7f57362f5f 100644 --- a/server/fleet/errors.go +++ b/server/fleet/errors.go @@ -211,6 +211,8 @@ type Error struct { const ( // ErrNoRoleNeeded is the error number for valid role needed ErrNoRoleNeeded = 1 + // ErrNoOneAdminNeeded is the error number when all admins are about to be removed + ErrNoOneAdminNeeded = 2 ) // NewError returns a fleet error with the code and message specified diff --git a/server/fleet/service.go b/server/fleet/service.go index 3ca357e4d9..ae96abaaf1 100644 --- a/server/fleet/service.go +++ b/server/fleet/service.go @@ -19,5 +19,6 @@ type Service interface { CarveService TeamService ActivitiesService + UserRolesService GlobalScheduleService } diff --git a/server/fleet/teams.go b/server/fleet/teams.go index 6c398c1c82..e10bc8b410 100644 --- a/server/fleet/teams.go +++ b/server/fleet/teams.go @@ -158,6 +158,10 @@ func ValidateRole(globalRole *string, teamUsers []UserTeam) error { return nil } + if len(teamUsers) > 0 { + return NewError(ErrNoRoleNeeded, "Cannot specify both Global Role and Team Roles") + } + if !ValidGlobalRole(*globalRole) { return NewError(ErrNoRoleNeeded, "GlobalRole role can only be admin, observer, or maintainer.") } diff --git a/server/fleet/user_roles.go b/server/fleet/user_roles.go new file mode 100644 index 0000000000..de10c12659 --- /dev/null +++ b/server/fleet/user_roles.go @@ -0,0 +1,26 @@ +package fleet + +import "context" + +const ( + UserRolesKind = "user_roles" +) + +type UsersRoleSpec struct { + Roles map[string]*UserRoleSpec `json:"roles"` +} + +type UserRoleSpec struct { + GlobalRole *string `json:"global_role"` + Teams []TeamRoleSpec `json:"teams"` +} + +type TeamRoleSpec struct { + Name string `json:"team"` + Role string `json:"role"` +} + +type UserRolesService interface { + // ApplyUserRolesSpecs applies a list of user global and team role changes + ApplyUserRolesSpecs(ctx context.Context, specs UsersRoleSpec) error +} diff --git a/server/fleet/users.go b/server/fleet/users.go index ade1271552..467ee54791 100644 --- a/server/fleet/users.go +++ b/server/fleet/users.go @@ -16,6 +16,7 @@ type UserStore interface { UserByEmail(email string) (*User, error) UserByID(id uint) (*User, error) SaveUser(user *User) error + SaveUsers(users []*User) error // DeleteUser permanently deletes the user identified by the provided ID. DeleteUser(id uint) error // PendingEmailChange creates a record with a pending email change for a user identified diff --git a/server/mock/datastore_users.go b/server/mock/datastore_users.go index bfce310deb..52dd579f0f 100644 --- a/server/mock/datastore_users.go +++ b/server/mock/datastore_users.go @@ -16,6 +16,8 @@ type UserByIDFunc func(id uint) (*fleet.User, error) type SaveUserFunc func(user *fleet.User) error +type SaveUsersFunc func(user []*fleet.User) error + type DeleteUserFunc func(id uint) error type PendingEmailChangeFunc func(userID uint, newEmail string, token string) error @@ -38,6 +40,9 @@ type UserStore struct { SaveUserFunc SaveUserFunc SaveUserFuncInvoked bool + SaveUsersFunc SaveUsersFunc + SaveUsersFuncInvoked bool + DeleteUserFunc DeleteUserFunc DeleteUserFuncInvoked bool @@ -73,6 +78,11 @@ func (s *UserStore) SaveUser(user *fleet.User) error { return s.SaveUserFunc(user) } +func (s *UserStore) SaveUsers(users []*fleet.User) error { + s.SaveUsersFuncInvoked = true + return s.SaveUsersFunc(users) +} + func (s *UserStore) DeleteUser(id uint) error { s.DeleteUserFuncInvoked = true return s.DeleteUserFunc(id) diff --git a/server/service/client.go b/server/service/client.go index 42332f7616..1b4bbeabd1 100644 --- a/server/service/client.go +++ b/server/service/client.go @@ -41,7 +41,7 @@ func NewClient(addr string, insecureSkipVerify bool, rootCA, urlPrefix string, o return nil, errors.Wrap(err, "parsing URL") } - if baseURL.Scheme != "https" && !strings.Contains(baseURL.Host, "localhost") { + if baseURL.Scheme != "https" && !strings.Contains(baseURL.Host, "localhost") && !strings.Contains(baseURL.Host, "127.0.0.1") { return nil, errors.New("address must start with https:// for remote connections") } @@ -209,3 +209,33 @@ func (l *logRoundTripper) RoundTrip(req *http.Request) (*http.Response, error) { return res, nil } + +func (c *Client) authenticatedRequest(params interface{}, verb string, path string, responseDest interface{}) error { + response, err := c.AuthenticatedDo(verb, path, "", params) + if err != nil { + return errors.Wrapf(err, "%s %s", verb, path) + } + defer response.Body.Close() + + if response.StatusCode != http.StatusOK { + return errors.Errorf( + "%s %s received status %d %s", + verb, path, + response.StatusCode, + extractServerErrorText(response.Body), + ) + } + + err = json.NewDecoder(response.Body).Decode(&responseDest) + if err != nil { + return errors.Wrapf(err, "decode %s %s response", verb, path) + } + + if e, ok := responseDest.(errorer); ok { + if e.error() != nil { + return errors.Errorf("%s %s error: %s", verb, path, e.error()) + } + } + + return nil +} diff --git a/server/service/client_users.go b/server/service/client_users.go index dfb4a9bb3d..83b98f3b16 100644 --- a/server/service/client_users.go +++ b/server/service/client_users.go @@ -1,39 +1,33 @@ package service import ( - "encoding/json" - "net/http" - "github.com/fleetdm/fleet/v4/server/fleet" - "github.com/pkg/errors" ) // CreateUser creates a new user, skipping the invitation process. func (c *Client) CreateUser(p fleet.UserPayload) error { verb, path := "POST", "/api/v1/fleet/users/admin" - response, err := c.AuthenticatedDo(verb, path, "", p) - if err != nil { - return errors.Wrapf(err, "%s %s", verb, path) - } - defer response.Body.Close() - - if response.StatusCode != http.StatusOK { - return errors.Errorf( - "create user received status %d %s", - response.StatusCode, - extractServerErrorText(response.Body), - ) - } - var responseBody createUserResponse - err = json.NewDecoder(response.Body).Decode(&responseBody) - if err != nil { - return errors.Wrap(err, "decode create user response") - } - if responseBody.Err != nil { - return errors.Errorf("create user: %s", responseBody.Err) + return c.authenticatedRequest(p, verb, path, &responseBody) +} + +// ListUsers retrieves the list of users. +func (c *Client) ListUsers() ([]fleet.User, error) { + verb, path := "GET", "/api/v1/fleet/users" + var responseBody listUsersResponse + + err := c.authenticatedRequest(nil, verb, path, &responseBody) + if err != nil { + return nil, err } + return responseBody.Users, nil +} - return nil +// ApplyUsersRoleSecretSpec applies the global and team roles for users. +func (c *Client) ApplyUsersRoleSecretSpec(spec *fleet.UsersRoleSpec) error { + req := applyUserRoleSpecsRequest{Spec: spec} + verb, path := "POST", "/api/v1/fleet/users/roles/spec" + var responseBody applyUserRoleSpecsResponse + return c.authenticatedRequest(req, verb, path, &responseBody) } diff --git a/server/service/endpoint_labels.go b/server/service/endpoint_labels.go index 79c0b6b291..9ed122939f 100644 --- a/server/service/endpoint_labels.go +++ b/server/service/endpoint_labels.go @@ -234,7 +234,7 @@ func makeDeleteLabelByIDEndpoint(svc fleet.Service) endpoint.Endpoint { } //////////////////////////////////////////////////////////////////////////////// -// Apply Label Specs +// Apply Label Spec //////////////////////////////////////////////////////////////////////////////// type applyLabelSpecsRequest struct { @@ -259,12 +259,12 @@ func makeApplyLabelSpecsEndpoint(svc fleet.Service) endpoint.Endpoint { } //////////////////////////////////////////////////////////////////////////////// -// Get Label Specs +// Get Label Spec //////////////////////////////////////////////////////////////////////////////// type getLabelSpecsResponse struct { Specs []*fleet.LabelSpec `json:"specs"` - Err error `json:"error,omitempty"` + Err error `json:"error,omitempty"` } func (r getLabelSpecsResponse) error() error { return r.Err } @@ -285,7 +285,7 @@ func makeGetLabelSpecsEndpoint(svc fleet.Service) endpoint.Endpoint { type getLabelSpecResponse struct { Spec *fleet.LabelSpec `json:"specs,omitempty"` - Err error `json:"error,omitempty"` + Err error `json:"error,omitempty"` } func (r getLabelSpecResponse) error() error { return r.Err } diff --git a/server/service/endpoint_packs.go b/server/service/endpoint_packs.go index e11fcbed5b..4588a6e1d2 100644 --- a/server/service/endpoint_packs.go +++ b/server/service/endpoint_packs.go @@ -237,7 +237,7 @@ func makeDeletePackByIDEndpoint(svc fleet.Service) endpoint.Endpoint { } //////////////////////////////////////////////////////////////////////////////// -// Apply Pack Specs +// Apply Pack Spec //////////////////////////////////////////////////////////////////////////////// type applyPackSpecsRequest struct { @@ -262,7 +262,7 @@ func makeApplyPackSpecsEndpoint(svc fleet.Service) endpoint.Endpoint { } //////////////////////////////////////////////////////////////////////////////// -// Get Pack Specs +// Get Pack Spec //////////////////////////////////////////////////////////////////////////////// type getPackSpecsResponse struct { diff --git a/server/service/endpoint_queries.go b/server/service/endpoint_queries.go index fcb90c48f8..f759b5a259 100644 --- a/server/service/endpoint_queries.go +++ b/server/service/endpoint_queries.go @@ -17,7 +17,7 @@ type getQueryRequest struct { type getQueryResponse struct { Query *fleet.Query `json:"query,omitempty"` - Err error `json:"error,omitempty"` + Err error `json:"error,omitempty"` } func (r getQueryResponse) error() error { return r.Err } @@ -42,7 +42,7 @@ type listQueriesRequest struct { type listQueriesResponse struct { Queries []fleet.Query `json:"queries"` - Err error `json:"error,omitempty"` + Err error `json:"error,omitempty"` } func (r listQueriesResponse) error() error { return r.Err } @@ -73,7 +73,7 @@ type createQueryRequest struct { type createQueryResponse struct { Query *fleet.Query `json:"query,omitempty"` - Err error `json:"error,omitempty"` + Err error `json:"error,omitempty"` } func (r createQueryResponse) error() error { return r.Err } @@ -100,7 +100,7 @@ type modifyQueryRequest struct { type modifyQueryResponse struct { Query *fleet.Query `json:"query,omitempty"` - Err error `json:"error,omitempty"` + Err error `json:"error,omitempty"` } func (r modifyQueryResponse) error() error { return r.Err } @@ -193,7 +193,7 @@ func makeDeleteQueriesEndpoint(svc fleet.Service) endpoint.Endpoint { } //////////////////////////////////////////////////////////////////////////////// -// Apply Query Specs +// Apply Query Spec //////////////////////////////////////////////////////////////////////////////// type applyQuerySpecsRequest struct { @@ -218,12 +218,12 @@ func makeApplyQuerySpecsEndpoint(svc fleet.Service) endpoint.Endpoint { } //////////////////////////////////////////////////////////////////////////////// -// Get Query Specs +// Get Query Spec //////////////////////////////////////////////////////////////////////////////// type getQuerySpecsResponse struct { Specs []*fleet.QuerySpec `json:"specs"` - Err error `json:"error,omitempty"` + Err error `json:"error,omitempty"` } func (r getQuerySpecsResponse) error() error { return r.Err } @@ -244,7 +244,7 @@ func makeGetQuerySpecsEndpoint(svc fleet.Service) endpoint.Endpoint { type getQuerySpecResponse struct { Spec *fleet.QuerySpec `json:"specs,omitempty"` - Err error `json:"error,omitempty"` + Err error `json:"error,omitempty"` } func (r getQuerySpecResponse) error() error { return r.Err } diff --git a/server/service/endpoint_utils.go b/server/service/endpoint_utils.go new file mode 100644 index 0000000000..607d33e591 --- /dev/null +++ b/server/service/endpoint_utils.go @@ -0,0 +1,34 @@ +package service + +import ( + "context" + "encoding/json" + "net/http" + "reflect" + + "github.com/fleetdm/fleet/v4/server/fleet" + "github.com/go-kit/kit/endpoint" +) + +type handlerFunc func(ctx context.Context, request interface{}, svc fleet.Service) (interface{}, error) + +func makeDecoderForType(v interface{}) func(ctx context.Context, r *http.Request) (interface{}, error) { + t := reflect.TypeOf(v) + return func(ctx context.Context, r *http.Request) (interface{}, error) { + req := reflect.New(t).Interface() + if err := json.NewDecoder(r.Body).Decode(&req); err != nil { + return nil, err + } + return req, nil + } +} + +func makeAuthenticatedServiceEndpoint(svc fleet.Service, f handlerFunc) endpoint.Endpoint { + return authenticatedUser(svc, makeServiceEndpoint(svc, f)) +} + +func makeServiceEndpoint(svc fleet.Service, f handlerFunc) endpoint.Endpoint { + return func(ctx context.Context, request interface{}) (interface{}, error) { + return f(ctx, request, svc) + } +} diff --git a/server/service/handler.go b/server/service/handler.go index edf24b04f5..48e7c4347f 100644 --- a/server/service/handler.go +++ b/server/service/handler.go @@ -521,6 +521,7 @@ func MakeHandler(svc fleet.Service, config config.FleetConfig, logger kitlog.Log r := mux.NewRouter() attachFleetAPIRoutes(r, fleetHandlers) + attachNewStyleFleetAPIRoutes(r, svc, fleetAPIOptions) // Results endpoint is handled different due to websockets use r.PathPrefix("/api/v1/fleet/results/"). @@ -665,6 +666,23 @@ func attachFleetAPIRoutes(r *mux.Router, h *fleetHandlers) { r.Handle("/api/v1/fleet/activities", h.ListActivities).Methods("GET").Name("list_activities") } +func attachNewStyleFleetAPIRoutes(r *mux.Router, svc fleet.Service, opts []kithttp.ServerOption) { + handle("POST", "/api/v1/fleet/users/roles/spec", makeApplyUserRoleSpecsEndpoint(svc, opts), "apply_user_roles_spec", r) +} + +func handle(verb, path string, handler http.Handler, name string, r *mux.Router) { + r.Handle( + path, + handler, + ).Methods(verb).Name(name) +} + +// TODO: this duplicates the one in makeKitHandler +func newServer(e endpoint.Endpoint, decodeFn kithttp.DecodeRequestFunc, opts []kithttp.ServerOption) http.Handler { + e = authzcheck.NewMiddleware().AuthzCheck()(e) + return kithttp.NewServer(e, decodeFn, encodeResponse, opts...) +} + // WithSetup is an http middleware that checks if setup procedures have been completed. // If setup hasn't been completed it serves the API with a setup middleware. // If the server is already configured, the default API handler is exposed. diff --git a/server/service/http_auth_test.go b/server/service/http_auth_test.go index b1bc783bb3..da22d5236b 100644 --- a/server/service/http_auth_test.go +++ b/server/service/http_auth_test.go @@ -4,12 +4,10 @@ import ( "bytes" "encoding/json" "fmt" - "github.com/go-kit/kit/transport" "io" "io/ioutil" "net/http" "net/http/httptest" - "os" "strconv" "testing" @@ -17,13 +15,8 @@ import ( "github.com/fleetdm/fleet/v4/server/datastore/inmem" "github.com/fleetdm/fleet/v4/server/fleet" - kitlog "github.com/go-kit/kit/log" - - kithttp "github.com/go-kit/kit/transport/http" - "github.com/gorilla/mux" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" - "github.com/throttled/throttled/v2/store/memstore" ) func TestLogin(t *testing.T) { @@ -115,35 +108,34 @@ func TestLogin(t *testing.T) { func setupAuthTest(t *testing.T) (*inmem.Datastore, map[string]fleet.User, *httptest.Server) { ds, _ := inmem.New(config.TestConfig()) - users, server := runServerForTestsWithDS(t, ds) + users, server := RunServerForTestsWithDS(t, ds) return ds, users, server } -func runServerForTestsWithDS(t *testing.T, ds fleet.Datastore) (map[string]fleet.User, *httptest.Server) { - svc := newTestService(ds, nil, nil) - users := createTestUsers(t, ds) - logger := kitlog.NewLogfmtLogger(os.Stdout) +func getTestAdminToken(t *testing.T, server *httptest.Server) string { + testUser := testUsers["admin1"] - opts := []kithttp.ServerOption{ - kithttp.ServerBefore( - setRequestsContexts(svc), - ), - kithttp.ServerErrorHandler(transport.NewLogErrorHandler(logger)), - kithttp.ServerAfter( - kithttp.SetContentType("application/json; charset=utf-8"), - ), + params := loginRequest{ + Email: testUser.Email, + Password: testUser.PlaintextPassword, } - r := mux.NewRouter() - limitStore, _ := memstore.New(0) - ke := MakeFleetServerEndpoints(svc, "", limitStore) - kh := makeKitHandlers(ke, opts) - attachFleetAPIRoutes(r, kh) - r.Handle("/", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - fmt.Fprint(w, "index") - })) + j, err := json.Marshal(¶ms) + assert.Nil(t, err) - server := httptest.NewServer(r) - return users, server + requestBody := &nopCloser{bytes.NewBuffer(j)} + resp, err := http.Post(server.URL+"/api/v1/fleet/login", "application/json", requestBody) + require.Nil(t, err) + assert.Equal(t, http.StatusOK, resp.StatusCode) + + var jsn = struct { + User *fleet.User `json:"user"` + Token string `json:"token"` + Err []map[string]string `json:"errors,omitempty"` + }{} + err = json.NewDecoder(resp.Body).Decode(&jsn) + require.Nil(t, err) + + return jsn.Token } func TestNoHeaderErrorsDifferently(t *testing.T) { diff --git a/server/service/integration_test.go b/server/service/integration_test.go index 7e2d8b8c9f..8fb86bf054 100644 --- a/server/service/integration_test.go +++ b/server/service/integration_test.go @@ -18,34 +18,8 @@ import ( "github.com/stretchr/testify/require" ) -func getTestAdminToken(t *testing.T, server *httptest.Server) string { - testUser := testUsers["admin1"] - - params := loginRequest{ - Email: testUser.Email, - Password: testUser.PlaintextPassword, - } - j, err := json.Marshal(¶ms) - assert.Nil(t, err) - - requestBody := &nopCloser{bytes.NewBuffer(j)} - resp, err := http.Post(server.URL+"/api/v1/fleet/login", "application/json", requestBody) - require.Nil(t, err) - assert.Equal(t, http.StatusOK, resp.StatusCode) - - var jsn = struct { - User *fleet.User `json:"user"` - Token string `json:"token"` - Err []map[string]string `json:"errors,omitempty"` - }{} - err = json.NewDecoder(resp.Body).Decode(&jsn) - require.Nil(t, err) - - return jsn.Token -} - func testDoubleUserCreationErrors(t *testing.T, ds fleet.Datastore) { - _, server := runServerForTestsWithDS(t, ds) + _, server := RunServerForTestsWithDS(t, ds) token := getTestAdminToken(t, server) params := fleet.UserPayload{ @@ -75,7 +49,7 @@ func testDoubleUserCreationErrors(t *testing.T, ds fleet.Datastore) { } func testUserWithoutRoleErrors(t *testing.T, ds fleet.Datastore) { - _, server := runServerForTestsWithDS(t, ds) + _, server := RunServerForTestsWithDS(t, ds) token := getTestAdminToken(t, server) params := fleet.UserPayload{ @@ -97,7 +71,7 @@ func testUserWithoutRoleErrors(t *testing.T, ds fleet.Datastore) { } func testUserWithWrongRoleErrors(t *testing.T, ds fleet.Datastore) { - _, server := runServerForTestsWithDS(t, ds) + _, server := RunServerForTestsWithDS(t, ds) token := getTestAdminToken(t, server) params := fleet.UserPayload{ @@ -120,7 +94,7 @@ func testUserWithWrongRoleErrors(t *testing.T, ds fleet.Datastore) { } func testUserCreationWrongTeamErrors(t *testing.T, ds fleet.Datastore) { - _, server := runServerForTestsWithDS(t, ds) + _, server := RunServerForTestsWithDS(t, ds) token := getTestAdminToken(t, server) teams := []fleet.UserTeam{ @@ -128,15 +102,15 @@ func testUserCreationWrongTeamErrors(t *testing.T, ds fleet.Datastore) { Team: fleet.Team{ ID: 9999, }, + Role: fleet.RoleObserver, }, } params := fleet.UserPayload{ - Name: ptr.String("user1"), - Email: ptr.String("email@asd.com"), - Password: ptr.String("pass"), - GlobalRole: ptr.String(fleet.RoleObserver), - Teams: &teams, + Name: ptr.String("user1"), + Email: ptr.String("email@asd.com"), + Password: ptr.String("pass"), + Teams: &teams, } method := "POST" path := "/api/v1/fleet/users/admin" @@ -191,7 +165,7 @@ func assertBodyContains(t *testing.T, resp *http.Response, expectedError string) } func testQueryCreationLogsActivity(t *testing.T, ds fleet.Datastore) { - _, server := runServerForTestsWithDS(t, ds) + _, server := RunServerForTestsWithDS(t, ds) token := getTestAdminToken(t, server) params := fleet.QueryPayload{ @@ -222,7 +196,7 @@ func assertErrorCodeAndMessage(t *testing.T, resp *http.Response, code int, mess } func testAppConfigAdditionalQueriesCanBeRemoved(t *testing.T, ds fleet.Datastore) { - _, server := runServerForTestsWithDS(t, ds) + _, server := RunServerForTestsWithDS(t, ds) token := getTestAdminToken(t, server) spec := []byte(` @@ -261,10 +235,50 @@ func getConfig(t *testing.T, server *httptest.Server, token string) *fleet.AppCo return responseBody } +func testUserRolesSpec(t *testing.T, ds fleet.Datastore) { + _, server := RunServerForTestsWithDS(t, ds) + _, err := ds.NewTeam(&fleet.Team{ + ID: 42, + Name: "team1", + Description: "desc team1", + }) + require.NoError(t, err) + token := getTestAdminToken(t, server) + + user, err := ds.UserByEmail("user1@example.com") + require.NoError(t, err) + assert.Len(t, user.Teams, 0) + + spec := []byte(` + roles: + user1@example.com: + global_role: null + teams: + - role: maintainer + team: team1 +`) + + var userRoleSpec applyUserRoleSpecsRequest + err = yaml.Unmarshal(spec, &userRoleSpec.Spec) + require.NoError(t, err) + + doReq(t, userRoleSpec, "POST", server, "/api/v1/fleet/users/roles/spec", token, http.StatusOK) + + user, err = ds.UserByEmail("user1@example.com") + require.NoError(t, err) + require.Len(t, user.Teams, 1) + assert.Equal(t, fleet.RoleMaintainer, user.Teams[0].Role) + + // But users are not deleted + users, err := ds.ListUsers(fleet.UserListOptions{}) + require.NoError(t, err) + assert.Len(t, users, 3) +} + func testGlobalSchedule(t *testing.T, ds fleet.Datastore) { test.AddAllHostsLabel(t, ds) - _, server := runServerForTestsWithDS(t, ds) + _, server := RunServerForTestsWithDS(t, ds) token := getTestAdminToken(t, server) gs := fleet.GlobalSchedulePayload{} @@ -329,6 +343,7 @@ func TestIntegration(t *testing.T) { testUserWithoutRoleErrors, testUserWithWrongRoleErrors, testAppConfigAdditionalQueriesCanBeRemoved, + testUserRolesSpec, testGlobalSchedule, }) } diff --git a/server/service/metrics.go b/server/service/metrics.go index bee0c9ff76..f18824742e 100644 --- a/server/service/metrics.go +++ b/server/service/metrics.go @@ -11,7 +11,7 @@ type metricsMiddleware struct { requestLatency metrics.Histogram } -// NewMetrics service takes an existing service and wraps it +// NewMetricsService service takes an existing service and wraps it // with instrumentation middleware. func NewMetricsService( svc fleet.Service, diff --git a/server/service/util_test.go b/server/service/testing_utils.go similarity index 74% rename from server/service/util_test.go rename to server/service/testing_utils.go index c31c2829ed..aec85a8075 100644 --- a/server/service/util_test.go +++ b/server/service/testing_utils.go @@ -1,6 +1,10 @@ package service import ( + "fmt" + "net/http" + "net/http/httptest" + "os" "strings" "testing" @@ -9,7 +13,11 @@ import ( "github.com/fleetdm/fleet/v4/server/fleet" "github.com/fleetdm/fleet/v4/server/ptr" kitlog "github.com/go-kit/kit/log" + "github.com/go-kit/kit/transport" + kithttp "github.com/go-kit/kit/transport/http" + "github.com/gorilla/mux" "github.com/stretchr/testify/require" + "github.com/throttled/throttled/v2/store/memstore" ) func newTestService(ds fleet.Datastore, rs fleet.QueryResultStore, lq fleet.LiveQueryStore) fleet.Service { @@ -105,3 +113,31 @@ func (svc *mockMailService) SendEmail(e fleet.Email) error { svc.Invoked = true return svc.SendEmailFn(e) } + +func RunServerForTestsWithDS(t *testing.T, ds fleet.Datastore) (map[string]fleet.User, *httptest.Server) { + svc := newTestService(ds, nil, nil) + users := createTestUsers(t, ds) + logger := kitlog.NewLogfmtLogger(os.Stdout) + + opts := []kithttp.ServerOption{ + kithttp.ServerBefore( + setRequestsContexts(svc), + ), + kithttp.ServerErrorHandler(transport.NewLogErrorHandler(logger)), + kithttp.ServerAfter( + kithttp.SetContentType("application/json; charset=utf-8"), + ), + } + r := mux.NewRouter() + limitStore, _ := memstore.New(0) + ke := MakeFleetServerEndpoints(svc, "", limitStore) + kh := makeKitHandlers(ke, opts) + attachFleetAPIRoutes(r, kh) + attachNewStyleFleetAPIRoutes(r, svc, opts) + r.Handle("/", http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + fmt.Fprint(w, "index") + })) + + server := httptest.NewServer(r) + return users, server +} diff --git a/server/service/transport_queries.go b/server/service/transport_queries.go index c91d5cd642..763ffb14a6 100644 --- a/server/service/transport_queries.go +++ b/server/service/transport_queries.go @@ -79,5 +79,4 @@ func decodeApplyQuerySpecsRequest(ctx context.Context, r *http.Request) (interfa return nil, err } return req, nil - } diff --git a/server/service/user_roles.go b/server/service/user_roles.go new file mode 100644 index 0000000000..e70240fd39 --- /dev/null +++ b/server/service/user_roles.go @@ -0,0 +1,104 @@ +package service + +import ( + "context" + "net/http" + "time" + + "github.com/fleetdm/fleet/v4/server/fleet" + kithttp "github.com/go-kit/kit/transport/http" + "gopkg.in/guregu/null.v3" +) + +type applyUserRoleSpecsRequest struct { + Spec *fleet.UsersRoleSpec `json:"spec"` +} + +type applyUserRoleSpecsResponse struct { + Err error `json:"error,omitempty"` +} + +func (r applyUserRoleSpecsResponse) error() error { return r.Err } + +func makeApplyUserRoleSpecsEndpoint(svc fleet.Service, opts []kithttp.ServerOption) http.Handler { + return newServer( + makeAuthenticatedServiceEndpoint(svc, applyUserRoleSpecsEndpoint), + makeDecoderForType(applyUserRoleSpecsRequest{}), + opts, + ) +} + +func applyUserRoleSpecsEndpoint(ctx context.Context, request interface{}, svc fleet.Service) (interface{}, error) { + req := request.(*applyUserRoleSpecsRequest) + err := svc.ApplyUserRolesSpecs(ctx, *req.Spec) + if err != nil { + return applyUserRoleSpecsResponse{Err: err}, nil + } + return applyUserRoleSpecsResponse{}, nil +} + +func (svc Service) ApplyUserRolesSpecs(ctx context.Context, specs fleet.UsersRoleSpec) error { + if err := svc.authz.Authorize(ctx, &fleet.User{}, fleet.ActionWrite); err != nil { + return err + } + + var users []*fleet.User + for email, spec := range specs.Roles { + user, err := svc.ds.UserByEmail(email) + if err != nil { + return err + } + // If an admin is downgraded, make sure there is at least one other admin + err = svc.checkAtLeastOneAdmin(user, spec, email) + if err != nil { + return err + } + user.GlobalRole = spec.GlobalRole + var teams []fleet.UserTeam + for _, team := range spec.Teams { + t, err := svc.ds.TeamByName(team.Name) + if err != nil { + return err + } + teams = append(teams, fleet.UserTeam{ + Team: *t, + Role: team.Role, + }) + } + user.Teams = teams + users = append(users, user) + } + + return svc.ds.SaveUsers(users) +} + +func (svc Service) checkAtLeastOneAdmin(user *fleet.User, spec *fleet.UserRoleSpec, email string) error { + if null.StringFromPtr(user.GlobalRole).ValueOrZero() == fleet.RoleAdmin && + null.StringFromPtr(spec.GlobalRole).ValueOrZero() != fleet.RoleAdmin { + users, err := svc.ds.ListUsers(fleet.UserListOptions{}) + if err != nil { + return err + } + adminsExceptCurrent := 0 + for _, u := range users { + if u.Email == email { + continue + } + if null.StringFromPtr(u.GlobalRole).ValueOrZero() == fleet.RoleAdmin { + adminsExceptCurrent++ + } + } + if adminsExceptCurrent == 0 { + return fleet.NewError(fleet.ErrNoOneAdminNeeded, "You need at least one admin") + } + } + return nil +} + +func (mw loggingMiddleware) ApplyUserRolesSpecs(ctx context.Context, specs fleet.UsersRoleSpec) (err error) { + defer func(begin time.Time) { + _ = mw.loggerDebug(err).Log("method", "ApplyUserRolesSpecs", "err", err, "took", time.Since(begin)) + }(time.Now()) + err = mw.Service.ApplyUserRolesSpecs(ctx, specs) + return err +}