Compare commits

..

7 Commits

Author SHA1 Message Date
vergil_on f772d08a70 style(tests): переименовать счётчики циклов i в iteration/index
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-08-09 00:33:46 +03:00
vergil_on a9f5a74162 refactor(resources): вынести embed-конфиги в корневой resources
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-08-09 00:33:32 +03:00
vergil_on 174b5f9048 style: инициализировать слайсы через make
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-08-09 00:32:47 +03:00
vergil_on 3950dd94ee style: заменить короткие имена переменных
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-08-09 00:32:32 +03:00
vergil_on 3f59a97844 style(storage): привести receiver-ы к первой букве типа
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-08-09 00:32:18 +03:00
vergil_on 9c6fc8cffa refactor(pipeline): ввести Runner интерфейс
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-08-09 00:32:07 +03:00
vergil_on 328f517543 refactor(healthz): скрыть Server за интерфейсом
Ultraworked with [Sisyphus](https://github.com/code-yeongyu/oh-my-openagent)

Co-authored-by: Sisyphus <clio-agent@sisyphuslabs.ai>
2026-08-09 00:31:53 +03:00
23 changed files with 240 additions and 214 deletions
+4 -13
View File
@@ -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() })
+2 -2
View File
@@ -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{}),
+4 -9
View File
@@ -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) {
+13 -11
View File
@@ -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 {
+11 -8
View File
@@ -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 {
+3 -1
View File
@@ -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
+7 -2
View File
@@ -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 {
+1 -1
View File
@@ -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
+4 -4
View File
@@ -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")
} }
+43 -43
View File
@@ -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)
} }
+15 -15
View File
@@ -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.
+31 -21
View File
@@ -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 {
+2 -2
View File
@@ -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
} }
+5 -5
View File
@@ -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)
} }
} }
} }
+23 -12
View File
@@ -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
+14 -27
View File
@@ -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)
+14 -2
View File
@@ -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)
+28 -28
View File
@@ -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
} }
+13
View File
@@ -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