Issue 1362 fleetctl user roles (#1397)
* WIP * Add get user_roles and apply for a user_roles spec to fleetctl * Uncomment other tests * Update test to check output * Update test with the new struct * Mock token so that it doesn't pick up the one in the local machine * Address review comments * Fix printJSON and printYaml * Fix merge conflict error * If both roles are specified, fail * Fix test * Switch arguments around * Update test with the new rule * Fix other tests that fell through the cracks
This commit is contained in:
@@ -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 {
|
||||
|
||||
+25
-12
@@ -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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
+154
-75
@@ -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()
|
||||
}
|
||||
|
||||
@@ -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"}))
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -19,5 +19,6 @@ type Service interface {
|
||||
CarveService
|
||||
TeamService
|
||||
ActivitiesService
|
||||
UserRolesService
|
||||
GlobalScheduleService
|
||||
}
|
||||
|
||||
@@ -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.")
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 }
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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.
|
||||
|
||||
@@ -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) {
|
||||
|
||||
@@ -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,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -79,5 +79,4 @@ func decodeApplyQuerySpecsRequest(ctx context.Context, r *http.Request) (interfa
|
||||
return nil, err
|
||||
}
|
||||
return req, nil
|
||||
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user