From c4cfd3fc74046caefffa63e2901bb593e995c4eb Mon Sep 17 00:00:00 2001 From: Eric Rakestraw Date: Tue, 2 Jun 2026 18:35:15 +0000 Subject: [PATCH] Unify backend open config construction --- internal/app/backends.go | 153 +++++++++------- internal/app/backends_test.go | 316 +++++++++++++++++++++++++++++++++- 2 files changed, 399 insertions(+), 70 deletions(-) diff --git a/internal/app/backends.go b/internal/app/backends.go index 33d5a46..4303533 100644 --- a/internal/app/backends.go +++ b/internal/app/backends.go @@ -37,6 +37,22 @@ type backendFactory struct { readOnlyKnownHosts bool } +type backendOpenSpec struct { + role string + backend string + path string + host string + user string + port int + ssh config.SSH + endpoint string + bucket string + prefix string + region string + forcePath *bool + credentials config.Credentials +} + func newBackendFactory() *backendFactory { return newBackendFactoryWithEnvironment(config.ProcessEnvironment()) } @@ -91,25 +107,11 @@ func newBackendFactoryWithEnvironment(environment config.Environment) *backendFa } func (f *backendFactory) openSource(ctx context.Context, source config.Backend) (storage.Backend, error) { - if source.Backend != config.BackendLocal && source.Backend != config.BackendSSH && source.Backend != config.BackendS3 { - return nil, fmt.Errorf("source backend %s is not implemented for execution", source.Backend) - } - openConfig, err := f.sourceOpenConfig(source) - if err != nil { - return nil, err - } - return f.registry.Open(ctx, source.Backend, openConfig) + return f.openBackend(ctx, backendOpenSpecFromSource(source)) } func (f *backendFactory) openDestination(ctx context.Context, destination config.Destination) (storage.Backend, error) { - if destination.Backend != config.BackendLocal && destination.Backend != config.BackendSSH && destination.Backend != config.BackendS3 { - return nil, fmt.Errorf("backend %s is not implemented for execution", destination.Backend) - } - openConfig, err := f.destinationOpenConfig(destination) - if err != nil { - return nil, err - } - return f.registry.Open(ctx, destination.Backend, openConfig) + return f.openBackend(ctx, backendOpenSpecFromDestination(destination)) } func (f *backendFactory) openLocalPath(ctx context.Context, path string) (storage.Backend, error) { @@ -120,40 +122,51 @@ func (f *backendFactory) resolveCredentials(creds config.Credentials) (config.Re return f.environment.ResolveCredentials(creds) } -func (f *backendFactory) sourceOpenConfig(source config.Backend) (storage.OpenConfig, error) { - cfg := sourceOpenConfig(source) - if source.Backend == config.BackendS3 { - if err := f.addS3Config(cfg, source.Endpoint, source.Bucket, source.Prefix, source.Region, source.ForcePath, source.Creds); err != nil { +func (f *backendFactory) openBackend(ctx context.Context, spec backendOpenSpec) (storage.Backend, error) { + if !backendExecutable(spec.backend) { + if spec.role == "source" { + return nil, fmt.Errorf("source backend %s is not implemented for execution", spec.backend) + } + return nil, fmt.Errorf("backend %s is not implemented for execution", spec.backend) + } + openConfig, err := f.openConfig(spec) + if err != nil { + return nil, err + } + return f.registry.Open(ctx, spec.backend, openConfig) +} + +func backendExecutable(name string) bool { + return name == config.BackendLocal || name == config.BackendSSH || name == config.BackendS3 +} + +func (f *backendFactory) openConfig(spec backendOpenSpec) (storage.OpenConfig, error) { + cfg := storage.OpenConfig{storagePathKey: spec.path} + switch spec.backend { + case config.BackendSSH: + cfg[sshHostKey] = spec.host + cfg[sshUserKey] = spec.user + cfg[sshPortKey] = strconv.Itoa(spec.port) + cfg[sshKeyFileKey] = spec.ssh.KeyFile + cfg[sshKnownHostsKey] = spec.ssh.KnownHosts + cfg[sshHostKeyPolicyKey] = string(spec.ssh.HostKeyPolicy) + cfg[sshReadOnlyHostsKey] = strconv.FormatBool(f.readOnlyKnownHosts) + case config.BackendS3: + if err := f.addS3Config(cfg, spec); err != nil { return nil, err } } - if source.Backend == config.BackendSSH { - cfg[sshReadOnlyHostsKey] = strconv.FormatBool(f.readOnlyKnownHosts) - } return cfg, nil } -func (f *backendFactory) destinationOpenConfig(destination config.Destination) (storage.OpenConfig, error) { - cfg := destinationOpenConfig(destination) - if destination.Backend == config.BackendS3 { - if err := f.addS3Config(cfg, destination.Endpoint, destination.Bucket, destination.Prefix, destination.Region, destination.ForcePath, destination.Creds); err != nil { - return nil, err - } - } - if destination.Backend == config.BackendSSH { - cfg[sshReadOnlyHostsKey] = strconv.FormatBool(f.readOnlyKnownHosts) - } - return cfg, nil -} - -func (f *backendFactory) addS3Config(cfg storage.OpenConfig, endpoint, bucket, prefix, region string, forcePath *bool, creds config.Credentials) error { - cfg[s3EndpointKey] = endpoint - cfg[s3BucketKey] = bucket - cfg[s3PrefixKey] = prefix - cfg[s3RegionKey] = region - cfg[s3ForcePathStyleKey] = strconv.FormatBool(config.ForcePathStyle(forcePath)) - if creds.AccessKeyIDEnv != "" || creds.SecretAccessKeyEnv != "" { - resolved, err := f.resolveCredentials(creds) +func (f *backendFactory) addS3Config(cfg storage.OpenConfig, spec backendOpenSpec) error { + cfg[s3EndpointKey] = spec.endpoint + cfg[s3BucketKey] = spec.bucket + cfg[s3PrefixKey] = spec.prefix + cfg[s3RegionKey] = spec.region + cfg[s3ForcePathStyleKey] = strconv.FormatBool(config.ForcePathStyle(spec.forcePath)) + if spec.credentials.AccessKeyIDEnv != "" || spec.credentials.SecretAccessKeyEnv != "" { + resolved, err := f.resolveCredentials(spec.credentials) if err != nil { return err } @@ -163,30 +176,38 @@ func (f *backendFactory) addS3Config(cfg storage.OpenConfig, endpoint, bucket, p return nil } -func sourceOpenConfig(source config.Backend) storage.OpenConfig { - cfg := storage.OpenConfig{storagePathKey: source.Path} - if source.Backend == config.BackendSSH { - cfg[sshHostKey] = source.Host - cfg[sshUserKey] = source.User - cfg[sshPortKey] = strconv.Itoa(source.Port) - cfg[sshKeyFileKey] = source.SSH.KeyFile - cfg[sshKnownHostsKey] = source.SSH.KnownHosts - cfg[sshHostKeyPolicyKey] = string(source.SSH.HostKeyPolicy) - cfg[sshReadOnlyHostsKey] = "false" +func backendOpenSpecFromSource(source config.Backend) backendOpenSpec { + return backendOpenSpec{ + role: "source", + backend: source.Backend, + path: source.Path, + host: source.Host, + user: source.User, + port: source.Port, + ssh: source.SSH, + endpoint: source.Endpoint, + bucket: source.Bucket, + prefix: source.Prefix, + region: source.Region, + forcePath: source.ForcePath, + credentials: source.Creds, } - return cfg } -func destinationOpenConfig(destination config.Destination) storage.OpenConfig { - cfg := storage.OpenConfig{storagePathKey: destination.Path} - if destination.Backend == config.BackendSSH { - cfg[sshHostKey] = destination.Host - cfg[sshUserKey] = destination.User - cfg[sshPortKey] = strconv.Itoa(destination.Port) - cfg[sshKeyFileKey] = destination.SSH.KeyFile - cfg[sshKnownHostsKey] = destination.SSH.KnownHosts - cfg[sshHostKeyPolicyKey] = string(destination.SSH.HostKeyPolicy) - cfg[sshReadOnlyHostsKey] = "false" +func backendOpenSpecFromDestination(destination config.Destination) backendOpenSpec { + return backendOpenSpec{ + role: "destination", + backend: destination.Backend, + path: destination.Path, + host: destination.Host, + user: destination.User, + port: destination.Port, + ssh: destination.SSH, + endpoint: destination.Endpoint, + bucket: destination.Bucket, + prefix: destination.Prefix, + region: destination.Region, + forcePath: destination.ForcePath, + credentials: destination.Creds, } - return cfg } diff --git a/internal/app/backends_test.go b/internal/app/backends_test.go index 443aaf9..eef2e0d 100644 --- a/internal/app/backends_test.go +++ b/internal/app/backends_test.go @@ -1,13 +1,16 @@ package app import ( + "bytes" "context" + "fmt" "strings" "testing" "gitea.maximumdirect.net/eric/distributor/internal/config" "gitea.maximumdirect.net/eric/distributor/internal/storage" "gitea.maximumdirect.net/eric/distributor/internal/storage/fake" + "gitea.maximumdirect.net/eric/distributor/internal/testutil" ) func TestBackendFactoryOpensLocalSource(t *testing.T) { @@ -275,7 +278,8 @@ func TestBackendFactoryResolvesCredentialsThroughEnvironment(t *testing.T) { } func TestBackendFactoryBuildsSSHSourceOpenConfig(t *testing.T) { - cfg := sourceOpenConfig(config.Backend{ + factory := &backendFactory{environment: config.NewEnvironment(nil, nil)} + cfg, err := factory.openConfig(backendOpenSpecFromSource(config.Backend{ Backend: config.BackendSSH, Host: "source.example.com", User: "reports", @@ -286,7 +290,10 @@ func TestBackendFactoryBuildsSSHSourceOpenConfig(t *testing.T) { KnownHosts: "/home/reports/.ssh/known_hosts", HostKeyPolicy: config.HostKeyPolicyStrict, }, - }) + })) + if err != nil { + t.Fatalf("openConfig() error = %v", err) + } assertOpenConfig(t, cfg, map[string]string{ storagePathKey: "/reports", @@ -300,7 +307,8 @@ func TestBackendFactoryBuildsSSHSourceOpenConfig(t *testing.T) { } func TestBackendFactoryBuildsSSHDestinationOpenConfig(t *testing.T) { - cfg := destinationOpenConfig(config.Destination{ + factory := &backendFactory{environment: config.NewEnvironment(nil, nil)} + cfg, err := factory.openConfig(backendOpenSpecFromDestination(config.Destination{ Backend: config.BackendSSH, Host: "destination.example.com", User: "deploy", @@ -309,7 +317,10 @@ func TestBackendFactoryBuildsSSHDestinationOpenConfig(t *testing.T) { SSH: config.SSH{ HostKeyPolicy: config.HostKeyPolicyAcceptNew, }, - }) + })) + if err != nil { + t.Fatalf("openConfig() error = %v", err) + } assertOpenConfig(t, cfg, map[string]string{ storagePathKey: "/srv/archive", @@ -320,6 +331,257 @@ func TestBackendFactoryBuildsSSHDestinationOpenConfig(t *testing.T) { }) } +func TestBackendFactoryBuildsEquivalentSourceAndDestinationOpenConfig(t *testing.T) { + forcePathStyle := false + tests := []struct { + name string + source config.Backend + destination config.Destination + }{ + { + name: "local", + source: config.Backend{Backend: config.BackendLocal, Path: "/reports"}, + destination: config.Destination{Backend: config.BackendLocal, Path: "/reports"}, + }, + { + name: "ssh", + source: config.Backend{ + Backend: config.BackendSSH, + Host: "reports.example.com", + User: "reports", + Port: 2222, + Path: "/reports", + SSH: config.SSH{ + KeyFile: "/home/reports/.ssh/id_ed25519", + KnownHosts: "/home/reports/.ssh/known_hosts", + HostKeyPolicy: config.HostKeyPolicyStrict, + }, + }, + destination: config.Destination{ + Backend: config.BackendSSH, + Host: "reports.example.com", + User: "reports", + Port: 2222, + Path: "/reports", + SSH: config.SSH{ + KeyFile: "/home/reports/.ssh/id_ed25519", + KnownHosts: "/home/reports/.ssh/known_hosts", + HostKeyPolicy: config.HostKeyPolicyStrict, + }, + }, + }, + { + name: "s3", + source: config.Backend{ + Backend: config.BackendS3, + Endpoint: "https://s3.example.com", + Bucket: "reports", + Prefix: "archive", + Region: "us-west-2", + ForcePath: &forcePathStyle, + }, + destination: config.Destination{ + Backend: config.BackendS3, + Endpoint: "https://s3.example.com", + Bucket: "reports", + Prefix: "archive", + Region: "us-west-2", + ForcePath: &forcePathStyle, + }, + }, + { + name: "s3 explicit credentials", + source: config.Backend{ + Backend: config.BackendS3, + Endpoint: "https://s3.example.com", + Bucket: "reports", + Region: config.DefaultS3Region, + Creds: config.Credentials{ + AccessKeyIDEnv: "ACCESS_KEY_ID", + SecretAccessKeyEnv: "SECRET_ACCESS_KEY", + }, + }, + destination: config.Destination{ + Backend: config.BackendS3, + Endpoint: "https://s3.example.com", + Bucket: "reports", + Region: config.DefaultS3Region, + Creds: config.Credentials{ + AccessKeyIDEnv: "ACCESS_KEY_ID", + SecretAccessKeyEnv: "SECRET_ACCESS_KEY", + }, + }, + }, + } + + factory := &backendFactory{ + environment: config.NewEnvironment(map[string]string{ + "ACCESS_KEY_ID": "secret-access", + "SECRET_ACCESS_KEY": "secret-secret", + }, func(string) (string, bool) { return "", false }), + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + sourceConfig, err := factory.openConfig(backendOpenSpecFromSource(tt.source)) + if err != nil { + t.Fatalf("source openConfig() error = %v", err) + } + destinationConfig, err := factory.openConfig(backendOpenSpecFromDestination(tt.destination)) + if err != nil { + t.Fatalf("destination openConfig() error = %v", err) + } + if !openConfigEqual(sourceConfig, destinationConfig) { + t.Fatalf("source open config = %#v, destination open config = %#v, want equivalent", sourceConfig, destinationConfig) + } + }) + } +} + +func TestBackendFactoryBuildsEquivalentDryRunSSHOpenConfig(t *testing.T) { + factory := &backendFactory{readOnlyKnownHosts: true} + sourceConfig, err := factory.openConfig(backendOpenSpecFromSource(config.Backend{ + Backend: config.BackendSSH, + Host: "reports.example.com", + Port: 22, + Path: "/reports", + SSH: config.SSH{HostKeyPolicy: config.HostKeyPolicyAcceptNew}, + })) + if err != nil { + t.Fatalf("source openConfig() error = %v", err) + } + destinationConfig, err := factory.openConfig(backendOpenSpecFromDestination(config.Destination{ + Backend: config.BackendSSH, + Host: "reports.example.com", + Port: 22, + Path: "/reports", + SSH: config.SSH{HostKeyPolicy: config.HostKeyPolicyAcceptNew}, + })) + if err != nil { + t.Fatalf("destination openConfig() error = %v", err) + } + if !openConfigEqual(sourceConfig, destinationConfig) { + t.Fatalf("source open config = %#v, destination open config = %#v, want equivalent", sourceConfig, destinationConfig) + } + if sourceConfig[sshReadOnlyHostsKey] != "true" { + t.Fatalf("open config %s = %q, want true", sshReadOnlyHostsKey, sourceConfig[sshReadOnlyHostsKey]) + } +} + +func TestConfiguredSourceValidationAndRunUseEquivalentSourceOpenConfig(t *testing.T) { + tests := []struct { + name string + source config.Backend + sourceKey string + dest config.Destination + destKey string + wantFields map[string]string + }{ + { + name: "s3", + source: config.Backend{ + Backend: config.BackendS3, + Endpoint: "https://s3.example.com", + Bucket: "source-bucket", + Prefix: "source-prefix", + Region: config.DefaultS3Region, + }, + sourceKey: "s3:source-bucket", + dest: config.Destination{ + ID: "archive", + Backend: config.BackendS3, + Endpoint: "https://s3.example.com", + Bucket: "destination-bucket", + Region: config.DefaultS3Region, + }, + destKey: "s3:destination-bucket", + wantFields: map[string]string{ + s3EndpointKey: "https://s3.example.com", + s3BucketKey: "source-bucket", + s3PrefixKey: "source-prefix", + s3RegionKey: config.DefaultS3Region, + s3ForcePathStyleKey: "true", + }, + }, + { + name: "ssh", + source: config.Backend{ + Backend: config.BackendSSH, + Host: "ssh.example.com", + User: "reports", + Port: 2222, + Path: "/source", + SSH: config.SSH{ + KeyFile: "/home/reports/.ssh/id_ed25519", + KnownHosts: "/home/reports/.ssh/known_hosts", + HostKeyPolicy: config.HostKeyPolicyStrict, + }, + }, + sourceKey: "ssh:/source", + dest: config.Destination{ + ID: "archive", + Backend: config.BackendSSH, + Host: "ssh.example.com", + Port: 2222, + Path: "/destination", + SSH: config.SSH{HostKeyPolicy: config.HostKeyPolicyStrict}, + }, + destKey: "ssh:/destination", + wantFields: map[string]string{ + storagePathKey: "/source", + sshHostKey: "ssh.example.com", + sshUserKey: "reports", + sshPortKey: "2222", + sshKeyFileKey: "/home/reports/.ssh/id_ed25519", + sshKnownHostsKey: "/home/reports/.ssh/known_hosts", + sshHostKeyPolicyKey: "strict", + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + sourceBackend := fake.New() + testutil.WriteFakeSourceBundle(t, sourceBackend, "", testutil.BundleOptions{ID: "reports.source"}) + destinationBackend := fake.New() + var validateSourceConfig storage.OpenConfig + var runSourceConfig storage.OpenConfig + validateProvider := recordingBackendFactoryProvider(t, map[string]storage.Backend{ + tt.sourceKey: sourceBackend, + tt.destKey: destinationBackend, + }, func(cfg storage.OpenConfig) { + validateSourceConfig = cfg + }) + runProvider := recordingBackendFactoryProvider(t, map[string]storage.Backend{ + tt.sourceKey: sourceBackend, + tt.destKey: destinationBackend, + }, func(cfg storage.OpenConfig) { + runSourceConfig = cfg + }) + cfg := config.Config{Pipelines: []config.Pipeline{{ + ID: "reports", + Source: tt.source, + Destinations: []config.Destination{tt.dest}, + }}} + config.ApplyDefaults(&cfg) + + var validateOutput bytes.Buffer + if err := validateConfigWithBackendFactory(context.Background(), cfg, ValidateOptions{ + PipelineID: "reports", + Stdout: &validateOutput, + }, validateProvider); err != nil { + t.Fatalf("validateConfigWithBackendFactory() error = %v", err) + } + if err := runConfigWithBackendFactory(context.Background(), cfg, RunOptions{}, runProvider); err != nil { + t.Fatalf("runConfigWithBackendFactory() error = %v", err) + } + if !openConfigEqual(validateSourceConfig, runSourceConfig) { + t.Fatalf("validate source config = %#v, run source config = %#v, want equivalent", validateSourceConfig, runSourceConfig) + } + assertOpenConfig(t, runSourceConfig, tt.wantFields) + }) + } +} + func assertOpenConfig(t *testing.T, got map[string]string, want map[string]string) { t.Helper() for key, wantValue := range want { @@ -328,3 +590,49 @@ func assertOpenConfig(t *testing.T, got map[string]string, want map[string]strin } } } + +func openConfigEqual(left, right storage.OpenConfig) bool { + if len(left) != len(right) { + return false + } + for key, leftValue := range left { + if right[key] != leftValue { + return false + } + } + return true +} + +func recordingBackendFactoryProvider(t *testing.T, remoteBackends map[string]storage.Backend, recordSource func(storage.OpenConfig)) backendFactoryProvider { + t.Helper() + return func(environment config.Environment) *backendFactory { + registry := storage.NewRegistry() + if err := registry.Register(config.BackendS3, func(ctx context.Context, cfg storage.OpenConfig) (storage.Backend, error) { + if cfg[s3BucketKey] == "source-bucket" { + recordSource(cfg) + } + key := "s3:" + cfg[s3BucketKey] + backend := remoteBackends[key] + if backend == nil { + return nil, fmt.Errorf("missing fake backend for %s", key) + } + return backend, nil + }); err != nil { + t.Fatalf("register s3 backend: %v", err) + } + if err := registry.Register(config.BackendSSH, func(ctx context.Context, cfg storage.OpenConfig) (storage.Backend, error) { + if cfg[storagePathKey] == "/source" { + recordSource(cfg) + } + key := "ssh:" + cfg[storagePathKey] + backend := remoteBackends[key] + if backend == nil { + return nil, fmt.Errorf("missing fake backend for %s", key) + } + return backend, nil + }); err != nil { + t.Fatalf("register ssh backend: %v", err) + } + return &backendFactory{registry: registry, environment: environment} + } +}