Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
43 changes: 35 additions & 8 deletions packages/cmd/root.go
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,10 @@ var RootCmd = &cobra.Command{
Version: util.CLI_VERSION,
}

// getCurrentLoggedInUserDetails is a seam for testing root pre-run behavior
// without reading the platform keyring.
var getCurrentLoggedInUserDetails = util.GetCurrentLoggedInUserDetails

// rootCmdStderrWriter is a writer wrapper that dynamically reads from RootCmd.ErrOrStderr()
// on each write. This allows the logger to automatically use RootCmd's stderr even if it's
// changed after logger initialization (e.g., in tests).
Expand Down Expand Up @@ -85,6 +89,31 @@ func isStructuredOutputRequested(cmd *cobra.Command) bool {
return false
}

// shouldReadSavedSession reports whether this invocation needs the local user
// session. Machine-authenticated and silent commands do not need it.
func shouldReadSavedSession(silent, hasExplicitToken, hasExplicitUniversalAuthCredentials bool) bool {
return !silent && !hasExplicitToken && !hasExplicitUniversalAuthCredentials
}

func hasExplicitUniversalAuthCredentials(cmd *cobra.Command) bool {
if cmd.Name() != "login" {
return false
}

method, err := cmd.Flags().GetString("method")
if err != nil || method != string(util.AuthStrategy.UNIVERSAL_AUTH) {
return false
}

clientID, err := util.GetCmdFlagOrEnv(cmd, "client-id", []string{util.INFISICAL_UNIVERSAL_AUTH_CLIENT_ID_NAME})
if err != nil || clientID == "" {
return false
}

clientSecret, err := util.GetCmdFlagOrEnv(cmd, "client-secret", []string{util.INFISICAL_UNIVERSAL_AUTH_CLIENT_SECRET_NAME})
return err == nil && clientSecret != ""
}

// Execute adds all child commands to the root command and sets flags appropriately.
// This is called by main.main(). It only needs to happen once to the RootCmd.
func Execute() {
Expand Down Expand Up @@ -145,14 +174,12 @@ func init() {
util.DisplayPackageRepoMigrationNoticeWithWriter(silent, cmd.ErrOrStderr())
}

loggedInDetails, err := util.GetCurrentLoggedInUserDetails(false)

if !silent && err == nil && loggedInDetails.IsUserLoggedIn && !loggedInDetails.LoginExpired {
token, err := util.GetInfisicalToken(cmd)

if err == nil && token != nil {
util.PrintWarningWithWriter(fmt.Sprintf("Your logged-in session is being overwritten by the token provided from the %s.", token.Source), cmd.ErrOrStderr())
}
token, err := util.GetInfisicalToken(cmd)
if err != nil {
token = nil
}
if shouldReadSavedSession(silent, token != nil, hasExplicitUniversalAuthCredentials(cmd)) {
_, _ = getCurrentLoggedInUserDetails(false)
}

}
Expand Down
161 changes: 161 additions & 0 deletions packages/cmd/root_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
package cmd

import (
"testing"

"github.com/Infisical/infisical-merge/packages/config"
"github.com/Infisical/infisical-merge/packages/util"
"github.com/spf13/cobra"
)

func TestShouldReadSavedSession(t *testing.T) {
cases := []struct {
name string
silent bool
hasExplicitToken bool
hasUniversalAuthCredentials bool
want bool
}{
{name: "interactive user command", want: true},
{name: "service token", hasExplicitToken: true, want: false},
{name: "universal auth login", hasUniversalAuthCredentials: true, want: false},
{name: "silent command", silent: true, want: false},
}

for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
if got := shouldReadSavedSession(tc.silent, tc.hasExplicitToken, tc.hasUniversalAuthCredentials); got != tc.want {
t.Errorf("shouldReadSavedSession(%v, %v, %v) = %v, want %v", tc.silent, tc.hasExplicitToken, tc.hasUniversalAuthCredentials, got, tc.want)
}
})
}
}

func TestHasExplicitUniversalAuthCredentials(t *testing.T) {
newLoginCommand := func() *cobra.Command {
cmd := &cobra.Command{Use: "login"}
cmd.Flags().String("method", "user", "")
cmd.Flags().String("client-id", "", "")
cmd.Flags().String("client-secret", "", "")
return cmd
}

t.Run("flags", func(t *testing.T) {
cmd := newLoginCommand()
if err := cmd.Flags().Set("method", string(util.AuthStrategy.UNIVERSAL_AUTH)); err != nil {
t.Fatal(err)
}
if err := cmd.Flags().Set("client-id", "test-client-id"); err != nil {
t.Fatal(err)
}
if err := cmd.Flags().Set("client-secret", "test-client-secret"); err != nil {
t.Fatal(err)
}

if !hasExplicitUniversalAuthCredentials(cmd) {
t.Fatal("expected universal auth flags to bypass the saved-session lookup")
}
})

t.Run("environment variables", func(t *testing.T) {
cmd := newLoginCommand()
if err := cmd.Flags().Set("method", string(util.AuthStrategy.UNIVERSAL_AUTH)); err != nil {
t.Fatal(err)
}
t.Setenv(util.INFISICAL_UNIVERSAL_AUTH_CLIENT_ID_NAME, "test-client-id")
t.Setenv(util.INFISICAL_UNIVERSAL_AUTH_CLIENT_SECRET_NAME, "test-client-secret")

if !hasExplicitUniversalAuthCredentials(cmd) {
t.Fatal("expected universal auth environment variables to bypass the saved-session lookup")
}
})
}

func TestRootPersistentPreRunSavedSessionLookup(t *testing.T) {
originalLookup := getCurrentLoggedInUserDetails
t.Cleanup(func() { getCurrentLoggedInUserDetails = originalLookup })
originalURL := config.INFISICAL_URL
t.Cleanup(func() { config.INFISICAL_URL = originalURL })
t.Setenv("INFISICAL_DISABLE_UPDATE_CHECK", "1")
t.Setenv("INFISICAL_DISABLE_MIGRATION_NOTICE", "1")

lookupCalls := 0
getCurrentLoggedInUserDetails = func(bool) (util.LoggedInUserDetails, error) {
lookupCalls++
return util.LoggedInUserDetails{}, nil
}

newCommand := func(use string) *cobra.Command {
cmd := &cobra.Command{Use: use}
cmd.Flags().Bool("silent", false, "")
cmd.Flags().String("token", "", "")
cmd.Flags().String("method", "user", "")
cmd.Flags().String("client-id", "", "")
cmd.Flags().String("client-secret", "", "")
return cmd
}

cases := []struct {
name string
command *cobra.Command
setup func(t *testing.T, cmd *cobra.Command)
want int
}{
{
name: "INFISICAL_TOKEN environment variable",
command: newCommand("run"),
setup: func(t *testing.T, cmd *cobra.Command) {
t.Setenv(util.INFISICAL_TOKEN_NAME, "st.test")
},
want: 0,
},
{
name: "universal auth flags",
command: newCommand("login"),
setup: func(t *testing.T, cmd *cobra.Command) {
for flag, value := range map[string]string{
"method": string(util.AuthStrategy.UNIVERSAL_AUTH),
"client-id": "test-client-id",
"client-secret": "test-client-secret",
} {
if err := cmd.Flags().Set(flag, value); err != nil {
t.Fatal(err)
}
}
},
want: 0,
},
{
name: "universal auth environment variables",
command: newCommand("login"),
setup: func(t *testing.T, cmd *cobra.Command) {
if err := cmd.Flags().Set("method", string(util.AuthStrategy.UNIVERSAL_AUTH)); err != nil {
t.Fatal(err)
}
t.Setenv(util.INFISICAL_UNIVERSAL_AUTH_CLIENT_ID_NAME, "test-client-id")
t.Setenv(util.INFISICAL_UNIVERSAL_AUTH_CLIENT_SECRET_NAME, "test-client-secret")
},
want: 0,
},
{
name: "normal saved-session command",
command: newCommand("run"),
setup: func(*testing.T, *cobra.Command) {},
want: 1,
},
}

for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
lookupCalls = 0
config.INFISICAL_URL = "https://app.infisical.com/api"
tc.setup(t, tc.command)

RootCmd.PersistentPreRun(tc.command, nil)

if lookupCalls != tc.want {
t.Fatalf("saved-session lookup called %d times, want %d", lookupCalls, tc.want)
}
})
}
}