Skip to content

Commit d10f3ba

Browse files
authored
Merge pull request #12 from SlyngDK/request-ip
Support enroll with groups and requesting nebula ip
2 parents 9bdd7e7 + 3018a35 commit d10f3ba

26 files changed

Lines changed: 887 additions & 342 deletions

cmd/agent/enrollment.go

Lines changed: 37 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -7,6 +7,7 @@ import (
77
"io"
88
"io/ioutil"
99
"os"
10+
"strings"
1011
"time"
1112

1213
"github.com/slackhq/nebula/cert"
@@ -31,8 +32,16 @@ var enrollCmd = &cobra.Command{
3132
if err != nil {
3233
l.WithError(err).Fatalln("failed to get token")
3334
}
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+
}
3443

35-
if err = enroll(agent, token); err != nil {
44+
if err = enroll(agent, token, groups, ip); err != nil {
3645
l.WithError(err).Fatalln("failed to enroll to server")
3746
}
3847
},
@@ -86,7 +95,16 @@ var enrollWaitCmd = &cobra.Command{
8695
if status == 0 {
8796
token, _ := cmd.Flags().GetString("token")
8897
if token != "" {
89-
if err := enroll(agent, token); err != nil {
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+
}
106+
107+
if err := enroll(agent, token, groups, ip); err != nil {
90108
l.WithError(err).Fatalln("failed to enroll to server")
91109
os.Exit(2)
92110
}
@@ -136,14 +154,19 @@ func getStatus(agent *agentClient) int8 {
136154

137155
func init() {
138156
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")
139159
enrollCmd.MarkFlagRequired("token")
140160
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")
141164

142165
enrollCmd.AddCommand(enrollStatusCmd)
143166
enrollCmd.AddCommand(enrollWaitCmd)
144167
}
145168

146-
func enroll(c *agentClient, enrollmentToken string) error {
169+
func enroll(c *agentClient, enrollmentToken, groups, ip string) error {
147170
if enrollmentToken == "" {
148171
return fmt.Errorf("requires enrollmentToken")
149172
}
@@ -158,6 +181,17 @@ func enroll(c *agentClient, enrollmentToken string) error {
158181
CsrPEM: string(csr),
159182
}
160183

184+
if groups != "" {
185+
g := strings.Split(groups, ",")
186+
if len(g) > 0 {
187+
enrollRequest.Groups = g
188+
}
189+
}
190+
191+
if ip != "" {
192+
enrollRequest.RequestedIP = ip
193+
}
194+
161195
hostname, err := os.Hostname()
162196
if err == nil {
163197
enrollRequest.Name = hostname

protocol/agent-service.pb.go

Lines changed: 79 additions & 68 deletions
Some generated files are not rendered by default. Learn more about customizing how changed files appear on GitHub.

protocol/agent-service.proto

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ message EnrollRequest {
2020
string csrPEM = 2;
2121
repeated string groups = 3;
2222
string name = 4;
23+
string requestedIP = 5;
2324
}
2425

2526
message EnrollResponse {

server/agent-service.go

Lines changed: 8 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -5,6 +5,7 @@ import (
55
"crypto/sha256"
66
"crypto/x509"
77
"fmt"
8+
"net"
89
"strings"
910

1011
"github.com/sirupsen/logrus"
@@ -58,6 +59,12 @@ func (a *agentService) Enroll(ctx context.Context, request *protocol.EnrollReque
5859
if request.CsrPEM == "" {
5960
return nil, status.Error(codes.InvalidArgument, "CsrPEM is required")
6061
}
62+
if request.RequestedIP != "" {
63+
ip := net.ParseIP(request.RequestedIP)
64+
if ip == nil {
65+
return nil, status.Error(codes.InvalidArgument, "RequestedIP is not valid")
66+
}
67+
}
6168

6269
_, _, err = cert.UnmarshalX25519PublicKey([]byte(request.CsrPEM))
6370
if err != nil {
@@ -68,7 +75,7 @@ func (a *agentService) Enroll(ctx context.Context, request *protocol.EnrollReque
6875
addr := p.Addr.String()
6976
ip := addr[0:strings.LastIndex(addr, ":")]
7077

71-
_, err = a.store.CreateEnrollmentRequest(fingerprint, request.Token, request.CsrPEM, ip, request.Name)
78+
_, err = a.store.CreateEnrollmentRequest(fingerprint, request.Token, request.CsrPEM, ip, request.Name, request.RequestedIP, request.Groups)
7279
if err != nil {
7380
return nil, status.Error(codes.Internal, fmt.Sprintf("%s", err))
7481
}

0 commit comments

Comments
 (0)