Skip to content

Commit e7e0f04

Browse files
committed
Refactored agent enroll, to create a new enrollment request if input is different from currently.
1 parent e4ac087 commit e7e0f04

2 files changed

Lines changed: 133 additions & 73 deletions

File tree

cmd/agent/enrollment.go

Lines changed: 117 additions & 73 deletions
Original file line numberDiff line numberDiff line change
@@ -6,8 +6,8 @@ import (
66
"fmt"
77
"io"
88
"io/ioutil"
9+
"net"
910
"os"
10-
"strings"
1111
"time"
1212

1313
"github.com/slackhq/nebula/cert"
@@ -28,22 +28,7 @@ var enrollCmd = &cobra.Command{
2828
}
2929
defer agent.Close()
3030

31-
token, err := cmd.Flags().GetString("token")
32-
if err != nil {
33-
l.WithError(err).Fatalln("failed to get token")
34-
}
35-
groups, err := cmd.Flags().GetString("groups")
36-
if err != nil {
37-
l.WithError(err).Fatalln("failed to get groups")
38-
}
39-
ip, err := cmd.Flags().GetString("ip")
40-
if err != nil {
41-
l.WithError(err).Fatalln("failed to get ip")
42-
}
43-
44-
if err = enroll(agent, token, groups, ip); err != nil {
45-
l.WithError(err).Fatalln("failed to enroll to server")
46-
}
31+
updateEnrollmentRequest(agent, cmd)
4732
},
4833
}
4934

@@ -91,82 +76,146 @@ var enrollWaitCmd = &cobra.Command{
9176
ticker := time.NewTicker(10 * time.Second)
9277
defer ticker.Stop()
9378

94-
status := getStatus(agent)
95-
if status == 0 {
96-
token, _ := cmd.Flags().GetString("token")
97-
if token != "" {
98-
groups, err := cmd.Flags().GetString("groups")
99-
if err != nil {
100-
l.WithError(err).Fatalln("failed to get groups")
101-
}
102-
ip, err := cmd.Flags().GetString("ip")
103-
if err != nil {
104-
l.WithError(err).Fatalln("failed to get ip")
105-
}
79+
updateEnrollmentRequest(agent, cmd)
10680

107-
if err := enroll(agent, token, groups, ip); err != nil {
108-
l.WithError(err).Fatalln("failed to enroll to server")
109-
os.Exit(2)
110-
}
111-
} else {
112-
l.Errorln("Require the enrollment has been started, you can provide enrollment token and doing it in one process.")
113-
os.Exit(1)
114-
}
115-
}
116-
if status == 2 {
81+
if isEnrollDone(agent) {
11782
return
11883
}
11984

12085
for {
12186
select {
12287
case _ = <-ticker.C:
123-
status := getStatus(agent)
124-
if status == 2 {
125-
l.Info("Agent is now enrolled")
88+
if isEnrollDone(agent) {
12689
return
12790
}
12891
}
12992
}
13093
},
13194
}
13295

133-
func getStatus(agent *agentClient) int8 {
96+
func init() {
97+
enrollCmd.Flags().StringP("token", "t", "", "Enrollment token")
98+
enrollCmd.Flags().StringSliceP("groups", "g", []string{}, "Comma separated list of groups")
99+
enrollCmd.Flags().StringP("ip", "i", "", "Requesting for this specific nebula ip")
100+
enrollCmd.MarkFlagRequired("token")
101+
enrollWaitCmd.Flags().StringP("token", "t", "", "Enrollment token")
102+
enrollWaitCmd.Flags().StringSliceP("groups", "g", []string{}, "Comma separated list of groups")
103+
enrollWaitCmd.Flags().StringP("ip", "i", "", "Requesting for this specific nebula ip")
104+
enrollWaitCmd.MarkFlagRequired("token")
105+
106+
enrollCmd.AddCommand(enrollStatusCmd)
107+
enrollCmd.AddCommand(enrollWaitCmd)
108+
}
109+
110+
func updateEnrollmentRequest(agent *agentClient, cmd *cobra.Command) {
134111
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
135112
defer cancel()
136113

137-
res, err := agent.client.GetEnrollStatus(ctx, &emptypb.Empty{})
114+
status, err := agent.client.GetEnrollStatus(ctx, &emptypb.Empty{})
138115
if err != nil {
139116
l.WithError(err).Fatalln("failed to get enrollment status")
117+
os.Exit(1)
140118
}
141119

142-
if res.IsEnrolled {
143-
l.Info("Agent is enrolled")
144-
l.Infof("IssuedAt: %s, ExpiresAt: %s\n", res.IssuedAt.AsTime().Format(time.RFC3339), res.ExpiresAt.AsTime().Format(time.RFC3339))
145-
return 2
146-
} else if res.IsEnrollmentRequested {
147-
l.Info("Agent has requested to be enrolled")
148-
return 1
120+
token, err := cmd.Flags().GetString("token")
121+
if err != nil {
122+
l.WithError(err).Fatalln("failed to get token")
123+
}
124+
if token == "" {
125+
l.Errorln("Require the enrollment has been started, you can provide enrollment token and doing it in one process.")
126+
os.Exit(1)
127+
}
128+
groups, err := cmd.Flags().GetStringSlice("groups")
129+
if err != nil {
130+
l.WithError(err).Fatalln("failed to get groups")
131+
}
132+
ip, err := cmd.Flags().GetString("ip")
133+
if err != nil {
134+
l.WithError(err).Fatalln("failed to get ip")
149135
} else {
150-
l.Info("Agent enrollment not started")
151-
return 0
136+
parseIP := net.ParseIP(ip)
137+
if parseIP == nil && ip != "" {
138+
l.WithError(err).Fatalln("ip is of invalid format")
139+
os.Exit(1)
140+
}
141+
}
142+
hostname, err := os.Hostname()
143+
if err != nil {
144+
l.WithError(err).Errorln("error when getting hostname")
145+
os.Exit(1)
146+
}
147+
148+
var diff = false
149+
150+
if status.EnrollmentRequest != nil {
151+
l.Debug("comparing against existing enrollment request")
152+
l.Debugf("hostname compare %s <=> %s", hostname, status.EnrollmentRequest.Name)
153+
if hostname != status.EnrollmentRequest.Name {
154+
diff = true
155+
l.Debugf("diff on hostname")
156+
}
157+
158+
l.Debugf("ip compare %s <=> %s", ip, status.EnrollmentRequest.RequestedIP)
159+
if ip != status.EnrollmentRequest.RequestedIP {
160+
diff = true
161+
l.Debugf("diff on ip")
162+
}
163+
164+
l.Debugf("groups compare %s <=> %s", groups, status.EnrollmentRequest.Groups)
165+
if !stringSlicesEqual(groups, status.EnrollmentRequest.Groups) {
166+
diff = true
167+
l.Debugf("diff on groups")
168+
}
169+
} else if status.IsEnrolled {
170+
l.Debug("comparing against enrolled agent")
171+
l.Debugf("hostname compare %s <=> %s", hostname, status.Name)
172+
if hostname != status.Name {
173+
diff = true
174+
l.Debugf("diff on hostname")
175+
}
176+
177+
l.Debugf("ip compare %s <=> %s", ip, status.AssignedIP)
178+
if ip != status.AssignedIP {
179+
diff = true
180+
l.Debugf("diff on ip")
181+
}
182+
183+
l.Debugf("groups compare %s <=> %s", groups, status.Groups)
184+
if !stringSlicesEqual(groups, status.Groups) {
185+
diff = true
186+
l.Debugf("diff on groups")
187+
}
188+
} else {
189+
diff = true
190+
}
191+
192+
if diff {
193+
l.Info("adding enrollment request")
194+
if err := enroll(agent, token, ip, groups); err != nil {
195+
l.WithError(err).Fatalln("failed to enroll agent")
196+
os.Exit(2)
197+
}
152198
}
153199
}
154200

155-
func init() {
156-
enrollCmd.Flags().StringP("token", "t", "", "Enrollment token")
157-
enrollCmd.Flags().StringP("groups", "g", "", "Comma separated list of groups")
158-
enrollCmd.Flags().StringP("ip", "i", "", "Requesting for this specific nebula ip")
159-
enrollCmd.MarkFlagRequired("token")
160-
enrollWaitCmd.Flags().StringP("token", "t", "", "Enrollment token")
161-
enrollWaitCmd.Flags().StringP("groups", "g", "", "Comma separated list of groups")
162-
enrollWaitCmd.Flags().StringP("ip", "i", "", "Requesting for this specific nebula ip")
163-
enrollWaitCmd.MarkFlagRequired("token")
201+
func isEnrollDone(agent *agentClient) bool {
202+
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
203+
defer cancel()
204+
status, err := agent.client.GetEnrollStatus(ctx, &emptypb.Empty{})
205+
if err != nil {
206+
l.WithError(err).Fatalln("failed to get enrollment status")
207+
os.Exit(1)
208+
}
164209

165-
enrollCmd.AddCommand(enrollStatusCmd)
166-
enrollCmd.AddCommand(enrollWaitCmd)
210+
if status.IsEnrolled && !status.IsEnrollmentRequested {
211+
l.Info("Agent is enrolled")
212+
return true
213+
}
214+
215+
return false
167216
}
168217

169-
func enroll(c *agentClient, enrollmentToken, groups, ip string) error {
218+
func enroll(c *agentClient, enrollmentToken, ip string, groups []string) error {
170219
if enrollmentToken == "" {
171220
return fmt.Errorf("requires enrollmentToken")
172221
}
@@ -181,11 +230,8 @@ func enroll(c *agentClient, enrollmentToken, groups, ip string) error {
181230
CsrPEM: string(csr),
182231
}
183232

184-
if groups != "" {
185-
g := strings.Split(groups, ",")
186-
if len(g) > 0 {
187-
enrollRequest.Groups = g
188-
}
233+
if len(groups) > 0 {
234+
enrollRequest.Groups = groups
189235
}
190236

191237
if ip != "" {
@@ -200,13 +246,11 @@ func enroll(c *agentClient, enrollmentToken, groups, ip string) error {
200246
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
201247
defer cancel()
202248

203-
res, err := c.client.Enroll(ctx, enrollRequest)
249+
_, err = c.client.Enroll(ctx, enrollRequest)
204250
if err != nil {
205251
return err
206252
}
207253

208-
c.l.Println(res)
209-
210254
return nil
211255
}
212256

cmd/agent/utils.go

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@ package main
33
import (
44
"os"
55
"path/filepath"
6+
"sort"
67
)
78

89
func fileExists(filepath string) (bool, os.FileInfo) {
@@ -52,3 +53,18 @@ func resolvePath(path string) string {
5253
}
5354
return path
5455
}
56+
57+
func stringSlicesEqual(a, b []string) bool {
58+
if len(a) != len(b) {
59+
return false
60+
}
61+
sort.Strings(a)
62+
sort.Strings(b)
63+
64+
for i, v := range a {
65+
if v != b[i] {
66+
return false
67+
}
68+
}
69+
return true
70+
}

0 commit comments

Comments
 (0)