keys, oidcIdP

This commit is contained in:
kkb0318
2024-03-26 21:25:25 +09:00
parent e33ea809f0
commit 022c561938
5 changed files with 195 additions and 56 deletions
+2 -2
View File
@@ -1,8 +1,8 @@
module github.com/kkb0318/irsa-manager
go 1.22
go 1.21
toolchain go1.22.1
toolchain go1.21.8
require (
github.com/go-jose/go-jose/v4 v4.0.1
+82
View File
@@ -0,0 +1,82 @@
package selfhosted
import (
"crypto"
"crypto/rsa"
"crypto/x509"
"encoding/base64"
"errors"
"fmt"
jose "github.com/go-jose/go-jose/v4"
"k8s.io/client-go/util/keyutil"
)
// keyIDFromPublicKey derives a key ID non-reversibly from a public key.
//
// The Key ID is field on a given on JWTs and JWKs that help relying parties
// pick the correct key for verification when the identity party advertises
// multiple keys.
//
// Making the derivation non-reversible makes it impossible for someone to
// accidentally obtain the real key from the key ID and use it for token
// validation.
// This method is copied from
// https://github.com/kubernetes/kubernetes/blob/v1.29.3/pkg/serviceaccount/jwt.go#L99
func keyIDFromPublicKey(publicKey interface{}) (string, error) {
publicKeyDERBytes, err := x509.MarshalPKIXPublicKey(publicKey)
if err != nil {
return "", fmt.Errorf("failed to serialize public key to DER format: %v", err)
}
hasher := crypto.SHA256.New()
hasher.Write(publicKeyDERBytes)
publicKeyDERHash := hasher.Sum(nil)
keyID := base64.RawURLEncoding.EncodeToString(publicKeyDERHash)
return keyID, nil
}
type JWK struct {
Keys []jose.JSONWebKey `json:"keys"`
}
func NewJWK(pub []byte) (*JWK, error) {
pubKeys, err := keyutil.ParsePublicKeysPEM(pub)
if err != nil {
return nil, err
}
pubKey := pubKeys[0]
var alg jose.SignatureAlgorithm
switch pubKey.(type) {
case *rsa.PublicKey:
alg = jose.RS256
default:
return nil, errors.New("public key is not RSA")
}
kid, err := keyIDFromPublicKey(pubKey)
if err != nil {
return nil, err
}
var keys []jose.JSONWebKey
keys = append(keys, jose.JSONWebKey{
Key: pubKey,
KeyID: kid,
Algorithm: string(alg),
Use: "sig",
})
keys = append(keys, jose.JSONWebKey{
Key: pubKey,
KeyID: "",
Algorithm: string(alg),
Use: "sig",
})
return &JWK{Keys: keys}, nil
}
func (j *JWK) GetKeys() []jose.JSONWebKey {
return j.Keys
}
+43
View File
@@ -0,0 +1,43 @@
package selfhosted
import (
"os"
"testing"
"github.com/stretchr/testify/assert"
)
const rsaKeyID = "JHJehTTTZlsspKHT-GaJxK7Kd1NQgZJu3fyK6K_QDYU"
func TestJWK(t *testing.T) {
tests := []struct {
name string
filename string
expected string
expectErr bool
}{
{
name: "rsa",
filename: "testdata/rsa.pub",
expected: rsaKeyID,
},
{
name: "rsa",
filename: "testdata/ecdsa.pub",
expectErr: true,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
content, err := os.ReadFile(tt.filename)
assert.NoError(t, err)
actual, err := NewJWK(content)
if tt.expectErr {
assert.Error(t, err, "")
} else {
assert.NoError(t, err)
assert.Equal(t, tt.expected, actual.Keys[0].KeyID)
}
})
}
}
+33 -53
View File
@@ -1,72 +1,52 @@
package selfhosted
import (
"crypto"
"crypto/rand"
"crypto/rsa"
"crypto/x509"
"encoding/base64"
"errors"
"fmt"
jose "github.com/go-jose/go-jose/v4"
"k8s.io/client-go/util/keyutil"
"encoding/pem"
"os"
)
// keyIDFromPublicKey derives a key ID non-reversibly from a public key.
//
// The Key ID is field on a given on JWTs and JWKs that help relying parties
// pick the correct key for verification when the identity party advertises
// multiple keys.
//
// Making the derivation non-reversible makes it impossible for someone to
// accidentally obtain the real key from the key ID and use it for token
// validation.
// This method is copied from
// https://github.com/kubernetes/kubernetes/blob/v1.29.3/pkg/serviceaccount/jwt.go#L99
func keyIDFromPublicKey(publicKey interface{}) (string, error) {
publicKeyDERBytes, err := x509.MarshalPKIXPublicKey(publicKey)
func createKeyPair() error {
// RSAキーペアの生成
privateKey, err := rsa.GenerateKey(rand.Reader, 2048)
if err != nil {
return "", fmt.Errorf("failed to serialize public key to DER format: %v", err)
return err
}
hasher := crypto.SHA256.New()
hasher.Write(publicKeyDERBytes)
publicKeyDERHash := hasher.Sum(nil)
keyID := base64.RawURLEncoding.EncodeToString(publicKeyDERHash)
return keyID, nil
// private keyをPEM形式で保存
privPem := pem.Block{
Type: "RSA PRIVATE KEY",
Bytes: x509.MarshalPKCS1PrivateKey(privateKey),
}
type JWK struct {
Keys []jose.JSONWebKey `json:"keys"`
}
func NewJWK(pub []byte) (*JWK, error){
pubKeys, err := keyutil.ParsePublicKeysPEM(pub)
privPemFile, err := os.Create("private_key.pem")
if err != nil {
return nil, err
return err
}
pubKey := pubKeys[0]
var alg jose.SignatureAlgorithm
switch pubKey.(type) {
case *rsa.PublicKey:
alg = jose.RS256
default:
return nil, errors.New("public key is not RSA")
defer privPemFile.Close()
if err := pem.Encode(privPemFile, &privPem); err != nil {
return err
}
kid, err := keyIDFromPublicKey(pubKey)
// 公開鍵をPKIX, ASN.1 DER形式に変換
pubASN1, err := x509.MarshalPKIXPublicKey(&privateKey.PublicKey)
if err != nil {
return nil, err
return err
}
var keys []jose.JSONWebKey
keys = append(keys, jose.JSONWebKey{
Key: pubKey,
KeyID: kid,
Algorithm: string(alg),
Use: "sig",
})
return &JWK{Keys: keys}, nil
// 公開鍵をPEM形式で保存
pubPem := pem.Block{
Type: "PUBLIC KEY",
Bytes: pubASN1,
}
pubPemFile, err := os.Create("public_key.pem")
if err != nil {
return err
}
defer pubPemFile.Close()
if err := pem.Encode(pubPemFile, &pubPem); err != nil {
return err
}
return nil
}
+34
View File
@@ -1 +1,35 @@
package selfhosted
import "fmt"
type OIDCIdProvider interface {
Discovery() string
JWK() string
Endpoint() string
}
type OIDCIdPCreator interface {
Upload(o OIDCIdProvider) error
CreateProvider() error
}
type S3IdPCreator struct {
region string
bucketName string
}
func (s *S3IdPCreator) CreateProvider() error {
// create S3 Bucket // bucketName, Region
// create OIDCProvider //
return nil
}
func (s *S3IdPCreator) Upload(o OIDCIdProvider) error {
o.Discovery() // Upload to Endpoint()/.well-known/openid-configuration
o.JWK() // Upload to Endpoint()/keys.json
return nil
}
func (s *S3IdPCreator) issuerHostPath() string {
hostName := fmt.Sprintf("s3-%s.amazonaws.com", s.region)
return fmt.Sprintf("%s/%s", hostName, s.bucketName)
}