Skip to content

Commit 30d90fb

Browse files
committed
feat(nnctl): use cobra for CLI help
Signed-off-by: Ville Vesilehto <ville@vesilehto.fi>
1 parent 7baf167 commit 30d90fb

12 files changed

Lines changed: 759 additions & 708 deletions

File tree

nnctl/go.mod

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,10 @@
11
module nnctl
22

33
go 1.26
4+
5+
require github.com/spf13/cobra v1.10.2
6+
7+
require (
8+
github.com/inconshreveable/mousetrap v1.1.0 // indirect
9+
github.com/spf13/pflag v1.0.9 // indirect
10+
)

nnctl/go.sum

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,10 @@
1+
github.com/cpuguy83/go-md2man/v2 v2.0.6/go.mod h1:oOW0eioCTA6cOiMLiUPZOpcVxMig6NIQQ7OS05n1F4g=
2+
github.com/inconshreveable/mousetrap v1.1.0 h1:wN+x4NVGpMsO7ErUn/mUI3vEoE6Jt13X2s0bqwp9tc8=
3+
github.com/inconshreveable/mousetrap v1.1.0/go.mod h1:vpF70FUmC8bwa3OWnCshd2FqLfsEA9PFc4w1p2J65bw=
4+
github.com/russross/blackfriday/v2 v2.1.0/go.mod h1:+Rmxgy9KzJVeS9/2gXHxylqXiyQDYRxCVz55jmeOWTM=
5+
github.com/spf13/cobra v1.10.2 h1:DMTTonx5m65Ic0GOoRY2c16WCbHxOOw6xxezuLaBpcU=
6+
github.com/spf13/cobra v1.10.2/go.mod h1:7C1pvHqHw5A4vrJfjNwvOdzYu0Gml16OCs2GRiTUUS4=
7+
github.com/spf13/pflag v1.0.9 h1:9exaQaMOCwffKiiiYk6/BndUBv+iRViNW+4lEMi0PvY=
8+
github.com/spf13/pflag v1.0.9/go.mod h1:McXfInJRrz4CZXVZOBLb0bTZqETkiAhM9Iw0y3An2Bg=
9+
go.yaml.in/yaml/v3 v3.0.4/go.mod h1:DhzuOOF2ATzADvBadXxruRBLzYTpT36CKvDb3+aBEFg=
10+
gopkg.in/check.v1 v0.0.0-20161208181325-20d25e280405/go.mod h1:Co6ibVJAznAaIkqp8huTwlJQCZ016jof/cbN4VW5Yz0=

nnctl/internal/cli/app.go

Lines changed: 3 additions & 82 deletions
Original file line numberDiff line numberDiff line change
@@ -2,15 +2,10 @@ package cli
22

33
import (
44
"context"
5-
"errors"
6-
"flag"
75
"fmt"
86
"io"
97
"os"
10-
"runtime"
118
"strings"
12-
13-
"nnctl/internal/repo"
149
)
1510

1611
const (
@@ -41,77 +36,9 @@ func Main(args []string) int {
4136
}
4237

4338
func (a *App) Run(ctx context.Context, args []string) error {
44-
global := flag.NewFlagSet("nnctl", flag.ContinueOnError)
45-
global.SetOutput(a.stderr())
46-
47-
repoRoot := ""
48-
zig := getenvDefault("ZIG", "zig")
49-
global.StringVar(&repoRoot, "repo", "", "repository root to operate on")
50-
global.StringVar(&repoRoot, "C", "", "repository root to operate on")
51-
global.StringVar(&zig, "zig", zig, "Zig executable to use")
52-
53-
if err := global.Parse(args); err != nil {
54-
if errors.Is(err, flag.ErrHelp) {
55-
printUsage(a.stdout())
56-
return nil
57-
}
58-
return err
59-
}
60-
61-
rest := global.Args()
62-
if len(rest) == 0 {
63-
printUsage(a.stdout())
64-
return nil
65-
}
66-
67-
command := rest[0]
68-
commandArgs := rest[1:]
69-
if command == "help" || command == "-h" || command == "--help" {
70-
printHelp(a.stdout(), commandArgs)
71-
return nil
72-
}
73-
if command == "version" {
74-
fmt.Fprintf(a.stdout(), "nnctl dev (%s)\n", runtime.Version())
75-
return nil
76-
}
77-
78-
root, err := repo.ResolveRoot(repoRoot)
79-
if err != nil {
80-
return err
81-
}
82-
a.repoRoot = root
83-
a.zig = zig
84-
85-
switch command {
86-
case "all":
87-
return a.cmdAll(ctx, commandArgs)
88-
case "build":
89-
return a.cmdBuild(ctx, commandArgs, false)
90-
case "release":
91-
return a.cmdBuild(ctx, commandArgs, true)
92-
case "test":
93-
return a.cmdTest(ctx, commandArgs)
94-
case "examples":
95-
return a.cmdExamples(ctx, commandArgs)
96-
case "run":
97-
return a.cmdRun(ctx, commandArgs)
98-
case "train":
99-
return a.cmdTrain(ctx, commandArgs)
100-
case "chat":
101-
return a.cmdChat(ctx, commandArgs)
102-
case "list":
103-
return a.cmdList(commandArgs)
104-
case "fmt", "format":
105-
return a.cmdFormat(ctx, commandArgs)
106-
case "clean":
107-
return a.cmdClean(commandArgs)
108-
case "data":
109-
return a.cmdData(ctx, commandArgs)
110-
case "doctor", "env":
111-
return a.cmdDoctor(ctx, commandArgs)
112-
default:
113-
return fmt.Errorf("unknown command %q; run nnctl help", command)
114-
}
39+
root := a.newRootCommand()
40+
root.SetArgs(args)
41+
return root.ExecuteContext(ctx)
11542
}
11643

11744
func (a *App) stdout() io.Writer {
@@ -135,12 +62,6 @@ func (a *App) stdin() io.Reader {
13562
return a.Stdin
13663
}
13764

138-
func newFlagSet(out io.Writer, name string) *flag.FlagSet {
139-
fs := flag.NewFlagSet(name, flag.ContinueOnError)
140-
fs.SetOutput(out)
141-
return fs
142-
}
143-
14465
func getenvDefault(key, fallback string) string {
14566
if value := os.Getenv(key); value != "" {
14667
return value

nnctl/internal/cli/chat.go

Lines changed: 49 additions & 68 deletions
Original file line numberDiff line numberDiff line change
@@ -13,96 +13,77 @@ import (
1313
"strconv"
1414
"time"
1515

16-
"nnctl/internal/flagx"
1716
"nnctl/internal/zig"
1817
)
1918

20-
func (a *App) cmdChat(ctx context.Context, args []string) error {
21-
fs := newFlagSet(a.stderr(), "chat")
22-
mode := getenvDefault("NNCTL_OPTIMIZE", getenvDefault("BUILD_MODE", defaultBuildMode))
23-
host := "127.0.0.1"
24-
port := 8090
25-
apiHost := "127.0.0.1"
26-
apiPort := 8080
27-
modelPath := ""
28-
modelName := "tiny-gpt-zig"
29-
maxTokens := 120
30-
temperature := 1.0
31-
topK := 8
32-
allowUntrained := false
33-
readyTimeoutText := "2m"
34-
35-
fs.StringVar(&mode, "mode", mode, "Zig optimization mode")
36-
fs.StringVar(&mode, "optimize", mode, "Zig optimization mode")
37-
fs.StringVar(&host, "host", host, "chat app bind host")
38-
fs.IntVar(&port, "port", port, "chat app bind port")
39-
fs.StringVar(&apiHost, "api-host", apiHost, "TinyGPT inference bind host")
40-
fs.IntVar(&apiPort, "api-port", apiPort, "TinyGPT inference bind port")
41-
fs.StringVar(&modelPath, "model", modelPath, "TinyGPT checkpoint to serve")
42-
fs.StringVar(&modelName, "model-name", modelName, "OpenAI-compatible model id")
43-
fs.IntVar(&maxTokens, "max-tokens", maxTokens, "default generated tokens per reply")
44-
fs.Float64Var(&temperature, "temperature", temperature, "default sampling temperature")
45-
fs.IntVar(&topK, "top-k", topK, "default TinyGPT top-k sampler cutoff")
46-
fs.BoolVar(&allowUntrained, "allow-untrained", allowUntrained, "serve a seeded untrained model when --model is omitted")
47-
fs.StringVar(&readyTimeoutText, "ready-timeout", readyTimeoutText, "time to wait for inference server readiness")
19+
type chatOptions struct {
20+
mode string
21+
host string
22+
port int
23+
apiHost string
24+
apiPort int
25+
modelPath string
26+
modelName string
27+
maxTokens int
28+
temperature float64
29+
topK int
30+
allowUntrained bool
31+
readyTimeoutText string
32+
}
4833

49-
if err := flagx.ParseInterspersed(fs, args, map[string]bool{
50-
"mode": true,
51-
"optimize": true,
52-
"host": true,
53-
"port": true,
54-
"api-host": true,
55-
"api-port": true,
56-
"model": true,
57-
"model-name": true,
58-
"max-tokens": true,
59-
"temperature": true,
60-
"top-k": true,
61-
"allow-untrained": false,
62-
"ready-timeout": true,
63-
}); err != nil {
64-
return err
34+
func defaultChatOptions() chatOptions {
35+
return chatOptions{
36+
mode: getenvDefault("NNCTL_OPTIMIZE", getenvDefault("BUILD_MODE", defaultBuildMode)),
37+
host: "127.0.0.1",
38+
port: 8090,
39+
apiHost: "127.0.0.1",
40+
apiPort: 8080,
41+
modelName: "tiny-gpt-zig",
42+
maxTokens: 120,
43+
temperature: 1.0,
44+
topK: 8,
45+
readyTimeoutText: "2m",
6546
}
66-
if fs.NArg() != 0 {
67-
return fmt.Errorf("chat does not accept positional arguments")
68-
}
69-
if modelPath == "" && !allowUntrained {
47+
}
48+
49+
func (a *App) runChat(ctx context.Context, opts chatOptions) error {
50+
if opts.modelPath == "" && !opts.allowUntrained {
7051
return fmt.Errorf("chat requires --model <checkpoint> or --allow-untrained")
7152
}
72-
if port <= 0 || port > 65535 {
53+
if opts.port <= 0 || opts.port > 65535 {
7354
return fmt.Errorf("--port must be between 1 and 65535")
7455
}
75-
if apiPort <= 0 || apiPort > 65535 {
56+
if opts.apiPort <= 0 || opts.apiPort > 65535 {
7657
return fmt.Errorf("--api-port must be between 1 and 65535")
7758
}
78-
if maxTokens < 0 {
59+
if opts.maxTokens < 0 {
7960
return fmt.Errorf("--max-tokens must be non-negative")
8061
}
81-
if topK < 0 {
62+
if opts.topK < 0 {
8263
return fmt.Errorf("--top-k must be non-negative")
8364
}
8465

85-
readyTimeout, err := time.ParseDuration(readyTimeoutText)
66+
readyTimeout, err := time.ParseDuration(opts.readyTimeoutText)
8667
if err != nil {
8768
return fmt.Errorf("parse --ready-timeout: %w", err)
8869
}
8970

90-
apiBaseURL := "http://" + net.JoinHostPort(apiHost, strconv.Itoa(apiPort))
71+
apiBaseURL := "http://" + net.JoinHostPort(opts.apiHost, strconv.Itoa(opts.apiPort))
9172
serverArgs := []string{
92-
"--host", apiHost,
93-
"--port", strconv.Itoa(apiPort),
94-
"--model-name", modelName,
95-
"--max-tokens", strconv.Itoa(maxTokens),
96-
"--temperature", strconv.FormatFloat(temperature, 'f', -1, 64),
97-
"--top-k", strconv.Itoa(topK),
73+
"--host", opts.apiHost,
74+
"--port", strconv.Itoa(opts.apiPort),
75+
"--model-name", opts.modelName,
76+
"--max-tokens", strconv.Itoa(opts.maxTokens),
77+
"--temperature", strconv.FormatFloat(opts.temperature, 'f', -1, 64),
78+
"--top-k", strconv.Itoa(opts.topK),
9879
}
99-
if modelPath != "" {
100-
serverArgs = append(serverArgs, "--model", modelPath)
80+
if opts.modelPath != "" {
81+
serverArgs = append(serverArgs, "--model", opts.modelPath)
10182
} else {
10283
serverArgs = append(serverArgs, "--allow-untrained")
10384
}
10485

105-
runArgs := zig.RunArgs("run_tiny_gpt_openai", zig.Options{Optimize: mode}, serverArgs)
86+
runArgs := zig.RunArgs("run_tiny_gpt_openai", zig.Options{Optimize: opts.mode}, serverArgs)
10687
fmt.Fprintf(a.stderr(), "==> %s\n", zig.CommandString(a.zig, runArgs))
10788
serverCmd := exec.CommandContext(ctx, a.zig, runArgs...)
10889
serverCmd.Dir = a.repoRoot
@@ -131,13 +112,13 @@ func (a *App) cmdChat(ctx context.Context, args []string) error {
131112
mux.HandleFunc("/", serveChatApp)
132113
mux.Handle("/api/chat", chatProxy{
133114
apiBaseURL: apiBaseURL,
134-
modelName: modelName,
135-
maxTokens: maxTokens,
136-
temperature: temperature,
115+
modelName: opts.modelName,
116+
maxTokens: opts.maxTokens,
117+
temperature: opts.temperature,
137118
client: &http.Client{Timeout: 2 * time.Minute},
138119
})
139120

140-
chatAddr := net.JoinHostPort(host, strconv.Itoa(port))
121+
chatAddr := net.JoinHostPort(opts.host, strconv.Itoa(opts.port))
141122
chatServer := &http.Server{
142123
Addr: chatAddr,
143124
Handler: mux,

0 commit comments

Comments
 (0)