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
0 commit comments