Skip to content

Commit 0451cb5

Browse files
committed
fix(cli): honour ax ssh --help and stop global flag parsing at --
ax ssh --help was taken as a task name and ended in a NotFound error after resolving the server. The ssh arguments are now parsed up front, so -h/--help prints ssh usage and a missing task name fails fast, before any server lookup or port-forward. The global flag parser also kept consuming -a, -n, --server and --context after "--", so ax ssh t -- grep -n foo file silently set the namespace to foo. Parsing now stops at "--" and hands the rest to the command unchanged. Fixes #373 Fixes #410
1 parent ac23328 commit 0451cb5

2 files changed

Lines changed: 202 additions & 54 deletions

File tree

‎cmd/ax/main.go‎

Lines changed: 112 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -45,53 +45,16 @@ func main() {
4545
os.Exit(1)
4646
}
4747

48+
opts := parseGlobalArgs(os.Args[1:])
4849
var (
49-
cmd string
50-
cleanArgs []string
51-
atespace = "default"
52-
explicitServer = ""
53-
kubeContext = ""
54-
axNamespace = "ax-system"
50+
cmd = opts.cmd
51+
cleanArgs = opts.args
52+
atespace = opts.atespace
53+
explicitServer = opts.server
54+
kubeContext = opts.kubeContext
55+
axNamespace = opts.namespace
5556
)
5657

57-
args := os.Args[1:]
58-
for i := 0; i < len(args); i++ {
59-
arg := args[i]
60-
if arg == "-a" || arg == "--atespace" {
61-
if i+1 < len(args) {
62-
atespace = args[i+1]
63-
i++
64-
}
65-
} else if strings.HasPrefix(arg, "--atespace=") {
66-
atespace = strings.TrimPrefix(arg, "--atespace=")
67-
} else if arg == "--server" {
68-
if i+1 < len(args) {
69-
explicitServer = args[i+1]
70-
i++
71-
}
72-
} else if strings.HasPrefix(arg, "--server=") {
73-
explicitServer = strings.TrimPrefix(arg, "--server=")
74-
} else if arg == "--context" {
75-
if i+1 < len(args) {
76-
kubeContext = args[i+1]
77-
i++
78-
}
79-
} else if strings.HasPrefix(arg, "--context=") {
80-
kubeContext = strings.TrimPrefix(arg, "--context=")
81-
} else if arg == "-n" || arg == "--namespace" {
82-
if i+1 < len(args) {
83-
axNamespace = args[i+1]
84-
i++
85-
}
86-
} else if strings.HasPrefix(arg, "--namespace=") {
87-
axNamespace = strings.TrimPrefix(arg, "--namespace=")
88-
} else if cmd == "" && !strings.HasPrefix(arg, "-") {
89-
cmd = arg
90-
} else {
91-
cleanArgs = append(cleanArgs, arg)
92-
}
93-
}
94-
9558
if cmd == "" {
9659
printUsage()
9760
os.Exit(1)
@@ -117,6 +80,17 @@ func main() {
11780
os.Exit(1)
11881
}
11982
return
83+
case "ssh":
84+
// Answer help and usage errors before resolving the server, which may
85+
// start a port-forward to the cluster.
86+
if _, _, help, err := parseSSHArgs(cleanArgs); help || err != nil {
87+
fmt.Println(sshUsage)
88+
if err != nil {
89+
fmt.Fprintf(os.Stderr, "Error: %v\n", err)
90+
os.Exit(1)
91+
}
92+
return
93+
}
12094
}
12195

12296
// Resolve the AX server URL (auto-tunneling to active kube context if not explicitly set)
@@ -159,6 +133,64 @@ func main() {
159133
}
160134
}
161135

136+
// globalArgs holds the command name, the flags shared by every command, and the
137+
// remaining arguments for the command itself.
138+
type globalArgs struct {
139+
cmd string
140+
args []string
141+
atespace string
142+
server string
143+
kubeContext string
144+
namespace string
145+
}
146+
147+
// parseGlobalArgs extracts the command name and the global flags from args. It
148+
// stops interpreting flags at "--", so everything after it (for example the
149+
// command run by `ax ssh <task> -- ...`) is passed through to the command intact.
150+
func parseGlobalArgs(args []string) globalArgs {
151+
g := globalArgs{atespace: "default", namespace: "ax-system"}
152+
for i := 0; i < len(args); i++ {
153+
arg := args[i]
154+
if arg == "--" {
155+
g.args = append(g.args, args[i:]...)
156+
break
157+
} else if arg == "-a" || arg == "--atespace" {
158+
if i+1 < len(args) {
159+
g.atespace = args[i+1]
160+
i++
161+
}
162+
} else if strings.HasPrefix(arg, "--atespace=") {
163+
g.atespace = strings.TrimPrefix(arg, "--atespace=")
164+
} else if arg == "--server" {
165+
if i+1 < len(args) {
166+
g.server = args[i+1]
167+
i++
168+
}
169+
} else if strings.HasPrefix(arg, "--server=") {
170+
g.server = strings.TrimPrefix(arg, "--server=")
171+
} else if arg == "--context" {
172+
if i+1 < len(args) {
173+
g.kubeContext = args[i+1]
174+
i++
175+
}
176+
} else if strings.HasPrefix(arg, "--context=") {
177+
g.kubeContext = strings.TrimPrefix(arg, "--context=")
178+
} else if arg == "-n" || arg == "--namespace" {
179+
if i+1 < len(args) {
180+
g.namespace = args[i+1]
181+
i++
182+
}
183+
} else if strings.HasPrefix(arg, "--namespace=") {
184+
g.namespace = strings.TrimPrefix(arg, "--namespace=")
185+
} else if g.cmd == "" && !strings.HasPrefix(arg, "-") {
186+
g.cmd = arg
187+
} else {
188+
g.args = append(g.args, arg)
189+
}
190+
}
191+
return g
192+
}
193+
162194
func printUsage() {
163195
fmt.Println(`AX CLI - Autonomous agent execution control
164196
@@ -959,24 +991,50 @@ func runTunnel(args []string) error {
959991
}
960992
}
961993

962-
func runSSH(serverURL, atespace, kubeContext string, args []string) error {
963-
if len(args) == 0 {
964-
return fmt.Errorf("usage: ax ssh <task-name> [-- command...]")
994+
const sshUsage = `Usage:
995+
ax ssh <task-name> [-- command...]
996+
997+
Run a command inside a running task's container, or /bin/sh when no command is
998+
given. The task must be Running and have spec.debug: true.
999+
1000+
Examples:
1001+
ax ssh task123
1002+
ax ssh task123 -- ls -la /workspace
1003+
ax ssh -a my-atespace task123 -- python3 main.py`
1004+
1005+
// parseSSHArgs splits the arguments of `ax ssh` into the task name and the command
1006+
// to run, defaulting the command to /bin/sh. help reports a -h or --help given in
1007+
// place of the task name. Arguments after the task name form the command, with or
1008+
// without a separating "--".
1009+
func parseSSHArgs(args []string) (taskName string, command []string, help bool, err error) {
1010+
if len(args) == 0 || args[0] == "--" {
1011+
return "", nil, false, errors.New("missing task name")
1012+
}
1013+
if args[0] == "-h" || args[0] == "--help" {
1014+
return "", nil, true, nil
1015+
}
1016+
if strings.HasPrefix(args[0], "-") {
1017+
return "", nil, false, fmt.Errorf("unknown flag %q", args[0])
9651018
}
9661019

967-
taskName := args[0]
968-
var cmdToRun []string
1020+
taskName = args[0]
9691021
for i := 1; i < len(args); i++ {
9701022
if args[i] == "--" {
971-
cmdToRun = args[i+1:]
1023+
command = append(command, args[i+1:]...)
9721024
break
973-
} else {
974-
cmdToRun = append(cmdToRun, args[i])
9751025
}
1026+
command = append(command, args[i])
1027+
}
1028+
if len(command) == 0 {
1029+
command = []string{"/bin/sh"}
9761030
}
1031+
return taskName, command, false, nil
1032+
}
9771033

978-
if len(cmdToRun) == 0 {
979-
cmdToRun = []string{"/bin/sh"}
1034+
func runSSH(serverURL, atespace, kubeContext string, args []string) error {
1035+
taskName, cmdToRun, _, err := parseSSHArgs(args)
1036+
if err != nil {
1037+
return err
9801038
}
9811039

9821040
client, conn, err := getAXClient(serverURL)

‎cmd/ax/main_test.go‎

Lines changed: 90 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -20,6 +20,7 @@ import (
2020
"net"
2121
"os"
2222
"path/filepath"
23+
"reflect"
2324
"strings"
2425
"testing"
2526
"time"
@@ -278,3 +279,92 @@ func TestRunDeleteTask(t *testing.T) {
278279
t.Fatalf("expected task to be NotFound after deletion, got %v", err)
279280
}
280281
}
282+
283+
func TestParseGlobalArgs(t *testing.T) {
284+
tests := []struct {
285+
name string
286+
args []string
287+
want globalArgs
288+
}{
289+
{
290+
name: "flags before and after the command",
291+
args: []string{"-a", "team", "get", "tasks", "--namespace=ax"},
292+
want: globalArgs{cmd: "get", args: []string{"tasks"}, atespace: "team", namespace: "ax"},
293+
},
294+
{
295+
name: "flags after -- belong to the remote command",
296+
args: []string{"ssh", "task123", "--", "grep", "-n", "foo", "-a", "file"},
297+
want: globalArgs{
298+
cmd: "ssh",
299+
args: []string{"task123", "--", "grep", "-n", "foo", "-a", "file"},
300+
atespace: "default",
301+
namespace: "ax-system",
302+
},
303+
},
304+
{
305+
name: "global flags before -- still apply",
306+
args: []string{"ssh", "-a", "team", "task123", "--", "ls", "--context=x"},
307+
want: globalArgs{
308+
cmd: "ssh",
309+
args: []string{"task123", "--", "ls", "--context=x"},
310+
atespace: "team",
311+
namespace: "ax-system",
312+
},
313+
},
314+
}
315+
for _, tt := range tests {
316+
t.Run(tt.name, func(t *testing.T) {
317+
if got := parseGlobalArgs(tt.args); !reflect.DeepEqual(got, tt.want) {
318+
t.Errorf("parseGlobalArgs(%q) = %+v, want %+v", tt.args, got, tt.want)
319+
}
320+
})
321+
}
322+
}
323+
324+
func TestParseSSHArgs(t *testing.T) {
325+
tests := []struct {
326+
name string
327+
args []string
328+
wantTask string
329+
wantCommand []string
330+
wantHelp bool
331+
wantErr bool
332+
}{
333+
{name: "long help", args: []string{"--help"}, wantHelp: true},
334+
{name: "short help", args: []string{"-h"}, wantHelp: true},
335+
{name: "no args", args: nil, wantErr: true},
336+
{name: "no task before --", args: []string{"--", "ls"}, wantErr: true},
337+
{name: "unknown flag", args: []string{"--verbose", "task123"}, wantErr: true},
338+
{name: "default shell", args: []string{"task123"}, wantTask: "task123", wantCommand: []string{"/bin/sh"}},
339+
{
340+
name: "command after --",
341+
args: []string{"task123", "--", "ls", "-la", "/workspace"},
342+
wantTask: "task123",
343+
wantCommand: []string{"ls", "-la", "/workspace"},
344+
},
345+
{
346+
name: "command without --",
347+
args: []string{"task123", "python3", "main.py"},
348+
wantTask: "task123",
349+
wantCommand: []string{"python3", "main.py"},
350+
},
351+
{
352+
name: "help after the task name goes to the remote command",
353+
args: []string{"task123", "--", "git", "--help"},
354+
wantTask: "task123",
355+
wantCommand: []string{"git", "--help"},
356+
},
357+
}
358+
for _, tt := range tests {
359+
t.Run(tt.name, func(t *testing.T) {
360+
task, command, help, err := parseSSHArgs(tt.args)
361+
if (err != nil) != tt.wantErr {
362+
t.Fatalf("parseSSHArgs(%q) error = %v, wantErr %v", tt.args, err, tt.wantErr)
363+
}
364+
if task != tt.wantTask || help != tt.wantHelp || !reflect.DeepEqual(command, tt.wantCommand) {
365+
t.Errorf("parseSSHArgs(%q) = (%q, %q, %v), want (%q, %q, %v)",
366+
tt.args, task, command, help, tt.wantTask, tt.wantCommand, tt.wantHelp)
367+
}
368+
})
369+
}
370+
}

0 commit comments

Comments
 (0)