Add native verified login and service-token commands
All checks were successful
Build and Publish Container Image / build-and-push (push) Successful in 41s
All checks were successful
Build and Publish Container Image / build-and-push (push) Successful in 41s
Assistant: codex Assistant-Model: gpt-6-astra Assistant-Session: 01a06e87-e039-7ed2-b85c-20ad37f8a21b
This commit is contained in:
parent
66df5fcf07
commit
b989de4e90
12 changed files with 928 additions and 14 deletions
|
|
@ -21,6 +21,7 @@ import (
|
|||
"keycape/internal/adapters/authelia"
|
||||
"keycape/internal/adapters/lldap"
|
||||
"keycape/internal/adapters/privacyidea"
|
||||
"keycape/internal/authclient"
|
||||
"keycape/internal/config"
|
||||
"keycape/internal/domain"
|
||||
servererrors "keycape/internal/server/errors"
|
||||
|
|
@ -31,6 +32,14 @@ import (
|
|||
const version = "0.1.0"
|
||||
|
||||
func main() {
|
||||
if len(os.Args) > 1 && (os.Args[1] == "login" || os.Args[1] == "service-token") {
|
||||
if err := authclient.Run(context.Background(), os.Args[1:], os.Stderr); err != nil {
|
||||
fmt.Fprintln(os.Stderr, err)
|
||||
os.Exit(1)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
log := zerolog.New(os.Stdout).With().Timestamp().Logger()
|
||||
|
||||
// -----------------------------------------------------------------
|
||||
|
|
|
|||
227
src/internal/authclient/cli.go
Normal file
227
src/internal/authclient/cli.go
Normal file
|
|
@ -0,0 +1,227 @@
|
|||
package authclient
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"flag"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Run executes a caller command. Output contains only instructions or status.
|
||||
func Run(ctx context.Context, args []string, stderr io.Writer) error {
|
||||
if len(args) == 0 {
|
||||
return errors.New("expected login or service-token")
|
||||
}
|
||||
mode := args[0]
|
||||
if mode != "login" && mode != "service-token" {
|
||||
return errors.New("unknown authentication command")
|
||||
}
|
||||
fs := flag.NewFlagSet(mode, flag.ContinueOnError)
|
||||
fs.SetOutput(stderr)
|
||||
issuer := fs.String("issuer", "", "HTTPS issuer")
|
||||
id := fs.String("client-id", "", "registered client ID")
|
||||
audience := fs.String("audience", "", "expected access audience (defaults to client ID)")
|
||||
scope := fs.String("scope", "", "space-separated registered scopes")
|
||||
secretEnv := fs.String("secret-env", "", "environment variable containing service client secret")
|
||||
out := fs.String("out", "", "new private JSON token file outside Git")
|
||||
redirect := fs.String("redirect-uri", "", "exact registered HTTP loopback callback (login only)")
|
||||
if err := fs.Parse(args[1:]); err != nil {
|
||||
return err
|
||||
}
|
||||
if fs.NArg() != 0 || *id == "" || *out == "" || strings.TrimSpace(*scope) == "" {
|
||||
return errors.New("client-id, scope and out are required; positional arguments are not accepted")
|
||||
}
|
||||
if *audience == "" {
|
||||
*audience = *id
|
||||
}
|
||||
if mode == "login" && (*secretEnv != "" || !hasScope(*scope, "openid")) {
|
||||
return errors.New("login requires openid scope and a public PKCE client")
|
||||
}
|
||||
if mode == "service-token" && (*secretEnv == "" || *redirect != "") {
|
||||
return errors.New("service-token requires secret-env and does not accept redirect-uri")
|
||||
}
|
||||
c, err := New(*issuer)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
file, err := reserveOutput(*out)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
success := false
|
||||
defer func() {
|
||||
file.Close()
|
||||
if !success {
|
||||
os.Remove(file.Name())
|
||||
}
|
||||
}()
|
||||
ctx, cancel := context.WithTimeout(ctx, 5*time.Minute)
|
||||
defer cancel()
|
||||
d, err := c.Discover(ctx)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
var tokens Tokens
|
||||
if mode == "service-token" {
|
||||
secret := os.Getenv(*secretEnv)
|
||||
if secret == "" {
|
||||
return errors.New("client secret environment variable is empty")
|
||||
}
|
||||
tokens, err = c.Exchange(ctx, d, url.Values{"grant_type": {"client_credentials"}, "scope": {*scope}}, *id, secret, *audience, "")
|
||||
} else {
|
||||
tokens, err = c.Login(ctx, d, *id, *audience, *scope, *redirect, stderr)
|
||||
}
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if err = json.NewEncoder(file).Encode(tokens); err != nil {
|
||||
return errors.New("could not write token file")
|
||||
}
|
||||
if err = file.Sync(); err != nil {
|
||||
return errors.New("could not sync token file")
|
||||
}
|
||||
if err = file.Close(); err != nil {
|
||||
return errors.New("could not close token file")
|
||||
}
|
||||
success = true
|
||||
fmt.Fprintln(stderr, "Verified tokens saved to the requested private file.")
|
||||
return nil
|
||||
}
|
||||
|
||||
func reserveOutput(path string) (*os.File, error) {
|
||||
absolute, err := filepath.Abs(path)
|
||||
if err != nil {
|
||||
return nil, errors.New("invalid output path")
|
||||
}
|
||||
parent, err := filepath.EvalSymlinks(filepath.Dir(absolute))
|
||||
if err != nil {
|
||||
return nil, errors.New("output directory must already exist")
|
||||
}
|
||||
for dir := parent; ; dir = filepath.Dir(dir) {
|
||||
marker := filepath.Join(dir, ".git")
|
||||
info, statErr := os.Lstat(marker)
|
||||
if statErr == nil {
|
||||
// A worktree uses a .git file; normal repositories have .git/HEAD.
|
||||
// An empty directory alone is not a Git repository.
|
||||
if !info.IsDir() {
|
||||
return nil, errors.New("token output must be outside Git worktrees")
|
||||
}
|
||||
if _, err := os.Lstat(filepath.Join(marker, "HEAD")); err == nil {
|
||||
return nil, errors.New("token output must be outside Git worktrees")
|
||||
} else if !os.IsNotExist(err) {
|
||||
return nil, errors.New("could not inspect Git directory")
|
||||
}
|
||||
} else if !os.IsNotExist(statErr) {
|
||||
return nil, errors.New("could not inspect output directory")
|
||||
}
|
||||
if filepath.Dir(dir) == dir {
|
||||
break
|
||||
}
|
||||
}
|
||||
file, err := os.OpenFile(filepath.Join(parent, filepath.Base(absolute)), os.O_WRONLY|os.O_CREATE|os.O_EXCL, 0600)
|
||||
if err != nil {
|
||||
return nil, errors.New("output file must be new and writable")
|
||||
}
|
||||
return file, nil
|
||||
}
|
||||
|
||||
// Login binds the exact registered IPv4/IPv6 loopback address before offering
|
||||
// the browser URL. State, nonce and PKCE are independently generated per attempt.
|
||||
func (c *Client) Login(ctx context.Context, d Discovery, id, audience, scope, redirect string, output io.Writer) (Tokens, error) {
|
||||
callback, err := url.Parse(redirect)
|
||||
if err != nil || callback.Scheme != "http" || callback.User != nil || callback.RawQuery != "" || callback.Fragment != "" || callback.Port() == "" || callback.Port() == "0" || callback.Path == "" {
|
||||
return Tokens{}, errors.New("redirect-uri must be an exact HTTP loopback URL with a fixed port and path")
|
||||
}
|
||||
ip := net.ParseIP(callback.Hostname())
|
||||
if ip == nil || !ip.IsLoopback() {
|
||||
return Tokens{}, errors.New("callback must use a literal loopback IP address")
|
||||
}
|
||||
listener, err := net.Listen("tcp", callback.Host)
|
||||
if err != nil {
|
||||
return Tokens{}, errors.New("could not bind registered callback")
|
||||
}
|
||||
defer listener.Close()
|
||||
state, err := randomValue()
|
||||
if err != nil {
|
||||
return Tokens{}, err
|
||||
}
|
||||
nonce, err := randomValue()
|
||||
if err != nil {
|
||||
return Tokens{}, err
|
||||
}
|
||||
verifier, err := randomValue()
|
||||
if err != nil {
|
||||
return Tokens{}, err
|
||||
}
|
||||
challenge := sha256.Sum256([]byte(verifier))
|
||||
type callbackResult struct {
|
||||
code string
|
||||
err error
|
||||
}
|
||||
result := make(chan callbackResult, 1)
|
||||
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Content-Type", "text/plain; charset=utf-8")
|
||||
if r.Method != http.MethodGet || r.URL.Path != callback.Path || r.Host != callback.Host {
|
||||
http.Error(w, "Invalid callback", 400)
|
||||
return
|
||||
}
|
||||
q, err := url.ParseQuery(r.URL.RawQuery)
|
||||
if err != nil || len(q["state"]) != 1 || subtle.ConstantTimeCompare([]byte(q.Get("state")), []byte(state)) != 1 {
|
||||
http.Error(w, "Invalid callback state", 400)
|
||||
return
|
||||
}
|
||||
var value callbackResult
|
||||
if q.Get("error") != "" {
|
||||
value.err = errors.New("login was declined by the identity provider")
|
||||
} else if len(q["code"]) != 1 || q.Get("code") == "" {
|
||||
http.Error(w, "Missing authorization code", 400)
|
||||
return
|
||||
} else {
|
||||
value.code = q.Get("code")
|
||||
}
|
||||
select {
|
||||
case result <- value:
|
||||
fmt.Fprintln(w, "Callback received. Return to the terminal.")
|
||||
default:
|
||||
http.Error(w, "Callback already received", 409)
|
||||
}
|
||||
})
|
||||
server := &http.Server{Handler: handler, ReadHeaderTimeout: 5 * time.Second, ReadTimeout: 10 * time.Second, WriteTimeout: 10 * time.Second, MaxHeaderBytes: 16384}
|
||||
defer server.Close()
|
||||
serveErrors := make(chan error, 1)
|
||||
go func() { serveErrors <- server.Serve(listener) }()
|
||||
authorize, err := url.Parse(d.Authorization)
|
||||
if err != nil {
|
||||
return Tokens{}, errors.New("invalid authorization endpoint")
|
||||
}
|
||||
q := authorize.Query()
|
||||
for name, value := range map[string]string{"client_id": id, "redirect_uri": redirect, "response_type": "code", "scope": scope, "state": state, "nonce": nonce, "code_challenge": base64.RawURLEncoding.EncodeToString(challenge[:]), "code_challenge_method": "S256"} {
|
||||
q.Set(name, value)
|
||||
}
|
||||
authorize.RawQuery = q.Encode()
|
||||
fmt.Fprintf(output, "Open this URL in your browser to authenticate and complete MFA:\n%s\n", authorize.String())
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return Tokens{}, errors.New("login timed out or was cancelled")
|
||||
case <-serveErrors:
|
||||
return Tokens{}, errors.New("callback listener stopped")
|
||||
case value := <-result:
|
||||
if value.err != nil {
|
||||
return Tokens{}, value.err
|
||||
}
|
||||
return c.Exchange(ctx, d, url.Values{"grant_type": {"authorization_code"}, "client_id": {id}, "code": {value.code}, "code_verifier": {verifier}, "redirect_uri": {redirect}, "scope": {scope}}, id, "", audience, nonce)
|
||||
}
|
||||
}
|
||||
221
src/internal/authclient/client.go
Normal file
221
src/internal/authclient/client.go
Normal file
|
|
@ -0,0 +1,221 @@
|
|||
// Package authclient implements caller-side KeyCape authentication. Credentials
|
||||
// are delivered to a private file; errors never include response bodies or tokens.
|
||||
package authclient
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"math/big"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Client struct {
|
||||
HTTP *http.Client
|
||||
Issuer string
|
||||
}
|
||||
type Discovery struct {
|
||||
Issuer string `json:"issuer"`
|
||||
Authorization string `json:"authorization_endpoint"`
|
||||
Token string `json:"token_endpoint"`
|
||||
JWKS string `json:"jwks_uri"`
|
||||
}
|
||||
type Tokens struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
IDToken string `json:"id_token,omitempty"`
|
||||
TokenType string `json:"token_type"`
|
||||
ExpiresIn int `json:"expires_in"`
|
||||
}
|
||||
|
||||
func New(issuer string) (*Client, error) {
|
||||
u, err := url.Parse(issuer)
|
||||
if err != nil || u.Scheme != "https" || u.Host == "" || u.User != nil || u.RawQuery != "" || u.Fragment != "" {
|
||||
return nil, errors.New("issuer must be an HTTPS URL without credentials, query or fragment")
|
||||
}
|
||||
return &Client{Issuer: issuer, HTTP: &http.Client{Timeout: 30 * time.Second, CheckRedirect: func(*http.Request, []*http.Request) error { return http.ErrUseLastResponse }}}, nil
|
||||
}
|
||||
|
||||
func (c *Client) request(ctx context.Context, method, endpoint string, form url.Values, id, secret string, out any) error {
|
||||
req, err := http.NewRequestWithContext(ctx, method, endpoint, strings.NewReader(form.Encode()))
|
||||
if err != nil {
|
||||
return errors.New("invalid provider endpoint")
|
||||
}
|
||||
if method == http.MethodPost {
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
}
|
||||
if secret != "" {
|
||||
req.SetBasicAuth(url.QueryEscape(id), url.QueryEscape(secret))
|
||||
}
|
||||
response, err := c.HTTP.Do(req)
|
||||
if err != nil {
|
||||
return errors.New("provider request failed")
|
||||
}
|
||||
defer response.Body.Close()
|
||||
if response.StatusCode != http.StatusOK {
|
||||
return fmt.Errorf("provider rejected request (HTTP %d)", response.StatusCode)
|
||||
}
|
||||
body, err := io.ReadAll(io.LimitReader(response.Body, 1024*1024+1))
|
||||
if err != nil || len(body) > 1024*1024 {
|
||||
return errors.New("invalid provider response size")
|
||||
}
|
||||
if err := json.Unmarshal(body, out); err != nil {
|
||||
return errors.New("invalid provider JSON response")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *Client) Discover(ctx context.Context) (Discovery, error) {
|
||||
var d Discovery
|
||||
if err := c.request(ctx, http.MethodGet, strings.TrimRight(c.Issuer, "/")+"/.well-known/openid-configuration", nil, "", "", &d); err != nil {
|
||||
return d, err
|
||||
}
|
||||
if d.Issuer != c.Issuer {
|
||||
return d, errors.New("discovery issuer mismatch")
|
||||
}
|
||||
base, _ := url.Parse(c.Issuer)
|
||||
for _, endpoint := range []string{d.Authorization, d.Token, d.JWKS} {
|
||||
u, err := url.Parse(endpoint)
|
||||
if err != nil || u.Scheme != "https" || u.Host != base.Host || u.User != nil || u.Fragment != "" {
|
||||
return d, errors.New("discovery endpoint must use the issuer HTTPS origin")
|
||||
}
|
||||
}
|
||||
return d, nil
|
||||
}
|
||||
|
||||
// Verify checks RS256 using the discovered issuer's JWKS and exact claim bindings.
|
||||
func (c *Client) Verify(ctx context.Context, d Discovery, token, audience, nonce string) (map[string]any, error) {
|
||||
fail := errors.New("token signature or claim validation failed")
|
||||
parts := strings.Split(token, ".")
|
||||
if len(parts) != 3 {
|
||||
return nil, fail
|
||||
}
|
||||
decode := base64.RawURLEncoding.DecodeString
|
||||
header, err := decode(parts[0])
|
||||
if err != nil {
|
||||
return nil, fail
|
||||
}
|
||||
var h struct {
|
||||
Alg string `json:"alg"`
|
||||
Kid string `json:"kid"`
|
||||
Crit []string `json:"crit"`
|
||||
}
|
||||
if json.Unmarshal(header, &h) != nil || h.Alg != "RS256" || h.Kid == "" || len(h.Crit) != 0 {
|
||||
return nil, fail
|
||||
}
|
||||
var jwks struct {
|
||||
Keys []struct{ Kty, Use, Alg, Kid, N, E string } `json:"keys"`
|
||||
}
|
||||
if err := c.request(ctx, http.MethodGet, d.JWKS, nil, "", "", &jwks); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
var pub *rsa.PublicKey
|
||||
for _, k := range jwks.Keys {
|
||||
if k.Kid != h.Kid {
|
||||
continue
|
||||
}
|
||||
if pub != nil || k.Kty != "RSA" || (k.Alg != "" && k.Alg != "RS256") || (k.Use != "" && k.Use != "sig") {
|
||||
return nil, fail
|
||||
}
|
||||
n, ne := decode(k.N)
|
||||
e, ee := decode(k.E)
|
||||
if ne != nil || ee != nil || len(e) == 0 || len(e) > 4 {
|
||||
return nil, fail
|
||||
}
|
||||
pub = &rsa.PublicKey{N: new(big.Int).SetBytes(n), E: int(new(big.Int).SetBytes(e).Int64())}
|
||||
if pub.N.BitLen() < 2048 || pub.E < 3 || pub.E%2 == 0 {
|
||||
return nil, fail
|
||||
}
|
||||
}
|
||||
if pub == nil {
|
||||
return nil, fail
|
||||
}
|
||||
sig, err := decode(parts[2])
|
||||
if err != nil {
|
||||
return nil, fail
|
||||
}
|
||||
hash := sha256.Sum256([]byte(parts[0] + "." + parts[1]))
|
||||
if rsa.VerifyPKCS1v15(pub, crypto.SHA256, hash[:], sig) != nil {
|
||||
return nil, fail
|
||||
}
|
||||
payload, err := decode(parts[1])
|
||||
if err != nil {
|
||||
return nil, fail
|
||||
}
|
||||
claims := map[string]any{}
|
||||
if json.Unmarshal(payload, &claims) != nil {
|
||||
return nil, fail
|
||||
}
|
||||
exp, okExp := claims["exp"].(float64)
|
||||
iat, okIat := claims["iat"].(float64)
|
||||
now := float64(time.Now().Unix())
|
||||
if claims["iss"] != c.Issuer || claims["aud"] != audience || !okExp || !okIat || exp <= now || iat > now+30 || iat >= exp {
|
||||
return nil, fail
|
||||
}
|
||||
if sub, ok := claims["sub"].(string); !ok || sub == "" {
|
||||
return nil, fail
|
||||
}
|
||||
if nonce != "" && claims["nonce"] != nonce {
|
||||
return nil, fail
|
||||
}
|
||||
if nbf, exists := claims["nbf"]; exists {
|
||||
if n, ok := nbf.(float64); !ok || n > now {
|
||||
return nil, fail
|
||||
}
|
||||
}
|
||||
return claims, nil
|
||||
}
|
||||
|
||||
func (c *Client) Exchange(ctx context.Context, d Discovery, form url.Values, id, secret, audience, nonce string) (Tokens, error) {
|
||||
var result Tokens
|
||||
if err := c.request(ctx, http.MethodPost, d.Token, form, id, secret, &result); err != nil {
|
||||
return Tokens{}, err
|
||||
}
|
||||
if !strings.EqualFold(result.TokenType, "Bearer") || result.ExpiresIn <= 0 {
|
||||
return Tokens{}, errors.New("invalid token response")
|
||||
}
|
||||
claims, err := c.Verify(ctx, d, result.AccessToken, audience, "")
|
||||
if err != nil {
|
||||
return Tokens{}, err
|
||||
}
|
||||
granted, _ := claims["scope"].(string)
|
||||
for _, s := range strings.Fields(form.Get("scope")) {
|
||||
if !hasScope(granted, s) {
|
||||
return Tokens{}, errors.New("required scope absent from token")
|
||||
}
|
||||
}
|
||||
if nonce != "" {
|
||||
idClaims, err := c.Verify(ctx, d, result.IDToken, id, nonce)
|
||||
if err != nil {
|
||||
return Tokens{}, err
|
||||
}
|
||||
if idClaims["sub"] != claims["sub"] {
|
||||
return Tokens{}, errors.New("token subject mismatch")
|
||||
}
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
func hasScope(scopes, want string) bool {
|
||||
for _, s := range strings.Fields(scopes) {
|
||||
if s == want {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
func randomValue() (string, error) {
|
||||
b := make([]byte, 32)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(b), nil
|
||||
}
|
||||
254
src/internal/authclient/client_test.go
Normal file
254
src/internal/authclient/client_test.go
Normal file
|
|
@ -0,0 +1,254 @@
|
|||
package authclient
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/rsa"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"keycape/internal/domain"
|
||||
"keycape/internal/server/oidc"
|
||||
"keycape/internal/server/telemetry"
|
||||
)
|
||||
|
||||
type users struct{}
|
||||
|
||||
func (users) LookupUser(context.Context, string) (*domain.User, error) {
|
||||
return &domain.User{ID: "user:test", Username: "test"}, nil
|
||||
}
|
||||
func (users) LookupGroups(context.Context, string) ([]domain.Group, error) { return nil, nil }
|
||||
func (users) ValidatePassword(context.Context, string, string) (bool, error) { return true, nil }
|
||||
func (users) ListUsers(context.Context) ([]domain.User, error) { return nil, nil }
|
||||
|
||||
func provider(t *testing.T) (*Client, Discovery, *oidc.TokenHandler) {
|
||||
t.Helper()
|
||||
key, err := rsa.GenerateKey(rand.Reader, 2048)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
mux := http.NewServeMux()
|
||||
server := httptest.NewTLSServer(mux)
|
||||
t.Cleanup(server.Close)
|
||||
sessions := oidc.NewSessionStore()
|
||||
h := &oidc.TokenHandler{Issuer: server.URL, SigningKey: key, TokenLifetime: 15 * time.Minute, Sessions: sessions, Users: users{}, Emitter: telemetry.NoopEmitter{}, ClientConfig: map[string]*domain.Client{
|
||||
"service:consumer": {ClientID: "service:consumer", ClientType: "confidential", ClientSecret: "special+%: secret", GrantTypes: []string{"client_credentials"}, AllowedScopes: []string{"approval:read"}, Audience: "approval-engine", ServiceSubject: "service:test", Tenant: "tenant:test"},
|
||||
"human": {ClientID: "human", AllowedScopes: []string{"openid", "approval:approve"}, Audience: "approval-engine"},
|
||||
}}
|
||||
mux.Handle("/token", h)
|
||||
keys := oidc.NewKeySet()
|
||||
keys.AddKey("key-1", &key.PublicKey)
|
||||
mux.Handle("/jwks", oidc.NewJWKSHandler(keys))
|
||||
d := Discovery{Issuer: server.URL, Authorization: server.URL + "/authorize", Token: server.URL + "/token", JWKS: server.URL + "/jwks"}
|
||||
mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, r *http.Request) { json.NewEncoder(w).Encode(d) })
|
||||
mux.HandleFunc("/authorize", func(w http.ResponseWriter, r *http.Request) {
|
||||
q := r.URL.Query()
|
||||
code := sessions.Create(&oidc.PKCESession{ClientID: q.Get("client_id"), Username: "test", Nonce: q.Get("nonce"), Scopes: strings.Fields(q.Get("scope")), PKCEChallenge: q.Get("code_challenge"), ExpiresAt: time.Now().Add(time.Minute)})
|
||||
target, _ := url.Parse(q.Get("redirect_uri"))
|
||||
params := target.Query()
|
||||
params.Set("state", q.Get("state"))
|
||||
params.Set("code", code)
|
||||
target.RawQuery = params.Encode()
|
||||
http.Redirect(w, r, target.String(), 302)
|
||||
})
|
||||
c, err := New(server.URL)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
c.HTTP.Transport = server.Client().Transport
|
||||
return c, d, h
|
||||
}
|
||||
|
||||
func TestServiceExchangeAndClaimValidation(t *testing.T) {
|
||||
c, _, h := provider(t)
|
||||
ctx := context.Background()
|
||||
d, err := c.Discover(ctx)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
form := url.Values{"grant_type": {"client_credentials"}, "scope": {"approval:read"}}
|
||||
token, err := c.Exchange(ctx, d, form, "service:consumer", "special+%: secret", "approval-engine", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = c.Verify(ctx, d, token.AccessToken, "other", ""); err == nil {
|
||||
t.Fatal("wrong audience accepted")
|
||||
}
|
||||
if _, err = c.Verify(ctx, d, token.AccessToken+"tampered", "approval-engine", ""); err == nil {
|
||||
t.Fatal("tampering accepted")
|
||||
}
|
||||
if _, err = c.Exchange(ctx, d, form, "service:consumer", "wrong", "approval-engine", ""); err == nil || strings.Contains(err.Error(), "special") {
|
||||
t.Fatal("wrong secret not safely rejected")
|
||||
}
|
||||
form.Set("scope", "approval:consume")
|
||||
if _, err = c.Exchange(ctx, d, form, "service:consumer", "special+%: secret", "approval-engine", ""); err == nil {
|
||||
t.Fatal("excess scope accepted")
|
||||
}
|
||||
form.Set("scope", "approval:read")
|
||||
h.TokenLifetime = -time.Minute
|
||||
if _, err = c.Exchange(ctx, d, form, "service:consumer", "special+%: secret", "approval-engine", ""); err == nil {
|
||||
t.Fatal("expired response accepted")
|
||||
}
|
||||
}
|
||||
|
||||
type urlWriter struct{ urls chan string }
|
||||
|
||||
func (w urlWriter) Write(p []byte) (int, error) {
|
||||
for _, line := range strings.Split(string(p), "\n") {
|
||||
if strings.HasPrefix(line, "https://") {
|
||||
w.urls <- line
|
||||
}
|
||||
}
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
func TestBrowserLoginPKCEAndState(t *testing.T) {
|
||||
c, d, _ := provider(t)
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
redirect := "http://" + listener.Addr().String() + "/callback"
|
||||
listener.Close()
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 10*time.Second)
|
||||
defer cancel()
|
||||
urls := make(chan string, 1)
|
||||
completed := make(chan error, 1)
|
||||
go func() {
|
||||
_, err := c.Login(ctx, d, "human", "approval-engine", "openid approval:approve", redirect, urlWriter{urls})
|
||||
completed <- err
|
||||
}()
|
||||
var address string
|
||||
select {
|
||||
case address = <-urls:
|
||||
case err := <-completed:
|
||||
t.Fatal(err)
|
||||
case <-ctx.Done():
|
||||
t.Fatal("no login URL")
|
||||
}
|
||||
// A forged callback must not consume the real login attempt.
|
||||
res, err := http.Get(redirect + "?state=forged&code=forged")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
res.Body.Close()
|
||||
if res.StatusCode != 400 {
|
||||
t.Fatal("forged state accepted")
|
||||
}
|
||||
browser := &http.Client{Transport: c.HTTP.Transport, Timeout: 5 * time.Second}
|
||||
res, err = browser.Get(address)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
io.Copy(io.Discard, res.Body)
|
||||
res.Body.Close()
|
||||
select {
|
||||
case err := <-completed:
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
case <-ctx.Done():
|
||||
t.Fatal("login did not finish")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginRejectsUnsafeCallbacks(t *testing.T) {
|
||||
c, d, _ := provider(t)
|
||||
for _, callback := range []string{"http://example.com:8000/callback", "http://localhost:8000/callback", "http://127.0.0.1:0/callback", "http://127.0.0.1:8000/callback?extra=yes"} {
|
||||
if _, err := c.Login(context.Background(), d, "human", "approval-engine", "openid", callback, io.Discard); err == nil {
|
||||
t.Fatalf("accepted %s", callback)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestOutputProtection(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
path := filepath.Join(dir, "token.json")
|
||||
file, err := reserveOutput(path)
|
||||
if err != nil {
|
||||
t.Fatalf("reserve %s: %v", path, err)
|
||||
}
|
||||
file.Close()
|
||||
info, _ := os.Stat(path)
|
||||
if info.Mode().Perm() != 0600 {
|
||||
t.Fatal("file not private")
|
||||
}
|
||||
if _, err = reserveOutput(path); err == nil {
|
||||
t.Fatal("overwrote existing file")
|
||||
}
|
||||
link := filepath.Join(dir, "link")
|
||||
if err = os.Symlink(path, link); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = reserveOutput(link); err == nil {
|
||||
t.Fatal("followed output symlink")
|
||||
}
|
||||
repo := filepath.Join(dir, "repo")
|
||||
os.Mkdir(repo, 0700)
|
||||
os.WriteFile(filepath.Join(repo, ".git"), []byte("gitdir: elsewhere"), 0600)
|
||||
if _, err = reserveOutput(filepath.Join(repo, "token")); err == nil {
|
||||
t.Fatal("allowed token in worktree")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDiscoveryAndRedirectBoundaries(t *testing.T) {
|
||||
for _, issuer := range []string{"http://example.com", "https://user:pass@example.com", "https://example.com?secret=value"} {
|
||||
if _, err := New(issuer); err == nil {
|
||||
t.Fatal("unsafe issuer accepted")
|
||||
}
|
||||
}
|
||||
mux := http.NewServeMux()
|
||||
server := httptest.NewTLSServer(mux)
|
||||
defer server.Close()
|
||||
c, _ := New(server.URL)
|
||||
c.HTTP.Transport = server.Client().Transport
|
||||
mux.HandleFunc("/.well-known/openid-configuration", func(w http.ResponseWriter, r *http.Request) {
|
||||
fmt.Fprintf(w, `{"issuer":%q,"authorization_endpoint":"https://evil.example/a","token_endpoint":"https://evil.example/t","jwks_uri":"https://evil.example/j"}`, server.URL)
|
||||
})
|
||||
if _, err := c.Discover(context.Background()); err == nil {
|
||||
t.Fatal("cross-origin discovery accepted")
|
||||
}
|
||||
mux.HandleFunc("/redirect", func(w http.ResponseWriter, r *http.Request) { http.Redirect(w, r, "https://evil.example", 307) })
|
||||
var out any
|
||||
if err := c.request(context.Background(), "POST", server.URL+"/redirect", nil, "id", "secret", &out); err == nil {
|
||||
t.Fatal("followed credential redirect")
|
||||
}
|
||||
}
|
||||
|
||||
func TestNonceMismatchAndCancelledLogin(t *testing.T) {
|
||||
c, d, _ := provider(t)
|
||||
form := url.Values{"grant_type": {"client_credentials"}, "scope": {"approval:read"}}
|
||||
token, err := c.Exchange(context.Background(), d, form, "service:consumer", "special+%: secret", "approval-engine", "")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err = c.Verify(context.Background(), d, token.AccessToken, "approval-engine", "required-nonce"); err == nil {
|
||||
t.Fatal("missing nonce accepted")
|
||||
}
|
||||
listener, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
redirect := "http://" + listener.Addr().String() + "/callback"
|
||||
listener.Close()
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
if _, err = c.Login(ctx, d, "human", "approval-engine", "openid", redirect, io.Discard); err == nil {
|
||||
t.Fatal("cancelled login succeeded")
|
||||
}
|
||||
listener, err = net.Listen("tcp", strings.TrimPrefix(strings.TrimSuffix(redirect, "/callback"), "http://"))
|
||||
if err != nil {
|
||||
t.Fatal("listener not released")
|
||||
}
|
||||
listener.Close()
|
||||
}
|
||||
|
|
@ -9,6 +9,7 @@ import (
|
|||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
|
|
@ -224,6 +225,12 @@ func (h *TokenHandler) serveClientCredentials(w http.ResponseWriter, r *http.Req
|
|||
Write(w, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
clientID, idErr := url.QueryUnescape(clientID)
|
||||
clientSecret, secretErr := url.QueryUnescape(clientSecret)
|
||||
if idErr != nil || secretErr != nil {
|
||||
profileerrors.InvalidProfileUsage("invalid client authentication encoding", "Authorization").Write(w, http.StatusUnauthorized)
|
||||
return
|
||||
}
|
||||
client, ok := h.ClientConfig[clientID]
|
||||
if !ok || client.ClientType != "confidential" || !containsString(client.GrantTypes, "client_credentials") {
|
||||
profileerrors.InvalidProfileUsage("invalid confidential client", "client_id").
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue