diff --git a/core/core.go b/core/core.go index 1aac20e..a8b16da 100644 --- a/core/core.go +++ b/core/core.go @@ -15,6 +15,12 @@ import ( "github.com/cloudflare/redoctober/passvault" ) +var ( + crypt cryptor.Cryptor + records passvault.Records + cache keycache.Cache +) + // Each of these structures corresponds to the JSON expected on the // correspondingly named URI (e.g. the delegate structure maps to the // JSON that should be sent on the /delegate URI and it is handled by @@ -51,11 +57,9 @@ type EncryptRequest struct { Name string Password string - Minimum int - Owners []string - LeftOwners []string - RightOwners []string - Data []byte + Minimum int + Owners []string + Data []byte Labels []string } @@ -102,7 +106,7 @@ func jsonStatusError(err error) ([]byte, error) { return json.Marshal(ResponseData{Status: err.Error()}) } func jsonSummary() ([]byte, error) { - return json.Marshal(SummaryData{Status: "ok", Live: keycache.GetSummary(), All: passvault.GetSummary()}) + return json.Marshal(SummaryData{Status: "ok", Live: cache.GetSummary(), All: records.GetSummary()}) } func jsonResponse(resp []byte) ([]byte, error) { return json.Marshal(ResponseData{Status: "ok", Response: resp}) @@ -111,11 +115,11 @@ func jsonResponse(resp []byte) ([]byte, error) { // validateAdmin checks that the username and password passed in are // correct and that the user is an admin func validateAdmin(name, password string) error { - if passvault.NumRecords() == 0 { + if records.NumRecords() == 0 { return errors.New("Vault is not created yet") } - pr, ok := passvault.GetRecord(name) + pr, ok := records.GetRecord(name) if !ok { return errors.New("User not present") } @@ -144,9 +148,13 @@ func validateUser(name, password string) error { // Init reads the records from disk from a given path func Init(path string) (err error) { - if err = passvault.InitFromDisk(path); err != nil { + if records, err = passvault.InitFrom(path); err != nil { err = fmt.Errorf("Failed to load password vault %s: %s", path, err) } + + cache = keycache.Cache{make(map[string]keycache.ActiveUser)} + crypt = cryptor.New(&records, &cache) + return } @@ -157,7 +165,7 @@ func Create(jsonIn []byte) ([]byte, error) { return jsonStatusError(err) } - if passvault.NumRecords() != 0 { + if records.NumRecords() != 0 { return jsonStatusError(errors.New("Vault is already created")) } @@ -166,7 +174,7 @@ func Create(jsonIn []byte) ([]byte, error) { return jsonStatusError(err) } - if _, err := passvault.AddNewRecord(s.Name, s.Password, true); err != nil { + if _, err := records.AddNewRecord(s.Name, s.Password, true, passvault.DefaultRecordType); err != nil { log.Printf("Error adding record for %s: %s\n", s.Name, err) return jsonStatusError(err) } @@ -177,13 +185,13 @@ func Create(jsonIn []byte) ([]byte, error) { // Summary processes a summary request. func Summary(jsonIn []byte) ([]byte, error) { var s SummaryRequest - keycache.Refresh() + cache.Refresh() if err := json.Unmarshal(jsonIn, &s); err != nil { return jsonStatusError(err) } - if passvault.NumRecords() == 0 { + if records.NumRecords() == 0 { return jsonStatusError(errors.New("Vault is not created yet")) } @@ -202,7 +210,7 @@ func Delegate(jsonIn []byte) ([]byte, error) { return jsonStatusError(err) } - if passvault.NumRecords() == 0 { + if records.NumRecords() == 0 { return jsonStatusError(errors.New("Vault is not created yet")) } @@ -214,21 +222,21 @@ func Delegate(jsonIn []byte) ([]byte, error) { // Find password record for user and verify that their password // matches. If not found then add a new entry for this user. - pr, found := passvault.GetRecord(s.Name) + pr, found := records.GetRecord(s.Name) if found { if err := pr.ValidatePassword(s.Password); err != nil { return jsonStatusError(err) } } else { var err error - if pr, err = passvault.AddNewRecord(s.Name, s.Password, false); err != nil { + if pr, err = records.AddNewRecord(s.Name, s.Password, false, passvault.DefaultRecordType); err != nil { log.Printf("Error adding record for %s: %s\n", s.Name, err) return jsonStatusError(err) } } // add signed-in record to active set - if err := keycache.AddKeyFromRecord(pr, s.Name, s.Password, s.Users, s.Labels, s.Uses, s.Time); err != nil { + if err := cache.AddKeyFromRecord(pr, s.Name, s.Password, s.Users, s.Labels, s.Uses, s.Time); err != nil { log.Printf("Error adding key to cache for %s: %s\n", s.Name, err) return jsonStatusError(err) } @@ -243,12 +251,12 @@ func Password(jsonIn []byte) ([]byte, error) { return jsonStatusError(err) } - if passvault.NumRecords() == 0 { + if records.NumRecords() == 0 { return jsonStatusError(errors.New("Vault is not created yet")) } // add signed-in record to active set - if err := passvault.ChangePassword(s.Name, s.Password, s.NewPassword); err != nil { + if err := records.ChangePassword(s.Name, s.Password, s.NewPassword); err != nil { log.Println("Error changing password:", err) return jsonStatusError(err) } @@ -268,13 +276,8 @@ func Encrypt(jsonIn []byte) ([]byte, error) { return jsonStatusError(err) } - if len(s.Owners) > 0 { - s.LeftOwners = s.Owners - s.RightOwners = s.Owners - } - // Encrypt file with list of owners - if resp, err := cryptor.Encrypt(s.Data, s.Labels, s.LeftOwners, s.RightOwners, s.Minimum); err != nil { + if resp, err := crypt.Encrypt(s.Data, s.Labels, s.Owners, s.Minimum); err != nil { log.Println("Error encrypting:", err) return jsonStatusError(err) } else { @@ -296,7 +299,7 @@ func Decrypt(jsonIn []byte) ([]byte, error) { return jsonStatusError(err) } - data, names, err := cryptor.Decrypt(s.Data, s.Name) + data, names, err := crypt.Decrypt(s.Data, s.Name) if err != nil { log.Println("Error decrypting:", err) return jsonStatusError(err) @@ -328,7 +331,7 @@ func Modify(jsonIn []byte) ([]byte, error) { return jsonStatusError(err) } - if _, ok := passvault.GetRecord(s.ToModify); !ok { + if _, ok := records.GetRecord(s.ToModify); !ok { return jsonStatusError(errors.New("Record to modify missing")) } @@ -339,11 +342,11 @@ func Modify(jsonIn []byte) ([]byte, error) { var err error switch s.Command { case "delete": - err = passvault.DeleteRecord(s.ToModify) + err = records.DeleteRecord(s.ToModify) case "revoke": - err = passvault.RevokeRecord(s.ToModify) + err = records.RevokeRecord(s.ToModify) case "admin": - err = passvault.MakeAdmin(s.ToModify) + err = records.MakeAdmin(s.ToModify) default: return jsonStatusError(errors.New("Unknown command")) } diff --git a/core/core_test.go b/core/core_test.go index 4acadcc..04e0f8d 100644 --- a/core/core_test.go +++ b/core/core_test.go @@ -10,7 +10,6 @@ import ( "os" "testing" - "github.com/cloudflare/redoctober/keycache" "github.com/cloudflare/redoctober/passvault" ) @@ -156,7 +155,7 @@ func TestSummary(t *testing.T) { dataLive, ok := s.Live["Bob"] if !ok { - t.Fatalf("Error in summary of account, record missing, %v", keycache.UserKeys) + t.Fatalf("Error in summary of account, record missing, %v", cache.UserKeys) } if dataLive.Admin != false { t.Fatalf("Error in summary of account, record missing") @@ -165,7 +164,7 @@ func TestSummary(t *testing.T) { t.Fatalf("Error in summary of account, record missing") } - keycache.FlushCache() + cache.FlushCache() os.Remove("/tmp/db1.json") } @@ -278,7 +277,7 @@ func TestPassword(t *testing.T) { t.Fatalf("Error in delegating account, %v", s.Status) } - keycache.FlushCache() + cache.FlushCache() os.Remove("/tmp/db1.json") } @@ -335,7 +334,7 @@ func TestEncryptDecrypt(t *testing.T) { } // check summary to see if none are delegated - keycache.Refresh() + cache.Refresh() respJson, err = Summary(summaryJson) if err != nil { t.Fatalf("Error in summary, %v", err) @@ -422,7 +421,7 @@ func TestEncryptDecrypt(t *testing.T) { } // verify the presence of the two delgations - keycache.Refresh() + cache.Refresh() var sum2 SummaryData respJson, err = Summary(summaryJson) if err != nil { @@ -467,7 +466,7 @@ func TestEncryptDecrypt(t *testing.T) { } } - keycache.FlushCache() + cache.FlushCache() os.Remove("/tmp/db1.json") } @@ -526,7 +525,7 @@ func TestModify(t *testing.T) { } // check summary to see if none are delegated - keycache.Refresh() + cache.Refresh() respJson, err = Summary(summaryJson) if err != nil { t.Fatalf("Error in summary, %v", err) @@ -654,7 +653,7 @@ func TestModify(t *testing.T) { t.Fatalf("Error in summary, %v", sum3.All) } - keycache.FlushCache() + cache.FlushCache() os.Remove("/tmp/db1.json") } @@ -731,7 +730,7 @@ func TestStatic(t *testing.T) { t.Fatalf("Error in summary, %v, %v", expected, r.Response) } - keycache.FlushCache() + cache.FlushCache() os.Remove("/tmp/db1.json") } diff --git a/cryptor/cryptor.go b/cryptor/cryptor.go index e7faa92..6978865 100644 --- a/cryptor/cryptor.go +++ b/cryptor/cryptor.go @@ -25,6 +25,15 @@ const ( DEFAULT_VERSION = 1 ) +type Cryptor struct { + records *passvault.Records + cache *keycache.Cache +} + +func New(records *passvault.Records, cache *keycache.Cache) Cryptor { + return Cryptor{records, cache} +} + // MultiWrappedKey is a structure containing a 16-byte key encrypted // once for each of the keys corresponding to the names of the users // in Name in order. @@ -53,174 +62,47 @@ type EncryptedData struct { Signature []byte } -// encryptKey encrypts data with the key associated with name inner, -// then name outer -func encryptKey(nameInner, nameOuter string, clearKey []byte, pubKeys map[string]SingleWrappedKey) (out MultiWrappedKey, err error) { - out.Name = []string{nameOuter, nameInner} - - recInner, ok := passvault.GetRecord(nameInner) - if !ok { - err = errors.New("Missing user on disk") - return - } - - recOuter, ok := passvault.GetRecord(nameOuter) - if !ok { - err = errors.New("Missing user on disk") - return - } - - if recInner.Type != recOuter.Type { - err = errors.New("Mismatched record types") - return - } - - var keyBytes []byte - var overrideInner SingleWrappedKey - var overrideOuter SingleWrappedKey - - // For AES records, use the live user key - // For RSA and ECC records, use the public key from the passvault - switch recInner.Type { - case passvault.RSARecord, passvault.ECCRecord: - if overrideInner, ok = pubKeys[nameInner]; !ok { - err = errors.New("Missing user in file") - return - } - - if overrideOuter, ok = pubKeys[nameOuter]; !ok { - err = errors.New("Missing user in file") - return - } - - default: - return out, errors.New("Unknown record type inner") - } - - // double-wrap the keys - if keyBytes, err = keycache.EncryptKey(clearKey, nameInner, overrideInner.aesKey); err != nil { - return out, err - } - if keyBytes, err = keycache.EncryptKey(keyBytes, nameOuter, overrideOuter.aesKey); err != nil { - return out, err - } - - out.Key = keyBytes - - return +type pair struct { + name string + key []byte } -// unwrapKey decrypts first key in keys whose encryption keys are in keycache -func unwrapKey(keys []MultiWrappedKey, pubKeys map[string]SingleWrappedKey, user string, labels []string) (unwrappedKey []byte, names []string, err error) { - var ( - keyFound error - fullMatch bool = false - nameSet = map[string]bool{} - ) +type mwkSlice []MultiWrappedKey +type swkSlice []pair - for _, mwKey := range keys { - if err != nil { - return nil, nil, err - } - - tmpKeyValue := mwKey.Key - - for _, mwName := range mwKey.Name { - pubEncrypted := pubKeys[mwName] - // if this is null, it's an AES encrypted key - if tmpKeyValue, keyFound = keycache.DecryptKey(tmpKeyValue, mwName, user, labels, pubEncrypted.Key); keyFound != nil { - break - } - nameSet[mwName] = true - } - if keyFound == nil { - fullMatch = true - // concatenate all the decrypted bytes - unwrappedKey = tmpKeyValue - break - } - } - - if !fullMatch { - err = errors.New("Need more delegated keys") - names = nil - } - - names = make([]string, 0, len(nameSet)) - for name := range nameSet { - names = append(names, name) - } - return -} - -// mwkSorter describes a slice of MultiWrappedKeys to be sorted. -type mwkSorter struct { - keySet []MultiWrappedKey -} - -// Len is part of sort.Interface. -func (s *mwkSorter) Len() int { - return len(s.keySet) -} - -// Swap is part of sort.Interface. -func (s *mwkSorter) Swap(i, j int) { - s.keySet[i], s.keySet[j] = s.keySet[j], s.keySet[i] -} - -// Less is part of sort.Interface, it sorts lexicographically -// based on the list of names -func (s *mwkSorter) Less(i, j int) bool { +func (s mwkSlice) Len() int { return len(s) } +func (s mwkSlice) Swap(i, j int) { s[i], s[j] = s[j], s[i] } +func (s mwkSlice) Less(i, j int) bool { // Alphabetic order var shorter = i - if len(s.keySet[i].Name) > len(s.keySet[j].Name) { + if len(s[i].Name) > len(s[j].Name) { shorter = j } - for index := range s.keySet[shorter].Name { - if s.keySet[i].Name[index] != s.keySet[j].Name[index] { - return s.keySet[i].Name[index] < s.keySet[j].Name[index] + + for index := range s[shorter].Name { + if s[i].Name[index] != s[j].Name[index] { + return s[i].Name[index] < s[j].Name[index] } } return false } -// swkSorter joins a slice of names with SingleWrappedKeys to be sorted. -type pair struct { - name string - key []byte -} - -type swkSorter []pair - -// Len is part of sort.Interface. -func (s swkSorter) Len() int { - return len(s) -} - -// Swap is part of sort.Interface. -func (s swkSorter) Swap(i, j int) { - s[i], s[j] = s[j], s[i] -} - -// Less is part of sort.Interface. -func (s swkSorter) Less(i, j int) bool { - return s[i].name < s[j].name -} +func (s swkSlice) Len() int { return len(s) } +func (s swkSlice) Swap(i, j int) { s[i], s[j] = s[j], s[i] } +func (s swkSlice) Less(i, j int) bool { return s[i].name < s[j].name } // computeHmac computes the signature of the encrypted data structure // the signature takes into account every element of the EncryptedData // structure, with all keys sorted alphabetically by name -func computeHmac(key []byte, encrypted EncryptedData) []byte { +func (encrypted *EncryptedData) computeHmac(key []byte) []byte { mac := hmac.New(sha1.New, key) // sort the multi-wrapped keys - mwks := &mwkSorter{ - keySet: encrypted.KeySet, - } + mwks := mwkSlice(encrypted.KeySet) sort.Sort(mwks) // sort the singly-wrapped keys - var swks swkSorter + var swks swkSlice for name, val := range encrypted.KeySetRSA { swks = append(swks, pair{name, val.Key}) } @@ -259,97 +141,141 @@ func computeHmac(key []byte, encrypted EncryptedData) []byte { return mac.Sum(nil) } +// wrapKey encrypts the clear key such that a minimum number of delegated keys +// are required to decrypt. NOTE: Currently the max value for min is 2. +func (encrypted *EncryptedData) wrapKey(records *passvault.Records, clearKey []byte, names []string, min int) (err error) { + // Generate a random AES key for each user and RSA/ECIES encrypt it + encrypted.KeySetRSA = make(map[string]SingleWrappedKey, len(names)) + + for _, name := range names { + rec, ok := records.GetRecord(name) + if !ok { + err = errors.New("Missing user on disk") + return + } + + var singleWrappedKey SingleWrappedKey + + if singleWrappedKey.aesKey, err = symcrypt.MakeRandom(16); err != nil { + return err + } + + if singleWrappedKey.Key, err = rec.EncryptKey(singleWrappedKey.aesKey); err != nil { + return err + } + + encrypted.KeySetRSA[name] = singleWrappedKey + } + + // encrypt file key with every combination of two keys + encrypted.KeySet = make([]MultiWrappedKey, 0) + + for i := 0; i < len(names); i++ { + for j := i + 1; j < len(names); j++ { + var outerCrypt, innerCrypt cipher.Block + keyBytes := make([]byte, 16) + + outerCrypt, err = aes.NewCipher(encrypted.KeySetRSA[names[i]].aesKey) + if err != nil { + return + } + + innerCrypt, err = aes.NewCipher(encrypted.KeySetRSA[names[j]].aesKey) + if err != nil { + return + } + + innerCrypt.Encrypt(keyBytes, clearKey) + outerCrypt.Encrypt(keyBytes, keyBytes) + + out := MultiWrappedKey{ + Name: []string{names[i], names[j]}, + Key: keyBytes, + } + + encrypted.KeySet = append(encrypted.KeySet, out) + } + } + + return nil +} + +// unwrapKey decrypts first key in keys whose encryption keys are in keycache +func (encrypted *EncryptedData) unwrapKey(cache *keycache.Cache, user string) (unwrappedKey []byte, names []string, err error) { + var ( + keyFound error + fullMatch bool = false + nameSet = map[string]bool{} + ) + + for _, mwKey := range encrypted.KeySet { + // validate the size of the keys + if len(mwKey.Key) != 16 { + err = errors.New("Invalid Input") + } + + if err != nil { + return nil, nil, err + } + + tmpKeyValue := mwKey.Key + + for _, mwName := range mwKey.Name { + pubEncrypted := encrypted.KeySetRSA[mwName] + // if this is null, it's an AES encrypted key + if tmpKeyValue, keyFound = cache.DecryptKey(tmpKeyValue, mwName, user, encrypted.Labels, pubEncrypted.Key); keyFound != nil { + break + } + nameSet[mwName] = true + } + if keyFound == nil { + fullMatch = true + // concatenate all the decrypted bytes + unwrappedKey = tmpKeyValue + break + } + } + + if !fullMatch { + err = errors.New("Need more delegated keys") + names = nil + } + + names = make([]string, 0, len(nameSet)) + for name := range nameSet { + names = append(names, name) + } + return +} + // Encrypt encrypts data with the keys associated with names. This // requires a minimum of min keys to decrypt. NOTE: as currently // implemented, the maximum value for min is 2. -func Encrypt(in []byte, labels, leftNames, rightNames []string, min int) (resp []byte, err error) { +func (c *Cryptor) Encrypt(in []byte, labels, names []string, min int) (resp []byte, err error) { if min > 2 { return nil, errors.New("Minimum restricted to 2") } var encrypted EncryptedData encrypted.Version = DEFAULT_VERSION - if encrypted.VaultId, err = passvault.GetVaultId(); err != nil { + if encrypted.VaultId, err = c.records.GetVaultId(); err != nil { return } // Generate random IV and encryption key - ivBytes, err := symcrypt.MakeRandom(16) + encrypted.IV, err = symcrypt.MakeRandom(16) if err != nil { return } - // append used here to make a new slice from ivBytes and assign to - // encrypted.IV - - encrypted.IV = append([]byte{}, ivBytes...) clearKey, err := symcrypt.MakeRandom(16) if err != nil { return } - var names = make(map[string]bool) - var overlap int - - // Count overlapping names, we don't want to double-encrypt - // with the same name - for _, n := range leftNames { - names[n] = true - } - for _, n := range rightNames { - used, ok := names[n] - if !used && ok { - names[n] = true - } else { - overlap++ - } - } - - // Allocate set of keys to be able to cover all unequal pairs of - // names with one from leftNames and one from rightNames - - // Combinatorially, the number of ordered pairs with one element - // from one set and one from another for which both elements of - // the pair is distinct is - // len(n) * len(k) - overlap - encrypted.KeySet = make([]MultiWrappedKey, (len(leftNames)*len(rightNames) - overlap)) - encrypted.KeySetRSA = make(map[string]SingleWrappedKey) - - var singleWrappedKey SingleWrappedKey - for name := range names { - rec, ok := passvault.GetRecord(name) - if !ok { - err = errors.New("Missing user on disk") - return - } - - if rec.GetType() == passvault.RSARecord || rec.GetType() == passvault.ECCRecord { - // only wrap key with RSA key if found - if singleWrappedKey.aesKey, err = symcrypt.MakeRandom(16); err != nil { - return nil, err - } - - if singleWrappedKey.Key, err = rec.EncryptKey(singleWrappedKey.aesKey); err != nil { - return nil, err - } - encrypted.KeySetRSA[name] = singleWrappedKey - } else { - err = nil - } - } - - // encrypt file key with every combination of two keys - var n int - for _, nameOuter := range leftNames { - for _, nameInner := range rightNames { - if nameInner != nameOuter { - encrypted.KeySet[n], err = encryptKey(nameInner, nameOuter, clearKey, encrypted.KeySetRSA) - n += 1 - } - if err != nil { - return - } - } + err = encrypted.wrapKey(c.records, clearKey, names, min) + if err != nil { + return } // encrypt file with clear key @@ -361,23 +287,23 @@ func Encrypt(in []byte, labels, leftNames, rightNames []string, min int) (resp [ clearFile := padding.AddPadding(in) encryptedFile := make([]byte, len(clearFile)) - aesCBC := cipher.NewCBCEncrypter(aesCrypt, ivBytes) + aesCBC := cipher.NewCBCEncrypter(aesCrypt, encrypted.IV) aesCBC.CryptBlocks(encryptedFile, clearFile) encrypted.Data = encryptedFile encrypted.Labels = labels - hmacKey, err := passvault.GetHmacKey() + hmacKey, err := c.records.GetHmacKey() if err != nil { return } - encrypted.Signature = computeHmac(hmacKey, encrypted) + encrypted.Signature = encrypted.computeHmac(hmacKey) return json.Marshal(encrypted) } // Decrypt decrypts a file using the keys in the key cache. -func Decrypt(in []byte, user string) (resp []byte, names []string, err error) { +func (c *Cryptor) Decrypt(in []byte, user string) (resp []byte, names []string, err error) { // unwrap encrypted file var encrypted EncryptedData if err = json.Unmarshal(in, &encrypted); err != nil { @@ -388,7 +314,7 @@ func Decrypt(in []byte, user string) (resp []byte, names []string, err error) { } // make sure file was encrypted with the active vault - vaultId, err := passvault.GetVaultId() + vaultId, err := c.records.GetVaultId() if err != nil { return } @@ -396,20 +322,12 @@ func Decrypt(in []byte, user string) (resp []byte, names []string, err error) { return nil, nil, errors.New("Wrong vault") } - // validate the size of the keys - for _, multiKey := range encrypted.KeySet { - if len(multiKey.Key) != 16 { - err = errors.New("Invalid Input") - return - } - } - // compute HMAC - hmacKey, err := passvault.GetHmacKey() + hmacKey, err := c.records.GetHmacKey() if err != nil { return } - expectedMAC := computeHmac(hmacKey, encrypted) + expectedMAC := encrypted.computeHmac(hmacKey) if !hmac.Equal(encrypted.Signature, expectedMAC) { err = errors.New("Signature mismatch") return @@ -417,7 +335,7 @@ func Decrypt(in []byte, user string) (resp []byte, names []string, err error) { // decrypt file key with delegate keys var unwrappedKey = make([]byte, 16) - unwrappedKey, names, err = unwrapKey(encrypted.KeySet, encrypted.KeySetRSA, user, encrypted.Labels) + unwrappedKey, names, err = encrypted.unwrapKey(c.cache, user) if err != nil { return } diff --git a/cryptor/cryptor_test.go b/cryptor/cryptor_test.go index d3d2d3a..011c036 100644 --- a/cryptor/cryptor_test.go +++ b/cryptor/cryptor_test.go @@ -22,7 +22,7 @@ func TestHash(t *testing.T) { var hmacKey, _ = base64.StdEncoding.DecodeString("Qugc5ZQ0vC7KQSgmDHTVgQ==") var signature = append([]byte{}, encrypted.Signature...) - expectedSig := computeHmac(hmacKey, encrypted) + expectedSig := encrypted.computeHmac(hmacKey) if diff := bytes.Compare(signature, expectedSig); diff != 0 { t.Fatalf("Error comparing signature %v", base64.StdEncoding.EncodeToString(expectedSig)) @@ -30,7 +30,7 @@ func TestHash(t *testing.T) { // change version and check hmac encrypted.Version = 2 - unexpectedSig := computeHmac(hmacKey, encrypted) + unexpectedSig := encrypted.computeHmac(hmacKey) if diff := bytes.Compare(signature, unexpectedSig); diff == 0 { t.Fatalf("Error comparing signature") @@ -39,7 +39,7 @@ func TestHash(t *testing.T) { // change vaultid and check hmac encrypted.VaultId = 529853896 - unexpectedSig = computeHmac(hmacKey, encrypted) + unexpectedSig = encrypted.computeHmac(hmacKey) if diff := bytes.Compare(signature, unexpectedSig); diff == 0 { t.Fatalf("Error comparing signature") @@ -48,7 +48,7 @@ func TestHash(t *testing.T) { // swap two records and check hmac encrypted.KeySet[0], encrypted.KeySet[1] = encrypted.KeySet[1], encrypted.KeySet[0] - unexpectedSig = computeHmac(hmacKey, encrypted) + unexpectedSig = encrypted.computeHmac(hmacKey) if diff := bytes.Compare(signature, unexpectedSig); diff != 0 { t.Fatalf("Error comparing signature %v, %v", @@ -59,7 +59,7 @@ func TestHash(t *testing.T) { // delete RSA key and check hmac encrypted.Version = 1 delete(encrypted.KeySetRSA, "Carol") - unexpectedSig = computeHmac(hmacKey, encrypted) + unexpectedSig = encrypted.computeHmac(hmacKey) if diff := bytes.Compare(signature, unexpectedSig); diff == 0 { t.Fatalf("Error comparing signature") diff --git a/keycache/keycache.go b/keycache/keycache.go index 496cc3d..2d1886a 100644 --- a/keycache/keycache.go +++ b/keycache/keycache.go @@ -19,9 +19,6 @@ import ( "github.com/cloudflare/redoctober/passvault" ) -// UserKeys is the set of decrypted keys in memory, indexed by name. -var UserKeys map[string]ActiveUser = make(map[string]ActiveUser) - // Usage holds the permissions of a delegated permission type Usage struct { Uses int // Number of uses delegated @@ -40,24 +37,8 @@ type ActiveUser struct { eccKey *ecdsa.PrivateKey } -// matchUser returns the matching active user if present -// and a boolean to indicate its presence. -func matchUser(name, user string, labels []string) (out ActiveUser, present bool) { - key, present := UserKeys[name] - if present { - if key.Usage.matches(user, labels) { - return key, true - } else { - present = false - } - } - - return -} - -// setUser takes an ActiveUser and adds it to the cache. -func setUser(in ActiveUser, name string) { - UserKeys[name] = in +type Cache struct { + UserKeys map[string]ActiveUser // Decrypted keys in memory, indexed by name. } // matchesLabel returns true if this usage applies the user and label @@ -96,42 +77,66 @@ func (usage Usage) matches(user string, labels []string) bool { return false } +func NewCache() Cache { + return Cache{make(map[string]ActiveUser)} +} + +// setUser takes an ActiveUser and adds it to the cache. +func (cache *Cache) setUser(in ActiveUser, name string) { + cache.UserKeys[name] = in +} + +// matchUser returns the matching active user if present +// and a boolean to indicate its presence. +func (cache *Cache) matchUser(name, user string, labels []string) (out ActiveUser, present bool) { + key, present := cache.UserKeys[name] + if present { + if key.Usage.matches(user, labels) { + return key, true + } else { + present = false + } + } + + return +} + // useKey decrements the counter on an active key // for decryption or symmetric encryption -func useKey(name, user string, labels []string) { - if val, present := matchUser(name, user, labels); present { +func (cache *Cache) useKey(name, user string, labels []string) { + if val, present := cache.matchUser(name, user, labels); present { val.Usage.Uses -= 1 - setUser(val, name) + cache.setUser(val, name) } } // GetSummary returns the list of active user keys. -func GetSummary() map[string]ActiveUser { - return UserKeys +func (cache *Cache) GetSummary() map[string]ActiveUser { + return cache.UserKeys } // FlushCache removes all delegated keys. -func FlushCache() { - for name := range UserKeys { - delete(UserKeys, name) +func (cache *Cache) FlushCache() { + for name := range cache.UserKeys { + delete(cache.UserKeys, name) } } // Refresh purges all expired or used up keys. -func Refresh() { - for name, active := range UserKeys { +func (cache *Cache) Refresh() { + for name, active := range cache.UserKeys { if active.Usage.Expiry.Before(time.Now()) || active.Usage.Uses <= 0 { log.Println("Record expired", name, active.Usage.Users, active.Usage.Labels, active.Usage.Expiry) - delete(UserKeys, name) + delete(cache.UserKeys, name) } } } // AddKeyFromRecord decrypts a key for a given record and adds it to the cache. -func AddKeyFromRecord(record passvault.PasswordRecord, name, password string, users, labels []string, uses int, durationString string) (err error) { +func (cache *Cache) AddKeyFromRecord(record passvault.PasswordRecord, name, password string, users, labels []string, uses int, durationString string) (err error) { var current ActiveUser - Refresh() + cache.Refresh() // compute exipiration duration, err := time.ParseDuration(durationString) @@ -162,22 +167,7 @@ func AddKeyFromRecord(record passvault.PasswordRecord, name, password string, us current.Admin = record.Admin // add current to map (overwriting previous for this name) - setUser(current, name) - - return -} - -// EncryptKey encrypts a 16 byte key using the cached key corresponding to name. -func EncryptKey(in []byte, name string, aesKey []byte) (out []byte, err error) { - Refresh() - - // encrypt - aesSession, err := aes.NewCipher(aesKey) - if err != nil { - return - } - out = make([]byte, 16) - aesSession.Encrypt(out, in) + cache.setUser(current, name) return } @@ -186,10 +176,10 @@ func EncryptKey(in []byte, name string, aesKey []byte) (out []byte, err error) { // For RSA and EC keys, the cached RSA/EC key is used to decrypt // the pubEncryptedKey which is then used to decrypt the input // buffer. -func DecryptKey(in []byte, name, user string, labels []string, pubEncryptedKey []byte) (out []byte, err error) { - Refresh() +func (cache *Cache) DecryptKey(in []byte, name, user string, labels []string, pubEncryptedKey []byte) (out []byte, err error) { + cache.Refresh() - decryptKey, ok := matchUser(name, user, labels) + decryptKey, ok := cache.matchUser(name, user, labels) if !ok { return nil, errors.New("Key not delegated") } @@ -223,7 +213,7 @@ func DecryptKey(in []byte, name, user string, labels []string, pubEncryptedKey [ out = make([]byte, 16) aesSession.Decrypt(out, in) - useKey(name, user, labels) + cache.useKey(name, user, labels) return } diff --git a/keycache/keycache_test.go b/keycache/keycache_test.go index 1a9ad7c..e70f640 100644 --- a/keycache/keycache_test.go +++ b/keycache/keycache_test.go @@ -4,230 +4,290 @@ package keycache import ( + "bytes" "github.com/cloudflare/redoctober/passvault" + "github.com/cloudflare/redoctober/symcrypt" "testing" "time" ) -var now = time.Now() -var nextYear = now.AddDate(1, 0, 0) -var emptyKey = make([]byte, 16) -var dummy = make([]byte, 16) - func TestUsesFlush(t *testing.T) { - singleUse := ActiveUser{ - Admin: true, - Type: passvault.AESRecord, - Usage: Usage{ - Expiry: nextYear, - Uses: 2, - }, - aesKey: emptyKey, + // Initialize passvault with one dummy user. + records, err := passvault.InitFrom("memory") + if err != nil { + t.Fatalf("%v", err) } - UserKeys["first"] = singleUse + pr, err := records.AddNewRecord("user", "weakpassword", true, passvault.DefaultRecordType) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 1 { + // Initialize keycache and delegate the user's key to it. + cache := NewCache() + + err = cache.AddKeyFromRecord(pr, "user", "weakpassword", nil, nil, 2, "1h") + if err != nil { + t.Fatalf("%v", err) + } + + cache.Refresh() + if len(cache.UserKeys) != 1 { t.Fatalf("Error in number of live keys") } - EncryptKey(dummy, "first", nil) - - Refresh() - if len(UserKeys) != 1 { - t.Fatalf("Error in number of live keys %v", UserKeys) + // Generate a random symmetric key, encrypt a blank block with it, and encrypt + // the key itself with the user's public key. + dummy := make([]byte, 16) + key, err := symcrypt.MakeRandom(16) + if err != nil { + t.Fatalf("%v", err) } - DecryptKey(dummy, "first", "", []string{}, nil) + encKey, err := symcrypt.EncryptCBC(dummy, dummy, key) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 0 { - t.Fatalf("Error in number of live keys %v", UserKeys) + pubEncryptedKey, err := pr.EncryptKey(key) + if err != nil { + t.Fatalf("%v", err) + } + + key2, err := cache.DecryptKey(encKey, "user", "anybody", []string{}, pubEncryptedKey) + if err != nil { + t.Fatalf("%v", err) + } + + if bytes.Equal(key, key2) { + t.Fatalf("cache.DecryptKey didnt decrypt the right key!") + } + + // Second decryption allowed. + _, err = cache.DecryptKey(encKey, "user", "anybody else", []string{}, pubEncryptedKey) + if err != nil { + t.Fatalf("%v", err) + } + + cache.Refresh() + if len(cache.UserKeys) != 0 { + t.Fatalf("Error in number of live keys %v", cache.UserKeys) } } func TestTimeFlush(t *testing.T) { - oneSec, _ := time.ParseDuration("1s") - one := now.Add(oneSec) - - singleUse := ActiveUser{ - Admin: true, - Type: passvault.AESRecord, - Usage: Usage{ - Expiry: one, - Uses: 10, - }, - aesKey: emptyKey, + // Initialize passvault and keycache. Delegate a key for 1s, wait a + // second and then make sure that it's gone. + records, err := passvault.InitFrom("memory") + if err != nil { + t.Fatalf("%v", err) } - UserKeys["first"] = singleUse + pr, err := records.AddNewRecord("user", "weakpassword", true, passvault.DefaultRecordType) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 1 { + cache := NewCache() + + err = cache.AddKeyFromRecord(pr, "user", "weakpassword", nil, nil, 10, "1s") + if err != nil { + t.Fatalf("%v", err) + } + + cache.Refresh() + if len(cache.UserKeys) != 1 { t.Fatalf("Error in number of live keys") } - EncryptKey(dummy, "first", nil) + time.Sleep(time.Second) - Refresh() - if len(UserKeys) != 1 { - t.Fatalf("Error in number of live keys") + dummy := make([]byte, 16) + pubEncryptedKey, err := pr.EncryptKey(dummy) + if err != nil { + t.Fatalf("%v", err) } - time.Sleep(oneSec) - - _, err := DecryptKey(dummy, "first", "", []string{}, nil) - + _, err = cache.DecryptKey(dummy, "user", "anybody", []string{}, pubEncryptedKey) if err == nil { t.Fatalf("Error in pruning expired key") } } func TestGoodLabel(t *testing.T) { - singleUse := ActiveUser{ - Admin: true, - Type: passvault.AESRecord, - Usage: Usage{ - Expiry: nextYear, - Uses: 2, - Labels: []string{"red"}, - }, - aesKey: emptyKey, + // Initialize passvault and keycache. Delegate a key with the tag "red" and + // verify that decryption with the tag "red" is allowed. + records, err := passvault.InitFrom("memory") + if err != nil { + t.Fatalf("%v", err) } - UserKeys["first"] = singleUse + pr, err := records.AddNewRecord("user", "weakpassword", true, passvault.DefaultRecordType) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 1 { + cache := NewCache() + + err = cache.AddKeyFromRecord(pr, "user", "weakpassword", nil, []string{"red"}, 1, "1h") + if err != nil { + t.Fatalf("%v", err) + } + + cache.Refresh() + if len(cache.UserKeys) != 1 { t.Fatalf("Error in number of live keys") } - EncryptKey(dummy, "first", nil) - - Refresh() - if len(UserKeys) != 1 { - t.Fatalf("Error in number of live keys") + dummy := make([]byte, 16) + pubEncryptedKey, err := pr.EncryptKey(dummy) + if err != nil { + t.Fatalf("%v", err) } - DecryptKey(dummy, "first", "", []string{"red"}, nil) + _, err = cache.DecryptKey(dummy, "user", "anybody", []string{"red"}, pubEncryptedKey) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 0 { - t.Fatalf("Error in number of live keys %v", UserKeys) + cache.Refresh() + if len(cache.UserKeys) != 0 { + t.Fatalf("Error in number of live keys %v", cache.UserKeys) } } func TestBadLabel(t *testing.T) { - singleUse := ActiveUser{ - Admin: true, - Type: passvault.AESRecord, - Usage: Usage{ - Expiry: nextYear, - Uses: 2, - Labels: []string{"red"}, - }, - aesKey: emptyKey, + // Initialize passvault and keycache. Delegate a key with the tag "red" and + // verify that decryption with the tag "blue" is disallowed. + records, err := passvault.InitFrom("memory") + if err != nil { + t.Fatalf("%v", err) } - UserKeys["first"] = singleUse + pr, err := records.AddNewRecord("user", "weakpassword", true, passvault.DefaultRecordType) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 1 { + cache := NewCache() + + err = cache.AddKeyFromRecord(pr, "user", "weakpassword", nil, []string{"red"}, 1, "1h") + if err != nil { + t.Fatalf("%v", err) + } + + cache.Refresh() + if len(cache.UserKeys) != 1 { t.Fatalf("Error in number of live keys") } - EncryptKey(dummy, "first", nil) - - Refresh() - if len(UserKeys) != 1 { - t.Fatalf("Error in number of live keys") + dummy := make([]byte, 16) + pubEncryptedKey, err := pr.EncryptKey(dummy) + if err != nil { + t.Fatalf("%v", err) } - _, err := DecryptKey(dummy, "first", "", []string{"blue"}, nil) - + _, err = cache.DecryptKey(dummy, "user", "anybody", []string{"blue"}, pubEncryptedKey) if err == nil { - t.Fatalf("Decryption of labeled key with no permission") + t.Fatalf("Decryption of labeled key allowed without permission.") } - Refresh() - if len(UserKeys) != 1 { - t.Fatalf("Error in number of live keys %v", UserKeys) + cache.Refresh() + if len(cache.UserKeys) != 1 { + t.Fatalf("Error in number of live keys %v", cache.UserKeys) } } func TestGoodUser(t *testing.T) { - singleUse := ActiveUser{ - Admin: true, - Type: passvault.AESRecord, - Usage: Usage{ - Expiry: nextYear, - Uses: 2, - Users: []string{"ci", "buildeng", "first"}, - Labels: []string{"red", "blue"}, - }, - aesKey: emptyKey, + // Initialize passvault and keycache. Delegate a key with tag and user + // restrictions and verify that permissible decryption is allowed. + records, err := passvault.InitFrom("memory") + if err != nil { + t.Fatalf("%v", err) } - UserKeys["first"] = singleUse + pr, err := records.AddNewRecord("user", "weakpassword", true, passvault.DefaultRecordType) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 1 { + cache := NewCache() + + err = cache.AddKeyFromRecord( + pr, "user", "weakpassword", + []string{"ci", "buildeng", "user"}, + []string{"red", "blue"}, + 1, "1h", + ) + if err != nil { + t.Fatalf("%v", err) + } + + cache.Refresh() + if len(cache.UserKeys) != 1 { t.Fatalf("Error in number of live keys") } - EncryptKey(dummy, "first", nil) - - Refresh() - if len(UserKeys) != 1 { - t.Fatalf("Error in number of live keys") + dummy := make([]byte, 16) + pubEncryptedKey, err := pr.EncryptKey(dummy) + if err != nil { + t.Fatalf("%v", err) } - DecryptKey(dummy, "first", "ci", []string{"red"}, nil) + _, err = cache.DecryptKey(dummy, "user", "ci", []string{"red"}, pubEncryptedKey) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 0 { - t.Fatalf("Error in number of live keys %v", UserKeys) + cache.Refresh() + if len(cache.UserKeys) != 0 { + t.Fatalf("Error in number of live keys %v", cache.UserKeys) } } func TestBadUser(t *testing.T) { - singleUse := ActiveUser{ - Admin: true, - Type: passvault.AESRecord, - Usage: Usage{ - Expiry: nextYear, - Uses: 2, - Users: []string{"ci", "buildeng", "first"}, - Labels: []string{"red", "blue"}, - }, - aesKey: emptyKey, + // Initialize passvault and keycache. Delegate a key with tag and user + // restrictions and verify that illegal decryption is disallowed. + records, err := passvault.InitFrom("memory") + if err != nil { + t.Fatalf("%v", err) } - UserKeys["first"] = singleUse + pr, err := records.AddNewRecord("user", "weakpassword", true, passvault.DefaultRecordType) + if err != nil { + t.Fatalf("%v", err) + } - Refresh() - if len(UserKeys) != 1 { + cache := NewCache() + + err = cache.AddKeyFromRecord( + pr, "user", "weakpassword", + []string{"ci", "buildeng", "user"}, + []string{"red", "blue"}, + 1, "1h", + ) + if err != nil { + t.Fatalf("%v", err) + } + + cache.Refresh() + if len(cache.UserKeys) != 1 { t.Fatalf("Error in number of live keys") } - // Note that the active user needs to be in the set of delegated - // users in the AES case only - EncryptKey(dummy, "first", nil) - - Refresh() - if len(UserKeys) != 1 { - t.Fatalf("Error in number of live keys") + dummy := make([]byte, 16) + pubEncryptedKey, err := pr.EncryptKey(dummy) + if err != nil { + t.Fatalf("%v", err) } - _, err := DecryptKey(dummy, "first", "", []string{"blue"}, nil) - + _, err = cache.DecryptKey(dummy, "user", "anybody", []string{"blue"}, pubEncryptedKey) if err == nil { - t.Fatalf("Decryption of labeled key by unauthorized user") + t.Fatalf("Decryption of labeled key allowed without permission.") } - Refresh() - if len(UserKeys) != 1 { - t.Fatalf("Error in number of live keys %v", UserKeys) + cache.Refresh() + if len(cache.UserKeys) != 1 { + t.Fatalf("Error in number of live keys %v", cache.UserKeys) } } diff --git a/passvault/passvault.go b/passvault/passvault.go index 8ae275a..a595cca 100644 --- a/passvault/passvault.go +++ b/passvault/passvault.go @@ -46,9 +46,6 @@ const ( DEFAULT_VERSION = 1 ) -// Path of current vault -var localPath string - type ECPublicKey struct { Curve *elliptic.CurveParams X, Y *big.Int @@ -92,16 +89,14 @@ type PasswordRecord struct { // diskRecords is the structure used to read and write a JSON file // containing the contents of a password vault -type diskRecords struct { +type Records struct { Version int VaultId int HmacKey []byte Passwords map[string]PasswordRecord -} -// records is the set of encrypted records read from disk and -// unmarshalled -var records diskRecords + localPath string // Path of current vault +} // Summary is a minmial account summary. type Summary struct { @@ -171,8 +166,8 @@ func encryptECCRecord(newRec *PasswordRecord, ecPriv *ecdsa.PrivateKey, passKey } // createPasswordRec creates a new record from a username and password -func createPasswordRec(password string, admin bool) (newRec PasswordRecord, err error) { - newRec.Type = DefaultRecordType +func createPasswordRec(password string, admin bool, userType string) (newRec PasswordRecord, err error) { + newRec.Type = userType if newRec.PasswordSalt, err = symcrypt.MakeRandom(16); err != nil { return @@ -192,7 +187,7 @@ func createPasswordRec(password string, admin bool) (newRec PasswordRecord, err } // generate a key pair - switch DefaultRecordType { + switch userType { case RSARecord: var rsaPriv *rsa.PrivateKey rsaPriv, err = rsa.GenerateKey(rand.Reader, 2048) @@ -217,6 +212,8 @@ func createPasswordRec(password string, admin bool) (newRec PasswordRecord, err newRec.ECKey.ECPublic.Curve = ecPriv.PublicKey.Curve.Params() newRec.ECKey.ECPublic.X = ecPriv.PublicKey.X newRec.ECKey.ECPublic.Y = ecPriv.PublicKey.Y + default: + err = errors.New("Unknown record type") } newRec.Admin = admin @@ -257,14 +254,18 @@ func encryptECB(data, key []byte) (encryptedData []byte, err error) { } // InitFromDisk reads the record from disk and initialize global context. -func InitFromDisk(path string) error { - jsonDiskRecord, err := ioutil.ReadFile(path) +func InitFrom(path string) (records Records, err error) { + var jsonDiskRecord []byte - // It's OK for the file to be missing, we'll create it later if - // anything is added. + if path != "memory" { + jsonDiskRecord, err = ioutil.ReadFile(path) - if err != nil && !os.IsNotExist(err) { - return err + // It's OK for the file to be missing, we'll create it later if + // anything is added. + + if err != nil && !os.IsNotExist(err) { + return + } } // Initialized so that we can determine later if anything was read @@ -274,47 +275,47 @@ func InitFromDisk(path string) error { if len(jsonDiskRecord) != 0 { if err = json.Unmarshal(jsonDiskRecord, &records); err != nil { - return err + return } } - formatErr := errors.New("Format error") + err = errors.New("Format error") for _, rec := range records.Passwords { if len(rec.PasswordSalt) != 16 { - return formatErr + return } if len(rec.HashedPassword) != 16 { - return formatErr + return } if len(rec.KeySalt) != 16 { - return formatErr + return } if rec.Type == RSARecord { if len(rec.RSAKey.RSAExp) == 0 || len(rec.RSAKey.RSAExp)%16 != 0 { - return formatErr + return } if len(rec.RSAKey.RSAPrimeP) == 0 || len(rec.RSAKey.RSAPrimeP)%16 != 0 { - return formatErr + return } if len(rec.RSAKey.RSAPrimeQ) == 0 || len(rec.RSAKey.RSAPrimeQ)%16 != 0 { - return formatErr + return } if len(rec.RSAKey.RSAExpIV) != 16 { - return formatErr + return } if len(rec.RSAKey.RSAPrimePIV) != 16 { - return formatErr + return } if len(rec.RSAKey.RSAPrimeQIV) != 16 { - return formatErr + return } } if rec.Type == ECCRecord { if len(rec.ECKey.ECPriv) == 0 || len(rec.ECKey.ECPriv)%16 != 0 { - return formatErr + return } if len(rec.ECKey.ECPrivIV) != 16 { - return formatErr + return } } } @@ -327,62 +328,53 @@ func InitFromDisk(path string) error { records.VaultId = int(mrand.Int31()) records.HmacKey, err = symcrypt.MakeRandom(16) if err != nil { - return err + return } records.Passwords = make(map[string]PasswordRecord) } - localPath = path + records.localPath = path - return nil + err = nil + return } // WriteRecordsToDisk saves the current state of the records to disk. -func WriteRecordsToDisk() error { - if !IsInitialized() { - return errors.New("Path not initialized") +func (records *Records) WriteRecordsToDisk() error { + if records.localPath == "memory" { + return nil + } else { + jsonDiskRecord, err := json.Marshal(records) + if err != nil { + return err + } + return ioutil.WriteFile(records.localPath, jsonDiskRecord, 0644) } - - jsonDiskRecord, err := json.Marshal(records) - if err != nil { - return err - } - return ioutil.WriteFile(localPath, jsonDiskRecord, 0644) } // AddNewRecord adds a new record for a given username and password. -func AddNewRecord(name, password string, admin bool) (PasswordRecord, error) { - pr, err := createPasswordRec(password, admin) +func (records *Records) AddNewRecord(name, password string, admin bool, userType string) (PasswordRecord, error) { + pr, err := createPasswordRec(password, admin, userType) if err != nil { return pr, err } - SetRecord(pr, name) - return pr, WriteRecordsToDisk() + records.SetRecord(pr, name) + return pr, records.WriteRecordsToDisk() } // ChangePassword changes the password for a given user. -func ChangePassword(name, password, newPassword string) (err error) { - pr, ok := GetRecord(name) +func (records *Records) ChangePassword(name, password, newPassword string) (err error) { + pr, ok := records.GetRecord(name) if !ok { err = errors.New("Record not present") return } - if err = pr.ValidatePassword(password); err != nil { - return - } - // add the password salt and hash - if pr.PasswordSalt, err = symcrypt.MakeRandom(16); err != nil { + var keySalt []byte + if keySalt, err = symcrypt.MakeRandom(16); err != nil { return } - if pr.HashedPassword, err = hashPassword(newPassword, pr.PasswordSalt); err != nil { - return - } - - if pr.KeySalt, err = symcrypt.MakeRandom(16); err != nil { - return - } - newPassKey, err := derivePasswordKey(newPassword, pr.KeySalt) + newPassKey, err := derivePasswordKey(newPassword, keySalt) if err != nil { return } @@ -417,84 +409,81 @@ func ChangePassword(name, password, newPassword string) (err error) { return } - SetRecord(pr, name) + // add the password salt and hash + if pr.PasswordSalt, err = symcrypt.MakeRandom(16); err != nil { + return + } + if pr.HashedPassword, err = hashPassword(newPassword, pr.PasswordSalt); err != nil { + return + } - return WriteRecordsToDisk() + pr.KeySalt = keySalt + + records.SetRecord(pr, name) + + return records.WriteRecordsToDisk() } // DeleteRecord deletes a given record. -func DeleteRecord(name string) error { - if _, ok := GetRecord(name); ok { +func (records *Records) DeleteRecord(name string) error { + if _, ok := records.GetRecord(name); ok { delete(records.Passwords, name) - return WriteRecordsToDisk() + return records.WriteRecordsToDisk() } else { return errors.New("Record missing") } } // RevokeRecord removes admin status from a record. -func RevokeRecord(name string) error { - if rec, ok := GetRecord(name); ok { +func (records *Records) RevokeRecord(name string) error { + if rec, ok := records.GetRecord(name); ok { rec.Admin = false - SetRecord(rec, name) - return WriteRecordsToDisk() + records.SetRecord(rec, name) + return records.WriteRecordsToDisk() } else { return errors.New("Record missing") } } // MakeAdmin adds admin status to a given record. -func MakeAdmin(name string) error { - if rec, ok := GetRecord(name); ok { +func (records *Records) MakeAdmin(name string) error { + if rec, ok := records.GetRecord(name); ok { rec.Admin = true - SetRecord(rec, name) - return WriteRecordsToDisk() + records.SetRecord(rec, name) + return records.WriteRecordsToDisk() } else { return errors.New("Record missing") } } // SetRecord puts a record into the global status. -func SetRecord(pr PasswordRecord, name string) { +func (records *Records) SetRecord(pr PasswordRecord, name string) { records.Passwords[name] = pr } // GetRecord returns a record given a name. -func GetRecord(name string) (PasswordRecord, bool) { +func (records *Records) GetRecord(name string) (PasswordRecord, bool) { dpr, found := records.Passwords[name] return dpr, found } // GetVaultId returns the id of the current vault. -func GetVaultId() (id int, err error) { - if !IsInitialized() { - return 0, errors.New("Path not initialized") - } else { - return records.VaultId, nil - } +func (records *Records) GetVaultId() (id int, err error) { + return records.VaultId, nil } // GetHmacKey returns the hmac key of the current vault. -func GetHmacKey() (key []byte, err error) { - if !IsInitialized() { - return nil, errors.New("Path not initialized") - } else { - return records.HmacKey, nil - } -} - -// IsInitialized returns true if the disk vault has been loaded. -func IsInitialized() bool { - return localPath != "" +func (records *Records) GetHmacKey() (key []byte, err error) { + return records.HmacKey, nil } // NumRecords returns the number of records in the vault. -func NumRecords() int { +func (records *Records) NumRecords() int { return len(records.Passwords) } // GetSummary returns a summary of the records on disk. -func GetSummary() (summary map[string]Summary) { +func (records *Records) GetSummary() (summary map[string]Summary) { summary = make(map[string]Summary) for name, pass := range records.Passwords { summary[name] = Summary{pass.Admin, pass.Type} @@ -503,17 +492,17 @@ func GetSummary() (summary map[string]Summary) { } // IsAdmin returns the admin status of the PasswordRecord. -func (pr PasswordRecord) IsAdmin() bool { +func (pr *PasswordRecord) IsAdmin() bool { return pr.Admin } // GetType returns the type status of the PasswordRecord. -func (pr PasswordRecord) GetType() string { +func (pr *PasswordRecord) GetType() string { return pr.Type } // EncryptKey encrypts a 16-byte key with the RSA or EC key of the record. -func (pr PasswordRecord) EncryptKey(in []byte) (out []byte, err error) { +func (pr *PasswordRecord) EncryptKey(in []byte) (out []byte, err error) { if pr.Type == RSARecord { return rsa.EncryptOAEP(sha1.New(), rand.Reader, &pr.RSAKey.RSAPublic, in, nil) } else if pr.Type == ECCRecord { @@ -524,7 +513,7 @@ func (pr PasswordRecord) EncryptKey(in []byte) (out []byte, err error) { } // GetKeyRSAPub returns the RSA public key of the record. -func (pr PasswordRecord) GetKeyRSAPub() (out *rsa.PublicKey, err error) { +func (pr *PasswordRecord) GetKeyRSAPub() (out *rsa.PublicKey, err error) { if pr.Type != RSARecord { return out, errors.New("Invalid function for record type") } else { @@ -533,7 +522,7 @@ func (pr PasswordRecord) GetKeyRSAPub() (out *rsa.PublicKey, err error) { } // GetKeyECCPub returns the ECDSA public key out of the record. -func (pr PasswordRecord) GetKeyECCPub() (out *ecdsa.PublicKey, err error) { +func (pr *PasswordRecord) GetKeyECCPub() (out *ecdsa.PublicKey, err error) { if pr.Type != ECCRecord { return out, errors.New("Invalid function for record type") } else { @@ -542,7 +531,7 @@ func (pr PasswordRecord) GetKeyECCPub() (out *ecdsa.PublicKey, err error) { } // GetKeyECC returns the ECDSA private key of the record given the correct password. -func (pr PasswordRecord) GetKeyECC(password string) (key *ecdsa.PrivateKey, err error) { +func (pr *PasswordRecord) GetKeyECC(password string) (key *ecdsa.PrivateKey, err error) { if pr.Type != ECCRecord { return key, errors.New("Invalid function for record type") } @@ -569,7 +558,7 @@ func (pr PasswordRecord) GetKeyECC(password string) (key *ecdsa.PrivateKey, err } // GetKeyRSA returns the RSA private key of the record given the correct password. -func (pr PasswordRecord) GetKeyRSA(password string) (key rsa.PrivateKey, err error) { +func (pr *PasswordRecord) GetKeyRSA(password string) (key rsa.PrivateKey, err error) { if pr.Type != RSARecord { return key, errors.New("Invalid function for record type") } @@ -626,7 +615,7 @@ func (pr PasswordRecord) GetKeyRSA(password string) (key rsa.PrivateKey, err err } // ValidatePassword returns an error if the password is incorrect. -func (pr PasswordRecord) ValidatePassword(password string) error { +func (pr *PasswordRecord) ValidatePassword(password string) error { h, err := hashPassword(password, pr.PasswordSalt) if err != nil { return err diff --git a/passvault/passvault_test.go b/passvault/passvault_test.go index 51f54e0..88b3daf 100644 --- a/passvault/passvault_test.go +++ b/passvault/passvault_test.go @@ -5,33 +5,41 @@ package passvault import ( - "testing" "os" + "testing" ) func TestStaticVault(t *testing.T) { - err := InitFromDisk("/tmp/redoctober.json") + // Initial create. + records, err := InitFrom("/tmp/redoctober.jso") if err != nil { t.Fatalf("Error reading record") } - _, err = AddNewRecord("test", "bad pass", true) + _, err = records.AddNewRecord("test", "bad pass", true, DefaultRecordType) if err != nil { t.Fatalf("Error creating record") } - err = InitFromDisk("/tmp/redoctober.json") + + // Reads data written last time. + _, err = InitFrom("/tmp/redoctober.json") if err != nil { t.Fatalf("Error reading record") } + + // Cleaning. os.Remove("/tmp/redoctober.json") } func TestRSAEncryptDecrypt(t *testing.T) { - oldDefaultRecordType := DefaultRecordType - DefaultRecordType = RSARecord - myRec, err := createPasswordRec("mypasswordisweak", true) + records, err := InitFrom("memory") if err != nil { - t.Fatalf("Error creating record") + t.Fatalf("%v", err) + } + + myRec, err := records.AddNewRecord("user", "weakpassword", true, RSARecord) + if err != nil { + t.Fatalf("%v", err) } _, err = myRec.GetKeyRSAPub() @@ -44,7 +52,7 @@ func TestRSAEncryptDecrypt(t *testing.T) { t.Fatalf("Incorrect password did not fail") } - rsaPriv, err = myRec.GetKeyRSA("mypasswordisweak") + rsaPriv, err = myRec.GetKeyRSA("weakpassword") if err != nil { t.Fatalf("Error decrypting RSA key") } @@ -53,20 +61,22 @@ func TestRSAEncryptDecrypt(t *testing.T) { if err != nil { t.Fatalf("Error validating RSA key") } - DefaultRecordType = oldDefaultRecordType } func TestECCEncryptDecrypt(t *testing.T) { - oldDefaultRecordType := DefaultRecordType - DefaultRecordType = ECCRecord - myRec, err := createPasswordRec("mypasswordisweak", true) + records, err := InitFrom("memory") if err != nil { - t.Fatalf("Error creating record") + t.Fatalf("%v", err) + } + + myRec, err := records.AddNewRecord("user", "weakpassword", true, ECCRecord) + if err != nil { + t.Fatalf("%v", err) } _, err = myRec.GetKeyECCPub() if err != nil { - t.Fatalf("Error extracting EC pub") + t.Fatalf("%v", err) } _, err = myRec.GetKeyECC("mypasswordiswrong") @@ -74,9 +84,8 @@ func TestECCEncryptDecrypt(t *testing.T) { t.Fatalf("Incorrect password did not fail") } - _, err = myRec.GetKeyECC("mypasswordisweak") + _, err = myRec.GetKeyECC("weakpassword") if err != nil { - t.Fatalf("Error decrypting EC key") + t.Fatalf("%v", err) } - DefaultRecordType = oldDefaultRecordType } diff --git a/redoctober.go b/redoctober.go index 90ec880..4feb6ac 100644 --- a/redoctober.go +++ b/redoctober.go @@ -89,7 +89,7 @@ func NewServer(process chan<- userRequest, staticPath, addr, certPath, keyPath, Rand: rand.Reader, PreferServerCipherSuites: true, SessionTicketsDisabled: true, - MinVersion: tls.VersionTLS10, + MinVersion: tls.VersionTLS10, } // If a caPath has been specified then a local CA is being used