package keymanager import ( "crypto/sha256" "encoding/pem" "errors" "fmt" "io" "os" "git.tswf.io/infra/go-synapse-backupper/pkg/domain/crypto" ) var ( // ErrPEMTypeMismatch is returned when the PEM block type does not match // the expected type for the given schemeID. ErrPEMTypeMismatch = errors.New("PEM type does not match scheme") // ErrInvalidPEM is returned when the file does not contain a valid PEM block. ErrInvalidPEM = errors.New("invalid PEM data") ) var pubPEMTypes = map[uint16]string{ 0x0006: "ML-KEM-768 PUBLIC KEY", 0x0007: "X25519 PUBLIC KEY", } var privPEMTypes = map[uint16]string{ 0x0006: "ML-KEM-768 PRIVATE KEY", 0x0007: "X25519 PRIVATE KEY", } // keyManager handles PEM encoding and decoding of recipient keys. type keyManager struct { registry crypto.Registry } // NewKeyManager creates a new KeyManager backed by the provided Registry. func NewKeyManager( registry crypto.Registry, ) crypto.KeyManager { return &keyManager{ registry: registry, } } // Generate creates a new key pair for the given schemeID and writes them // as PEM blocks to pubOut and privOut. func (k *keyManager) Generate( schemeID uint16, pubOut io.Writer, privOut io.Writer, rand io.Reader, ) error { factory, err := k.registry.Lookup(schemeID) if err != nil { return err } kem := factory() pub, priv, err := kem.GenerateKeyPair(rand) if err != nil { return err } pubType, ok := pubPEMTypes[schemeID] if !ok { return fmt.Errorf( "unsupported scheme 0x%04x for public key PEM", schemeID, ) } privType, ok := privPEMTypes[schemeID] if !ok { return fmt.Errorf( "unsupported scheme 0x%04x for private key PEM", schemeID, ) } pubBlock := &pem.Block{ Type: pubType, Bytes: pub.Raw(), } if err := pem.Encode(pubOut, pubBlock); err != nil { return fmt.Errorf("encode public key PEM: %w", err) } privBlock := &pem.Block{ Type: privType, Bytes: priv.Raw(), } if err := pem.Encode(privOut, privBlock); err != nil { return fmt.Errorf("encode private key PEM: %w", err) } return nil } // LoadPub reads a PEM-encoded public key from path and validates that its // type matches the expected type for schemeID. func (k *keyManager) LoadPub( path string, schemeID uint16, ) ( crypto.RecipientPub, error, ) { data, err := os.ReadFile(path) if err != nil { return nil, fmt.Errorf("read public key file: %w", err) } block, _ := pem.Decode(data) if block == nil { return nil, fmt.Errorf("%w: no valid PEM block found", ErrInvalidPEM) } expectedType, ok := pubPEMTypes[schemeID] if !ok { return nil, fmt.Errorf( "unsupported scheme 0x%04x for public key PEM", schemeID, ) } if block.Type != expectedType { return nil, fmt.Errorf( "expected PEM type %q, got %q: %w", expectedType, block.Type, ErrPEMTypeMismatch, ) } return newRecipientPub(schemeID, block.Bytes), nil } // LoadPriv reads a PEM-encoded private key from path and validates that its // type matches the expected type for schemeID. func (k *keyManager) LoadPriv( path string, schemeID uint16, ) ( crypto.RecipientPriv, error, ) { data, err := os.ReadFile(path) if err != nil { return nil, fmt.Errorf("read private key file: %w", err) } block, _ := pem.Decode(data) if block == nil { return nil, fmt.Errorf("%w: no valid PEM block found", ErrInvalidPEM) } expectedType, ok := privPEMTypes[schemeID] if !ok { return nil, fmt.Errorf( "unsupported scheme 0x%04x for private key PEM", schemeID, ) } if block.Type != expectedType { return nil, fmt.Errorf( "expected PEM type %q, got %q: %w", expectedType, block.Type, ErrPEMTypeMismatch, ) } factory, err := k.registry.Lookup(schemeID) if err != nil { return nil, err } kem := factory() return kem.LoadPriv(block.Bytes) } // recipientPub is a generic RecipientPub implementation backed by raw bytes. type recipientPub struct { schemeID uint16 raw []byte keyID []byte } func newRecipientPub( schemeID uint16, raw []byte, ) crypto.RecipientPub { var keyID []byte switch schemeID { case 0x0006: h := sha256.Sum256(raw[:8]) keyID = h[:8] case 0x0007: h := sha256.Sum256(raw) keyID = h[:8] } return &recipientPub{ schemeID: schemeID, raw: append([]byte(nil), raw...), keyID: keyID, } } func (r *recipientPub) SchemeID() uint16 { return r.schemeID } func (r *recipientPub) KeyID() []byte { return r.keyID } func (r *recipientPub) Raw() []byte { return r.raw }