Skip to content

Commit 6757fab

Browse files
svnltoaaearon
andauthored
fix: AWS elevation shows Azure CLI message (#37)
* fix: use CSP to determine post-elevation message instead of AccessCredentials AWS elevations with nil AccessCredentials incorrectly showed the Azure CLI message. Switch on target CSP so AWS always gets AWS-appropriate guidance and Azure gets the az CLI message. Also fix time-sensitive test in session_tracker_test.go that used a hardcoded past date, causing failures after 24h elapsed. * fix: use explicit CSPAzure case instead of default in post-elevation switch The default case would incorrectly show the Azure CLI message for any future CSP. Using an explicit case means unknown CSPs simply show no post-elevation guidance, which is the safer behavior. --------- Co-authored-by: Tim Schindler <tim@iosharp.com>
1 parent 2853db5 commit 6757fab

3 files changed

Lines changed: 89 additions & 15 deletions

File tree

cmd/root.go

Lines changed: 14 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -811,16 +811,21 @@ func runElevateWithDeps(
811811
fmt.Fprintf(cmd.OutOrStdout(), " Session ID: %s\n", res.result.SessionID)
812812

813813
// CSP-aware post-elevation guidance
814-
if res.result.AccessCredentials != nil {
815-
awsCreds, err := models.ParseAWSCredentials(*res.result.AccessCredentials)
816-
if err != nil {
817-
return fmt.Errorf("failed to parse access credentials: %w", err)
814+
switch res.target.CSP {
815+
case models.CSPAWS:
816+
if res.result.AccessCredentials != nil {
817+
awsCreds, err := models.ParseAWSCredentials(*res.result.AccessCredentials)
818+
if err != nil {
819+
return fmt.Errorf("failed to parse access credentials: %w", err)
820+
}
821+
fmt.Fprintf(cmd.OutOrStdout(), "\n export AWS_ACCESS_KEY_ID='%s'\n", awsCreds.AccessKeyID)
822+
fmt.Fprintf(cmd.OutOrStdout(), " export AWS_SECRET_ACCESS_KEY='%s'\n", awsCreds.SecretAccessKey)
823+
fmt.Fprintf(cmd.OutOrStdout(), " export AWS_SESSION_TOKEN='%s'\n", awsCreds.SessionToken)
824+
fmt.Fprintf(cmd.OutOrStdout(), "\n Or run: eval $(grant env --provider aws)\n")
825+
} else {
826+
fmt.Fprintf(cmd.OutOrStdout(), "\n Run: eval $(grant env --provider aws) to get credentials.\n")
818827
}
819-
fmt.Fprintf(cmd.OutOrStdout(), "\n export AWS_ACCESS_KEY_ID='%s'\n", awsCreds.AccessKeyID)
820-
fmt.Fprintf(cmd.OutOrStdout(), " export AWS_SECRET_ACCESS_KEY='%s'\n", awsCreds.SecretAccessKey)
821-
fmt.Fprintf(cmd.OutOrStdout(), " export AWS_SESSION_TOKEN='%s'\n", awsCreds.SessionToken)
822-
fmt.Fprintf(cmd.OutOrStdout(), "\n Or run: eval $(grant env --provider aws)\n")
823-
} else {
828+
case models.CSPAzure:
824829
fmt.Fprintf(cmd.OutOrStdout(), "\n Your az CLI session now has the elevated permissions.\n")
825830
}
826831

cmd/root_elevate_test.go

Lines changed: 74 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -17,11 +17,12 @@ import (
1717

1818
func TestRootElevate_InteractiveMode(t *testing.T) {
1919
tests := []struct {
20-
name string
21-
setupMocks func() (*mockAuthLoader, *mockEligibilityLister, *mockElevateService, *mockUnifiedSelector, *config.Config)
22-
args []string
23-
wantContain []string
24-
wantErr bool
20+
name string
21+
setupMocks func() (*mockAuthLoader, *mockEligibilityLister, *mockElevateService, *mockUnifiedSelector, *config.Config)
22+
args []string
23+
wantContain []string
24+
wantNotContain []string
25+
wantErr bool
2526
}{
2627
{
2728
name: "interactive mode success",
@@ -314,6 +315,69 @@ func TestRootElevate_InteractiveMode(t *testing.T) {
314315
},
315316
wantErr: false,
316317
},
318+
{
319+
name: "AWS elevation without credentials should not show Azure message",
320+
setupMocks: func() (*mockAuthLoader, *mockEligibilityLister, *mockElevateService, *mockUnifiedSelector, *config.Config) {
321+
authLoader := &mockAuthLoader{
322+
token: &authmodels.IdsecToken{
323+
Token: "test-jwt",
324+
Username: "test@example.com",
325+
ExpiresIn: commonmodels.IdsecRFC3339Time(time.Now().Add(1 * time.Hour)),
326+
},
327+
}
328+
329+
awsTarget := models.EligibleTarget{
330+
CSP: models.CSPAWS,
331+
OrganizationID: "o-abc123",
332+
WorkspaceID: "123456789012",
333+
WorkspaceName: "test-aws-account-7x9k",
334+
WorkspaceType: models.WorkspaceTypeAccount,
335+
RoleInfo: models.RoleInfo{
336+
ID: "arn:aws:iam::123456789012:role/Edit",
337+
Name: "Edit",
338+
},
339+
}
340+
341+
eligibilityLister := &mockEligibilityLister{
342+
listFunc: func(ctx context.Context, csp models.CSP) (*models.EligibilityResponse, error) {
343+
if csp == models.CSPAWS {
344+
return &models.EligibilityResponse{Response: []models.EligibleTarget{awsTarget}, Total: 1}, nil
345+
}
346+
return &models.EligibilityResponse{}, nil
347+
},
348+
}
349+
350+
// AWS elevation succeeds but API returns nil AccessCredentials
351+
elevateService := &mockElevateService{
352+
response: &models.ElevateResponse{
353+
Response: models.ElevateAccessResult{
354+
CSP: models.CSPAWS,
355+
OrganizationID: "o-abc123",
356+
Results: []models.ElevateTargetResult{
357+
{
358+
WorkspaceID: "123456789012",
359+
RoleID: "Edit",
360+
SessionID: "session-aws-nocreds",
361+
AccessCredentials: nil,
362+
},
363+
},
364+
},
365+
},
366+
}
367+
368+
selector := &mockUnifiedSelector{
369+
item: &selectionItem{kind: selectionCloud, cloud: &awsTarget},
370+
}
371+
372+
cfg := config.DefaultConfig()
373+
374+
return authLoader, eligibilityLister, elevateService, selector, cfg
375+
},
376+
args: []string{},
377+
wantContain: []string{"Elevated to Edit on test-aws-account-7x9k"},
378+
wantNotContain: []string{"az CLI session"},
379+
wantErr: false,
380+
},
317381
{
318382
name: "no eligible targets found across all providers",
319383
setupMocks: func() (*mockAuthLoader, *mockEligibilityLister, *mockElevateService, *mockUnifiedSelector, *config.Config) {
@@ -360,6 +424,11 @@ func TestRootElevate_InteractiveMode(t *testing.T) {
360424
t.Errorf("output missing %q\ngot:\n%s", want, output)
361425
}
362426
}
427+
for _, notWant := range tt.wantNotContain {
428+
if strings.Contains(output, notWant) {
429+
t.Errorf("output should not contain %q\ngot:\n%s", notWant, output)
430+
}
431+
}
363432
})
364433
}
365434
}

internal/cache/session_tracker_test.go

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@ import (
1010
func TestRecordSession_AndLookup(t *testing.T) {
1111
t.Parallel()
1212
s := NewStore(t.TempDir(), 25*time.Hour)
13-
now := time.Date(2026, 2, 21, 12, 0, 0, 0, time.UTC)
13+
now := time.Now().UTC().Truncate(time.Second)
1414

1515
if err := RecordSession(s, "sess-1", now); err != nil {
1616
t.Fatalf("RecordSession() error = %v", err)

0 commit comments

Comments
 (0)