diff --git a/adapter/provider/provider.go b/adapter/provider/provider.go index 2cb159b6..b92f6c3f 100644 --- a/adapter/provider/provider.go +++ b/adapter/provider/provider.go @@ -366,20 +366,15 @@ func NewProxiesParser(pdName string, tunnel C.Tunnel, filter string, excludeFilt filterRegs = append(filterRegs, filterReg) } - var identities []age.Identity - if ageSecretKey != "" { - var err error - identities, err = age.ParseIdentities(ageSecretKey) - if err != nil { - return nil, fmt.Errorf("parse age-secret-key error: %w", err) - } + if err := age.VeritySecretKeys(ageSecretKey); err != nil { + return nil, fmt.Errorf("invalid age-secret-key: %w", err) } return func(buf []byte) ([]C.Proxy, error) { schema := &ProxySchema{} // decrypt config - buf, err := age.DecryptBytes(buf, identities...) + buf, err := age.DecryptBytes(buf, ageSecretKey) if err != nil { return nil, fmt.Errorf("decrypt config error: %w", err) } diff --git a/component/age/age.go b/component/age/age.go index c16ec510..60334dbb 100644 --- a/component/age/age.go +++ b/component/age/age.go @@ -14,30 +14,95 @@ import ( const FileHeader = armor.Header -type Identity = age.Identity -type Recipient = age.Recipient +var globalSecretKeys []string -var globalIdentities []Identity - -func ParseIdentities(secretKey string) ([]Identity, error) { +// parseIdentities parse age-secret-key to age.Identity +func parseIdentities(secretKey string) ([]age.Identity, error) { return age.ParseIdentities(strings.NewReader(secretKey)) } -func ParseRecipients(publicKey string) ([]Recipient, error) { +// parseRecipients parse age-public-key to age.Recipient +func parseRecipients(publicKey string) ([]age.Recipient, error) { return age.ParseRecipients(strings.NewReader(publicKey)) } -func SetGlobalIdentities(id []Identity) { - globalIdentities = append(globalIdentities[:0], id...) +// convertToRecipient convert age.Identity to age.Recipient +func convertToRecipient(identity age.Identity) (age.Recipient, error) { + switch identity := identity.(type) { + case *age.X25519Identity: + return identity.Recipient(), nil + case *age.HybridIdentity: + return identity.Recipient(), nil + default: + return nil, fmt.Errorf("unexpected identity type: %T", identity) + } +} + +// ToPublicKeys convert age-secret-key to age-public-key +func ToPublicKeys(secretKeys ...string) (publicKeys []string, err error) { + for _, secretKey := range secretKeys { + identities, err := parseIdentities(secretKey) + if err != nil { + return nil, err + } + for _, identity := range identities { + recipient, err := convertToRecipient(identity) + if err != nil { + return nil, err + } + publicKeys = append(publicKeys, fmt.Sprint(recipient)) + } + } + return +} + +// SetGlobalSecretKeys set global secret keys, which will be used when decrypting +func SetGlobalSecretKeys(secretKeys ...string) { + globalSecretKeys = append(globalSecretKeys[:0], secretKeys...) +} + +// VeritySecretKeys check if the secret key is valid +func VeritySecretKeys(secretKeys ...string) error { + for _, secretKey := range secretKeys { + if _, err := parseIdentities(secretKey); err != nil { + return err + } + } + return nil +} + +// VerityPublicKeys check if the public key is valid +func VerityPublicKeys(publicKeys ...string) error { + for _, publicKey := range publicKeys { + if _, err := parseRecipients(publicKey); err != nil { + return err + } + } + return nil } // DecryptBytes decrypt age armor format encrypted data // if not the age armor format, return original data -func DecryptBytes(data []byte, identities ...Identity) ([]byte, error) { +func DecryptBytes(data []byte, secretKeys ...string) ([]byte, error) { if !strings.HasPrefix(string(data), FileHeader) { // not age armor format return data, nil } - identities = append(identities[:len(identities):len(identities)], globalIdentities...) + var identities []age.Identity + for _, secretKey := range secretKeys { + identity, err := parseIdentities(secretKey) + if err != nil { + return nil, err + } + identities = append(identities, identity...) + } + for _, secretKey := range globalSecretKeys { + identity, err := parseIdentities(secretKey) + if err != nil { + return nil, err + } + identities = append(identities, identity...) + } + r, err := age.Decrypt(armor.NewReader(bytes.NewReader(data)), identities...) if err != nil { return nil, err @@ -46,7 +111,15 @@ func DecryptBytes(data []byte, identities ...Identity) ([]byte, error) { } // EncryptBytes encrypt data with age armor format -func EncryptBytes(data []byte, recipients ...Recipient) ([]byte, error) { +func EncryptBytes(data []byte, publicKeys ...string) ([]byte, error) { + var recipients []age.Recipient + for _, publicKey := range publicKeys { + recipient, err := parseRecipients(publicKey) + if err != nil { + return nil, err + } + recipients = append(recipients, recipient...) + } buf := &bytes.Buffer{} armorWriter := armor.NewWriter(buf) w, err := age.Encrypt(armorWriter, recipients...) @@ -68,19 +141,8 @@ func EncryptBytes(data []byte, recipients ...Recipient) ([]byte, error) { return buf.Bytes(), nil } -// ConvertToRecipient convert age.Identity to age.Recipient -func ConvertToRecipient(identity Identity) (Recipient, error) { - switch identity := identity.(type) { - case *age.X25519Identity: - return identity.Recipient(), nil - case *age.HybridIdentity: - return identity.Recipient(), nil - default: - return nil, fmt.Errorf("unexpected identity type: %T", identity) - } -} - -func GenX25519KeyPair() (string, string, error) { +// GenX25519KeyPair generate x25519 recipient type age-secret-key and age-public-key +func GenX25519KeyPair() (secretKey string, publicKey string, err error) { identity, err := age.GenerateX25519Identity() if err != nil { return "", "", err @@ -88,7 +150,8 @@ func GenX25519KeyPair() (string, string, error) { return identity.String(), identity.Recipient().String(), nil } -func GenHybridKeyPair() (string, string, error) { +// GenHybridKeyPair generate mlkem768-x25519 hybrid post-quantum recipient type age-secret-key and age-public-key +func GenHybridKeyPair() (secretKey string, publicKey string, err error) { identity, err := age.GenerateHybridIdentity() if err != nil { return "", "", err @@ -121,29 +184,22 @@ func Main(args []string) { if len(args) < 1 { panic("Using: age convert ") } - identities, err := ParseIdentities(args[1]) + publicKeys, err := ToPublicKeys(args[1]) if err != nil { panic(err) } - if len(identities) == 0 { - panic("no identities found in the input") + if len(publicKeys) == 0 { + panic("no public keys found in the input") } - for _, identity := range identities { - recipient, err := ConvertToRecipient(identity) - if err != nil { - panic(err) - } - fmt.Println(recipient) + for _, publicKey := range publicKeys { + fmt.Println(publicKey) } case "decrypt": if len(args) < 3 { panic("Using: age decrypt ") } - identities, err := ParseIdentities(args[1]) - if err != nil { - panic(err) - } var data []byte + var err error if args[2] == "-" { data, err = io.ReadAll(os.Stdin) } else { @@ -152,7 +208,7 @@ func Main(args []string) { if err != nil { panic(err) } - result, err := DecryptBytes(data, identities...) + result, err := DecryptBytes(data, args[1]) if err != nil { panic(err) } @@ -168,11 +224,8 @@ func Main(args []string) { if len(args) < 3 { panic("Using: age encrypt ") } - recipients, err := ParseRecipients(args[1]) - if err != nil { - panic(err) - } var data []byte + var err error if args[2] == "-" { data, err = io.ReadAll(os.Stdin) } else { @@ -181,7 +234,7 @@ func Main(args []string) { if err != nil { panic(err) } - result, err := EncryptBytes(data, recipients...) + result, err := EncryptBytes(data, args[1]) if err != nil { panic(err) } diff --git a/component/age/age_test.go b/component/age/age_test.go index 54ece94a..b71b4558 100644 --- a/component/age/age_test.go +++ b/component/age/age_test.go @@ -1,7 +1,6 @@ package age_test import ( - "fmt" "testing" "github.com/metacubex/mihomo/component/age" @@ -23,33 +22,23 @@ func TestAge(t *testing.T) { t.Fatal(err) } t.Log(secretKey, publicKey) - identities, err := age.ParseIdentities(secretKey) + publicKeys, err := age.ToPublicKeys(secretKey) if err != nil { t.Fatal(err) } - recipients, err := age.ParseRecipients(publicKey) - if err != nil { - t.Fatal(err) + if len(publicKeys) != 1 { + t.Fatal("public keys length is not equal to 1") } - if len(identities) != len(recipients) { - t.Fatal("identities and recipients are not equal") - } - for i, identity := range identities { - recipient, err := age.ConvertToRecipient(identity) - if err != nil { - t.Fatal(err) - } - if fmt.Sprint(recipient) != fmt.Sprint(recipients[i]) { - t.Fatal("recipient is not equal to recipients: ", recipient, " != ", recipients[i], "") - } + if publicKeys[0] != publicKey { + t.Fatal("public key is not equal") } rawData := []byte("hello world") - encryptData, err := age.EncryptBytes(rawData, recipients...) + encryptData, err := age.EncryptBytes(rawData, publicKey) if err != nil { t.Fatal(err) } t.Log(string(encryptData)) - decryptData, err := age.DecryptBytes(encryptData, identities...) + decryptData, err := age.DecryptBytes(encryptData, secretKey) if err != nil { t.Fatal(err) } @@ -58,5 +47,4 @@ func TestAge(t *testing.T) { } }) } - } diff --git a/main.go b/main.go index 8ef7471c..1ad9127c 100644 --- a/main.go +++ b/main.go @@ -125,11 +125,10 @@ func main() { } if ageSecretKey != "" { - identities, err := age.ParseIdentities(ageSecretKey) - if err != nil { + if err := age.VeritySecretKeys(ageSecretKey); err != nil { log.Errorln("Parse age-secret-key error: %s", err.Error()) } - age.SetGlobalIdentities(identities) + age.SetGlobalSecretKeys(ageSecretKey) } if configString != "" {