Compare commits
7 Commits
bf2bceb520
...
f772d08a70
| Author | SHA1 | Date | |
|---|---|---|---|
| f772d08a70 | |||
| a9f5a74162 | |||
| 174b5f9048 | |||
| 3950dd94ee | |||
| 3f59a97844 | |||
| 9c6fc8cffa | |||
| 328f517543 |
@@ -35,17 +35,7 @@ var (
|
|||||||
outputWriter io.Writer = os.Stderr
|
outputWriter io.Writer = os.Stderr
|
||||||
)
|
)
|
||||||
|
|
||||||
type pipelineRunner interface {
|
func defaultNewRunner(options ...pipeline.Option) pipeline.Runner {
|
||||||
Run(
|
|
||||||
ctx context.Context,
|
|
||||||
pgDumpOpts pgdump.Options,
|
|
||||||
recipients []crypto.RecipientPub,
|
|
||||||
sink domain.Sink,
|
|
||||||
rand io.Reader,
|
|
||||||
) error
|
|
||||||
}
|
|
||||||
|
|
||||||
func defaultNewRunner(options ...pipeline.Option) pipelineRunner {
|
|
||||||
return pipeline.NewRunner(options...)
|
return pipeline.NewRunner(options...)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -91,7 +81,8 @@ var backupCmd = &cobra.Command{
|
|||||||
return fmt.Errorf("load classical public key: %w", err)
|
return fmt.Errorf("load classical public key: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
recipients := []crypto.RecipientPub{pqPub, classicalPub}
|
recipients := make([]crypto.RecipientPub, 0, 2)
|
||||||
|
recipients = append(recipients, pqPub, classicalPub)
|
||||||
|
|
||||||
pgDumpOpts := pgdump.Options{
|
pgDumpOpts := pgdump.Options{
|
||||||
Host: cfg.PG.Host,
|
Host: cfg.PG.Host,
|
||||||
@@ -110,7 +101,7 @@ var backupCmd = &cobra.Command{
|
|||||||
slog.String("output_path", finalPath),
|
slog.String("output_path", finalPath),
|
||||||
)
|
)
|
||||||
|
|
||||||
sink := newLocalSink(cfg.Backup.Dir)
|
var sink domain.Sink = newLocalSink(cfg.Backup.Dir)
|
||||||
|
|
||||||
registry := crypto.NewRegistry()
|
registry := crypto.NewRegistry()
|
||||||
_ = registry.Register(0x0006, func() crypto.KEM { return mlkem768.New() })
|
_ = registry.Register(0x0006, func() crypto.KEM { return mlkem768.New() })
|
||||||
|
|||||||
@@ -183,7 +183,7 @@ func TestBackupCmd_Success(t *testing.T) {
|
|||||||
|
|
||||||
testData := []byte("test backup payload")
|
testData := []byte("test backup payload")
|
||||||
dumper := &successDumper{data: testData}
|
dumper := &successDumper{data: testData}
|
||||||
newRunner = func(...pipeline.Option) pipelineRunner {
|
newRunner = func(...pipeline.Option) pipeline.Runner {
|
||||||
return pipeline.NewRunner(
|
return pipeline.NewRunner(
|
||||||
pipeline.WithDumper(dumper),
|
pipeline.WithDumper(dumper),
|
||||||
pipeline.WithEncryptor(&passthroughEncryptor{}),
|
pipeline.WithEncryptor(&passthroughEncryptor{}),
|
||||||
@@ -320,7 +320,7 @@ func TestBackupCmd_PgDumpFailure(t *testing.T) {
|
|||||||
return mockKM
|
return mockKM
|
||||||
}
|
}
|
||||||
|
|
||||||
newRunner = func(...pipeline.Option) pipelineRunner {
|
newRunner = func(...pipeline.Option) pipeline.Runner {
|
||||||
return pipeline.NewRunner(
|
return pipeline.NewRunner(
|
||||||
pipeline.WithDumper(&failDumper{}),
|
pipeline.WithDumper(&failDumper{}),
|
||||||
pipeline.WithEncryptor(&passthroughEncryptor{}),
|
pipeline.WithEncryptor(&passthroughEncryptor{}),
|
||||||
|
|||||||
@@ -1,19 +1,14 @@
|
|||||||
package main
|
package main
|
||||||
|
|
||||||
import (
|
import (
|
||||||
_ "embed"
|
|
||||||
"fmt"
|
"fmt"
|
||||||
"os"
|
"os"
|
||||||
|
|
||||||
"github.com/spf13/cobra"
|
"github.com/spf13/cobra"
|
||||||
|
|
||||||
|
"git.tswf.io/infra/go-synapse-backupper/resources"
|
||||||
)
|
)
|
||||||
|
|
||||||
//go:embed resources/config.en.yaml
|
|
||||||
var configEnYAML []byte
|
|
||||||
|
|
||||||
//go:embed resources/config.ru.yaml
|
|
||||||
var configRuYAML []byte
|
|
||||||
|
|
||||||
func generateConfigCmd() *cobra.Command {
|
func generateConfigCmd() *cobra.Command {
|
||||||
var lang string
|
var lang string
|
||||||
var output string
|
var output string
|
||||||
@@ -25,9 +20,9 @@ func generateConfigCmd() *cobra.Command {
|
|||||||
var tmpl []byte
|
var tmpl []byte
|
||||||
switch lang {
|
switch lang {
|
||||||
case "en":
|
case "en":
|
||||||
tmpl = configEnYAML
|
tmpl = resources.ConfigEnYAML
|
||||||
case "ru":
|
case "ru":
|
||||||
tmpl = configRuYAML
|
tmpl = resources.ConfigRuYAML
|
||||||
default:
|
default:
|
||||||
return fmt.Errorf("unsupported language: %q (must be \"en\" or \"ru\")", lang)
|
return fmt.Errorf("unsupported language: %q (must be \"en\" or \"ru\")", lang)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -140,13 +140,8 @@ func TestGenerateConfig_AllKeys(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
out := buf.String()
|
out := buf.String()
|
||||||
requiredKeys := []string{
|
requiredKeys := make([]string, 0, 16)
|
||||||
"host", "port", "user", "password", "database", "sslmode", "exclude_tables",
|
requiredKeys = append(requiredKeys, "host", "port", "user", "password", "database", "sslmode", "exclude_tables", "dir", "retention_days", "cron", "pq_scheme", "classical_scheme", "pq_public_key_path", "classical_public_key_path", "healthz", "log")
|
||||||
"dir", "retention_days", "cron",
|
|
||||||
"pq_scheme", "classical_scheme",
|
|
||||||
"pq_public_key_path", "classical_public_key_path",
|
|
||||||
"healthz", "log",
|
|
||||||
}
|
|
||||||
|
|
||||||
for _, key := range requiredKeys {
|
for _, key := range requiredKeys {
|
||||||
if !strings.Contains(out, key) {
|
if !strings.Contains(out, key) {
|
||||||
|
|||||||
@@ -68,12 +68,14 @@ func newKeygenCmdWithDeps(reg crypto.Registry) *cobra.Command {
|
|||||||
}
|
}
|
||||||
|
|
||||||
if !force {
|
if !force {
|
||||||
for _, s := range schemes {
|
for _, scheme := range schemes {
|
||||||
pubPath := outPrefix + "." + s.name + ".pub.pem"
|
pubPath := outPrefix + "." + scheme.name + ".pub.pem"
|
||||||
privPath := outPrefix + "." + s.name + ".priv.pem"
|
privPath := outPrefix + "." + scheme.name + ".priv.pem"
|
||||||
for _, p := range []string{pubPath, privPath} {
|
paths := make([]string, 0, 2)
|
||||||
if _, err := os.Stat(p); err == nil {
|
paths = append(paths, pubPath, privPath)
|
||||||
return fmt.Errorf("file already exists: %s (use --force to overwrite)", p)
|
for _, path := range paths {
|
||||||
|
if _, err := os.Stat(path); err == nil {
|
||||||
|
return fmt.Errorf("file already exists: %s (use --force to overwrite)", path)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -81,9 +83,9 @@ func newKeygenCmdWithDeps(reg crypto.Registry) *cobra.Command {
|
|||||||
|
|
||||||
km := keymanager.NewKeyManager(reg)
|
km := keymanager.NewKeyManager(reg)
|
||||||
|
|
||||||
for _, s := range schemes {
|
for _, scheme := range schemes {
|
||||||
pubPath := outPrefix + "." + s.name + ".pub.pem"
|
pubPath := outPrefix + "." + scheme.name + ".pub.pem"
|
||||||
privPath := outPrefix + "." + s.name + ".priv.pem"
|
privPath := outPrefix + "." + scheme.name + ".priv.pem"
|
||||||
|
|
||||||
pubFile, err := os.OpenFile(pubPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o644)
|
pubFile, err := os.OpenFile(pubPath, os.O_WRONLY|os.O_CREATE|os.O_TRUNC, 0o644)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -96,10 +98,10 @@ func newKeygenCmdWithDeps(reg crypto.Registry) *cobra.Command {
|
|||||||
return fmt.Errorf("create private key file %s: %w", privPath, err)
|
return fmt.Errorf("create private key file %s: %w", privPath, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := km.Generate(s.schemeID, pubFile, privFile, rand.Reader); err != nil {
|
if err := km.Generate(scheme.schemeID, pubFile, privFile, rand.Reader); err != nil {
|
||||||
_ = pubFile.Close()
|
_ = pubFile.Close()
|
||||||
_ = privFile.Close()
|
_ = privFile.Close()
|
||||||
return fmt.Errorf("generate %s keys: %w", s.name, err)
|
return fmt.Errorf("generate %s keys: %w", scheme.name, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := pubFile.Close(); err != nil {
|
if err := pubFile.Close(); err != nil {
|
||||||
|
|||||||
@@ -74,12 +74,9 @@ func TestKeygenCmdBoth(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Verify all 4 files exist.
|
// Verify all 4 files exist.
|
||||||
files := []string{
|
files := make([]string, 0, 4)
|
||||||
prefix + ".pq.pub.pem",
|
files = append(files, prefix+".pq.pub.pem", prefix+".pq.priv.pem", prefix+".classical.pub.pem", prefix+".classical.priv.pem")
|
||||||
prefix + ".pq.priv.pem",
|
|
||||||
prefix + ".classical.pub.pem",
|
|
||||||
prefix + ".classical.priv.pem",
|
|
||||||
}
|
|
||||||
for _, f := range files {
|
for _, f := range files {
|
||||||
if _, err := os.Stat(f); err != nil {
|
if _, err := os.Stat(f); err != nil {
|
||||||
t.Errorf("expected file %s to exist: %v", f, err)
|
t.Errorf("expected file %s to exist: %v", f, err)
|
||||||
@@ -274,7 +271,10 @@ func TestKeygenCmdPermissions(t *testing.T) {
|
|||||||
t.Fatalf("Execute failed: %v", err)
|
t.Fatalf("Execute failed: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
pubFiles := []string{prefix + ".pq.pub.pem", prefix + ".classical.pub.pem"}
|
pubFiles := make([]string, 0, 2)
|
||||||
|
|
||||||
|
pubFiles = append(pubFiles, prefix+".pq.pub.pem", prefix+".classical.pub.pem")
|
||||||
|
|
||||||
for _, f := range pubFiles {
|
for _, f := range pubFiles {
|
||||||
info, err := os.Stat(f)
|
info, err := os.Stat(f)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -286,7 +286,10 @@ func TestKeygenCmdPermissions(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
privFiles := []string{prefix + ".pq.priv.pem", prefix + ".classical.priv.pem"}
|
privFiles := make([]string, 0, 2)
|
||||||
|
|
||||||
|
privFiles = append(privFiles, prefix+".pq.priv.pem", prefix+".classical.priv.pem")
|
||||||
|
|
||||||
for _, f := range privFiles {
|
for _, f := range privFiles {
|
||||||
info, err := os.Stat(f)
|
info, err := os.Stat(f)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
|
|||||||
@@ -94,9 +94,11 @@ var restoreCmd = &cobra.Command{
|
|||||||
}
|
}
|
||||||
|
|
||||||
decryptor := newRestoreDecryptor(registry)
|
decryptor := newRestoreDecryptor(registry)
|
||||||
|
privateKeys := make([]crypto.RecipientPriv, 0, 2)
|
||||||
|
privateKeys = append(privateKeys, pqPriv, classicalPriv)
|
||||||
if err := decryptor.Decrypt(
|
if err := decryptor.Decrypt(
|
||||||
inFile,
|
inFile,
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
privateKeys,
|
||||||
out,
|
out,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
return err
|
return err
|
||||||
|
|||||||
@@ -153,7 +153,10 @@ func TestRestoreCmd_Structure(t *testing.T) {
|
|||||||
t.Fatalf("expected Use='restore', got %q", restoreCmd.Use)
|
t.Fatalf("expected Use='restore', got %q", restoreCmd.Use)
|
||||||
}
|
}
|
||||||
|
|
||||||
requiredFlags := []string{"in", "privkey-pq", "privkey-classical"}
|
requiredFlags := make([]string, 0, 3)
|
||||||
|
|
||||||
|
requiredFlags = append(requiredFlags, "in", "privkey-pq", "privkey-classical")
|
||||||
|
|
||||||
for _, f := range requiredFlags {
|
for _, f := range requiredFlags {
|
||||||
if restoreCmd.Flag(f) == nil {
|
if restoreCmd.Flag(f) == nil {
|
||||||
t.Fatalf("missing required --%s flag", f)
|
t.Fatalf("missing required --%s flag", f)
|
||||||
@@ -172,7 +175,9 @@ func TestRestoreCmd_Structure(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestRestoreCmd_RequiredFlags(t *testing.T) {
|
func TestRestoreCmd_RequiredFlags(t *testing.T) {
|
||||||
required := []string{"in", "privkey-pq", "privkey-classical"}
|
required := make([]string, 0, 3)
|
||||||
|
required = append(required, "in", "privkey-pq", "privkey-classical")
|
||||||
|
|
||||||
for _, name := range required {
|
for _, name := range required {
|
||||||
flag := restoreCmd.Flag(name)
|
flag := restoreCmd.Flag(name)
|
||||||
if flag == nil {
|
if flag == nil {
|
||||||
|
|||||||
@@ -29,7 +29,7 @@ var (
|
|||||||
) {
|
) {
|
||||||
return cron.NewCronScheduler(expr, job)
|
return cron.NewCronScheduler(expr, job)
|
||||||
}
|
}
|
||||||
newHealthz = func(port int) (*healthz.Server, error) {
|
newHealthz = func(port int) (healthz.Server, error) {
|
||||||
return healthz.New(port)
|
return healthz.New(port)
|
||||||
}
|
}
|
||||||
runOnceFunc = backup.RunOnce
|
runOnceFunc = backup.RunOnce
|
||||||
|
|||||||
@@ -158,7 +158,7 @@ func TestRunCmd_SIGTERM_forcesExitAfterTimeout(t *testing.T) {
|
|||||||
return fake, nil
|
return fake, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
newHealthz = func(port int) (*healthz.Server, error) {
|
newHealthz = func(port int) (healthz.Server, error) {
|
||||||
return healthz.New(port)
|
return healthz.New(port)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -245,7 +245,7 @@ func TestRunCmd_Healthz503DuringShutdown(t *testing.T) {
|
|||||||
t.Fatalf("failed to create healthz server: %v", err)
|
t.Fatalf("failed to create healthz server: %v", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
newHealthz = func(port int) (*healthz.Server, error) {
|
newHealthz = func(port int) (healthz.Server, error) {
|
||||||
return srv, nil
|
return srv, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -292,7 +292,7 @@ func TestRunCmd_Healthz503DuringShutdown(t *testing.T) {
|
|||||||
cancel()
|
cancel()
|
||||||
|
|
||||||
var status503 bool
|
var status503 bool
|
||||||
for i := 0; i < 20; i++ {
|
for iteration := 0; iteration < 20; iteration++ {
|
||||||
resp, err = http.Get("http://" + addr + "/healthz")
|
resp, err = http.Get("http://" + addr + "/healthz")
|
||||||
if err == nil {
|
if err == nil {
|
||||||
resp.Body.Close()
|
resp.Body.Close()
|
||||||
@@ -340,7 +340,7 @@ func TestRunCmd_HealthzBindFailure(t *testing.T) {
|
|||||||
return fake, nil
|
return fake, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
newHealthz = func(port int) (*healthz.Server, error) {
|
newHealthz = func(port int) (healthz.Server, error) {
|
||||||
return nil, errors.New("bind failed")
|
return nil, errors.New("bind failed")
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -74,88 +74,88 @@ func RegisterFlags(cmd *cobra.Command) {
|
|||||||
// launch arguments (Cobra flags) → APP_* env variables → config file → defaults.
|
// launch arguments (Cobra flags) → APP_* env variables → config file → defaults.
|
||||||
func Load(cmd *cobra.Command) (*Config, error) {
|
func Load(cmd *cobra.Command) (*Config, error) {
|
||||||
// Use a fresh Viper instance so that successive calls do not leak state.
|
// Use a fresh Viper instance so that successive calls do not leak state.
|
||||||
v := viper.New()
|
viperInstance := viper.New()
|
||||||
|
|
||||||
// Defaults (lowest priority).
|
// Defaults (lowest priority).
|
||||||
v.SetDefault("pg.port", 5432)
|
viperInstance.SetDefault("pg.port", 5432)
|
||||||
v.SetDefault("pg.sslmode", "prefer")
|
viperInstance.SetDefault("pg.sslmode", "prefer")
|
||||||
v.SetDefault("pg.exclude_tables", []string{"e2e_one_time_keys_json"})
|
viperInstance.SetDefault("pg.exclude_tables", []string{"e2e_one_time_keys_json"})
|
||||||
v.SetDefault("backup.retention_days", 180)
|
viperInstance.SetDefault("backup.retention_days", 180)
|
||||||
v.SetDefault("backup.cron", "0 0 3 * * *")
|
viperInstance.SetDefault("backup.cron", "0 0 3 * * *")
|
||||||
v.SetDefault("pq_scheme", uint16(0x0006))
|
viperInstance.SetDefault("pq_scheme", uint16(0x0006))
|
||||||
v.SetDefault("classical_scheme", uint16(0x0007))
|
viperInstance.SetDefault("classical_scheme", uint16(0x0007))
|
||||||
v.SetDefault("healthz.port", 8080)
|
viperInstance.SetDefault("healthz.port", 8080)
|
||||||
v.SetDefault("log.level", "info")
|
viperInstance.SetDefault("log.level", "info")
|
||||||
v.SetDefault("shutdown_timeout", 30*time.Second)
|
viperInstance.SetDefault("shutdown_timeout", 30*time.Second)
|
||||||
|
|
||||||
// Bind parsed Cobra flags to Viper keys.
|
// Bind parsed Cobra flags to Viper keys.
|
||||||
if cmd != nil {
|
if cmd != nil {
|
||||||
_ = v.BindPFlag("pg.host", cmd.Flags().Lookup("pg-host"))
|
_ = viperInstance.BindPFlag("pg.host", cmd.Flags().Lookup("pg-host"))
|
||||||
_ = v.BindPFlag("pg.port", cmd.Flags().Lookup("pg-port"))
|
_ = viperInstance.BindPFlag("pg.port", cmd.Flags().Lookup("pg-port"))
|
||||||
_ = v.BindPFlag("pg.user", cmd.Flags().Lookup("pg-user"))
|
_ = viperInstance.BindPFlag("pg.user", cmd.Flags().Lookup("pg-user"))
|
||||||
_ = v.BindPFlag("pg.database", cmd.Flags().Lookup("pg-database"))
|
_ = viperInstance.BindPFlag("pg.database", cmd.Flags().Lookup("pg-database"))
|
||||||
_ = v.BindPFlag("pg.sslmode", cmd.Flags().Lookup("pg-sslmode"))
|
_ = viperInstance.BindPFlag("pg.sslmode", cmd.Flags().Lookup("pg-sslmode"))
|
||||||
_ = v.BindPFlag("pg.exclude_tables", cmd.Flags().Lookup("pg-exclude-tables"))
|
_ = viperInstance.BindPFlag("pg.exclude_tables", cmd.Flags().Lookup("pg-exclude-tables"))
|
||||||
_ = v.BindPFlag("backup.dir", cmd.Flags().Lookup("backup-dir"))
|
_ = viperInstance.BindPFlag("backup.dir", cmd.Flags().Lookup("backup-dir"))
|
||||||
_ = v.BindPFlag("backup.retention_days", cmd.Flags().Lookup("backup-retention-days"))
|
_ = viperInstance.BindPFlag("backup.retention_days", cmd.Flags().Lookup("backup-retention-days"))
|
||||||
_ = v.BindPFlag("backup.cron", cmd.Flags().Lookup("backup-cron"))
|
_ = viperInstance.BindPFlag("backup.cron", cmd.Flags().Lookup("backup-cron"))
|
||||||
_ = v.BindPFlag("pq_scheme", cmd.Flags().Lookup("pq-scheme"))
|
_ = viperInstance.BindPFlag("pq_scheme", cmd.Flags().Lookup("pq-scheme"))
|
||||||
_ = v.BindPFlag("classical_scheme", cmd.Flags().Lookup("classical-scheme"))
|
_ = viperInstance.BindPFlag("classical_scheme", cmd.Flags().Lookup("classical-scheme"))
|
||||||
_ = v.BindPFlag("pq_public_key_path", cmd.Flags().Lookup("pq-public-key-path"))
|
_ = viperInstance.BindPFlag("pq_public_key_path", cmd.Flags().Lookup("pq-public-key-path"))
|
||||||
_ = v.BindPFlag("classical_public_key_path", cmd.Flags().Lookup("classical-public-key-path"))
|
_ = viperInstance.BindPFlag("classical_public_key_path", cmd.Flags().Lookup("classical-public-key-path"))
|
||||||
_ = v.BindPFlag("healthz.port", cmd.Flags().Lookup("healthz-port"))
|
_ = viperInstance.BindPFlag("healthz.port", cmd.Flags().Lookup("healthz-port"))
|
||||||
_ = v.BindPFlag("log.level", cmd.Flags().Lookup("log-level"))
|
_ = viperInstance.BindPFlag("log.level", cmd.Flags().Lookup("log-level"))
|
||||||
_ = v.BindPFlag("shutdown_timeout", cmd.Flags().Lookup("shutdown-timeout"))
|
_ = viperInstance.BindPFlag("shutdown_timeout", cmd.Flags().Lookup("shutdown-timeout"))
|
||||||
}
|
}
|
||||||
|
|
||||||
// Config file search with logging.
|
// Config file search with logging.
|
||||||
v.SetConfigName("config")
|
viperInstance.SetConfigName("config")
|
||||||
v.SetConfigType("yaml")
|
viperInstance.SetConfigType("yaml")
|
||||||
|
|
||||||
if envLoc := os.Getenv("APP_CONFIG_LOCATION"); envLoc != "" {
|
if envLoc := os.Getenv("APP_CONFIG_LOCATION"); envLoc != "" {
|
||||||
v.SetConfigFile(envLoc)
|
viperInstance.SetConfigFile(envLoc)
|
||||||
if err := v.ReadInConfig(); err == nil {
|
if err := viperInstance.ReadInConfig(); err == nil {
|
||||||
logf("Config file found: %s (from APP_CONFIG_LOCATION)", v.ConfigFileUsed())
|
logf("Config file found: %s (from APP_CONFIG_LOCATION)", viperInstance.ConfigFileUsed())
|
||||||
} else {
|
} else {
|
||||||
return nil, fmt.Errorf("config file specified in APP_CONFIG_LOCATION not found: %s", envLoc)
|
return nil, fmt.Errorf("config file specified in APP_CONFIG_LOCATION not found: %s", envLoc)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
v.AddConfigPath(".")
|
viperInstance.AddConfigPath(".")
|
||||||
|
|
||||||
homeDir, _ := os.UserHomeDir()
|
homeDir, _ := os.UserHomeDir()
|
||||||
appName := "synapse-backupper"
|
appName := "synapse-backupper"
|
||||||
userConfigPath := filepath.Join(homeDir, ".config", appName)
|
userConfigPath := filepath.Join(homeDir, ".config", appName)
|
||||||
v.AddConfigPath(userConfigPath)
|
viperInstance.AddConfigPath(userConfigPath)
|
||||||
|
|
||||||
exePath, _ := os.Executable()
|
exePath, _ := os.Executable()
|
||||||
exeDir := filepath.Dir(exePath)
|
exeDir := filepath.Dir(exePath)
|
||||||
v.AddConfigPath(exeDir)
|
viperInstance.AddConfigPath(exeDir)
|
||||||
|
|
||||||
etcPath := filepath.Join("/etc", appName)
|
etcPath := filepath.Join("/etc", appName)
|
||||||
v.AddConfigPath(etcPath)
|
viperInstance.AddConfigPath(etcPath)
|
||||||
|
|
||||||
if err := v.ReadInConfig(); err != nil {
|
if err := viperInstance.ReadInConfig(); err != nil {
|
||||||
if _, ok := err.(viper.ConfigFileNotFoundError); ok {
|
if _, ok := err.(viper.ConfigFileNotFoundError); ok {
|
||||||
logf("Configuration file not found. Using defaults and env/args.")
|
logf("Configuration file not found. Using defaults and env/args.")
|
||||||
} else {
|
} else {
|
||||||
return nil, fmt.Errorf("error reading config: %w", err)
|
return nil, fmt.Errorf("error reading config: %w", err)
|
||||||
}
|
}
|
||||||
} else {
|
} else {
|
||||||
logf("Config file found: %s", v.ConfigFileUsed())
|
logf("Config file found: %s", viperInstance.ConfigFileUsed())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Environment variables.
|
// Environment variables.
|
||||||
v.SetEnvPrefix("APP")
|
viperInstance.SetEnvPrefix("APP")
|
||||||
v.SetEnvKeyReplacer(strings.NewReplacer(".", "_", "-", "_"))
|
viperInstance.SetEnvKeyReplacer(strings.NewReplacer(".", "_", "-", "_"))
|
||||||
v.AutomaticEnv()
|
viperInstance.AutomaticEnv()
|
||||||
// pg.password is intentionally not exposed as a CLI flag (CWE-214), but
|
// pg.password is intentionally not exposed as a CLI flag (CWE-214), but
|
||||||
// must still be loadable from APP_PG_PASSWORD. Viper needs an explicit
|
// must still be loadable from APP_PG_PASSWORD. Viper needs an explicit
|
||||||
// BindEnv for a nested key that has no bound flag.
|
// BindEnv for a nested key that has no bound flag.
|
||||||
_ = v.BindEnv("pg.password")
|
_ = viperInstance.BindEnv("pg.password")
|
||||||
|
|
||||||
var cfg Config
|
var cfg Config
|
||||||
if err := v.Unmarshal(&cfg); err != nil {
|
if err := viperInstance.Unmarshal(&cfg); err != nil {
|
||||||
return nil, fmt.Errorf("config unmarshal failed: %w", err)
|
return nil, fmt.Errorf("config unmarshal failed: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -233,40 +233,40 @@ func writeHeader(
|
|||||||
wrapNonce, wrappedCek, firstPayloadNonce []byte,
|
wrapNonce, wrappedCek, firstPayloadNonce []byte,
|
||||||
) error {
|
) error {
|
||||||
var buf bytes.Buffer
|
var buf bytes.Buffer
|
||||||
var b4 [4]byte
|
var scratchBytes [4]byte
|
||||||
|
|
||||||
binary.BigEndian.PutUint32(b4[:], magic)
|
binary.BigEndian.PutUint32(scratchBytes[:], magic)
|
||||||
buf.Write(b4[:]) // [0:4] magic
|
buf.Write(scratchBytes[:]) // [0:4] magic
|
||||||
|
|
||||||
binary.BigEndian.PutUint16(b4[:2], version)
|
binary.BigEndian.PutUint16(scratchBytes[:2], version)
|
||||||
buf.Write(b4[:2]) // [4:6] version
|
buf.Write(scratchBytes[:2]) // [4:6] version
|
||||||
|
|
||||||
binary.BigEndian.PutUint32(b4[:], flags)
|
binary.BigEndian.PutUint32(scratchBytes[:], flags)
|
||||||
buf.Write(b4[:]) // [6:10] flags
|
buf.Write(scratchBytes[:]) // [6:10] flags
|
||||||
|
|
||||||
// [10] nRecipients — composite v2 always carries exactly two slots.
|
// [10] nRecipients — composite v2 always carries exactly two slots.
|
||||||
buf.WriteByte(byte(maxRecipients))
|
buf.WriteByte(byte(maxRecipients))
|
||||||
|
|
||||||
// Slot 0 (PQ).
|
// Slot 0 (PQ).
|
||||||
binary.BigEndian.PutUint16(b4[:2], pqPub.SchemeID())
|
binary.BigEndian.PutUint16(scratchBytes[:2], pqPub.SchemeID())
|
||||||
buf.Write(b4[:2])
|
buf.Write(scratchBytes[:2])
|
||||||
if len(pqPub.KeyID()) != 8 {
|
if len(pqPub.KeyID()) != 8 {
|
||||||
return ErrMalformedHeader
|
return ErrMalformedHeader
|
||||||
}
|
}
|
||||||
buf.Write(pqPub.KeyID())
|
buf.Write(pqPub.KeyID())
|
||||||
binary.BigEndian.PutUint32(b4[:], uint32(len(pqCt)))
|
binary.BigEndian.PutUint32(scratchBytes[:], uint32(len(pqCt)))
|
||||||
buf.Write(b4[:])
|
buf.Write(scratchBytes[:])
|
||||||
buf.Write(pqCt)
|
buf.Write(pqCt)
|
||||||
|
|
||||||
// Slot 1 (classical).
|
// Slot 1 (classical).
|
||||||
binary.BigEndian.PutUint16(b4[:2], classicalPub.SchemeID())
|
binary.BigEndian.PutUint16(scratchBytes[:2], classicalPub.SchemeID())
|
||||||
buf.Write(b4[:2])
|
buf.Write(scratchBytes[:2])
|
||||||
if len(classicalPub.KeyID()) != 8 {
|
if len(classicalPub.KeyID()) != 8 {
|
||||||
return ErrMalformedHeader
|
return ErrMalformedHeader
|
||||||
}
|
}
|
||||||
buf.Write(classicalPub.KeyID())
|
buf.Write(classicalPub.KeyID())
|
||||||
binary.BigEndian.PutUint32(b4[:], uint32(len(classicalCt)))
|
binary.BigEndian.PutUint32(scratchBytes[:], uint32(len(classicalCt)))
|
||||||
buf.Write(b4[:])
|
buf.Write(scratchBytes[:])
|
||||||
buf.Write(classicalCt)
|
buf.Write(classicalCt)
|
||||||
|
|
||||||
// wrapNonce + wrappedCEK + firstPayloadNonce.
|
// wrapNonce + wrappedCEK + firstPayloadNonce.
|
||||||
|
|||||||
@@ -18,6 +18,16 @@ import (
|
|||||||
"git.tswf.io/infra/go-synapse-backupper/pkg/domain/crypto"
|
"git.tswf.io/infra/go-synapse-backupper/pkg/domain/crypto"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func makeRecipientPubs(pqPub, classicalPub crypto.RecipientPub) []crypto.RecipientPub {
|
||||||
|
pubs := make([]crypto.RecipientPub, 0, 2)
|
||||||
|
return append(pubs, pqPub, classicalPub)
|
||||||
|
}
|
||||||
|
|
||||||
|
func makeRecipientPrivs(pqPriv, classicalPriv crypto.RecipientPriv) []crypto.RecipientPriv {
|
||||||
|
privs := make([]crypto.RecipientPriv, 0, 2)
|
||||||
|
return append(privs, pqPriv, classicalPriv)
|
||||||
|
}
|
||||||
|
|
||||||
// ---------------------------------------------------------------------------
|
// ---------------------------------------------------------------------------
|
||||||
// Test KEM harness
|
// Test KEM harness
|
||||||
//
|
//
|
||||||
@@ -286,7 +296,7 @@ func TestGoldenFormat(
|
|||||||
out := &bytes.Buffer{}
|
out := &bytes.Buffer{}
|
||||||
if err := dec.Decrypt(
|
if err := dec.Decrypt(
|
||||||
bytes.NewReader(goldenBytes),
|
bytes.NewReader(goldenBytes),
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
out,
|
out,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
t.Fatalf("Decrypt(golden) failed: %v", err)
|
t.Fatalf("Decrypt(golden) failed: %v", err)
|
||||||
@@ -304,8 +314,8 @@ func TestRoundTrip(
|
|||||||
for _, size := range sizes {
|
for _, size := range sizes {
|
||||||
t.Run(fmt.Sprintf("size=%d", size), func(t *testing.T) {
|
t.Run(fmt.Sprintf("size=%d", size), func(t *testing.T) {
|
||||||
plaintext := make([]byte, size)
|
plaintext := make([]byte, size)
|
||||||
for i := 0; i < size; i++ {
|
for index := 0; index < size; index++ {
|
||||||
plaintext[i] = byte(i)
|
plaintext[index] = byte(index)
|
||||||
}
|
}
|
||||||
|
|
||||||
reg := fakeRegistry(t)
|
reg := fakeRegistry(t)
|
||||||
@@ -318,7 +328,7 @@ func TestRoundTrip(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader(plaintext),
|
bytes.NewReader(plaintext),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -328,7 +338,7 @@ func TestRoundTrip(
|
|||||||
out := &bytes.Buffer{}
|
out := &bytes.Buffer{}
|
||||||
if err := dec.Decrypt(
|
if err := dec.Decrypt(
|
||||||
bytes.NewReader(encrypted.Bytes()),
|
bytes.NewReader(encrypted.Bytes()),
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
out,
|
out,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
t.Fatalf("Decrypt: %v", err)
|
t.Fatalf("Decrypt: %v", err)
|
||||||
@@ -354,7 +364,7 @@ func TestEmptyPlaintextSingleFinalChunk(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader(nil),
|
bytes.NewReader(nil),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -394,7 +404,7 @@ func TestExactly64KiBTwoChunks(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader(plaintext),
|
bytes.NewReader(plaintext),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -486,7 +496,7 @@ func TestTamperPayload(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader(plaintext),
|
bytes.NewReader(plaintext),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -504,7 +514,7 @@ func TestTamperPayload(
|
|||||||
out := &bytes.Buffer{}
|
out := &bytes.Buffer{}
|
||||||
err := dec.Decrypt(
|
err := dec.Decrypt(
|
||||||
bytes.NewReader(buf),
|
bytes.NewReader(buf),
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
out,
|
out,
|
||||||
)
|
)
|
||||||
if !errors.Is(err, ErrTamperingDetected) {
|
if !errors.Is(err, ErrTamperingDetected) {
|
||||||
@@ -526,7 +536,7 @@ func TestTamperWrappedCEK(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader([]byte{0xAA}),
|
bytes.NewReader([]byte{0xAA}),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -545,7 +555,7 @@ func TestTamperWrappedCEK(
|
|||||||
out := &bytes.Buffer{}
|
out := &bytes.Buffer{}
|
||||||
err := dec.Decrypt(
|
err := dec.Decrypt(
|
||||||
bytes.NewReader(buf),
|
bytes.NewReader(buf),
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
out,
|
out,
|
||||||
)
|
)
|
||||||
if !errors.Is(err, ErrTamperingDetected) {
|
if !errors.Is(err, ErrTamperingDetected) {
|
||||||
@@ -575,7 +585,7 @@ func TestWrongPrivKey(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader([]byte{0x11, 0x22, 0x33}),
|
bytes.NewReader([]byte{0x11, 0x22, 0x33}),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -605,7 +615,7 @@ func TestFormatConformance(
|
|||||||
var out bytes.Buffer
|
var out bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader([]byte{0x42}),
|
bytes.NewReader([]byte{0x42}),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&out,
|
&out,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -649,7 +659,7 @@ func TestUnsupportedVersionNoGCM(
|
|||||||
out := &bytes.Buffer{}
|
out := &bytes.Buffer{}
|
||||||
err := dec.Decrypt(
|
err := dec.Decrypt(
|
||||||
reader,
|
reader,
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
out,
|
out,
|
||||||
)
|
)
|
||||||
if !errors.Is(err, ErrUnsupportedVersion) {
|
if !errors.Is(err, ErrUnsupportedVersion) {
|
||||||
@@ -723,7 +733,7 @@ func TestMalformedHeaderCtLenOverflow(
|
|||||||
_, classicalPriv := generateFakeKeyPair(t, fakeRegistry(t), fakeClassicalSchemeID, rand.Reader)
|
_, classicalPriv := generateFakeKeyPair(t, fakeRegistry(t), fakeClassicalSchemeID, rand.Reader)
|
||||||
err := dec.Decrypt(
|
err := dec.Decrypt(
|
||||||
reader,
|
reader,
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
&bytes.Buffer{},
|
&bytes.Buffer{},
|
||||||
)
|
)
|
||||||
if !errors.Is(err, ErrMalformedHeader) {
|
if !errors.Is(err, ErrMalformedHeader) {
|
||||||
@@ -754,7 +764,7 @@ func TestZeroLengthNonFinalChunk(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader([]byte{0xAA}),
|
bytes.NewReader([]byte{0xAA}),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -770,7 +780,7 @@ func TestZeroLengthNonFinalChunk(
|
|||||||
out := &bytes.Buffer{}
|
out := &bytes.Buffer{}
|
||||||
err := dec.Decrypt(
|
err := dec.Decrypt(
|
||||||
bytes.NewReader(corrupt.Bytes()),
|
bytes.NewReader(corrupt.Bytes()),
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
out,
|
out,
|
||||||
)
|
)
|
||||||
if !errors.Is(err, ErrMalformedChunk) {
|
if !errors.Is(err, ErrMalformedChunk) {
|
||||||
@@ -792,7 +802,7 @@ func TestOversizedChunk(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader([]byte{0xAA}),
|
bytes.NewReader([]byte{0xAA}),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -808,7 +818,7 @@ func TestOversizedChunk(
|
|||||||
out := &bytes.Buffer{}
|
out := &bytes.Buffer{}
|
||||||
err := dec.Decrypt(
|
err := dec.Decrypt(
|
||||||
bytes.NewReader(corrupt.Bytes()),
|
bytes.NewReader(corrupt.Bytes()),
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
out,
|
out,
|
||||||
)
|
)
|
||||||
if !errors.Is(err, ErrMalformedChunk) {
|
if !errors.Is(err, ErrMalformedChunk) {
|
||||||
@@ -833,7 +843,7 @@ func TestPrematureEOF(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader(plaintext),
|
bytes.NewReader(plaintext),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rand.Reader,
|
rand.Reader,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
@@ -853,7 +863,7 @@ func TestPrematureEOF(
|
|||||||
out := &bytes.Buffer{}
|
out := &bytes.Buffer{}
|
||||||
err := dec.Decrypt(
|
err := dec.Decrypt(
|
||||||
bytes.NewReader(encrypted.Bytes()[:truncatedLen]),
|
bytes.NewReader(encrypted.Bytes()[:truncatedLen]),
|
||||||
[]crypto.RecipientPriv{pqPriv, classicalPriv},
|
makeRecipientPrivs(pqPriv, classicalPriv),
|
||||||
out,
|
out,
|
||||||
)
|
)
|
||||||
if !errors.Is(err, ErrUnexpectedEOF) {
|
if !errors.Is(err, ErrUnexpectedEOF) {
|
||||||
|
|||||||
@@ -88,7 +88,7 @@ func TestGenerateGoldenFixture(
|
|||||||
var encrypted bytes.Buffer
|
var encrypted bytes.Buffer
|
||||||
if err := enc.Encrypt(
|
if err := enc.Encrypt(
|
||||||
bytes.NewReader([]byte{0xAA}),
|
bytes.NewReader([]byte{0xAA}),
|
||||||
[]crypto.RecipientPub{pqPub, classicalPub},
|
makeRecipientPubs(pqPub, classicalPub),
|
||||||
&encrypted,
|
&encrypted,
|
||||||
rng,
|
rng,
|
||||||
); err != nil {
|
); err != nil {
|
||||||
|
|||||||
@@ -136,8 +136,8 @@ func (k *kemAdapter) Decapsulate(
|
|||||||
|
|
||||||
// Reject the all-zero public key (identity point), which yields an all-zero shared secret.
|
// Reject the all-zero public key (identity point), which yields an all-zero shared secret.
|
||||||
allZero := true
|
allZero := true
|
||||||
for _, b := range ciphertext {
|
for _, byteValue := range ciphertext {
|
||||||
if b != 0 {
|
if byteValue != 0 {
|
||||||
allZero = false
|
allZero = false
|
||||||
break
|
break
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -114,24 +114,24 @@ func TestRoundTrip(t *testing.T) {
|
|||||||
func TestRoundTripMany(t *testing.T) {
|
func TestRoundTripMany(t *testing.T) {
|
||||||
adapter := New()
|
adapter := New()
|
||||||
|
|
||||||
for i := 0; i < 1000; i++ {
|
for iteration := 0; iteration < 1000; iteration++ {
|
||||||
pub, priv, err := adapter.GenerateKeyPair(rand.Reader)
|
pub, priv, err := adapter.GenerateKeyPair(rand.Reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("iteration %d: GenerateKeyPair failed: %v", i, err)
|
t.Fatalf("iteration %d: GenerateKeyPair failed: %v", iteration, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ct, ssEnc, err := adapter.Encapsulate(pub, rand.Reader)
|
ct, ssEnc, err := adapter.Encapsulate(pub, rand.Reader)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("iteration %d: Encapsulate failed: %v", i, err)
|
t.Fatalf("iteration %d: Encapsulate failed: %v", iteration, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
ssDec, err := adapter.Decapsulate(priv, ct)
|
ssDec, err := adapter.Decapsulate(priv, ct)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatalf("iteration %d: Decapsulate failed: %v", i, err)
|
t.Fatalf("iteration %d: Decapsulate failed: %v", iteration, err)
|
||||||
}
|
}
|
||||||
|
|
||||||
if !bytes.Equal(ssEnc, ssDec) {
|
if !bytes.Equal(ssEnc, ssDec) {
|
||||||
t.Fatalf("iteration %d: shared secret mismatch", i)
|
t.Fatalf("iteration %d: shared secret mismatch", iteration)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -9,23 +9,34 @@ import (
|
|||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
// Server is a minimal HTTP health check server.
|
// Server is the public interface of the minimal HTTP health check server.
|
||||||
type Server struct {
|
type Server interface {
|
||||||
|
// Addr returns the bound network address (e.g. "127.0.0.1:8080").
|
||||||
|
Addr() string
|
||||||
|
// Start begins serving HTTP requests. It blocks until Stop is called.
|
||||||
|
Start() error
|
||||||
|
// Stop initiates graceful shutdown. After Stop is called the /healthz
|
||||||
|
// endpoint returns 503 while in-flight requests complete.
|
||||||
|
Stop(ctx context.Context) error
|
||||||
|
}
|
||||||
|
|
||||||
|
// server is the private implementation of Server.
|
||||||
|
type server struct {
|
||||||
listener net.Listener
|
listener net.Listener
|
||||||
server *http.Server
|
httpServer *http.Server
|
||||||
shuttingDown atomic.Bool
|
shuttingDown atomic.Bool
|
||||||
}
|
}
|
||||||
|
|
||||||
// New creates a health check server listening on the given port.
|
// New creates a health check server listening on the given port.
|
||||||
// Passing port 0 binds to an available ephemeral port.
|
// Passing port 0 binds to an available ephemeral port.
|
||||||
func New(port int) (*Server, error) {
|
func New(port int) (Server, error) {
|
||||||
listener, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", port))
|
listener, err := net.Listen("tcp", fmt.Sprintf("127.0.0.1:%d", port))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, fmt.Errorf("failed to create listener: %w", err)
|
return nil, fmt.Errorf("failed to create listener: %w", err)
|
||||||
}
|
}
|
||||||
|
|
||||||
s := &Server{listener: listener}
|
s := &server{listener: listener}
|
||||||
s.server = &http.Server{
|
s.httpServer = &http.Server{
|
||||||
Handler: http.HandlerFunc(s.handleHealthz),
|
Handler: http.HandlerFunc(s.handleHealthz),
|
||||||
ReadHeaderTimeout: 5 * time.Second,
|
ReadHeaderTimeout: 5 * time.Second,
|
||||||
}
|
}
|
||||||
@@ -33,7 +44,7 @@ func New(port int) (*Server, error) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Addr returns the bound network address (e.g. "127.0.0.1:8080").
|
// Addr returns the bound network address (e.g. "127.0.0.1:8080").
|
||||||
func (s *Server) Addr() string {
|
func (s *server) Addr() string {
|
||||||
if s.listener == nil {
|
if s.listener == nil {
|
||||||
return ""
|
return ""
|
||||||
}
|
}
|
||||||
@@ -41,21 +52,21 @@ func (s *Server) Addr() string {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Start begins serving HTTP requests. It blocks until Stop is called.
|
// Start begins serving HTTP requests. It blocks until Stop is called.
|
||||||
func (s *Server) Start() error {
|
func (s *server) Start() error {
|
||||||
return s.server.Serve(s.listener)
|
return s.httpServer.Serve(s.listener)
|
||||||
}
|
}
|
||||||
|
|
||||||
// Stop initiates graceful shutdown. After Stop is called the /healthz
|
// Stop initiates graceful shutdown. After Stop is called the /healthz
|
||||||
// endpoint returns 503 while in-flight requests complete.
|
// endpoint returns 503 while in-flight requests complete.
|
||||||
func (s *Server) Stop(ctx context.Context) error {
|
func (s *server) Stop(ctx context.Context) error {
|
||||||
s.shuttingDown.Store(true)
|
s.shuttingDown.Store(true)
|
||||||
// Grace period so that health checks can observe the 503 state
|
// Grace period so that health checks can observe the 503 state
|
||||||
// before the listener is closed.
|
// before the listener is closed.
|
||||||
time.Sleep(100 * time.Millisecond)
|
time.Sleep(100 * time.Millisecond)
|
||||||
return s.server.Shutdown(ctx)
|
return s.httpServer.Shutdown(ctx)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (s *Server) handleHealthz(w http.ResponseWriter, r *http.Request) {
|
func (s *server) handleHealthz(w http.ResponseWriter, r *http.Request) {
|
||||||
if r.URL.Path != "/healthz" {
|
if r.URL.Path != "/healthz" {
|
||||||
http.NotFound(w, r)
|
http.NotFound(w, r)
|
||||||
return
|
return
|
||||||
|
|||||||
@@ -44,10 +44,8 @@ func mockCommandContext(
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDump_Success(t *testing.T) {
|
func TestDump_Success(t *testing.T) {
|
||||||
wantArgs := []string{
|
wantArgs := make([]string, 0, 2)
|
||||||
"--format=custom",
|
wantArgs = append(wantArgs, "--format=custom", "--exclude-table=e2e_one_time_keys_json")
|
||||||
"--exclude-table=e2e_one_time_keys_json",
|
|
||||||
}
|
|
||||||
|
|
||||||
adapter := &adapter{
|
adapter := &adapter{
|
||||||
commandContext: mockCommandContext(
|
commandContext: mockCommandContext(
|
||||||
@@ -78,10 +76,8 @@ func TestDump_Success(t *testing.T) {
|
|||||||
func TestDump_WaitAfterStdoutEOF(t *testing.T) {
|
func TestDump_WaitAfterStdoutEOF(t *testing.T) {
|
||||||
// This test verifies that after io.Copy returns (stdout EOF),
|
// This test verifies that after io.Copy returns (stdout EOF),
|
||||||
// cmd.Wait() is called and the exit code is verified before returning.
|
// cmd.Wait() is called and the exit code is verified before returning.
|
||||||
wantArgs := []string{
|
wantArgs := make([]string, 0, 2)
|
||||||
"--format=custom",
|
wantArgs = append(wantArgs, "--format=custom", "--exclude-table=e2e_one_time_keys_json")
|
||||||
"--exclude-table=e2e_one_time_keys_json",
|
|
||||||
}
|
|
||||||
|
|
||||||
adapter := &adapter{
|
adapter := &adapter{
|
||||||
commandContext: mockCommandContext(
|
commandContext: mockCommandContext(
|
||||||
@@ -110,10 +106,8 @@ func TestDump_WaitAfterStdoutEOF(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDump_NonZeroExitCode(t *testing.T) {
|
func TestDump_NonZeroExitCode(t *testing.T) {
|
||||||
wantArgs := []string{
|
wantArgs := make([]string, 0, 2)
|
||||||
"--format=custom",
|
wantArgs = append(wantArgs, "--format=custom", "--exclude-table=e2e_one_time_keys_json")
|
||||||
"--exclude-table=e2e_one_time_keys_json",
|
|
||||||
}
|
|
||||||
|
|
||||||
adapter := &adapter{
|
adapter := &adapter{
|
||||||
commandContext: mockCommandContext(
|
commandContext: mockCommandContext(
|
||||||
@@ -200,10 +194,8 @@ func TestDump_ContextCancellation(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDump_DefaultExcludeTables(t *testing.T) {
|
func TestDump_DefaultExcludeTables(t *testing.T) {
|
||||||
wantArgs := []string{
|
wantArgs := make([]string, 0, 2)
|
||||||
"--format=custom",
|
wantArgs = append(wantArgs, "--format=custom", "--exclude-table=e2e_one_time_keys_json")
|
||||||
"--exclude-table=e2e_one_time_keys_json",
|
|
||||||
}
|
|
||||||
|
|
||||||
adapter := &adapter{
|
adapter := &adapter{
|
||||||
commandContext: mockCommandContext(
|
commandContext: mockCommandContext(
|
||||||
@@ -227,11 +219,8 @@ func TestDump_DefaultExcludeTables(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
func TestDump_CustomExcludeTables(t *testing.T) {
|
func TestDump_CustomExcludeTables(t *testing.T) {
|
||||||
wantArgs := []string{
|
wantArgs := make([]string, 0, 3)
|
||||||
"--format=custom",
|
wantArgs = append(wantArgs, "--format=custom", "--exclude-table=table_a", "--exclude-table=table_b")
|
||||||
"--exclude-table=table_a",
|
|
||||||
"--exclude-table=table_b",
|
|
||||||
}
|
|
||||||
|
|
||||||
adapter := &adapter{
|
adapter := &adapter{
|
||||||
commandContext: mockCommandContext(
|
commandContext: mockCommandContext(
|
||||||
@@ -288,12 +277,10 @@ func TestDump_EnvVars(t *testing.T) {
|
|||||||
t.Errorf("PGPASSWORD must not be passed to pg_dump subprocess")
|
t.Errorf("PGPASSWORD must not be passed to pg_dump subprocess")
|
||||||
}
|
}
|
||||||
|
|
||||||
wantEnvVars := []string{
|
wantEnvVars := make([]string, 0, 4)
|
||||||
"PGHOST=myhost",
|
|
||||||
"PGPORT=5433",
|
wantEnvVars = append(wantEnvVars, "PGHOST=myhost", "PGPORT=5433", "PGUSER=myuser", "PGDATABASE=mydb")
|
||||||
"PGUSER=myuser",
|
|
||||||
"PGDATABASE=mydb",
|
|
||||||
}
|
|
||||||
for _, wantEnv := range wantEnvVars {
|
for _, wantEnv := range wantEnvVars {
|
||||||
if !strings.Contains(envStr, wantEnv) {
|
if !strings.Contains(envStr, wantEnv) {
|
||||||
t.Errorf("env missing %q", wantEnv)
|
t.Errorf("env missing %q", wantEnv)
|
||||||
|
|||||||
@@ -9,6 +9,18 @@ import (
|
|||||||
"git.tswf.io/infra/go-synapse-backupper/pkg/domain/pgdump"
|
"git.tswf.io/infra/go-synapse-backupper/pkg/domain/pgdump"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// Runner orchestrates the dump → encrypt → sink pipeline.
|
||||||
|
type Runner interface {
|
||||||
|
// Run executes the full backup pipeline: pg_dump → encrypt → sink.
|
||||||
|
Run(
|
||||||
|
ctx context.Context,
|
||||||
|
pgDumpOpts pgdump.Options,
|
||||||
|
recipients []crypto.RecipientPub,
|
||||||
|
sink domain.Sink,
|
||||||
|
rand io.Reader,
|
||||||
|
) error
|
||||||
|
}
|
||||||
|
|
||||||
// Option configures a Runner.
|
// Option configures a Runner.
|
||||||
type Option func(*runner)
|
type Option func(*runner)
|
||||||
|
|
||||||
@@ -26,14 +38,14 @@ func WithEncryptor(encryptor crypto.Encryptor) Option {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// Runner orchestrates the dump → encrypt → sink pipeline.
|
// runner is the private implementation of Runner.
|
||||||
type runner struct {
|
type runner struct {
|
||||||
dumper pgdump.Dumper
|
dumper pgdump.Dumper
|
||||||
encryptor crypto.Encryptor
|
encryptor crypto.Encryptor
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewRunner creates a pipeline runner with the given functional options.
|
// NewRunner creates a pipeline runner with the given functional options.
|
||||||
func NewRunner(options ...Option) *runner {
|
func NewRunner(options ...Option) Runner {
|
||||||
r := &runner{}
|
r := &runner{}
|
||||||
for _, option := range options {
|
for _, option := range options {
|
||||||
option(r)
|
option(r)
|
||||||
|
|||||||
@@ -19,9 +19,9 @@ type localSink struct {
|
|||||||
dir string
|
dir string
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sink *localSink) Begin(key string) (domain.SinkTx, error) {
|
func (l *localSink) Begin(key string) (domain.SinkTx, error) {
|
||||||
tmpPath := filepath.Join(sink.dir, key+".tmp")
|
tmpPath := filepath.Join(l.dir, key+".tmp")
|
||||||
finalPath := filepath.Join(sink.dir, key)
|
finalPath := filepath.Join(l.dir, key)
|
||||||
|
|
||||||
file, err := os.Create(tmpPath)
|
file, err := os.Create(tmpPath)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -35,8 +35,8 @@ func (sink *localSink) Begin(key string) (domain.SinkTx, error) {
|
|||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sink *localSink) List(prefix string) ([]string, error) {
|
func (l *localSink) List(prefix string) ([]string, error) {
|
||||||
entries, err := os.ReadDir(sink.dir)
|
entries, err := os.ReadDir(l.dir)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return nil, err
|
return nil, err
|
||||||
}
|
}
|
||||||
@@ -56,8 +56,8 @@ func (sink *localSink) List(prefix string) ([]string, error) {
|
|||||||
return keys, nil
|
return keys, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (sink *localSink) Remove(key string) error {
|
func (l *localSink) Remove(key string) error {
|
||||||
path := filepath.Join(sink.dir, key)
|
path := filepath.Join(l.dir, key)
|
||||||
return os.Remove(path)
|
return os.Remove(path)
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -70,39 +70,39 @@ type localSinkTx struct {
|
|||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
}
|
}
|
||||||
|
|
||||||
func (transaction *localSinkTx) Write(p []byte) (int, error) {
|
func (t *localSinkTx) Write(p []byte) (int, error) {
|
||||||
transaction.mu.Lock()
|
t.mu.Lock()
|
||||||
defer transaction.mu.Unlock()
|
defer t.mu.Unlock()
|
||||||
|
|
||||||
if transaction.committed || transaction.aborted {
|
if t.committed || t.aborted {
|
||||||
return 0, errors.New("transaction already finished")
|
return 0, errors.New("transaction already finished")
|
||||||
}
|
}
|
||||||
|
|
||||||
return transaction.file.Write(p)
|
return t.file.Write(p)
|
||||||
}
|
}
|
||||||
|
|
||||||
func (transaction *localSinkTx) Commit() error {
|
func (t *localSinkTx) Commit() error {
|
||||||
transaction.mu.Lock()
|
t.mu.Lock()
|
||||||
defer transaction.mu.Unlock()
|
defer t.mu.Unlock()
|
||||||
|
|
||||||
if transaction.committed || transaction.aborted {
|
if t.committed || t.aborted {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
transaction.committed = true
|
t.committed = true
|
||||||
|
|
||||||
if err := transaction.file.Sync(); err != nil {
|
if err := t.file.Sync(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
if err := transaction.file.Close(); err != nil {
|
if err := t.file.Close(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
if err := os.Rename(transaction.tmpPath, transaction.finalPath); err != nil {
|
if err := os.Rename(t.tmpPath, t.finalPath); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
parent, err := os.Open(filepath.Dir(transaction.tmpPath))
|
parent, err := os.Open(filepath.Dir(t.tmpPath))
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
@@ -115,18 +115,18 @@ func (transaction *localSinkTx) Commit() error {
|
|||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (transaction *localSinkTx) Abort() error {
|
func (t *localSinkTx) Abort() error {
|
||||||
transaction.mu.Lock()
|
t.mu.Lock()
|
||||||
defer transaction.mu.Unlock()
|
defer t.mu.Unlock()
|
||||||
|
|
||||||
if transaction.committed || transaction.aborted {
|
if t.committed || t.aborted {
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
transaction.aborted = true
|
t.aborted = true
|
||||||
|
|
||||||
_ = transaction.file.Close()
|
_ = t.file.Close()
|
||||||
_ = os.Remove(transaction.tmpPath)
|
_ = os.Remove(t.tmpPath)
|
||||||
|
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -0,0 +1,13 @@
|
|||||||
|
package resources
|
||||||
|
|
||||||
|
import _ "embed"
|
||||||
|
|
||||||
|
// ConfigEnYAML is the embedded English example configuration file.
|
||||||
|
//
|
||||||
|
//go:embed config.en.yaml
|
||||||
|
var ConfigEnYAML []byte
|
||||||
|
|
||||||
|
// ConfigRuYAML is the embedded Russian example configuration file.
|
||||||
|
//
|
||||||
|
//go:embed config.ru.yaml
|
||||||
|
var ConfigRuYAML []byte
|
||||||
Reference in New Issue
Block a user