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:
Tomas Touceda
2021-07-16 15:28:13 -03:00
committed by GitHub
parent a38a7f4ad4
commit 545b3f396e
28 changed files with 902 additions and 251 deletions
+11
View File
@@ -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
View File
@@ -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
+97
View File
@@ -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)
}
+13 -2
View File
@@ -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
View File
@@ -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()
}
+83
View File
@@ -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"}))
}
+59
View File
@@ -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()
}
+17 -4
View File
@@ -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)
+67 -48
View File
@@ -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
}
+2
View File
@@ -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
+1
View File
@@ -19,5 +19,6 @@ type Service interface {
CarveService
TeamService
ActivitiesService
UserRolesService
GlobalScheduleService
}
+4
View File
@@ -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.")
}
+26
View File
@@ -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
}
+1
View File
@@ -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
+10
View File
@@ -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)
+31 -1
View File
@@ -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
}
+19 -25
View File
@@ -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)
}
+4 -4
View File
@@ -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 }
+2 -2
View File
@@ -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 {
+8 -8
View File
@@ -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 }
+34
View File
@@ -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)
}
}
+18
View File
@@ -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.
+22 -30
View File
@@ -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(&params)
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) {
+53 -38
View File
@@ -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(&params)
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,
})
}
+1 -1
View File
@@ -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
}
-1
View File
@@ -79,5 +79,4 @@ func decodeApplyQuerySpecsRequest(ctx context.Context, r *http.Request) (interfa
return nil, err
}
return req, nil
}
+104
View File
@@ -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
}