Skip to content
Open
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
225 changes: 225 additions & 0 deletions cmd/cli/README.md

Large diffs are not rendered by default.

951 changes: 951 additions & 0 deletions cmd/cli/main.go

Large diffs are not rendered by default.

65 changes: 65 additions & 0 deletions internal/cli/auth/discovery.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
package auth

import (
"context"
"encoding/json"
"fmt"
"io"
"net/http"
"strings"
)

// oidcDiscoveryDocument mirrors the subset of an OpenID Connect discovery
// document (OIDC Discovery 1.0 / RFC 8414) this CLI needs.
type oidcDiscoveryDocument struct {
AuthorizationEndpoint string `json:"authorization_endpoint"`
TokenEndpoint string `json:"token_endpoint"`
}

// maxDiscoveryResponseBytes bounds how much of a discovery response gets
// read into memory - real discovery documents are a few KB, so this is
// generous headroom against a misbehaving or malicious server sending an
// oversized or non-terminating response.
const maxDiscoveryResponseBytes = 1 << 20 // 1 MiB

// DiscoverEndpoints fetches the OIDC discovery document at
// issuer + "/.well-known/openid-configuration" and returns its authorization
// and token endpoints. This lets callers configure just an issuer/base URL,
// as most standards-compliant identity providers support discovery, instead
// of every individual endpoint.
func DiscoverEndpoints(ctx context.Context, httpClient *http.Client, issuer string) (authURL, tokenURL string, err error) {
if httpClient == nil {
httpClient = http.DefaultClient
}
discoveryURL := strings.TrimRight(issuer, "/") + "/.well-known/openid-configuration"

req, err := http.NewRequestWithContext(ctx, http.MethodGet, discoveryURL, nil)
if err != nil {
return "", "", fmt.Errorf("failed to create discovery request: %w", err)
}
resp, err := httpClient.Do(req)
if err != nil {
return "", "", fmt.Errorf("failed to reach discovery endpoint %s: %w", discoveryURL, err)
}
defer func() { _ = resp.Body.Close() }()

body, err := io.ReadAll(io.LimitReader(resp.Body, maxDiscoveryResponseBytes+1))
if err != nil {
return "", "", fmt.Errorf("failed to read discovery response: %w", err)
}
if len(body) > maxDiscoveryResponseBytes {
return "", "", fmt.Errorf("discovery endpoint %s returned a response larger than %d bytes", discoveryURL, maxDiscoveryResponseBytes)
}
if resp.StatusCode != http.StatusOK {
return "", "", fmt.Errorf("discovery endpoint %s returned status %d: %s", discoveryURL, resp.StatusCode, string(body))
}

var doc oidcDiscoveryDocument
if err := json.Unmarshal(body, &doc); err != nil {
return "", "", fmt.Errorf("failed to parse discovery document from %s: %w", discoveryURL, err)
}
if doc.AuthorizationEndpoint == "" || doc.TokenEndpoint == "" {
return "", "", fmt.Errorf("discovery document from %s is missing authorization_endpoint or token_endpoint", discoveryURL)
}
return doc.AuthorizationEndpoint, doc.TokenEndpoint, nil
}
80 changes: 80 additions & 0 deletions internal/cli/auth/discovery_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
package auth

import (
"context"
"net/http"
"net/http/httptest"
"testing"

"github.com/stretchr/testify/assert"
)

func TestDiscoverEndpoints_Success(t *testing.T) {
var requestedPath string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requestedPath = r.URL.Path
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"authorization_endpoint":"https://idp.example.com/oauth2/authorize","token_endpoint":"https://idp.example.com/oauth2/token"}`))
}))
defer server.Close()

authURL, tokenURL, err := DiscoverEndpoints(context.Background(), server.Client(), server.URL)
assert.NoError(t, err)
assert.Equal(t, "https://idp.example.com/oauth2/authorize", authURL)
assert.Equal(t, "https://idp.example.com/oauth2/token", tokenURL)
assert.Equal(t, "/.well-known/openid-configuration", requestedPath)
}

func TestDiscoverEndpoints_TrailingSlashOnIssuer(t *testing.T) {
var requestedPath string
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
requestedPath = r.URL.Path
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"authorization_endpoint":"https://idp.example.com/oauth2/authorize","token_endpoint":"https://idp.example.com/oauth2/token"}`))
}))
defer server.Close()

_, _, err := DiscoverEndpoints(context.Background(), server.Client(), server.URL+"/")
assert.NoError(t, err)
assert.Equal(t, "/.well-known/openid-configuration", requestedPath)
}

func TestDiscoverEndpoints_NotFound(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusNotFound)
_, _ = w.Write([]byte("not found"))
}))
defer server.Close()

_, _, err := DiscoverEndpoints(context.Background(), server.Client(), server.URL)
assert.Error(t, err)
assert.ErrorContains(t, err, "status 404")
}

func TestDiscoverEndpoints_ResponseTooLarge(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
oversized := make([]byte, maxDiscoveryResponseBytes+1)
for i := range oversized {
oversized[i] = ' '
}
_, _ = w.Write(oversized)
}))
defer server.Close()

_, _, err := DiscoverEndpoints(context.Background(), server.Client(), server.URL)
assert.Error(t, err)
assert.ErrorContains(t, err, "larger than")
}

func TestDiscoverEndpoints_MissingEndpoints(t *testing.T) {
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"issuer":"https://idp.example.com"}`))
}))
defer server.Close()

_, _, err := DiscoverEndpoints(context.Background(), server.Client(), server.URL)
assert.Error(t, err)
assert.ErrorContains(t, err, "missing authorization_endpoint or token_endpoint")
}
Loading
Loading