diff --git a/cmd/ate-setup/internal/config/config.go b/cmd/ate-setup/internal/config/config.go index 96b6af4130..99d828cd79 100644 --- a/cmd/ate-setup/internal/config/config.go +++ b/cmd/ate-setup/internal/config/config.go @@ -122,6 +122,10 @@ type Config struct { // follows neither form. ExpectedJWTIssuer string + // ActorJWTAlgorithm is the signing algorithm of the key in a new actor JWT + // pool (ACTOR_JWT_ALGORITHM): ES256 or RS256. + ActorJWTAlgorithm string + // BucketName is the snapshot bucket demos are templated with. BucketName string @@ -345,6 +349,7 @@ func Load(opts Options) (*Config, error) { ClusterName: env["CLUSTER_NAME"], ClusterLocation: env["CLUSTER_LOCATION"], ExpectedJWTIssuer: env["EXPECTED_JWT_ISSUER"], + ActorJWTAlgorithm: firstNonEmpty(env["ACTOR_JWT_ALGORITHM"], "ES256"), BucketName: env["BUCKET_NAME"], KODockerRepo: env["KO_DOCKER_REPO"], KODefaultPlatforms: env["KO_DEFAULTPLATFORMS"], @@ -449,6 +454,11 @@ func validate(cfg *Config) error { return fmt.Errorf("ATE_API_POSTGRES_CLOUDSQL_IP_TYPE must be %s, %s, or %s, got %q", CloudSQLIPTypePrivate, CloudSQLIPTypePublic, CloudSQLIPTypePSC, cfg.CloudSQL.IPType) } + switch cfg.ActorJWTAlgorithm { + case "ES256", "RS256": + default: + return fmt.Errorf("ACTOR_JWT_ALGORITHM must be ES256 or RS256, got %q", cfg.ActorJWTAlgorithm) + } switch cfg.ClusterSize { case ClusterSizeSize0, ClusterSizeSize10: default: diff --git a/cmd/ate-setup/internal/config/config_test.go b/cmd/ate-setup/internal/config/config_test.go index f70be4ce1a..7fca3b15e6 100644 --- a/cmd/ate-setup/internal/config/config_test.go +++ b/cmd/ate-setup/internal/config/config_test.go @@ -38,6 +38,7 @@ func loadEnv(t *testing.T) { t.Helper() t.Setenv("NO_DEV_ENV", "1") for _, name := range []string{ + "ACTOR_JWT_ALGORITHM", "ANTHROPIC_API_KEY", "ATE_ADDITIONAL_EGRESS_EXTPROC_SERVICE", "ATE_API_POSTGRES_CLOUDSQL_GSA", @@ -380,6 +381,38 @@ func TestLoadExpectedJWTIssuer(t *testing.T) { } } +func TestLoadActorJWTAlgorithm(t *testing.T) { + for _, tt := range []struct { + name string + env string + want string + wantErr bool + }{ + {name: "unset", env: "", want: "ES256"}, + {name: "RS256", env: "RS256", want: "RS256"}, + {name: "unsupported", env: "HS256", wantErr: true}, + } { + t.Run(tt.name, func(t *testing.T) { + loadEnv(t) + t.Setenv("ACTOR_JWT_ALGORITHM", tt.env) + + cfg, err := Load(Options{}) + if tt.wantErr { + if err == nil { + t.Fatalf("Load() with ACTOR_JWT_ALGORITHM=%q returned nil error", tt.env) + } + return + } + if err != nil { + t.Fatalf("Load() error = %v", err) + } + if cfg.ActorJWTAlgorithm != tt.want { + t.Errorf("ActorJWTAlgorithm = %q, want %q", cfg.ActorJWTAlgorithm, tt.want) + } + }) + } +} + // The endpoint has to reach both the Go steps and the shell scripts ate-setup // still delegates to, or the two halves of an install export different // collectors. diff --git a/cmd/ate-setup/internal/steps/create.go b/cmd/ate-setup/internal/steps/create.go index 4b69cfe191..f86e0e6780 100644 --- a/cmd/ate-setup/internal/steps/create.go +++ b/cmd/ate-setup/internal/steps/create.go @@ -40,8 +40,8 @@ const ( SecretPostgresServerCA = "postgres-server-ca" ConfigMapAPIEnvVars = "ate-api-server-envvars" ConfigMapAPIAuthn = "ate-api-authentication" - // poolKeyID is the identifier given to the first CA and JWT key in a new - // pool, matching the --ca-id/--key-id the shell scripts passed. + // poolKeyID is the identifier given to the first CA in a new pool, + // matching the --ca-id the shell scripts passed. poolKeyID = "1" ) @@ -259,18 +259,21 @@ func (e *Env) createJWTPool(ctx context.Context, namespace, name string) error { return nil } - authority, err := localjwtauthority.GenerateECDSAP256Authority(poolKeyID) + data, err := newJWTPoolSecretData(e.Cfg.ActorJWTAlgorithm) if err != nil { - return fmt.Errorf("while generating the JWT authority for %s/%s: %w", namespace, name, err) + return fmt.Errorf("while building the JWT pool for %s/%s: %w", namespace, name, err) } - poolBytes, err := localjwtauthority.Marshal(&localjwtauthority.ConcretePool{ - Authorities: []*localjwtauthority.Authority{authority}, - ActiveForSigning: poolKeyID, - }) + return e.createPoolSecret(ctx, namespace, name, corev1.SecretTypeOpaque, data) +} + +// newJWTPoolSecretData generates a pool with one active authority for +// algorithm, keyed by its thumbprint. +func newJWTPoolSecretData(algorithm string) (map[string][]byte, error) { + poolBytes, _, err := localjwtauthority.GeneratePool(algorithm, "") if err != nil { - return fmt.Errorf("while marshaling the JWT pool for %s/%s: %w", namespace, name, err) + return nil, fmt.Errorf("while generating the JWT pool: %w", err) } - return e.createPoolSecret(ctx, namespace, name, corev1.SecretTypeOpaque, map[string][]byte{"pool": poolBytes}) + return map[string][]byte{"pool": poolBytes}, nil } // createPoolSecret writes pool state. diff --git a/cmd/ate-setup/internal/steps/create_test.go b/cmd/ate-setup/internal/steps/create_test.go index ffab153750..32dcfad4c1 100644 --- a/cmd/ate-setup/internal/steps/create_test.go +++ b/cmd/ate-setup/internal/steps/create_test.go @@ -25,6 +25,8 @@ import ( corev1 "k8s.io/api/core/v1" "github.com/agent-substrate/substrate/internal/localca" + "github.com/agent-substrate/substrate/internal/localjwtauthority" + "github.com/agent-substrate/substrate/internal/oidcdiscovery" ) // ate-api-server requires both connection strings and its schema in the @@ -160,3 +162,36 @@ func TestNewCAPoolSecretData(t *testing.T) { }) } } + +func TestNewJWTPoolSecretData(t *testing.T) { + for _, alg := range []string{"ES256", "RS256"} { + t.Run(alg, func(t *testing.T) { + data, err := newJWTPoolSecretData(alg) + if err != nil { + t.Fatalf("newJWTPoolSecretData() error = %v", err) + } + if diff := cmp.Diff([]string{"pool"}, slices.Sorted(maps.Keys(data))); diff != "" { + t.Errorf("secret keys differ (-want +got):\n%s", diff) + } + + pool, err := localjwtauthority.Unmarshal(data["pool"]) + if err != nil { + t.Fatalf("Unmarshal() error = %v", err) + } + if len(pool.Authorities) != 1 { + t.Fatalf("pool has %d authorities, want 1", len(pool.Authorities)) + } + authority := pool.Authorities[0] + if authority.Algorithm != alg { + t.Errorf("Algorithm = %q, want %q", authority.Algorithm, alg) + } + thumbprint, err := oidcdiscovery.Thumbprint(authority.SigningKey.Public()) + if err != nil { + t.Fatalf("Thumbprint() error = %v", err) + } + if authority.ID != thumbprint || pool.ActiveForSigning != thumbprint { + t.Errorf("key ID %q, active %q; want both to be the thumbprint %q", authority.ID, pool.ActiveForSigning, thumbprint) + } + }) + } +} diff --git a/cmd/ateapi/internal/controlapi/functionaltest/common_test.go b/cmd/ateapi/internal/controlapi/functionaltest/common_test.go index 71d4b2b40f..45e66cef04 100644 --- a/cmd/ateapi/internal/controlapi/functionaltest/common_test.go +++ b/cmd/ateapi/internal/controlapi/functionaltest/common_test.go @@ -186,7 +186,7 @@ func setupTestWithVolumePlugins(t *testing.T, ns string, plugins map[string]volu } } - actorJWTAuthority, err := localjwtauthority.GenerateECDSAP256Authority("1") + actorJWTAuthority, err := localjwtauthority.GenerateAuthority("ES256", "1") if err != nil { t.Fatalf("Error generating actor JWT authority: %v", err) } diff --git a/cmd/kubectl-ate/README.md b/cmd/kubectl-ate/README.md index 6fa2150b5c..cde6d3649e 100644 --- a/cmd/kubectl-ate/README.md +++ b/cmd/kubectl-ate/README.md @@ -312,9 +312,10 @@ kubectl ate admin make-ca-pool \ --secret-namespace ate-system \ --ca-id "1" -# Generate a new JWT authority pool and push it to a Kubernetes Secret +# Generate a new JWT authority pool and push it to a Kubernetes Secret. The +# key is ES256 (--alg RS256 for relying parties that don't support ES256) and +# its ID defaults to the RFC 7638 thumbprint (--key-id to override). kubectl ate admin make-jwt-pool \ --name actor-id-jwt-pool \ - --secret-namespace ate-system \ - --key-id "1" + --secret-namespace ate-system ``` diff --git a/cmd/kubectl-ate/internal/cmd/admin.go b/cmd/kubectl-ate/internal/cmd/admin.go index b2fc6df38a..abe96139ce 100644 --- a/cmd/kubectl-ate/internal/cmd/admin.go +++ b/cmd/kubectl-ate/internal/cmd/admin.go @@ -34,6 +34,7 @@ var ( makeCaPoolIDFlag string makeCaPoolKeyTypeFlag string makeCaPoolValidityFlag time.Duration + makeJwtPoolAlgFlag string makeJwtPoolKeyIDFlag string ) @@ -134,29 +135,9 @@ var makeJwtPoolCmd = &cobra.Command{ return fmt.Errorf("while creating Kubernetes client: %w", err) } - authority, err := localjwtauthority.GenerateECDSAP256Authority(makeJwtPoolKeyIDFlag) + secret, keyID, err := newJWTPoolSecret(poolSecretNamespaceFlag, poolSecretNameFlag, makeJwtPoolAlgFlag, makeJwtPoolKeyIDFlag) if err != nil { - return fmt.Errorf("while generating JWT authority: %w", err) - } - - pool := &localjwtauthority.ConcretePool{ - Authorities: []*localjwtauthority.Authority{authority}, - ActiveForSigning: makeJwtPoolKeyIDFlag, - } - - poolBytes, err := localjwtauthority.Marshal(pool) - if err != nil { - return fmt.Errorf("while marshaling pool: %w", err) - } - - secret := &corev1.Secret{ - ObjectMeta: metav1.ObjectMeta{ - Namespace: poolSecretNamespaceFlag, - Name: poolSecretNameFlag, - }, - Data: map[string][]byte{ - "pool": poolBytes, - }, + return err } _, err = kc.CoreV1().Secrets(poolSecretNamespaceFlag).Create(ctx, secret, metav1.CreateOptions{}) @@ -164,11 +145,30 @@ var makeJwtPoolCmd = &cobra.Command{ return fmt.Errorf("while uploading pool state to secret: %w", err) } - fmt.Printf("Successfully created JWT authority pool secret %s/%s\n", poolSecretNamespaceFlag, poolSecretNameFlag) + fmt.Printf("Successfully created JWT authority pool secret %s/%s with %s key %s\n", poolSecretNamespaceFlag, poolSecretNameFlag, makeJwtPoolAlgFlag, keyID) return nil }, } +// newJWTPoolSecret builds a Secret holding a pool with one active authority, +// and returns the authority's key ID. +func newJWTPoolSecret(namespace, name, algorithm, keyID string) (*corev1.Secret, string, error) { + poolBytes, id, err := localjwtauthority.GeneratePool(algorithm, keyID) + if err != nil { + return nil, "", fmt.Errorf("while generating JWT authority pool: %w", err) + } + + return &corev1.Secret{ + ObjectMeta: metav1.ObjectMeta{ + Namespace: namespace, + Name: name, + }, + Data: map[string][]byte{ + "pool": poolBytes, + }, + }, id, nil +} + func init() { rootCmd.AddCommand(adminCmd) @@ -180,7 +180,8 @@ func init() { _ = makeCaPoolCmd.MarkFlagRequired("name") adminCmd.AddCommand(makeCaPoolCmd) - makeJwtPoolCmd.Flags().StringVar(&makeJwtPoolKeyIDFlag, "key-id", "1", "The ID of the initial JWT signing key in the pool") + makeJwtPoolCmd.Flags().StringVar(&makeJwtPoolAlgFlag, "alg", "ES256", "Signing algorithm of the initial key. One of [ES256, RS256]") + makeJwtPoolCmd.Flags().StringVar(&makeJwtPoolKeyIDFlag, "key-id", "", "The ID of the initial JWT signing key in the pool. Defaults to the key's RFC 7638 thumbprint") makeJwtPoolCmd.Flags().StringVar(&poolSecretNamespaceFlag, "secret-namespace", "default", "Create the secret in this namespace") makeJwtPoolCmd.Flags().StringVar(&poolSecretNameFlag, "name", "", "Create the secret with this name") _ = makeJwtPoolCmd.MarkFlagRequired("name") diff --git a/cmd/kubectl-ate/internal/cmd/admin_test.go b/cmd/kubectl-ate/internal/cmd/admin_test.go new file mode 100644 index 0000000000..656594aecb --- /dev/null +++ b/cmd/kubectl-ate/internal/cmd/admin_test.go @@ -0,0 +1,73 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package cmd + +import ( + "testing" + + "github.com/agent-substrate/substrate/internal/localjwtauthority" + "github.com/agent-substrate/substrate/internal/oidcdiscovery" +) + +func TestNewJWTPoolSecretWithFlagDefaults(t *testing.T) { + alg := makeJwtPoolCmd.Flags().Lookup("alg").DefValue + keyID := makeJwtPoolCmd.Flags().Lookup("key-id").DefValue + + secret, gotKeyID, err := newJWTPoolSecret("ate-system", "actor-id-jwt-pool", alg, keyID) + if err != nil { + t.Fatal(err) + } + if secret.Namespace != "ate-system" || secret.Name != "actor-id-jwt-pool" { + t.Errorf("secret = %s/%s, want ate-system/actor-id-jwt-pool", secret.Namespace, secret.Name) + } + pool, err := localjwtauthority.Unmarshal(secret.Data["pool"]) + if err != nil { + t.Fatal(err) + } + if len(pool.Authorities) != 1 { + t.Fatalf("pool has %d authorities, want 1", len(pool.Authorities)) + } + authority := pool.Authorities[0] + if authority.Algorithm != "ES256" { + t.Errorf("Algorithm = %q, want ES256", authority.Algorithm) + } + thumbprint, err := oidcdiscovery.Thumbprint(authority.SigningKey.Public()) + if err != nil { + t.Fatal(err) + } + if authority.ID != thumbprint || gotKeyID != thumbprint || pool.ActiveForSigning != thumbprint { + t.Errorf("key ID %q, returned %q, active %q; want all to be the thumbprint %q", authority.ID, gotKeyID, pool.ActiveForSigning, thumbprint) + } +} + +func TestNewJWTPoolSecretExplicitKey(t *testing.T) { + secret, keyID, err := newJWTPoolSecret("ate-system", "actor-id-jwt-pool", "RS256", "1") + if err != nil { + t.Fatal(err) + } + pool, err := localjwtauthority.Unmarshal(secret.Data["pool"]) + if err != nil { + t.Fatal(err) + } + if keyID != "1" || pool.ActiveForSigning != "1" || pool.Authorities[0].Algorithm != "RS256" { + t.Errorf("got key %q, active %q, algorithm %q; want 1, 1, RS256", keyID, pool.ActiveForSigning, pool.Authorities[0].Algorithm) + } +} + +func TestNewJWTPoolSecretRejectsUnsupportedAlgorithm(t *testing.T) { + if _, _, err := newJWTPoolSecret("ate-system", "actor-id-jwt-pool", "HS256", ""); err == nil { + t.Error("newJWTPoolSecret(HS256) returned nil error") + } +} diff --git a/go.mod b/go.mod index d5699f40dd..c64675509c 100644 --- a/go.mod +++ b/go.mod @@ -23,6 +23,7 @@ require ( github.com/envoyproxy/go-control-plane v0.14.0 github.com/envoyproxy/go-control-plane/envoy v1.37.1-0.20260812071801-353463cc7248 github.com/fsnotify/fsnotify v1.9.0 + github.com/go-jose/go-jose/v4 v4.1.4 github.com/go-logr/logr v1.4.4 github.com/google/go-cmp v0.7.0 github.com/google/go-containerregistry v0.21.7 @@ -137,7 +138,6 @@ require ( github.com/felixge/httpsnoop v1.1.0 // indirect github.com/fxamacker/cbor/v2 v2.9.1 // indirect github.com/go-errors/errors v1.4.2 // indirect - github.com/go-jose/go-jose/v4 v4.1.4 // indirect github.com/go-logr/stdr v1.2.2 // indirect github.com/go-ole/go-ole v1.3.0 // indirect github.com/go-openapi/jsonpointer v1.0.0 // indirect diff --git a/hack/install-ate.sh b/hack/install-ate.sh index a43b218b60..d9c2d94f12 100755 --- a/hack/install-ate.sh +++ b/hack/install-ate.sh @@ -152,6 +152,8 @@ usage() { echo " Default: derived from PROJECT_ID/CLUSTER_LOCATION/CLUSTER_NAME" echo " (https://container.googleapis.com/v1/projects/.../clusters/...)," echo " else the cluster's OIDC discovery document" + echo " ACTOR_JWT_ALGORITHM Signing algorithm of a newly created actor JWT pool: ES256 (default) | RS256." + echo " Use RS256 for relying parties that don't support ES256. Has no effect once the pool exists." echo "" echo "Benchmarks (see benchmarking/README.md for details and customization):" echo "" diff --git a/internal/localjwtauthority/localjwtauthority.go b/internal/localjwtauthority/localjwtauthority.go index a50f8dbd26..9c29dc134e 100644 --- a/internal/localjwtauthority/localjwtauthority.go +++ b/internal/localjwtauthority/localjwtauthority.go @@ -31,6 +31,7 @@ import ( "time" "github.com/agent-substrate/substrate/internal/actoridjwt" + "github.com/agent-substrate/substrate/internal/oidcdiscovery" ) // Pool is the interface for a JWT signing pool. @@ -167,8 +168,6 @@ func (p *ConcretePool) SignJWT(claims *actoridjwt.Claims) (string, error) { selectedAuthority = p.Authorities[0] } - // TODO(identity): The key IDs should probably be SHA256 of the key, to - // prevent user misuse. jwt, err := sign(payloadBytes, selectedAuthority.SigningKey, selectedAuthority.Algorithm, selectedAuthority.ID) if err != nil { return "", fmt.Errorf("while signing JWT: %w", err) @@ -338,16 +337,50 @@ func Unmarshal(wireBytes []byte) (*ConcretePool, error) { return pool, nil } -// GenerateECDSAP256Authority generates an ECDSA P256 JWT signing key. -func GenerateECDSAP256Authority(id string) (*Authority, error) { - privKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader) +// GenerateAuthority generates a JWT signing key for algorithm, which must be +// RS256 or ES256. An empty id defaults to the RFC 7638 thumbprint of the +// public key. +func GenerateAuthority(algorithm, id string) (*Authority, error) { + var key crypto.Signer + var err error + switch algorithm { + case "RS256": + key, err = rsa.GenerateKey(rand.Reader, 2048) + case "ES256": + key, err = ecdsa.GenerateKey(elliptic.P256(), rand.Reader) + default: + return nil, fmt.Errorf("unsupported algorithm %q, want RS256 or ES256", algorithm) + } if err != nil { return nil, fmt.Errorf("while generating key: %w", err) } - + if id == "" { + id, err = oidcdiscovery.Thumbprint(key.Public()) + if err != nil { + return nil, fmt.Errorf("while computing key thumbprint: %w", err) + } + } return &Authority{ ID: id, - Algorithm: "ES256", - SigningKey: privKey, + Algorithm: algorithm, + SigningKey: key, }, nil } + +// GeneratePool generates a pool holding one authority, active for signing, and +// returns the serialized pool and the authority's ID. algorithm and id are as +// for GenerateAuthority. +func GeneratePool(algorithm, id string) ([]byte, string, error) { + authority, err := GenerateAuthority(algorithm, id) + if err != nil { + return nil, "", err + } + wire, err := Marshal(&ConcretePool{ + Authorities: []*Authority{authority}, + ActiveForSigning: authority.ID, + }) + if err != nil { + return nil, "", err + } + return wire, authority.ID, nil +} diff --git a/internal/localjwtauthority/localjwtauthority_test.go b/internal/localjwtauthority/localjwtauthority_test.go index 31cac14760..2ff51f2f46 100644 --- a/internal/localjwtauthority/localjwtauthority_test.go +++ b/internal/localjwtauthority/localjwtauthority_test.go @@ -15,6 +15,9 @@ package localjwtauthority import ( + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rsa" "encoding/base64" "encoding/json" "os" @@ -27,10 +30,11 @@ import ( "github.com/google/go-cmp/cmp" "github.com/agent-substrate/substrate/internal/actoridjwt" + "github.com/agent-substrate/substrate/internal/oidcdiscovery" ) func TestRefreshingPool(t *testing.T) { - ca1, err := GenerateECDSAP256Authority("1") + ca1, err := GenerateAuthority("ES256", "1") if err != nil { t.Fatalf("Unexpected error generating CA 1: %v", err) } @@ -43,7 +47,7 @@ func TestRefreshingPool(t *testing.T) { t.Fatalf("Unexpected error marshaling pool 1: %v", err) } - ca2, err := GenerateECDSAP256Authority("2") + ca2, err := GenerateAuthority("ES256", "2") if err != nil { t.Fatalf("Unexpected error generating CA 2: %v", err) } @@ -154,7 +158,7 @@ func TestRefreshingPool(t *testing.T) { } func TestSignJWTHeader(t *testing.T) { - authority, err := GenerateECDSAP256Authority("key-1") + authority, err := GenerateAuthority("ES256", "key-1") if err != nil { t.Fatalf("Unexpected error generating authority: %v", err) } @@ -186,3 +190,107 @@ func TestSignJWTHeader(t *testing.T) { t.Errorf("Wrong JWT header; diff (-got +want)\n%s", diff) } } + +func TestGenerateAuthority(t *testing.T) { + for _, alg := range []string{"RS256", "ES256"} { + t.Run(alg, func(t *testing.T) { + authority, err := GenerateAuthority(alg, "") + if err != nil { + t.Fatal(err) + } + if authority.Algorithm != alg { + t.Errorf("Algorithm = %q, want %q", authority.Algorithm, alg) + } + thumbprint, err := oidcdiscovery.Thumbprint(authority.SigningKey.Public()) + if err != nil { + t.Fatal(err) + } + if authority.ID != thumbprint { + t.Errorf("ID = %q, want the key thumbprint %q", authority.ID, thumbprint) + } + switch key := authority.SigningKey.(type) { + case *rsa.PrivateKey: + if alg != "RS256" || key.N.BitLen() != 2048 { + t.Errorf("got a %d-bit RSA key for %s, want 2048-bit for RS256", key.N.BitLen(), alg) + } + case *ecdsa.PrivateKey: + if alg != "ES256" || key.Curve != elliptic.P256() { + t.Errorf("got an EC key on %s for %s, want P-256 for ES256", key.Curve.Params().Name, alg) + } + default: + t.Errorf("unexpected key type %T", key) + } + + pool := &ConcretePool{Authorities: []*Authority{authority}, ActiveForSigning: authority.ID} + poolBytes, err := Marshal(pool) + if err != nil { + t.Fatal(err) + } + loaded, err := Unmarshal(poolBytes) + if err != nil { + t.Fatal(err) + } + jwt, err := loaded.SignJWT(&actoridjwt.Claims{Subject: "actor/a/b", Audiences: []string{"aud"}}) + if err != nil { + t.Fatal(err) + } + headerB64, _, _ := strings.Cut(jwt, ".") + headerBytes, err := base64.RawURLEncoding.DecodeString(headerB64) + if err != nil { + t.Fatal(err) + } + var header map[string]string + if err := json.Unmarshal(headerBytes, &header); err != nil { + t.Fatal(err) + } + if header["alg"] != alg || header["kid"] != thumbprint { + t.Errorf("header = %v, want alg %s and kid %s", header, alg, thumbprint) + } + }) + } +} + +func TestGenerateAuthorityExplicitID(t *testing.T) { + authority, err := GenerateAuthority("RS256", "my-key") + if err != nil { + t.Fatal(err) + } + if authority.ID != "my-key" { + t.Errorf("ID = %q, want %q", authority.ID, "my-key") + } +} + +func TestGeneratePool(t *testing.T) { + for _, tc := range []struct{ alg, id string }{{"ES256", ""}, {"RS256", "my-key"}} { + wire, id, err := GeneratePool(tc.alg, tc.id) + if err != nil { + t.Fatalf("GeneratePool(%q, %q): %v", tc.alg, tc.id, err) + } + pool, err := Unmarshal(wire) + if err != nil { + t.Fatal(err) + } + if len(pool.Authorities) != 1 { + t.Fatalf("pool has %d authorities, want 1", len(pool.Authorities)) + } + authority := pool.Authorities[0] + if authority.Algorithm != tc.alg { + t.Errorf("Algorithm = %q, want %q", authority.Algorithm, tc.alg) + } + if tc.id != "" && id != tc.id { + t.Errorf("returned ID %q, want %q", id, tc.id) + } + if authority.ID != id || pool.ActiveForSigning != id { + t.Errorf("authority %q, active %q; want both to be the returned ID %q", authority.ID, pool.ActiveForSigning, id) + } + } + if _, _, err := GeneratePool("HS256", ""); err == nil { + t.Error("GeneratePool(HS256) returned nil error") + } +} + +func TestGenerateAuthorityRejectsUnsupportedAlgorithm(t *testing.T) { + if _, err := GenerateAuthority("HS256", ""); err == nil { + t.Error("GenerateAuthority(HS256) returned nil error") + } +} diff --git a/internal/oidcdiscovery/thumbprint.go b/internal/oidcdiscovery/thumbprint.go new file mode 100644 index 0000000000..43233d834d --- /dev/null +++ b/internal/oidcdiscovery/thumbprint.go @@ -0,0 +1,45 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package oidcdiscovery + +import ( + "crypto" + "crypto/ecdsa" + "crypto/elliptic" + "crypto/rsa" + "encoding/base64" + "fmt" + + jose "github.com/go-jose/go-jose/v4" +) + +// Thumbprint returns the RFC 7638 SHA-256 thumbprint of an RSA or P-256 EC +// public key, base64url-encoded without padding. +func Thumbprint(pub crypto.PublicKey) (string, error) { + switch k := pub.(type) { + case *rsa.PublicKey: + case *ecdsa.PublicKey: + if k.Curve != elliptic.P256() { + return "", fmt.Errorf("unsupported EC curve %s", k.Curve.Params().Name) + } + default: + return "", fmt.Errorf("unsupported public key type %T", pub) + } + sum, err := (&jose.JSONWebKey{Key: pub}).Thumbprint(crypto.SHA256) + if err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(sum), nil +} diff --git a/internal/oidcdiscovery/thumbprint_test.go b/internal/oidcdiscovery/thumbprint_test.go new file mode 100644 index 0000000000..e4f6bb0cfe --- /dev/null +++ b/internal/oidcdiscovery/thumbprint_test.go @@ -0,0 +1,100 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +package oidcdiscovery + +import ( + "crypto/ecdsa" + "crypto/ed25519" + "crypto/elliptic" + "crypto/rand" + "crypto/rsa" + "encoding/base64" + "math/big" + "testing" +) + +func mustDecode(t *testing.T, s string) []byte { + t.Helper() + b, err := base64.RawURLEncoding.DecodeString(s) + if err != nil { + t.Fatalf("decoding %q: %v", s, err) + } + return b +} + +// The example RSA key from RFC 7638 section 3.1. +const ( + rfcRSAN = "0vx7agoebGcQSuuPiLJXZptN9nndrQmbXEps2aiAFbWhM78LhWx4cbbfAAtVT86zwu1RK7aPFFxuhDR1L6tSoc_BJECPebWKRXjBZCiFV4n3oknjhMstn64tZ_2W-5JsGY4Hc5n9yBXArwl93lqt7_RN5w6Cf0h4QyQ5v-65YGjQR0_FDW2QvzqY368QQMicAtaSqzs8KJZgnYb9c7d0zgdAZHzu6qMQvRL5hajrn1n91CbOpbISD08qNLyrdkt-bFTWhAI4vMQFh6WeZu0fM4lFd2NcRwr3XPksINHaQ-G_xBniIqbw0Ls1jF44-csFCur-kEgU8awapJzKnqDKgw" + rfcRSAE = "AQAB" +) + +// The P-256 key from RFC 7517 appendix A.1. +const ( + rfcECX = "MKBCTNIcKUSDii11ySs3526iDZ8AiTo7Tu6KPAqv7D4" + rfcECY = "4Etl6SRW2YiLUrN5vfvVHuhp7x8PxltmWWlbbM4IFyM" +) + +func rfcRSAKey(t *testing.T) *rsa.PublicKey { + t.Helper() + return &rsa.PublicKey{N: new(big.Int).SetBytes(mustDecode(t, rfcRSAN)), E: 65537} +} + +func rfcECKey(t *testing.T) *ecdsa.PublicKey { + t.Helper() + point := append([]byte{0x04}, mustDecode(t, rfcECX)...) + pub, err := ecdsa.ParseUncompressedPublicKey(elliptic.P256(), append(point, mustDecode(t, rfcECY)...)) + if err != nil { + t.Fatal(err) + } + return pub +} + +func TestThumbprintRFC7638Example(t *testing.T) { + got, err := Thumbprint(rfcRSAKey(t)) + if err != nil { + t.Fatal(err) + } + if want := "NzbLsXh8uDCcd-6MNwXF4W_7noWXFZAfHkxZsRGC9Xs"; got != want { + t.Errorf("Thumbprint = %q, want %q", got, want) + } +} + +func TestThumbprintEC(t *testing.T) { + // RFC 7638 has no EC example; the expected value is go-jose's thumbprint + // of the RFC 7517 key. + got, err := Thumbprint(rfcECKey(t)) + if err != nil { + t.Fatal(err) + } + if want := "cn-I_WNMClehiVp51i_0VpOENW1upEerA8sEam5hn-s"; got != want { + t.Errorf("Thumbprint = %q, want %q", got, want) + } +} + +func TestThumbprintRejectsUnsupportedKeys(t *testing.T) { + p384, err := ecdsa.GenerateKey(elliptic.P384(), rand.Reader) + if err != nil { + t.Fatal(err) + } + edPub, _, err := ed25519.GenerateKey(rand.Reader) + if err != nil { + t.Fatal(err) + } + for _, pub := range []any{&p384.PublicKey, edPub} { + if got, err := Thumbprint(pub); err == nil { + t.Errorf("Thumbprint(%T) = %q, want error", pub, got) + } + } +}