84 lines
1.8 KiB
Go
84 lines
1.8 KiB
Go
package vault
|
|
|
|
import (
|
|
"os"
|
|
"path/filepath"
|
|
"testing"
|
|
)
|
|
|
|
func TestSaveLoadRoundTrip(t *testing.T) {
|
|
dir := t.TempDir()
|
|
v, err := Open(dir)
|
|
if err != nil {
|
|
t.Fatalf("open: %v", err)
|
|
}
|
|
secret := []byte(`{"cookies":"session=abc","token":"xyz"}`)
|
|
if err = v.Save("acc-1", secret); err != nil {
|
|
t.Fatalf("save: %v", err)
|
|
}
|
|
// 文件不应包含明文
|
|
raw, _ := os.ReadFile(filepath.Join(dir, "vault.enc"))
|
|
if stringContains(string(raw), "session=abc") {
|
|
t.Fatal("vault file must not contain plaintext")
|
|
}
|
|
got, err := v.Load("acc-1")
|
|
if err != nil {
|
|
t.Fatalf("load: %v", err)
|
|
}
|
|
if string(got) != string(secret) {
|
|
t.Fatalf("roundtrip mismatch: %s vs %s", got, secret)
|
|
}
|
|
}
|
|
|
|
func TestReopenPersists(t *testing.T) {
|
|
dir := t.TempDir()
|
|
v1, _ := Open(dir)
|
|
_ = v1.Save("acc-2", []byte("data-2"))
|
|
v2, err := Open(dir)
|
|
if err != nil {
|
|
t.Fatalf("reopen: %v", err)
|
|
}
|
|
got, err := v2.Load("acc-2")
|
|
if err != nil || string(got) != "data-2" {
|
|
t.Fatalf("reopen mismatch: %s %v", got, err)
|
|
}
|
|
}
|
|
|
|
func TestTamperFails(t *testing.T) {
|
|
dir := t.TempDir()
|
|
v, _ := Open(dir)
|
|
_ = v.Save("acc-3", []byte("data-3"))
|
|
path := filepath.Join(dir, "vault.enc")
|
|
raw, _ := os.ReadFile(path)
|
|
raw[len(raw)/2] ^= 0xff
|
|
_ = os.WriteFile(path, raw, 0o600)
|
|
if _, err := v.Load("acc-3"); err == nil {
|
|
t.Fatal("tampered ciphertext should fail auth")
|
|
}
|
|
}
|
|
|
|
func TestDelete(t *testing.T) {
|
|
dir := t.TempDir()
|
|
v, _ := Open(dir)
|
|
_ = v.Save("acc-4", []byte("d4"))
|
|
if err := v.Delete("acc-4"); err != nil {
|
|
t.Fatalf("delete: %v", err)
|
|
}
|
|
if _, err := v.Load("acc-4"); err == nil {
|
|
t.Fatal("deleted entry should not load")
|
|
}
|
|
}
|
|
|
|
func stringContains(s, sub string) bool {
|
|
return len(s) >= len(sub) && indexOf(s, sub) >= 0
|
|
}
|
|
|
|
func indexOf(s, sub string) int {
|
|
for i := 0; i+len(sub) <= len(s); i++ {
|
|
if s[i:i+len(sub)] == sub {
|
|
return i
|
|
}
|
|
}
|
|
return -1
|
|
}
|