← back to Cli Printing Press
feat(generator): OAuth2 auth flow for generated CLIs
dacda31f8213dcd73f748805f325d67b4259d2a6 · 2026-03-24 00:04:48 -0700 · Matt Van Horn
Extract OAuth2 authorization URL, token URL, and scopes from OpenAPI specs.
Generate auth login/status/logout commands with browser-based OAuth2 flow,
local callback server, token storage in config, and auto-refresh in client.
Only rendered when spec has OAuth2 security scheme.
Files touched
M internal/generator/generator.goM internal/generator/generator_test.goA internal/generator/templates/auth.go.tmplM internal/generator/templates/client.go.tmplM internal/generator/templates/config.go.tmplM internal/generator/templates/root.go.tmplM internal/openapi/parser.goM internal/openapi/parser_test.goM internal/spec/spec.go
Diff
commit dacda31f8213dcd73f748805f325d67b4259d2a6
Author: Matt Van Horn <455140+mvanhorn@users.noreply.github.com>
Date: Tue Mar 24 00:04:48 2026 -0700
feat(generator): OAuth2 auth flow for generated CLIs
Extract OAuth2 authorization URL, token URL, and scopes from OpenAPI specs.
Generate auth login/status/logout commands with browser-based OAuth2 flow,
local callback server, token storage in config, and auto-refresh in client.
Only rendered when spec has OAuth2 security scheme.
---
internal/generator/generator.go | 64 +++----
internal/generator/generator_test.go | 32 ++++
internal/generator/templates/auth.go.tmpl | 248 ++++++++++++++++++++++++++++
internal/generator/templates/client.go.tmpl | 101 ++++++++++-
internal/generator/templates/config.go.tmpl | 42 +++++
internal/generator/templates/root.go.tmpl | 5 +-
internal/openapi/parser.go | 26 +++
internal/openapi/parser_test.go | 16 ++
internal/spec/spec.go | 15 +-
9 files changed, 506 insertions(+), 43 deletions(-)
diff --git a/internal/generator/generator.go b/internal/generator/generator.go
index c16a784f..23ec9115 100644
--- a/internal/generator/generator.go
+++ b/internal/generator/generator.go
@@ -24,23 +24,23 @@ type Generator struct {
func New(s *spec.APISpec, outputDir string) *Generator {
g := &Generator{Spec: s, OutputDir: outputDir}
g.funcs = template.FuncMap{
- "title": strings.Title,
- "lower": strings.ToLower,
- "upper": strings.ToUpper,
- "camel": toCamel,
- "snake": toSnake,
- "goType": goType,
- "cobraFlagFunc": cobraFlagFunc,
- "defaultVal": defaultVal,
- "zeroVal": zeroVal,
- "positionalArgs": positionalArgs,
- "configTag": configTag,
- "envVarField": envVarField,
+ "title": strings.Title,
+ "lower": strings.ToLower,
+ "upper": strings.ToUpper,
+ "camel": toCamel,
+ "snake": toSnake,
+ "goType": goType,
+ "cobraFlagFunc": cobraFlagFunc,
+ "defaultVal": defaultVal,
+ "zeroVal": zeroVal,
+ "positionalArgs": positionalArgs,
+ "configTag": configTag,
+ "envVarField": envVarField,
"envVarPlaceholder": envVarPlaceholder,
- "add": func(a, b int) int { return a + b },
- "oneline": oneline,
- "flagName": flagName,
- "exampleLine": g.exampleLine,
+ "add": func(a, b int) int { return a + b },
+ "oneline": oneline,
+ "flagName": flagName,
+ "exampleLine": g.exampleLine,
}
return g
}
@@ -62,18 +62,18 @@ func (g *Generator) Generate() error {
// Generate single files
singleFiles := map[string]string{
- "main.go.tmpl": filepath.Join("cmd", g.Spec.Name+"-cli", "main.go"),
- "root.go.tmpl": filepath.Join("internal", "cli", "root.go"),
- "helpers.go.tmpl": filepath.Join("internal", "cli", "helpers.go"),
- "doctor.go.tmpl": filepath.Join("internal", "cli", "doctor.go"),
- "config.go.tmpl": filepath.Join("internal", "config", "config.go"),
- "client.go.tmpl": filepath.Join("internal", "client", "client.go"),
- "types.go.tmpl": filepath.Join("internal", "types", "types.go"),
- "go.mod.tmpl": "go.mod",
- "goreleaser.yaml.tmpl": ".goreleaser.yaml",
- "golangci.yml.tmpl": ".golangci.yml",
- "makefile.tmpl": "Makefile",
- "readme.md.tmpl": "README.md",
+ "main.go.tmpl": filepath.Join("cmd", g.Spec.Name+"-cli", "main.go"),
+ "root.go.tmpl": filepath.Join("internal", "cli", "root.go"),
+ "helpers.go.tmpl": filepath.Join("internal", "cli", "helpers.go"),
+ "doctor.go.tmpl": filepath.Join("internal", "cli", "doctor.go"),
+ "config.go.tmpl": filepath.Join("internal", "config", "config.go"),
+ "client.go.tmpl": filepath.Join("internal", "client", "client.go"),
+ "types.go.tmpl": filepath.Join("internal", "types", "types.go"),
+ "go.mod.tmpl": "go.mod",
+ "goreleaser.yaml.tmpl": ".goreleaser.yaml",
+ "golangci.yml.tmpl": ".golangci.yml",
+ "makefile.tmpl": "Makefile",
+ "readme.md.tmpl": "README.md",
}
for tmplName, outPath := range singleFiles {
@@ -124,6 +124,14 @@ func (g *Generator) Generate() error {
}
}
+ // Conditionally render auth command when OAuth2 is detected
+ if g.Spec.Auth.AuthorizationURL != "" {
+ authPath := filepath.Join("internal", "cli", "auth.go")
+ if err := g.renderTemplate("auth.go.tmpl", authPath, g.Spec); err != nil {
+ return fmt.Errorf("rendering auth: %w", err)
+ }
+ }
+
return nil
}
diff --git a/internal/generator/generator_test.go b/internal/generator/generator_test.go
index d41c1380..e6166328 100644
--- a/internal/generator/generator_test.go
+++ b/internal/generator/generator_test.go
@@ -6,6 +6,7 @@ import (
"path/filepath"
"testing"
+ "github.com/mvanhorn/cli-printing-press/internal/openapi"
"github.com/mvanhorn/cli-printing-press/internal/spec"
"github.com/stretchr/testify/require"
)
@@ -48,6 +49,37 @@ func TestGenerateProjectsCompile(t *testing.T) {
}
}
+func TestGenerateOAuth2AuthTemplateConditionally(t *testing.T) {
+ t.Parallel()
+
+ t.Run("oauth2 spec includes auth command", func(t *testing.T) {
+ data, err := os.ReadFile(filepath.Join("..", "..", "testdata", "openapi", "gmail.yaml"))
+ require.NoError(t, err)
+
+ apiSpec, err := openapi.Parse(data)
+ require.NoError(t, err)
+
+ outputDir := filepath.Join(t.TempDir(), apiSpec.Name+"-cli")
+ gen := New(apiSpec, outputDir)
+ require.NoError(t, gen.Generate())
+
+ _, err = os.Stat(filepath.Join(outputDir, "internal", "cli", "auth.go"))
+ require.NoError(t, err)
+ })
+
+ t.Run("non-oauth2 spec omits auth command", func(t *testing.T) {
+ apiSpec, err := spec.Parse(filepath.Join("..", "..", "testdata", "stytch.yaml"))
+ require.NoError(t, err)
+
+ outputDir := filepath.Join(t.TempDir(), apiSpec.Name+"-cli")
+ gen := New(apiSpec, outputDir)
+ require.NoError(t, gen.Generate())
+
+ _, err = os.Stat(filepath.Join(outputDir, "internal", "cli", "auth.go"))
+ require.True(t, os.IsNotExist(err))
+ })
+}
+
func countFiles(t *testing.T, root string) int {
t.Helper()
diff --git a/internal/generator/templates/auth.go.tmpl b/internal/generator/templates/auth.go.tmpl
new file mode 100644
index 00000000..66b6a2cf
--- /dev/null
+++ b/internal/generator/templates/auth.go.tmpl
@@ -0,0 +1,248 @@
+package cli
+
+import (
+ "context"
+ "crypto/rand"
+ "encoding/hex"
+ "encoding/json"
+ "fmt"
+ "net"
+ "net/http"
+ "net/url"
+ "os"
+ "os/exec"
+ "runtime"
+ "strings"
+ "time"
+
+ "github.com/USER/{{.Name}}-cli/internal/config"
+ "github.com/spf13/cobra"
+)
+
+func newAuthCmd(flags *rootFlags) *cobra.Command {
+ cmd := &cobra.Command{
+ Use: "auth",
+ Short: "Manage authentication",
+ }
+
+ cmd.AddCommand(newAuthLoginCmd(flags))
+ cmd.AddCommand(newAuthStatusCmd(flags))
+ cmd.AddCommand(newAuthLogoutCmd(flags))
+
+ return cmd
+}
+
+func newAuthLoginCmd(flags *rootFlags) *cobra.Command {
+ var clientID string
+ var clientSecret string
+ var port int
+
+ cmd := &cobra.Command{
+ Use: "login",
+ Short: "Authenticate via OAuth2",
+ RunE: func(cmd *cobra.Command, args []string) error {
+ if clientID == "" {
+ return fmt.Errorf("--client-id is required")
+ }
+
+ cfg, err := config.Load(flags.configPath)
+ if err != nil {
+ return err
+ }
+
+ stateBytes := make([]byte, 16)
+ if _, err := rand.Read(stateBytes); err != nil {
+ return fmt.Errorf("generating state: %w", err)
+ }
+ state := hex.EncodeToString(stateBytes)
+
+ listener, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", port))
+ if err != nil {
+ return fmt.Errorf("starting callback server: %w", err)
+ }
+ defer listener.Close()
+
+ redirectURI := fmt.Sprintf("http://localhost:%d/callback", listener.Addr().(*net.TCPAddr).Port)
+
+ authURL := "{{.Auth.AuthorizationURL}}"
+ params := url.Values{
+ "client_id": {clientID},
+ "redirect_uri": {redirectURI},
+ "response_type": {"code"},
+ "state": {state},
+ "access_type": {"offline"},
+ "prompt": {"consent"},
+ }
+ scopes := []string{ {{- range .Auth.Scopes}}"{{.}}", {{end}} }
+ if len(scopes) > 0 {
+ params.Set("scope", strings.Join(scopes, " "))
+ }
+
+ fullURL := authURL + "?" + params.Encode()
+ fmt.Fprintf(os.Stderr, "Opening browser for authentication...\n")
+ fmt.Fprintf(os.Stderr, "If the browser doesn't open, visit:\n%s\n\n", fullURL)
+ openBrowser(fullURL)
+
+ codeCh := make(chan string, 1)
+ errCh := make(chan error, 1)
+
+ mux := http.NewServeMux()
+ mux.HandleFunc("/callback", func(w http.ResponseWriter, r *http.Request) {
+ if r.URL.Query().Get("state") != state {
+ errCh <- fmt.Errorf("state mismatch")
+ http.Error(w, "State mismatch", http.StatusBadRequest)
+ return
+ }
+ if errMsg := r.URL.Query().Get("error"); errMsg != "" {
+ errCh <- fmt.Errorf("auth error: %s", errMsg)
+ http.Error(w, errMsg, http.StatusBadRequest)
+ return
+ }
+ code := r.URL.Query().Get("code")
+ if code == "" {
+ errCh <- fmt.Errorf("no code in callback")
+ http.Error(w, "No code", http.StatusBadRequest)
+ return
+ }
+ w.Header().Set("Content-Type", "text/html")
+ fmt.Fprint(w, "<html><body><h2>Authentication successful!</h2><p>You can close this tab.</p></body></html>")
+ codeCh <- code
+ })
+
+ server := &http.Server{Handler: mux}
+ go server.Serve(listener)
+
+ var code string
+ select {
+ case code = <-codeCh:
+ case err := <-errCh:
+ return err
+ case <-time.After(2 * time.Minute):
+ return fmt.Errorf("authentication timed out after 2 minutes")
+ }
+
+ server.Shutdown(context.Background())
+
+ tokenURL := "{{.Auth.TokenURL}}"
+ tokenParams := url.Values{
+ "grant_type": {"authorization_code"},
+ "code": {code},
+ "redirect_uri": {redirectURI},
+ "client_id": {clientID},
+ }
+ if clientSecret != "" {
+ tokenParams.Set("client_secret", clientSecret)
+ }
+
+ resp, err := http.PostForm(tokenURL, tokenParams)
+ if err != nil {
+ return fmt.Errorf("exchanging code for token: %w", err)
+ }
+ defer resp.Body.Close()
+
+ if resp.StatusCode >= 400 {
+ var body map[string]any
+ if err := json.NewDecoder(resp.Body).Decode(&body); err == nil {
+ return fmt.Errorf("exchanging code for token: HTTP %d: %v", resp.StatusCode, body)
+ }
+ return fmt.Errorf("exchanging code for token: HTTP %d", resp.StatusCode)
+ }
+
+ var tokenResp struct {
+ AccessToken string `json:"access_token"`
+ RefreshToken string `json:"refresh_token"`
+ ExpiresIn int `json:"expires_in"`
+ TokenType string `json:"token_type"`
+ }
+ if err := json.NewDecoder(resp.Body).Decode(&tokenResp); err != nil {
+ return fmt.Errorf("parsing token response: %w", err)
+ }
+ if tokenResp.AccessToken == "" {
+ return fmt.Errorf("no access token in response")
+ }
+
+ expiry := time.Now().Add(time.Duration(tokenResp.ExpiresIn) * time.Second)
+ if err := cfg.SaveTokens(clientID, clientSecret, tokenResp.AccessToken, tokenResp.RefreshToken, expiry); err != nil {
+ return fmt.Errorf("saving tokens: %w", err)
+ }
+
+ fmt.Fprintf(os.Stderr, "%s Authentication successful! Token expires at %s\n", green("OK"), expiry.Format(time.RFC3339))
+ return nil
+ },
+ }
+
+ cmd.Flags().StringVar(&clientID, "client-id", os.Getenv("{{upper .Name}}_CLIENT_ID"), "OAuth2 client ID")
+ cmd.Flags().StringVar(&clientSecret, "client-secret", os.Getenv("{{upper .Name}}_CLIENT_SECRET"), "OAuth2 client secret")
+ cmd.Flags().IntVar(&port, "port", 8085, "Local callback server port")
+
+ return cmd
+}
+
+func newAuthStatusCmd(flags *rootFlags) *cobra.Command {
+ return &cobra.Command{
+ Use: "status",
+ Short: "Show authentication status",
+ RunE: func(cmd *cobra.Command, args []string) error {
+ cfg, err := config.Load(flags.configPath)
+ if err != nil {
+ return err
+ }
+
+ w := cmd.OutOrStdout()
+ if cfg.AccessToken == "" {
+ fmt.Fprintf(w, " %s Not authenticated. Run 'auth login' to authenticate.\n", red("FAIL"))
+ return nil
+ }
+
+ if cfg.TokenExpiry.IsZero() {
+ fmt.Fprintf(w, " %s Authenticated (no expiry info)\n", green("OK"))
+ } else if time.Now().Before(cfg.TokenExpiry) {
+ fmt.Fprintf(w, " %s Authenticated (expires %s)\n", green("OK"), cfg.TokenExpiry.Format(time.RFC3339))
+ } else {
+ if cfg.RefreshToken != "" {
+ fmt.Fprintf(w, " %s Token expired (will auto-refresh on next request)\n", yellow("WARN"))
+ } else {
+ fmt.Fprintf(w, " %s Token expired. Run 'auth login' to re-authenticate.\n", red("FAIL"))
+ }
+ }
+
+ if cfg.AuthSource != "" {
+ fmt.Fprintf(w, " source: %s\n", cfg.AuthSource)
+ }
+ return nil
+ },
+ }
+}
+
+func newAuthLogoutCmd(flags *rootFlags) *cobra.Command {
+ return &cobra.Command{
+ Use: "logout",
+ Short: "Remove stored authentication tokens",
+ RunE: func(cmd *cobra.Command, args []string) error {
+ cfg, err := config.Load(flags.configPath)
+ if err != nil {
+ return err
+ }
+ if err := cfg.ClearTokens(); err != nil {
+ return fmt.Errorf("clearing tokens: %w", err)
+ }
+ fmt.Fprintf(os.Stderr, "Logged out. Tokens removed.\n")
+ return nil
+ },
+ }
+}
+
+func openBrowser(url string) {
+ var cmd *exec.Cmd
+ switch runtime.GOOS {
+ case "darwin":
+ cmd = exec.Command("open", url)
+ case "linux":
+ cmd = exec.Command("xdg-open", url)
+ case "windows":
+ cmd = exec.Command("rundll32", "url.dll,FileProtocolHandler", url)
+ default:
+ return
+ }
+ cmd.Start()
+}
diff --git a/internal/generator/templates/client.go.tmpl b/internal/generator/templates/client.go.tmpl
index 8521e5ff..95e8ce72 100644
--- a/internal/generator/templates/client.go.tmpl
+++ b/internal/generator/templates/client.go.tmpl
@@ -6,15 +6,18 @@ import (
"io"
"math"
"net/http"
+ "net/url"
"os"
"strconv"
"strings"
"time"
+
+ "github.com/USER/{{.Name}}-cli/internal/config"
)
type Client struct {
BaseURL string
- AuthHeader string
+ Config *config.Config
HTTPClient *http.Client
DryRun bool
}
@@ -31,10 +34,10 @@ func (e *APIError) Error() string {
return fmt.Sprintf("%s %s returned HTTP %d: %s", e.Method, e.Path, e.StatusCode, e.Body)
}
-func New(baseURL, authHeader string, timeout time.Duration) *Client {
+func New(cfg *config.Config, timeout time.Duration) *Client {
return &Client{
- BaseURL: strings.TrimRight(baseURL, "/"),
- AuthHeader: authHeader,
+ BaseURL: strings.TrimRight(cfg.BaseURL, "/"),
+ Config: cfg,
HTTPClient: &http.Client{Timeout: timeout},
}
}
@@ -90,8 +93,12 @@ func (c *Client) do(method, path string, params map[string]string, body any) (js
return nil, fmt.Errorf("creating request: %w", err)
}
- if c.AuthHeader != "" {
- req.Header.Set("Authorization", c.AuthHeader)
+ authHeader, err := c.authHeader()
+ if err != nil {
+ return nil, err
+ }
+ if authHeader != "" {
+ req.Header.Set("Authorization", authHeader)
}
if bodyBytes != nil {
req.Header.Set("Content-Type", "application/json")
@@ -166,9 +173,13 @@ func (c *Client) dryRun(method, url string, params map[string]string, body []byt
}
}
}
- if c.AuthHeader != "" {
+ authHeader, err := c.authHeader()
+ if err != nil {
+ return nil, err
+ }
+ if authHeader != "" {
// Mask token for safety
- auth := c.AuthHeader
+ auth := authHeader
if len(auth) > 20 {
auth = auth[:15] + "..."
}
@@ -187,6 +198,80 @@ func (c *Client) dryRun(method, url string, params map[string]string, body []byt
return json.RawMessage(`{"dry_run": true}`), nil
}
+func (c *Client) authHeader() (string, error) {
+ if c.Config == nil {
+ return "", nil
+ }
+ if c.Config.AccessToken != "" && !c.Config.TokenExpiry.IsZero() && time.Now().After(c.Config.TokenExpiry) && c.Config.RefreshToken != "" {
+ if err := c.refreshAccessToken(); err != nil {
+ return "", err
+ }
+ }
+ return c.Config.AuthHeader(), nil
+}
+
+func (c *Client) refreshAccessToken() error {
+ if c.Config == nil {
+ return nil
+ }
+ if c.Config.RefreshToken == "" {
+ return nil
+ }
+
+ tokenURL := "{{.Auth.TokenURL}}"
+ if tokenURL == "" {
+ return nil
+ }
+
+ params := url.Values{
+ "grant_type": {"refresh_token"},
+ "refresh_token": {c.Config.RefreshToken},
+ "client_id": {c.Config.ClientID},
+ }
+ if c.Config.ClientSecret != "" {
+ params.Set("client_secret", c.Config.ClientSecret)
+ }
+
+ resp, err := c.HTTPClient.PostForm(tokenURL, params)
+ if err != nil {
+ return fmt.Errorf("refreshing access token: %w", err)
+ }
+ defer resp.Body.Close()
+
+ if resp.StatusCode >= 400 {
+ body, _ := io.ReadAll(resp.Body)
+ return fmt.Errorf("refreshing access token: HTTP %d: %s", resp.StatusCode, truncateBody(body))
+ }
+
+ var tokenResp struct {
+ AccessToken string `json:"access_token"`
+ RefreshToken string `json:"refresh_token"`
+ ExpiresIn int `json:"expires_in"`
+ }
+ if err := json.NewDecoder(resp.Body).Decode(&tokenResp); err != nil {
+ return fmt.Errorf("parsing refresh response: %w", err)
+ }
+ if tokenResp.AccessToken == "" {
+ return fmt.Errorf("refreshing access token: no access token in response")
+ }
+
+ refreshToken := c.Config.RefreshToken
+ if tokenResp.RefreshToken != "" {
+ refreshToken = tokenResp.RefreshToken
+ }
+
+ expiry := time.Time{}
+ if tokenResp.ExpiresIn > 0 {
+ expiry = time.Now().Add(time.Duration(tokenResp.ExpiresIn) * time.Second)
+ }
+
+ if err := c.Config.SaveTokens(c.Config.ClientID, c.Config.ClientSecret, tokenResp.AccessToken, refreshToken, expiry); err != nil {
+ return fmt.Errorf("saving refreshed token: %w", err)
+ }
+
+ return nil
+}
+
func retryAfter(resp *http.Response) time.Duration {
header := resp.Header.Get("Retry-After")
if header == "" {
diff --git a/internal/generator/templates/config.go.tmpl b/internal/generator/templates/config.go.tmpl
index 3f1550a4..5329aaac 100644
--- a/internal/generator/templates/config.go.tmpl
+++ b/internal/generator/templates/config.go.tmpl
@@ -5,6 +5,7 @@ import (
"os"
"path/filepath"
"strings"
+ "time"
{{- if eq .Config.Format "toml"}}
"github.com/pelletier/go-toml/v2"
@@ -15,6 +16,11 @@ type Config struct {
BaseURL string `{{configTag .Config.Format}}:"base_url"`
AuthHeaderVal string `{{configTag .Config.Format}}:"auth_header"`
AuthSource string `{{configTag .Config.Format}}:"-"`
+ AccessToken string `{{configTag .Config.Format}}:"access_token"`
+ RefreshToken string `{{configTag .Config.Format}}:"refresh_token"`
+ TokenExpiry time.Time `{{configTag .Config.Format}}:"token_expiry"`
+ ClientID string `{{configTag .Config.Format}}:"client_id"`
+ ClientSecret string `{{configTag .Config.Format}}:"client_secret"`
Path string `{{configTag .Config.Format}}:"-"`
{{- range .Auth.EnvVars}}
{{envVarField .}} string `{{configTag $.Config.Format}}:"{{envVarPlaceholder .}}"`
@@ -79,6 +85,10 @@ func (c *Config) AuthHeader() string {
}
return format
{{- else if eq .Auth.Type "bearer_token"}}
+ if c.AccessToken != "" {
+ c.AuthSource = "oauth2"
+ return "Bearer " + c.AccessToken
+ }
{{- if gt (len .Auth.EnvVars) 0}}
if c.{{envVarField (index .Auth.EnvVars 0)}} != "" {
return "Bearer " + c.{{envVarField (index .Auth.EnvVars 0)}}
@@ -90,5 +100,37 @@ func (c *Config) AuthHeader() string {
{{- end}}
}
+func (c *Config) SaveTokens(clientID, clientSecret, accessToken, refreshToken string, expiry time.Time) error {
+ c.ClientID = clientID
+ c.ClientSecret = clientSecret
+ c.AccessToken = accessToken
+ c.RefreshToken = refreshToken
+ c.TokenExpiry = expiry
+ return c.save()
+}
+
+func (c *Config) ClearTokens() error {
+ c.AccessToken = ""
+ c.RefreshToken = ""
+ c.TokenExpiry = time.Time{}
+ return c.save()
+}
+
+func (c *Config) save() error {
+ dir := filepath.Dir(c.Path)
+ if err := os.MkdirAll(dir, 0o700); err != nil {
+ return fmt.Errorf("creating config dir: %w", err)
+ }
+{{- if eq .Config.Format "toml"}}
+ data, err := toml.Marshal(c)
+ if err != nil {
+ return fmt.Errorf("marshaling config: %w", err)
+ }
+ return os.WriteFile(c.Path, data, 0o600)
+{{- else}}
+ return fmt.Errorf("config format %q does not support writing", "{{.Config.Format}}")
+{{- end}}
+}
+
// Ensure strings import is used
var _ = strings.ReplaceAll
diff --git a/internal/generator/templates/root.go.tmpl b/internal/generator/templates/root.go.tmpl
index 95cecc4d..96e87562 100644
--- a/internal/generator/templates/root.go.tmpl
+++ b/internal/generator/templates/root.go.tmpl
@@ -46,6 +46,9 @@ func Execute() error {
rootCmd.AddCommand(new{{camel $name}}Cmd(&flags)) {{/* FuncPrefix matches resource name for top-level */}}
{{- end}}
rootCmd.AddCommand(newDoctorCmd(&flags))
+{{- if .Auth.AuthorizationURL}}
+ rootCmd.AddCommand(newAuthCmd(&flags))
+{{- end}}
rootCmd.AddCommand(newVersionCliCmd())
return rootCmd.Execute()
@@ -64,7 +67,7 @@ func (f *rootFlags) newClient() (*client.Client, error) {
if err != nil {
return nil, configErr(err)
}
- c := client.New(cfg.BaseURL, cfg.AuthHeader(), f.timeout)
+ c := client.New(cfg, f.timeout)
c.DryRun = f.dryRun
return c, nil
}
diff --git a/internal/openapi/parser.go b/internal/openapi/parser.go
index 68587f8f..033e0ee8 100644
--- a/internal/openapi/parser.go
+++ b/internal/openapi/parser.go
@@ -118,6 +118,22 @@ func mapAuth(doc *openapi3.T, name string) spec.AuthConfig {
case "oauth2":
auth.Type = "bearer_token"
auth.Header = "Authorization"
+ if scheme.Flows != nil {
+ if ac := scheme.Flows.AuthorizationCode; ac != nil {
+ auth.AuthorizationURL = ac.AuthorizationURL
+ auth.TokenURL = ac.TokenURL
+ for scope := range ac.Scopes {
+ auth.Scopes = append(auth.Scopes, scope)
+ }
+ sort.Strings(auth.Scopes)
+ } else if ic := scheme.Flows.Implicit; ic != nil {
+ auth.AuthorizationURL = ic.AuthorizationURL
+ for scope := range ic.Scopes {
+ auth.Scopes = append(auth.Scopes, scope)
+ }
+ sort.Strings(auth.Scopes)
+ }
+ }
}
envPrefix := strings.ToUpper(strings.ReplaceAll(name, "-", "_"))
@@ -143,6 +159,16 @@ func selectSecurityScheme(doc *openapi3.T) (string, *openapi3.SecurityScheme) {
}
orderedNames := orderedSecuritySchemeNames(doc)
+ for _, name := range orderedNames {
+ scheme := securitySchemeValue(doc.Components.SecuritySchemes[name])
+ if scheme == nil || !strings.EqualFold(scheme.Type, "oauth2") || scheme.Flows == nil {
+ continue
+ }
+ if ac := scheme.Flows.AuthorizationCode; ac != nil && strings.TrimSpace(ac.AuthorizationURL) != "" && strings.TrimSpace(ac.TokenURL) != "" {
+ return name, scheme
+ }
+ }
+
for _, name := range orderedNames {
scheme := securitySchemeValue(doc.Components.SecuritySchemes[name])
if scheme == nil {
diff --git a/internal/openapi/parser_test.go b/internal/openapi/parser_test.go
index fd588868..e1f5ada2 100644
--- a/internal/openapi/parser_test.go
+++ b/internal/openapi/parser_test.go
@@ -62,6 +62,22 @@ func TestParseStytchOpenAPI(t *testing.T) {
assert.Greater(t, totalEndpoints, 10)
}
+func TestParseGmailOAuth2(t *testing.T) {
+ t.Parallel()
+
+ data, err := os.ReadFile(filepath.Join("..", "..", "testdata", "openapi", "gmail.yaml"))
+ require.NoError(t, err)
+
+ parsed, err := Parse(data)
+ require.NoError(t, err)
+
+ assert.Equal(t, "bearer_token", parsed.Auth.Type)
+ assert.Equal(t, "Authorization", parsed.Auth.Header)
+ assert.Equal(t, "https://accounts.google.com/o/oauth2/auth", parsed.Auth.AuthorizationURL)
+ assert.Equal(t, "https://accounts.google.com/o/oauth2/token", parsed.Auth.TokenURL)
+ assert.NotEmpty(t, parsed.Auth.Scopes)
+}
+
func TestSkipUnderscoreFields(t *testing.T) {
spec := []byte(`
openapi: "3.0.0"
diff --git a/internal/spec/spec.go b/internal/spec/spec.go
index 10f5ccc7..2d44b49b 100644
--- a/internal/spec/spec.go
+++ b/internal/spec/spec.go
@@ -20,12 +20,15 @@ type APISpec struct {
}
type AuthConfig struct {
- Type string `yaml:"type"` // api_key, oauth2, bearer_token, none
- Header string `yaml:"header"`
- Format string `yaml:"format"`
- EnvVars []string `yaml:"env_vars"`
- Scheme string `yaml:"scheme,omitempty"` // OpenAPI security scheme name
- In string `yaml:"in,omitempty"` // header, query, cookie
+ Type string `yaml:"type"` // api_key, oauth2, bearer_token, none
+ Header string `yaml:"header"`
+ Format string `yaml:"format"`
+ EnvVars []string `yaml:"env_vars"`
+ Scheme string `yaml:"scheme,omitempty"` // OpenAPI security scheme name
+ In string `yaml:"in,omitempty"` // header, query, cookie
+ AuthorizationURL string `yaml:"authorization_url,omitempty"`
+ TokenURL string `yaml:"token_url,omitempty"`
+ Scopes []string `yaml:"scopes,omitempty"`
}
type ConfigSpec struct {
← 2e70bc46 feat(cli): accept URLs for --spec with local caching
·
back to Cli Printing Press
·
feat(cli): multi-spec composition with --spec repetition bb12a605 →