diff --git a/authn/authenticate.go b/authn/authenticate.go index 48dc5a823..81a87ed5c 100644 --- a/authn/authenticate.go +++ b/authn/authenticate.go @@ -122,6 +122,13 @@ func NewAuth(logger logger.Logger, url string, scopes []string, authURL, tokenUR return auth, nil } +// CleanOAuthConfig returns a's oauthConfig without the client secret +func (a Auth) CleanOAuthConfig() oauth2.Config { + b := *a.oAuthConfig + b.ClientSecret = "" + return b +} + // SecretKey is a convenient function to get the SecretKey from an Auth struct func (a Auth) SecretKey() []byte { return a.secretKey diff --git a/authn/authenticate_internal_test.go b/authn/authenticate_internal_test.go index 727180aa7..caae2350e 100644 --- a/authn/authenticate_internal_test.go +++ b/authn/authenticate_internal_test.go @@ -651,6 +651,22 @@ func (s *ServerTransportStream) SetTrailer(md metadata.MD) error { return nil } +func TestCleanOAuthConfig(t *testing.T) { + a := NewTestAuth(t) + res := a.CleanOAuthConfig() + assertEqual("", res.ClientSecret, t) + assertEqual(a.oAuthConfig.ClientID, res.ClientID, t) + assertEqual(a.oAuthConfig.RedirectURL, res.RedirectURL, t) + assertEqual(a.oAuthConfig.Scopes, res.Scopes, t) + assertEqual(a.oAuthConfig.Endpoint, res.Endpoint, t) +} + +func assertEqual(exp, got interface{}, t *testing.T) { + if !reflect.DeepEqual(exp, got) { + t.Fatalf("expected %v, got %v", exp, got) + } +} + func TestCheckAllowedNetworks(t *testing.T) { tests := []struct { diff --git a/cmd/auth_token.go b/cmd/auth_token.go new file mode 100644 index 000000000..894c0507e --- /dev/null +++ b/cmd/auth_token.go @@ -0,0 +1,29 @@ +// Copyright 2022 Molecula Corp. All rights reserved. +package cmd + +import ( + "context" + "io" + + "github.com/molecula/featurebase/v3/ctl" + "github.com/spf13/cobra" +) + +func newAuthTokenCommand(stdin io.Reader, stdout io.Writer, stderr io.Writer) *cobra.Command { + cmd := ctl.NewAuthTokenCommand(stdin, stdout, stderr) + ccmd := &cobra.Command{ + Use: "auth-token", + Short: "Get an auth-token", + Long: ` +Retrieves an auth-token for use in authenticating with FeatureBase from the configured identity provider. +`, + RunE: func(c *cobra.Command, args []string) error { + return cmd.Run(context.Background()) + }, + } + + flags := ccmd.Flags() + flags.StringVar(&cmd.Host, "host", "https://localhost:10101", "The address (host:port) of FeatureBase (HTTPs).") + ctl.SetTLSConfig(flags, "", &cmd.TLS.CertificatePath, &cmd.TLS.CertificateKeyPath, &cmd.TLS.CACertPath, &cmd.TLS.SkipVerify, &cmd.TLS.EnableClientVerification) + return ccmd +} diff --git a/cmd/root.go b/cmd/root.go index 5feee57d0..67b240fb9 100644 --- a/cmd/root.go +++ b/cmd/root.go @@ -57,6 +57,7 @@ at https://docs.molecula.cloud/. rc.AddCommand(newExportCommand(stdin, stdout, stderr)) rc.AddCommand(newGenerateConfigCommand(stdin, stdout, stderr)) rc.AddCommand(newImportCommand(stdin, stdout, stderr)) + rc.AddCommand(newAuthTokenCommand(stdin, stdout, stderr)) rc.AddCommand(newRBFCommand(stdin, stdout, stderr)) rc.AddCommand(newServeCmd(stdin, stdout, stderr)) rc.AddCommand(newHolderCmd(stdin, stdout, stderr)) diff --git a/ctl/auth_token.go b/ctl/auth_token.go new file mode 100644 index 000000000..a7a2de138 --- /dev/null +++ b/ctl/auth_token.go @@ -0,0 +1,273 @@ +// Copyright 2022 Molecula Corp. All rights reserved. +package ctl + +import ( + "context" + "crypto/tls" + "encoding/json" + "fmt" + "io" + "net/http" + "net/url" + "strings" + "time" + + pilosa "github.com/molecula/featurebase/v3" + "github.com/molecula/featurebase/v3/server" + "golang.org/x/oauth2" +) + +// AuthTokenCommand represents a command for retrieving an auth-token from a +// FeatureBase node. +type AuthTokenCommand struct { // nolint: maligned + // TLS configuration. + TLS server.TLSConfig + tlsConfig *tls.Config + + // Host is the host and port of the FeatureBase node to authenticate with. + Host string `json:"host"` + + // Reusable client. + client *pilosa.InternalClient + + // Standard input/output. + *pilosa.CmdIO +} + +// NewAuthTokenCommand returns a new instance of AuthTokenCommand. +func NewAuthTokenCommand(stdin io.Reader, stdout, stderr io.Writer) *AuthTokenCommand { + return &AuthTokenCommand{ + CmdIO: pilosa.NewCmdIO(stdin, stdout, stderr), + } +} + +type deviceAuthResponse struct { + DeviceCode string `json:"device_code"` + UserCode string `json:"user_code"` + VerificationURI string `json:"verification_uri"` + VerificationURIComplete string `json:"verification_uri_complete"` + ExpiresIn int `json:"expires_in"` + Interval int `json:"interval"` +} + +// deviceAuthRequest makes a Device Authorization Request and parses the Device +// Authorization Response as defined in RFC 8628 sections 3.1 and 3.2. +func deviceAuthRequest(cli *http.Client, config oauth2.Config) (rsp *deviceAuthResponse, err error) { + // build the request + req, err := http.NewRequest( + http.MethodPost, + pilosa.ReplaceFirstFromBack(config.Endpoint.AuthURL, "authorize", "devicecode"), + strings.NewReader( + url.Values{ + "client_id": {config.ClientID}, + "scope": {strings.Join(config.Scopes, " ")}, + }.Encode(), + ), + ) + if err != nil { + return nil, fmt.Errorf("building request: %w", err) + } + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + + // do the request + dar := &deviceAuthResponse{} + if resp, err := cli.Do(req); err != nil { + return nil, fmt.Errorf("making request with status %d: %v", resp.StatusCode, err) + } else if !(resp.StatusCode >= 200 && resp.StatusCode < 300) { + body, _ := io.ReadAll(resp.Body) // ignore the error bc it's ok if there's no body + return nil, fmt.Errorf( + "unsuccessful with status %d: %s", + resp.StatusCode, + strings.TrimSpace(string(body)), + ) + } else if err := json.NewDecoder(resp.Body).Decode(dar); err != nil { + return nil, fmt.Errorf("decoding device auth response body: %w", err) + } + + return dar, nil +} + +// successResponse represents a successful device auth token reponse. +type successResponse struct { + Access string `json:"access_token"` + Refresh string `json:"refresh_token"` + Type string `json:"token_type"` + ExpiresIn int `json:"expires_in"` + Scope string `json:"scope"` +} + +// errorResponse represents a device auth token error reponse. +type errorResponse struct { + Err error + Error string `json:"error"` + Description string `json:"error_description"` + URI string `json:"error_uri"` +} + +// waitResponse represents a device auth token reponse that indicates the client +// should continue polling. +type waitResponse struct { + SlowDown bool +} + +// parseResponse parses a device auth grant response from the /token endpoint. It +// returns a successResponse on success, an errorResponse on error, or a +// waitResponse if the authorization is still pending or we are requested to slow +// down. +func parseResponse(resp *http.Response) interface{} { + if s := resp.StatusCode; s >= 200 && s < 300 { + rsp := &successResponse{} + if err := json.NewDecoder(resp.Body).Decode(&rsp); err != nil { + return errorResponse{Err: fmt.Errorf("reading response body: %w", err)} + } + return *rsp + } + + rsp := &errorResponse{} + if err := json.NewDecoder(resp.Body).Decode(&rsp); err != nil { + return errorResponse{Err: fmt.Errorf("reading response body: %w", err)} + } + + switch rsp.Error { + case "slow_down": + return waitResponse{SlowDown: true} + case "authorization_pending": + return waitResponse{} + default: + rsp.Err = fmt.Errorf("error: %s, description: %s, uri: %s", rsp.Error, rsp.Description, rsp.URI) + return *rsp + } +} + +// Run executes the main program execution. +func (cmd *AuthTokenCommand) Run(ctx context.Context) (err error) { + // Parse TLS configuration for node-specific clients. + tls := cmd.TLSConfiguration() + if cmd.tlsConfig, err = server.GetTLSConfig(&tls, cmd.Logger()); err != nil { + return fmt.Errorf("parsing tls config: %w", err) + } + + // Create an internal client. + client, err := commandClient(cmd) + if err != nil { + return fmt.Errorf("creating client: %w", err) + } + + // Get the OAuth config. + var access, refresh string + config, err := client.OAuthConfig() + if err != nil { + return fmt.Errorf("getting oauth config: %w", err) + } + cli := &http.Client{} + + // Make the device authorization request. + dar, err := deviceAuthRequest(cli, config) + if err != nil { + return fmt.Errorf("making device auth request: %w", err) + } + + // Prompt the user to visit verification_uri and enter code. + fmt.Printf(formatPromptBox(dar.VerificationURI, dar.UserCode)) + + // Request a token until success or error response, slowing down if requested. + interval := dar.Interval + var exit bool + + for !exit { + time.Sleep(time.Duration(interval) * time.Second) + + // Make a device access token request to the token endpoint. + req, err := http.NewRequest( + http.MethodPost, + config.Endpoint.TokenURL, + strings.NewReader( + url.Values{ + "grant_type": {"urn:ietf:params:oauth:grant-type:device_code"}, + "client_id": {config.ClientID}, + "device_code": {dar.DeviceCode}, + }.Encode(), + ), + ) + if err != nil { + return fmt.Errorf("building request to %s: %w", config.Endpoint.TokenURL, err) + } + + resp, err := cli.Do(req) + if err != nil { + return fmt.Errorf("making request to %s: %w", config.Endpoint.TokenURL, err) + } + + output := parseResponse(resp) + + switch o := output.(type) { + case successResponse: + access = o.Access + refresh = o.Refresh + exit = true + case errorResponse: + return o.Err + case waitResponse: + if o.SlowDown { + interval += 5 + } + } + } + + // Convey the response to the user. + fmt.Printf("\nauth-token: %s\n", access) + fmt.Printf("\nrefresh-token: %s\n", refresh) + return nil +} + +// TLSHost implements the CommandWithTLSSupport interface. +func (cmd *AuthTokenCommand) TLSHost() string { return cmd.Host } + +// TLSConfiguration implements the CommandWithTLSSupport interface. +func (cmd *AuthTokenCommand) TLSConfiguration() server.TLSConfig { return cmd.TLS } + +// formatPromptBox makes a nice little prompt box to hold the url and code. +func formatPromptBox(url, code string) string { + top := "+-----------------------------------------------" + buffer := "| " + width := len(top) + 1 + if len(url) >= width { + width = len(url) + 5 + } + if len(code) >= width { + width = len(code) + 5 + } + + urlStr := fmt.Sprintf("| %s", url) + for i := len(urlStr); i < width-1; i++ { + urlStr += " " + } + urlStr += "|" + + codeStr := fmt.Sprintf("| %s", code) + for i := len(codeStr); i < width-1; i++ { + codeStr += " " + } + codeStr += "|" + + for i := len(top); i < width-1; i++ { + top += "-" + buffer += " " + } + top += "+\n" + buffer += "|\n" + + pls := "| Please visit:" + for i := len(pls); i < width-1; i++ { + pls += " " + } + pls += "|\n" + + and := "| And enter the code:" + for i := len(and); i < width-1; i++ { + and += " " + } + and += "|\n" + + return fmt.Sprintf("%s%s%s%s\n%s%s%s\n%s%s", top, buffer, pls, urlStr, buffer, and, codeStr, buffer, top) +} diff --git a/ctl/auth_token_internal_test.go b/ctl/auth_token_internal_test.go new file mode 100644 index 000000000..e850a6a2c --- /dev/null +++ b/ctl/auth_token_internal_test.go @@ -0,0 +1,191 @@ +package ctl + +import ( + "bytes" + "encoding/json" + "fmt" + "io" + "net/http" + "net/http/httptest" + "reflect" + "testing" + + "golang.org/x/oauth2" +) + +func TestFormatPromptBox(t *testing.T) { + golden := `+----------------------------------------------------+ +| | +| Please visit: | +| testingtestingtestingtestingtestingtestingtesting | +| | +| And enter the code: | +| blahblah | +| | ++----------------------------------------------------+ +` + got := formatPromptBox("testingtestingtestingtestingtestingtestingtesting", "blahblah") + if got != golden { + t.Fatalf("expected:\n%s, got:\n%s", golden, got) + } +} + +type authReqTest struct { + config oauth2.Config + expRsp *deviceAuthResponse + expErr error +} + +func TestDeviceAuthRequest(t *testing.T) { + goodClientID := "ring-a-ding-dillo" + goodResponse := deviceAuthResponse{ + DeviceCode: "Old knives are long enough as swords for hobbit-people.", + UserCode: "Sharp blades are good to have, if Shire-folk go walking, east, south, or far away into dark and danger.", + VerificationURI: "I am no weather-master, nor is aught that goes on two legs.", + VerificationURIComplete: "Hey! Come merry dol! derry dol! My hearties!", + ExpiresIn: -486846000, + Interval: 9, + } + srv := httptest.NewServer( + http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + clientID := r.PostFormValue("client_id") + if clientID != goodClientID { + http.Error( + w, + "Get out, you old Wight! Vanish in the sunlight!", + http.StatusBadRequest, + ) + return + } + b := bytes.Buffer{} + if err := json.NewEncoder(&b).Encode(goodResponse); err != nil { + t.Fatalf("unexpected error encoding goodResponse: %v", err) + } + w.Write(b.Bytes()) + w.WriteHeader(http.StatusOK) + }), + ) + cli := &http.Client{} + for name, test := range map[string]authReqTest{ + "badRequest": { + config: oauth2.Config{ + ClientID: "Iarwain Ben-adar", + Endpoint: oauth2.Endpoint{AuthURL: srv.URL}, + }, + expRsp: nil, + expErr: fmt.Errorf("unsuccessful with status 400: Get out, you old Wight! Vanish in the sunlight!"), + }, + "goodRequest": { + config: oauth2.Config{ + ClientID: "ring-a-ding-dillo", + Endpoint: oauth2.Endpoint{AuthURL: srv.URL}, + }, + expRsp: &goodResponse, + expErr: nil, + }, + } { + t.Run(name, func(t *testing.T) { + got, err := deviceAuthRequest(cli, test.config) + if !errEqual(err, test.expErr) { + t.Errorf("expected '%v', got '%v'", test.expErr, err) + } + if !reflect.DeepEqual(got, test.expRsp) { + t.Errorf("expected '%v', got '%v'", test.expRsp, got) + } + }) + } +} + +type respTest struct { + resp *http.Response + exp interface{} +} + +func TestParseResponse(t *testing.T) { + sr := successResponse{ + Access: "ACCESS", + Refresh: "REFRESH", + Type: "access", + ExpiresIn: 10000, + Scope: "scopity-scope-scopity-scope-pope", + } + good := bytes.Buffer{} + if err := json.NewEncoder(&good).Encode(sr); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + sd := errorResponse{Error: "slow_down"} + slowDown := bytes.Buffer{} + if err := json.NewEncoder(&slowDown).Encode(sd); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + sd.Error = "authorization_pending" + pending := bytes.Buffer{} + if err := json.NewEncoder(&pending).Encode(sd); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + sd.Error = "invalid_client" + genErr := bytes.Buffer{} + if err := json.NewEncoder(&genErr).Encode(sd); err != nil { + t.Fatalf("unexpected error: %v", err) + } + + for name, test := range map[string]respTest{ + "success": { + resp: &http.Response{ + StatusCode: 200, + Body: io.NopCloser(&good), + ContentLength: int64(good.Len()), + }, + exp: sr, + }, + "slowDown": { + resp: &http.Response{ + StatusCode: 400, + Body: io.NopCloser(&slowDown), + ContentLength: int64(slowDown.Len()), + }, + exp: waitResponse{SlowDown: true}, + }, + "pending": { + resp: &http.Response{ + StatusCode: 400, + Body: io.NopCloser(&pending), + ContentLength: int64(pending.Len()), + }, + exp: waitResponse{}, + }, + "generalError": { + resp: &http.Response{ + StatusCode: 400, + Body: io.NopCloser(&genErr), + ContentLength: int64(pending.Len()), + }, + exp: errorResponse{ + Error: "invalid_client", + Description: "", + URI: "", + Err: fmt.Errorf("error: %s, description: %s, uri: %s", "invalid_client", + "", ""), + }, + }, + } { + t.Run(name, func(t *testing.T) { + if got := parseResponse(test.resp); !reflect.DeepEqual(got, test.exp) { + t.Errorf("expected: '%v', got '%v'", test.exp, got) + } + }) + } +} + +func errEqual(a, b error) bool { + if a == nil { + return b == nil + } + if b == nil { + return a == nil + } + return a.Error() == b.Error() +} diff --git a/http_handler.go b/http_handler.go index 7fcba2816..c43ebaeda 100644 --- a/http_handler.go +++ b/http_handler.go @@ -533,6 +533,7 @@ func newRouter(handler *Handler) http.Handler { router.HandleFunc("/redirect", handler.handleRedirect).Methods("GET").Name("Redirect") router.HandleFunc("/auth", handler.handleCheckAuthentication).Methods("GET").Name("CheckAuthentication") router.HandleFunc("/userinfo", handler.handleUserInfo).Methods("GET").Name("UserInfo") + router.HandleFunc("/internal/oauth-config", handler.handleOAuthConfig).Methods("GET").Name("GetOAuthConfig") // Endpoints to support lattice UI embedded via statik. // The messiness here reflects the fact that assets live in a nontrivial @@ -3793,6 +3794,26 @@ func (h *Handler) handleRedirect(w http.ResponseWriter, r *http.Request) { h.auth.Redirect(w, r) } +// handleOAuthConfig handles requests for a cleaned version of our oAuthConfig. We +// use this endpoint /internal/oauth-config in the `featurebase auth-token` +// subcommand to create a RedirectURL on the fly. +func (h *Handler) handleOAuthConfig(w http.ResponseWriter, r *http.Request) { + if !validHeaderAcceptJSON(r.Header) { + http.Error(w, "JSON only acceptable response", http.StatusNotAcceptable) + return + } + + if h.auth == nil { + http.Error(w, "auth not enabled: no OAuthConfig", http.StatusNotFound) + return + } + + config := h.auth.CleanOAuthConfig() + if err := json.NewEncoder(w).Encode(config); err != nil { + h.logger.Errorf("writing oauth-config info: %s", err) + } +} + func (h *Handler) handleCheckAuthentication(w http.ResponseWriter, r *http.Request) { if !validHeaderAcceptJSON(r.Header) { http.Error(w, "JSON only acceptable response", http.StatusNotAcceptable) diff --git a/http_handler_internal_test.go b/http_handler_internal_test.go index 1ce0d3584..2ddd27d7c 100644 --- a/http_handler_internal_test.go +++ b/http_handler_internal_test.go @@ -827,6 +827,37 @@ func TestHandleGetDiskUsage(t *testing.T) { } } +func TestHandleOAuthConfig(t *testing.T) { + h := Handler{ + logger: logger.NewStandardLogger(os.Stdout), + queryLogger: logger.NewStandardLogger(os.Stdout), + auth: NewTestAuth(t), + api: &API{ + server: &Server{ + dataDir: t.TempDir(), + }, + }, + } + w := httptest.NewRecorder() + r := httptest.NewRequest("GET", "/whatever", nil) + + h.handleOAuthConfig(w, r) + + resp := w.Result() + if resp.StatusCode != http.StatusOK { + t.Fatalf("expected %v, got %v", http.StatusOK, resp.StatusCode) + } + + defer resp.Body.Close() + var rsp oauth2.Config + if err := json.NewDecoder(resp.Body).Decode(&rsp); err != nil { + t.Fatalf("unexpected error decoding body: %v", err) + } + if exp := h.auth.CleanOAuthConfig(); !reflect.DeepEqual(exp, rsp) { + t.Fatalf("expected %v, got %v", exp, rsp) + } +} + func TestAuthzAllowedIPs(t *testing.T) { tests := []struct { configuredIPs []string diff --git a/internal_client.go b/internal_client.go index 99774f015..e529e33f0 100644 --- a/internal_client.go +++ b/internal_client.go @@ -27,6 +27,7 @@ import ( "github.com/molecula/featurebase/v3/topology" "github.com/molecula/featurebase/v3/tracing" "github.com/pkg/errors" + "golang.org/x/oauth2" ) // InternalClient represents a client to the Pilosa cluster. @@ -2453,3 +2454,26 @@ func (c *InternalClient) PartitionNodes(ctx context.Context, partitionID int) ([ func (c *InternalClient) SetInternalAPI(api *API) { c.api = api } + +func (c *InternalClient) OAuthConfig() (rsp oauth2.Config, err error) { + u := uriPathToURL(c.defaultURI, "/internal/oauth-config") + + req, err := http.NewRequest("GET", u.String(), nil) + if err != nil { + return rsp, errors.Wrap(err, "creating request") + } + + req.Header.Set("User-Agent", "pilosa/"+Version) + req.Header.Set("Accept", "application/json") + resp, err := c.executeRequest(req) + if err != nil { + return rsp, fmt.Errorf("getting config: %w", err) + } + defer resp.Body.Close() + + if err := json.NewDecoder(resp.Body).Decode(&rsp); err != nil { + return rsp, fmt.Errorf("json decode: %s", err) + } + + return rsp, nil +} diff --git a/util.go b/util.go index 88287a91c..35c5974b9 100644 --- a/util.go +++ b/util.go @@ -124,3 +124,26 @@ func EtcdUnixSocket(tb testing.TB) string { }) return fmt.Sprintf("unix://%s", addr) } + +// Rev reverses a string +func Rev(input string) string { + n := 0 + runes := make([]rune, len(input)) + for _, r := range input { + runes[n] = r + n++ + } + runes = runes[0:n] + + for i := 0; i < n/2; i++ { + runes[i], runes[n-1-i] = runes[n-1-i], runes[i] + } + + return string(runes) +} + +// ReplaceFirstFromBack replaces the first instance of toReplace from the back of +// the string s +func ReplaceFirstFromBack(s, toReplace, replacement string) string { + return Rev(strings.Replace(Rev(s), Rev(toReplace), Rev(replacement), 1)) +} diff --git a/utils_internal_test.go b/utils_internal_test.go index 072fd5b1c..80667d71d 100644 --- a/utils_internal_test.go +++ b/utils_internal_test.go @@ -106,3 +106,37 @@ func NewTestClusterWithReplication(tb testing.TB, nNodes, nReplicas, partitionN c.close() } } + +func TestReplaceFirstFromBack(t *testing.T) { + for name, test := range map[string]struct { + input string + exp string + toReplace string + replacement string + }{ + "url": { + input: "https://login.microsoftonline.com/4a137d66-d161-4ae4-b1e6-07e9920874b8/oauth2/v2.0/authorize", + exp: "https://login.microsoftonline.com/4a137d66-d161-4ae4-b1e6-07e9920874b8/oauth2/v2.0/devicecode", + toReplace: "authorize", + replacement: "devicecode", + }, + "unicode": { + input: "那不是兽人号角", + exp: "那是一只兽人号角", + toReplace: "不是", + replacement: "是一只", + }, + "multiple": { + input: "cowscowscowscowscows", + exp: "cowscowscowscowscats", + toReplace: "cows", + replacement: "cats", + }, + } { + t.Run(name, func(t *testing.T) { + if got := ReplaceFirstFromBack(test.input, test.toReplace, test.replacement); got != test.exp { + t.Fatalf("expected %v, got %v", test.exp, got) + } + }) + } +}