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) } }