499 lines
12 KiB
Go
499 lines
12 KiB
Go
package main
|
|
|
|
import (
|
|
"bytes"
|
|
"errors"
|
|
"io"
|
|
"os"
|
|
"path/filepath"
|
|
"strings"
|
|
"testing"
|
|
|
|
"github.com/spf13/cobra"
|
|
|
|
"git.tswf.io/infra/go-synapse-backupper/pkg/adapters/crypto/composite"
|
|
"git.tswf.io/infra/go-synapse-backupper/pkg/domain/crypto"
|
|
)
|
|
|
|
// --- Mocks ---
|
|
|
|
type mockRecipientPriv struct {
|
|
schemeID uint16
|
|
keyID []byte
|
|
raw []byte
|
|
}
|
|
|
|
func (m *mockRecipientPriv) SchemeID() uint16 { return m.schemeID }
|
|
func (m *mockRecipientPriv) KeyID() []byte { return m.keyID }
|
|
func (m *mockRecipientPriv) Raw() []byte { return m.raw }
|
|
|
|
type mockRestoreKeyManager struct {
|
|
loadPrivCalls []loadPrivCall
|
|
pqPriv crypto.RecipientPriv
|
|
classicalPriv crypto.RecipientPriv
|
|
loadPrivErr error
|
|
}
|
|
|
|
type loadPrivCall struct {
|
|
path string
|
|
schemeID uint16
|
|
}
|
|
|
|
func (m *mockRestoreKeyManager) LoadPub(
|
|
path string,
|
|
schemeID uint16,
|
|
) (
|
|
crypto.RecipientPub,
|
|
error,
|
|
) {
|
|
return nil, nil
|
|
}
|
|
|
|
func (m *mockRestoreKeyManager) LoadPriv(
|
|
path string,
|
|
schemeID uint16,
|
|
) (
|
|
crypto.RecipientPriv,
|
|
error,
|
|
) {
|
|
m.loadPrivCalls = append(
|
|
m.loadPrivCalls,
|
|
loadPrivCall{path: path, schemeID: schemeID},
|
|
)
|
|
if m.loadPrivErr != nil {
|
|
return nil, m.loadPrivErr
|
|
}
|
|
if schemeID == 0x0006 {
|
|
return m.pqPriv, nil
|
|
}
|
|
return m.classicalPriv, nil
|
|
}
|
|
|
|
func (m *mockRestoreKeyManager) Generate(
|
|
schemeID uint16,
|
|
pubOut io.Writer,
|
|
privOut io.Writer,
|
|
rand io.Reader,
|
|
) error {
|
|
return nil
|
|
}
|
|
|
|
type mockDecryptor struct {
|
|
called bool
|
|
src io.Reader
|
|
privs []crypto.RecipientPriv
|
|
out io.Writer
|
|
err error
|
|
}
|
|
|
|
func (d *mockDecryptor) Decrypt(
|
|
src io.Reader,
|
|
privs []crypto.RecipientPriv,
|
|
plaintext io.Writer,
|
|
) error {
|
|
d.called = true
|
|
d.src = src
|
|
d.privs = privs
|
|
d.out = plaintext
|
|
if d.err != nil {
|
|
return d.err
|
|
}
|
|
_, writeErr := plaintext.Write([]byte("decrypted pg_dump data"))
|
|
return writeErr
|
|
}
|
|
|
|
// --- Helpers ---
|
|
|
|
func setRestoreFlags(
|
|
cmd *cobra.Command,
|
|
inPath string,
|
|
pqPath string,
|
|
classicalPath string,
|
|
) {
|
|
_ = cmd.Flags().Set("in", inPath)
|
|
_ = cmd.Flags().Set("privkey-pq", pqPath)
|
|
_ = cmd.Flags().Set("privkey-classical", classicalPath)
|
|
}
|
|
|
|
func makeRestoreTestCmd() *cobra.Command {
|
|
cmd := &cobra.Command{
|
|
SilenceUsage: true,
|
|
SilenceErrors: true,
|
|
}
|
|
cmd.SetArgs([]string{})
|
|
cmd.Flags().String("in", "", "Input .pqenc file path")
|
|
cmd.Flags().String("privkey-pq", "", "Path to PQ private key PEM")
|
|
cmd.Flags().String("privkey-classical", "", "Path to classical private key PEM")
|
|
cmd.Flags().String("out", "", "Output file path (empty = stdout)")
|
|
cmd.PreRunE = restoreCmd.PreRunE
|
|
cmd.RunE = restoreCmd.RunE
|
|
return cmd
|
|
}
|
|
|
|
func restoreTestGlobals(t *testing.T) {
|
|
originalNewRestoreKeyManager := newRestoreKeyManager
|
|
originalNewRestoreDecryptor := newRestoreDecryptor
|
|
originalRestoreOutput := restoreOutput
|
|
originalRestoreOsOpen := restoreOsOpen
|
|
t.Cleanup(func() {
|
|
newRestoreKeyManager = originalNewRestoreKeyManager
|
|
newRestoreDecryptor = originalNewRestoreDecryptor
|
|
restoreOutput = originalRestoreOutput
|
|
restoreOsOpen = originalRestoreOsOpen
|
|
})
|
|
}
|
|
|
|
// --- Tests ---
|
|
|
|
func TestRestoreCmd_Structure(t *testing.T) {
|
|
if restoreCmd == nil {
|
|
t.Fatal("restoreCmd is nil")
|
|
}
|
|
if restoreCmd.Use != "restore" {
|
|
t.Fatalf("expected Use='restore', got %q", restoreCmd.Use)
|
|
}
|
|
|
|
requiredFlags := []string{"in", "privkey-pq", "privkey-classical"}
|
|
for _, f := range requiredFlags {
|
|
if restoreCmd.Flag(f) == nil {
|
|
t.Fatalf("missing required --%s flag", f)
|
|
}
|
|
}
|
|
|
|
if restoreCmd.Flag("out") == nil {
|
|
t.Fatal("missing --out flag")
|
|
}
|
|
if restoreCmd.Flag("out").DefValue != "" {
|
|
t.Fatalf(
|
|
"expected --out default empty, got %q",
|
|
restoreCmd.Flag("out").DefValue,
|
|
)
|
|
}
|
|
}
|
|
|
|
func TestRestoreCmd_RequiredFlags(t *testing.T) {
|
|
required := []string{"in", "privkey-pq", "privkey-classical"}
|
|
for _, name := range required {
|
|
flag := restoreCmd.Flag(name)
|
|
if flag == nil {
|
|
t.Fatalf("missing required --%s flag", name)
|
|
}
|
|
ann, ok := flag.Annotations[cobra.BashCompOneRequiredFlag]
|
|
if !ok || len(ann) == 0 || ann[0] != "true" {
|
|
t.Fatalf("flag --%s is not marked required", name)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestRestoreCmd_PreRunE_MissingPrivkeyClassical(t *testing.T) {
|
|
restoreTestGlobals(t)
|
|
|
|
dir := t.TempDir()
|
|
inFile := filepath.Join(dir, "backup.pqenc")
|
|
_ = os.WriteFile(inFile, []byte("data"), 0o644)
|
|
pqFile := filepath.Join(dir, "pq.priv.pem")
|
|
_ = os.WriteFile(pqFile, []byte("pq"), 0o600)
|
|
classicalFile := filepath.Join(dir, "classical.priv.pem") // does not exist
|
|
|
|
openCount := 0
|
|
restoreOsOpen = func(name string) (*os.File, error) {
|
|
openCount++
|
|
return os.Open(name)
|
|
}
|
|
|
|
cmd := makeRestoreTestCmd()
|
|
setRestoreFlags(cmd, inFile, pqFile, classicalFile)
|
|
|
|
err := cmd.Execute()
|
|
if err == nil {
|
|
t.Fatal("expected non-nil error")
|
|
}
|
|
if !strings.Contains(
|
|
err.Error(),
|
|
"--privkey-pq and --privkey-classical are both required (AND model)",
|
|
) {
|
|
t.Fatalf("expected AND model error, got: %v", err)
|
|
}
|
|
if openCount != 0 {
|
|
t.Fatalf("expected no os.Open calls on .pqenc, got %d", openCount)
|
|
}
|
|
}
|
|
|
|
func TestRestoreCmd_PreRunE_MissingPrivkeyPQ(t *testing.T) {
|
|
restoreTestGlobals(t)
|
|
|
|
dir := t.TempDir()
|
|
inFile := filepath.Join(dir, "backup.pqenc")
|
|
_ = os.WriteFile(inFile, []byte("data"), 0o644)
|
|
pqFile := filepath.Join(dir, "pq.priv.pem") // does not exist
|
|
classicalFile := filepath.Join(dir, "classical.priv.pem")
|
|
_ = os.WriteFile(classicalFile, []byte("classical"), 0o600)
|
|
|
|
openCount := 0
|
|
restoreOsOpen = func(name string) (*os.File, error) {
|
|
openCount++
|
|
return os.Open(name)
|
|
}
|
|
|
|
cmd := makeRestoreTestCmd()
|
|
setRestoreFlags(cmd, inFile, pqFile, classicalFile)
|
|
|
|
err := cmd.Execute()
|
|
if err == nil {
|
|
t.Fatal("expected non-nil error")
|
|
}
|
|
if !strings.Contains(
|
|
err.Error(),
|
|
"--privkey-pq and --privkey-classical are both required (AND model)",
|
|
) {
|
|
t.Fatalf("expected AND model error, got: %v", err)
|
|
}
|
|
if openCount != 0 {
|
|
t.Fatalf("expected no os.Open calls on .pqenc, got %d", openCount)
|
|
}
|
|
}
|
|
|
|
func TestRestoreCmd_OnlyPrivkeyPQ(t *testing.T) {
|
|
restoreTestGlobals(t)
|
|
|
|
dir := t.TempDir()
|
|
inFile := filepath.Join(dir, "backup.pqenc")
|
|
_ = os.WriteFile(inFile, []byte("data"), 0o644)
|
|
pqFile := filepath.Join(dir, "pq.priv.pem")
|
|
_ = os.WriteFile(pqFile, []byte("pq"), 0o600)
|
|
|
|
openCount := 0
|
|
restoreOsOpen = func(name string) (*os.File, error) {
|
|
openCount++
|
|
return os.Open(name)
|
|
}
|
|
|
|
cmd := makeRestoreTestCmd()
|
|
_ = cmd.Flags().Set("in", inFile)
|
|
_ = cmd.Flags().Set("privkey-pq", pqFile)
|
|
// --privkey-classical intentionally omitted
|
|
|
|
err := cmd.Execute()
|
|
if err == nil {
|
|
t.Fatal("expected non-nil error")
|
|
}
|
|
if openCount != 0 {
|
|
t.Fatalf("expected no os.Open calls on .pqenc, got %d", openCount)
|
|
}
|
|
}
|
|
|
|
func TestRestoreCmd_HappyPath_Stdout(t *testing.T) {
|
|
restoreTestGlobals(t)
|
|
|
|
dir := t.TempDir()
|
|
inFile := filepath.Join(dir, "backup.pqenc")
|
|
_ = os.WriteFile(inFile, []byte("encrypted"), 0o644)
|
|
pqFile := filepath.Join(dir, "pq.priv.pem")
|
|
_ = os.WriteFile(pqFile, []byte("pq"), 0o600)
|
|
classicalFile := filepath.Join(dir, "classical.priv.pem")
|
|
_ = os.WriteFile(classicalFile, []byte("classical"), 0o600)
|
|
|
|
pqPriv := &mockRecipientPriv{
|
|
schemeID: 0x0006,
|
|
keyID: []byte{1, 2, 3, 4, 5, 6, 7, 8},
|
|
raw: make([]byte, 32),
|
|
}
|
|
classicalPriv := &mockRecipientPriv{
|
|
schemeID: 0x0007,
|
|
keyID: []byte{8, 7, 6, 5, 4, 3, 2, 1},
|
|
raw: make([]byte, 32),
|
|
}
|
|
|
|
mockKM := &mockRestoreKeyManager{pqPriv: pqPriv, classicalPriv: classicalPriv}
|
|
newRestoreKeyManager = func(crypto.Registry) crypto.KeyManager {
|
|
return mockKM
|
|
}
|
|
|
|
mockDec := &mockDecryptor{}
|
|
newRestoreDecryptor = func(crypto.Registry) crypto.Decryptor {
|
|
return mockDec
|
|
}
|
|
|
|
var outBuf bytes.Buffer
|
|
restoreOutput = &outBuf
|
|
|
|
cmd := makeRestoreTestCmd()
|
|
setRestoreFlags(cmd, inFile, pqFile, classicalFile)
|
|
|
|
err := cmd.Execute()
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
if len(mockKM.loadPrivCalls) != 2 {
|
|
t.Fatalf("expected 2 LoadPriv calls, got %d", len(mockKM.loadPrivCalls))
|
|
}
|
|
if mockKM.loadPrivCalls[0].path != pqFile ||
|
|
mockKM.loadPrivCalls[0].schemeID != 0x0006 {
|
|
t.Fatalf("unexpected PQ LoadPriv call: %+v", mockKM.loadPrivCalls[0])
|
|
}
|
|
if mockKM.loadPrivCalls[1].path != classicalFile ||
|
|
mockKM.loadPrivCalls[1].schemeID != 0x0007 {
|
|
t.Fatalf(
|
|
"unexpected classical LoadPriv call: %+v",
|
|
mockKM.loadPrivCalls[1],
|
|
)
|
|
}
|
|
|
|
if !mockDec.called {
|
|
t.Fatal("expected decryptor.Decrypt to be called")
|
|
}
|
|
if len(mockDec.privs) != 2 {
|
|
t.Fatalf("expected 2 privs, got %d", len(mockDec.privs))
|
|
}
|
|
if mockDec.privs[0] != pqPriv {
|
|
t.Fatal("expected pqPriv in positional slot 0")
|
|
}
|
|
if mockDec.privs[1] != classicalPriv {
|
|
t.Fatal("expected classicalPriv in positional slot 1")
|
|
}
|
|
|
|
if !bytes.Equal(outBuf.Bytes(), []byte("decrypted pg_dump data")) {
|
|
t.Fatalf("unexpected stdout content: %q", outBuf.Bytes())
|
|
}
|
|
}
|
|
|
|
func TestRestoreCmd_HappyPath_FileOut(t *testing.T) {
|
|
restoreTestGlobals(t)
|
|
|
|
dir := t.TempDir()
|
|
inFile := filepath.Join(dir, "backup.pqenc")
|
|
_ = os.WriteFile(inFile, []byte("encrypted"), 0o644)
|
|
pqFile := filepath.Join(dir, "pq.priv.pem")
|
|
_ = os.WriteFile(pqFile, []byte("pq"), 0o600)
|
|
classicalFile := filepath.Join(dir, "classical.priv.pem")
|
|
_ = os.WriteFile(classicalFile, []byte("classical"), 0o600)
|
|
outFile := filepath.Join(dir, "restored.dump")
|
|
|
|
pqPriv := &mockRecipientPriv{
|
|
schemeID: 0x0006,
|
|
keyID: []byte{1, 2, 3, 4, 5, 6, 7, 8},
|
|
raw: make([]byte, 32),
|
|
}
|
|
classicalPriv := &mockRecipientPriv{
|
|
schemeID: 0x0007,
|
|
keyID: []byte{8, 7, 6, 5, 4, 3, 2, 1},
|
|
raw: make([]byte, 32),
|
|
}
|
|
|
|
mockKM := &mockRestoreKeyManager{pqPriv: pqPriv, classicalPriv: classicalPriv}
|
|
newRestoreKeyManager = func(crypto.Registry) crypto.KeyManager {
|
|
return mockKM
|
|
}
|
|
|
|
mockDec := &mockDecryptor{}
|
|
newRestoreDecryptor = func(crypto.Registry) crypto.Decryptor {
|
|
return mockDec
|
|
}
|
|
|
|
cmd := makeRestoreTestCmd()
|
|
setRestoreFlags(cmd, inFile, pqFile, classicalFile)
|
|
_ = cmd.Flags().Set("out", outFile)
|
|
|
|
err := cmd.Execute()
|
|
if err != nil {
|
|
t.Fatalf("unexpected error: %v", err)
|
|
}
|
|
|
|
data, err := os.ReadFile(outFile)
|
|
if err != nil {
|
|
t.Fatalf("read output file: %v", err)
|
|
}
|
|
if !bytes.Equal(data, []byte("decrypted pg_dump data")) {
|
|
t.Fatalf("unexpected output file content: %q", data)
|
|
}
|
|
}
|
|
|
|
func TestRestoreCmd_TamperingDetected(t *testing.T) {
|
|
restoreTestGlobals(t)
|
|
|
|
dir := t.TempDir()
|
|
inFile := filepath.Join(dir, "backup.pqenc")
|
|
_ = os.WriteFile(inFile, []byte("encrypted"), 0o644)
|
|
pqFile := filepath.Join(dir, "pq.priv.pem")
|
|
_ = os.WriteFile(pqFile, []byte("pq"), 0o600)
|
|
classicalFile := filepath.Join(dir, "classical.priv.pem")
|
|
_ = os.WriteFile(classicalFile, []byte("classical"), 0o600)
|
|
|
|
pqPriv := &mockRecipientPriv{
|
|
schemeID: 0x0006,
|
|
keyID: []byte{1, 2, 3, 4, 5, 6, 7, 8},
|
|
raw: make([]byte, 32),
|
|
}
|
|
classicalPriv := &mockRecipientPriv{
|
|
schemeID: 0x0007,
|
|
keyID: []byte{8, 7, 6, 5, 4, 3, 2, 1},
|
|
raw: make([]byte, 32),
|
|
}
|
|
|
|
mockKM := &mockRestoreKeyManager{pqPriv: pqPriv, classicalPriv: classicalPriv}
|
|
newRestoreKeyManager = func(crypto.Registry) crypto.KeyManager {
|
|
return mockKM
|
|
}
|
|
|
|
mockDec := &mockDecryptor{err: composite.ErrTamperingDetected}
|
|
newRestoreDecryptor = func(crypto.Registry) crypto.Decryptor {
|
|
return mockDec
|
|
}
|
|
|
|
cmd := makeRestoreTestCmd()
|
|
setRestoreFlags(cmd, inFile, pqFile, classicalFile)
|
|
|
|
err := cmd.Execute()
|
|
if err == nil {
|
|
t.Fatal("expected non-nil error")
|
|
}
|
|
if !errors.Is(err, composite.ErrTamperingDetected) {
|
|
t.Fatalf("expected ErrTamperingDetected, got: %v", err)
|
|
}
|
|
}
|
|
|
|
func TestRestoreCmd_WrongKeys(t *testing.T) {
|
|
restoreTestGlobals(t)
|
|
|
|
dir := t.TempDir()
|
|
inFile := filepath.Join(dir, "backup.pqenc")
|
|
_ = os.WriteFile(inFile, []byte("encrypted"), 0o644)
|
|
pqFile := filepath.Join(dir, "pq.priv.pem")
|
|
_ = os.WriteFile(pqFile, []byte("pq"), 0o600)
|
|
classicalFile := filepath.Join(dir, "classical.priv.pem")
|
|
_ = os.WriteFile(classicalFile, []byte("classical"), 0o600)
|
|
|
|
pqPriv := &mockRecipientPriv{
|
|
schemeID: 0x0006,
|
|
keyID: []byte{1, 2, 3, 4, 5, 6, 7, 8},
|
|
raw: make([]byte, 32),
|
|
}
|
|
classicalPriv := &mockRecipientPriv{
|
|
schemeID: 0x0007,
|
|
keyID: []byte{8, 7, 6, 5, 4, 3, 2, 1},
|
|
raw: make([]byte, 32),
|
|
}
|
|
|
|
mockKM := &mockRestoreKeyManager{pqPriv: pqPriv, classicalPriv: classicalPriv}
|
|
newRestoreKeyManager = func(crypto.Registry) crypto.KeyManager {
|
|
return mockKM
|
|
}
|
|
|
|
mockDec := &mockDecryptor{err: composite.ErrWrongKeys}
|
|
newRestoreDecryptor = func(crypto.Registry) crypto.Decryptor {
|
|
return mockDec
|
|
}
|
|
|
|
cmd := makeRestoreTestCmd()
|
|
setRestoreFlags(cmd, inFile, pqFile, classicalFile)
|
|
|
|
err := cmd.Execute()
|
|
if err == nil {
|
|
t.Fatal("expected non-nil error")
|
|
}
|
|
if !errors.Is(err, composite.ErrWrongKeys) {
|
|
t.Fatalf("expected ErrWrongKeys, got: %v", err)
|
|
}
|
|
}
|