Skip to content
Merged
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
66 changes: 56 additions & 10 deletions internal/cmd/jwtinfo.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,7 +17,8 @@ import (
"github.com/spf13/cobra"
"github.com/xenos76/https-wrench/internal/errdisp"
"github.com/xenos76/https-wrench/internal/jwtinfo"
"github.com/xenos76/https-wrench/internal/style"
"github.com/xenos76/https-wrench/internal/view"
"golang.org/x/term"
)

var (
Expand All @@ -36,6 +37,7 @@ var (
refresh bool
tokenOutputFile string
renewThreshold float64
jwtinfoFmt string
keyfuncDefOverride keyfunc.Override

// requestSteps tracks the sequence of request-related flags as they appear on the command line.
Expand Down Expand Up @@ -99,6 +101,9 @@ Examples:
--request-values-json $REQ_VALUES \
--validation-url $VALIDATION_URL

# Emit machine-readable JSON (agents / MCP)
https-wrench jwtinfo --token-file ./token.jwt --format json

# Request a JWT token, write it to a file and refresh it before expiration
https-wrench jwtinfo \
--request-url $REQ_URL \
Expand All @@ -114,8 +119,15 @@ Examples:
requestValuesMap = make(map[string]string)
)

switch jwtinfoFmt {
case "", "text", "json":
default:
cmd.Printf("Error: unsupported --format %q (use text or json)\n", jwtinfoFmt)
return
}

if refresh && requestURL == "" {
fmt.Fprintln(cmd.OutOrStdout(), style.LgSprintf(style.Error, "Error: --refresh requires --request-url"))
cmd.Print("Error: --refresh requires --request-url\n")
return
}

Expand Down Expand Up @@ -187,14 +199,41 @@ Examples:
}
}

err = jwtinfo.PrintTokenInfo(tokenData, cmd.OutOrStdout())
if err != nil {
cmd.Printf("error while printing token data: %s\n", errdisp.FormatCause(err))
result, buildErr := tokenData.BuildResult()
if buildErr != nil {
cmd.Printf("error building Jwtinfo result: %s\n", errdisp.FormatCause(buildErr))
return
}

out := cmd.OutOrStdout()
statusOut := out

if jwtinfoFmt == "json" {
statusOut = cmd.ErrOrStderr()
}

if jwtinfoFmt == "json" {
payload, encErr := jwtinfo.EncodeJSON(result)
if encErr != nil {
cmd.Printf("error encoding Jwtinfo JSON: %s\n", errdisp.FormatCause(encErr))
return
}

cmd.Println(string(payload))
Comment thread
coderabbitai[bot] marked this conversation as resolved.
} else {
opts := view.Options{}
if f, ok := out.(*os.File); ok && term.IsTerminal(int(f.Fd())) {
opts.ForceColor = true
}

if err = view.Render(out, jwtinfo.BuildDoc(result), opts); err != nil {
cmd.Printf("error while printing token data: %s\n", errdisp.FormatCause(err))
return
}
}

if tokenOutputFile != "" {
tokenData.WriteTokenToFile(tokenOutputFile, cmd.OutOrStdout())
tokenData.WriteTokenToFile(tokenOutputFile, statusOut)
}

if refresh {
Expand All @@ -210,7 +249,7 @@ Examples:
cancel()
}()

cmd.Printf("Starting refresh loop...\n")
fmt.Fprintln(statusOut, "Starting refresh loop...")

err := tokenData.RefreshLoop(
ctx,
Expand All @@ -220,12 +259,12 @@ Examples:
io.ReadAll,
renewThreshold,
tokenOutputFile,
cmd.OutOrStdout(),
statusOut,
)
if err != nil {
cmd.Printf("Refresh loop exited with error: %s\n", errdisp.FormatCause(err))
fmt.Fprintf(statusOut, "Refresh loop exited with error: %s\n", errdisp.FormatCause(err))
} else {
cmd.Printf("Refresh loop stopped gracefully.\n")
fmt.Fprintln(statusOut, "Refresh loop stopped gracefully.")
}
}
} else {
Expand Down Expand Up @@ -297,6 +336,13 @@ func init() {
"Percentage of token lifetime to wait before refreshing",
)

jwtinfoCmd.Flags().StringVar(
&jwtinfoFmt,
"format",
"text",
"Output format: text (default) or json",
)

// Either read a token from a file or request it from an HTTP address
jwtinfoCmd.MarkFlagsMutuallyExclusive(flagNameTokenFile, flagNameRequestURL)
jwtinfoCmd.MarkFlagsOneRequired(flagNameTokenFile, flagNameRequestURL)
Expand Down
137 changes: 131 additions & 6 deletions internal/cmd/jwtinfo_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -3,12 +3,16 @@ package cmd
import (
"bytes"
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"path/filepath"
"testing"
"time"

"github.com/spf13/pflag"
"github.com/stretchr/testify/require"
"github.com/xenos76/https-wrench/internal/jwtinfo"
)

func TestJwtinfoCmd_Errors(t *testing.T) {
Expand All @@ -33,6 +37,14 @@ func TestJwtinfoCmd_Errors(t *testing.T) {
},
expected: []string{"Error: --refresh requires --request-url"},
},
{
name: "unsupported format",
setup: func() {
tokenFile = "some.jwt"
jwtinfoFmt = "yaml"
},
expected: []string{"Error: unsupported --format \"yaml\" (use text or json)"},
},
}

for _, tt := range tests {
Expand Down Expand Up @@ -68,21 +80,133 @@ func TestJwtinfoCmd_Success(t *testing.T) {
}))
defer ts.Close()

t.Run("default text format", func(t *testing.T) {
resetFlags()

requestURL = ts.URL
requestSteps = []requestValueStep{{kind: "kv", value: "key=val"}}

out := new(bytes.Buffer)
jwtinfoCmd.SetOut(out)
jwtinfoCmd.SetErr(out)
jwtinfoCmd.SetContext(context.Background())

jwtinfoCmd.Run(jwtinfoCmd, nil)

got := out.String()
require.Contains(t, got, "JwtInfo")
require.Contains(t, got, "AccessToken")
require.Contains(t, got, "\"sub\"")
require.Contains(t, got, "\"1234567890\"")
})

t.Run("json format", func(t *testing.T) {
resetFlags()

requestURL = ts.URL
requestSteps = []requestValueStep{{kind: "kv", value: "key=val"}}
jwtinfoFmt = "json"

out := new(bytes.Buffer)
jwtinfoCmd.SetOut(out)
jwtinfoCmd.SetErr(out)
jwtinfoCmd.SetContext(context.Background())

jwtinfoCmd.Run(jwtinfoCmd, nil)

got := out.String()
require.NotContains(t, got, "\x1b[")

var res jwtinfo.Result
require.NoError(t, json.Unmarshal(out.Bytes(), &res))
require.Equal(t, jwtinfo.ResultSchemaVersion, res.SchemaVersion)
require.Equal(t, "jwtinfo", res.Command)
require.NotNil(t, res.AccessToken)
require.Contains(t, string(res.AccessToken.Claims), "1234567890")
})
}

func TestJwtinfoCmd_JSONStatusRouting_OutputFile(t *testing.T) {
ts := newMockTokenServer()
defer ts.Close()

resetFlags()

tmpFile := filepath.Join(t.TempDir(), "out.jwt")
tokenOutputFile = tmpFile
requestURL = ts.URL
requestSteps = []requestValueStep{{kind: "kv", value: "key=val"}}
jwtinfoFmt = "json"

stdout := new(bytes.Buffer)
stderr := new(bytes.Buffer)

out := new(bytes.Buffer)
jwtinfoCmd.SetOut(out)
jwtinfoCmd.SetErr(out)
jwtinfoCmd.SetOut(stdout)
jwtinfoCmd.SetErr(stderr)
jwtinfoCmd.SetContext(context.Background())

jwtinfoCmd.Run(jwtinfoCmd, nil)

got := out.String()
require.Contains(t, got, "\"sub\"")
require.Contains(t, got, "\"1234567890\"")
gotOut := stdout.String()
require.NotContains(t, gotOut, "\x1b[")
require.NotContains(t, gotOut, "Token persisted to")

var res jwtinfo.Result
require.NoError(t, json.Unmarshal(stdout.Bytes(), &res))
require.Equal(t, jwtinfo.ResultSchemaVersion, res.SchemaVersion)

require.Contains(t, stderr.String(), "Token persisted to")
}

func TestJwtinfoCmd_JSONStatusRouting_Refresh(t *testing.T) {
ts := newMockTokenServer()
defer ts.Close()

resetFlags()

requestURL = ts.URL
requestSteps = []requestValueStep{{kind: "kv", value: "key=val"}}
refresh = true
jwtinfoFmt = "json"

stdout := new(bytes.Buffer)
stderr := new(bytes.Buffer)

jwtinfoCmd.SetOut(stdout)
jwtinfoCmd.SetErr(stderr)

ctx, cancel := context.WithCancel(context.Background())

go func() {
time.Sleep(50 * time.Millisecond)
cancel()
}()

jwtinfoCmd.SetContext(ctx)

jwtinfoCmd.Run(jwtinfoCmd, nil)

gotOut := stdout.String()
require.NotContains(t, gotOut, "\x1b[")
require.NotContains(t, gotOut, "Starting refresh loop")
require.NotContains(t, gotOut, "Refresh loop")

var res jwtinfo.Result
require.NoError(t, json.Unmarshal(stdout.Bytes(), &res))
require.Equal(t, jwtinfo.ResultSchemaVersion, res.SchemaVersion)

require.Contains(t, stderr.String(), "Starting refresh loop")
}

func newMockTokenServer() *httptest.Server {
return httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
w.Header().Set("Content-Type", "application/json")

token := "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9." +
"eyJzdWIiOiIxMjM0NTY3ODkwIiwibmFtZSI6IkpvaG4gRG9lIiwiaWF0IjoxNTE2M" +
"jM5MDIyLCJleHAiOjE1MTYyNDkwMjJ9.c2lnbmF0dXJl"
_, _ = w.Write([]byte(`{"access_token": "` + token + `"}`))
}))
}

func resetFlags() {
Expand All @@ -98,4 +222,5 @@ func resetFlags() {
jwksURL = ""
tokenOutputFile = ""
renewThreshold = 80.0
jwtinfoFmt = "text"
}
Loading
Loading