diff --git a/CLAUDE.md b/CLAUDE.md index 2e286619..6cb0c9ab 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -79,7 +79,7 @@ UNCORS follows a clean layered architecture with middleware composition: **`internal/infra`** - Infrastructure services - HTTP client with connection pooling and proxy support -- Logger setup (logs to stderr or file with debug flag) +- Logger setup (logs to stderr or file based on UNCORS_LOGGING env var) - TLS certificate generation and handling **`internal/tui`** - Terminal UI and logging @@ -142,8 +142,7 @@ Key test flags: **Key Config Options** - `proxy`: Upstream proxy URL (optional) -- `interactive`: Enable TUI mode -- `debug`: Enable debug logging +- `interactive`: Enable TUI mode (default: true) - `port`: Listen port (default: 3000) - `mappings`: Array of request mappings (from/to hosts) @@ -177,7 +176,7 @@ Key test flags: 4. Update CONTRIBUTING.md if user-facing ### Debugging -- Enable debug logs: `./uncors -d` (writes to `uncors.log`) +- Enable logging: Set `UNCORS_LOGGING=/path/to/logfile` environment variable - Run single test: `go test -run TestName ./internal/handler/proxy/` - Race detector: Already enabled in `make test` and `make test-cover` - Integration tests: `make test-integration` (slower, real network) diff --git a/internal/cli/generate_certs.go b/internal/cli/generate_certs.go new file mode 100644 index 00000000..d5baf26b --- /dev/null +++ b/internal/cli/generate_certs.go @@ -0,0 +1,28 @@ +package cli + +import ( + "errors" + + "github.com/evg4b/uncors/internal/di" + "github.com/spf13/pflag" +) + +const GenerateCertsCmd = "generate-certs" + +func GenerateCerts(container *di.Container) error { + cmd := container.GenerateCertsCommand() + + flags := pflag.NewFlagSet(GenerateCertsCmd, pflag.ContinueOnError) + cmd.DefineFlags(flags, container.Version()) + + err := flags.Parse(container.Args()) + if err != nil { + if !errors.Is(err, pflag.ErrHelp) { + return err + } + + return nil + } + + return cmd.Execute() +} diff --git a/internal/cli/generate_certs_test.go b/internal/cli/generate_certs_test.go new file mode 100644 index 00000000..da4c33e1 --- /dev/null +++ b/internal/cli/generate_certs_test.go @@ -0,0 +1,29 @@ +package cli_test + +import ( + "testing" + + "github.com/evg4b/uncors/internal/cli" + "github.com/evg4b/uncors/internal/di" + "github.com/stretchr/testify/require" +) + +func TestGenerateCerts(t *testing.T) { + t.Run("returns error for unknown flag", func(t *testing.T) { + err := cli.GenerateCerts(di.NewContainer(di.WithArgs([]string{"--unknown-flag"}))) + require.Error(t, err) + }) + + t.Run("generates CA certificate with valid args", func(t *testing.T) { + // Point HOME to a temp dir so certs go there, not ~/.config/uncors. + t.Setenv("HOME", t.TempDir()) + + err := cli.GenerateCerts(di.NewContainer(di.WithArgs([]string{"--validity-days=7"}))) + require.NoError(t, err) + }) + + t.Run("returns nil for --help flag", func(t *testing.T) { + err := cli.GenerateCerts(di.NewContainer(di.WithArgs([]string{"--help"}))) + require.NoError(t, err) + }) +} diff --git a/internal/cli/run_ineractive.go b/internal/cli/run_ineractive.go new file mode 100644 index 00000000..733b1721 --- /dev/null +++ b/internal/cli/run_ineractive.go @@ -0,0 +1,33 @@ +package cli + +import ( + "context" + + tea "charm.land/bubbletea/v2" + "github.com/evg4b/uncors/internal/config" + "github.com/evg4b/uncors/internal/di" + uncor "github.com/evg4b/uncors/internal/uncors_app" +) + +func runIneractive( + ctx context.Context, + container *di.Container, + cfg *config.UncorsConfig, + cfgPath string, +) error { + app := uncor.NewUncorsApp( + container, + cfgPath, + cfg, + func() *config.UncorsConfig { + reloaded, _, _ := config.LoadConfiguration(container.Fs(), container.Version(), container.Args()) + + return reloaded + }, + ) + + _, err := tea.NewProgram(app, tea.WithContext(ctx)). + Run() + + return err +} diff --git a/internal/cli/run_non_ineractive.go b/internal/cli/run_non_ineractive.go new file mode 100644 index 00000000..571e22f8 --- /dev/null +++ b/internal/cli/run_non_ineractive.go @@ -0,0 +1,124 @@ +package cli + +import ( + "context" + "log" + "os" + "os/signal" + "syscall" + "time" + + "github.com/evg4b/uncors/internal/config" + "github.com/evg4b/uncors/internal/di" + "github.com/evg4b/uncors/internal/server" + "github.com/evg4b/uncors/internal/tui" +) + +const shutdownTimeout = 15 * time.Second + +func runNonIneractive( + ctx context.Context, + container *di.Container, + cfg *config.UncorsConfig, + cfgPath string, +) error { + output := container.CliOutput() + tui.PrintLogo(output, container.Version()) + output.Print("") + output.WarnBox(tui.DisclaimerMessage) + output.Print("") + output.InfoBox(cfg.Mappings.String()) + output.Print("") + + targets, err := container.Targets(cfg) + if err != nil { + return err + } + + srv := container.Server() + + err = srv.Start(ctx, targets) + if err != nil { + return err + } + + go startVersionChecker(ctx, container, cfg.Proxy) + + go func() { + watcher := config.NewWatcher(cfgPath) + + err := watcher.Watch(ctx, func() { reloadServer(ctx, container, srv) }) + if err != nil { + output.Error(err) + } + }() + + go func() { //nolint:gosec // G118: shutdown needs a fresh context because parent ctx is being cancelled + stop := make(chan os.Signal, 1) + signal.Notify(stop, syscall.SIGINT, syscall.SIGTERM, syscall.SIGHUP) + + defer signal.Stop(stop) + + select { + case sig := <-stop: + if sig == syscall.SIGINT { + _, _ = os.Stdout.WriteString("\n") + } + + log.Println("shutdown signal received") + case <-ctx.Done(): + } + + shutdownCtx, cancel := context.WithTimeout(context.Background(), shutdownTimeout) + defer cancel() + + _ = srv.Shutdown(shutdownCtx) + }() + + srv.Wait() + output.Info("Server was stopped") + + return nil +} + +func reloadServer(ctx context.Context, container *di.Container, srv *server.Server) { + output := container.CliOutput() + + newUncorsConfig, _, err := config.LoadConfiguration(container.Fs(), container.Version(), container.Args()) + if err != nil { + output.Error(err) + + return + } + + output.Info("Restarting server....") + + targets, err := container.Targets(newUncorsConfig) + if err != nil { + output.Error(err) + + return + } + + err = srv.Restart(ctx, targets) + if err != nil { + output.Error(err) + + return + } + + output.InfoBox( + "Server restarted", + newUncorsConfig.Mappings.String(), + ) +} + +// startVersionChecker waits for a short delay then checks for a newer release. +func startVersionChecker(ctx context.Context, container *di.Container, proxy string) { + const checkDelay = 50 * time.Millisecond + + time.Sleep(checkDelay) + + container.VersionChecker(proxy). + CheckNewVersion(ctx) +} diff --git a/internal/cli/run_uncors.go b/internal/cli/run_uncors.go new file mode 100644 index 00000000..135fc1f5 --- /dev/null +++ b/internal/cli/run_uncors.go @@ -0,0 +1,42 @@ +package cli + +import ( + "context" + "errors" + "fmt" + "os" + + "github.com/evg4b/uncors/internal/config" + "github.com/evg4b/uncors/internal/di" + "github.com/spf13/pflag" +) + +func RunUncors(ctx context.Context, container *di.Container) error { + uncorsConfig, path, err := config.LoadConfiguration(container.Fs(), container.Version(), container.Args()) + if err != nil { + if errors.Is(err, config.ErrVersionRequested) { + fmt.Fprintln(os.Stdout, container.Version()) + + return nil + } + + if errors.Is(err, pflag.ErrHelp) { + return nil + } + + return err + } + + var runError error + if uncorsConfig.Interactive { + runError = runIneractive(ctx, container, uncorsConfig, path) + } else { + runError = runNonIneractive(ctx, container, uncorsConfig, path) + } + + if runError != nil && !errors.Is(runError, pflag.ErrHelp) { + return runError + } + + return nil +} diff --git a/internal/cli/run_uncors_test.go b/internal/cli/run_uncors_test.go new file mode 100644 index 00000000..14101e4d --- /dev/null +++ b/internal/cli/run_uncors_test.go @@ -0,0 +1,233 @@ +package cli_test + +import ( + "context" + "net" + "net/http" + "os" + "path/filepath" + "strconv" + "testing" + "time" + + "github.com/evg4b/uncors/internal/cli" + "github.com/evg4b/uncors/internal/config" + "github.com/evg4b/uncors/internal/di" + "github.com/evg4b/uncors/testing/hosts" + "github.com/evg4b/uncors/testing/testutils" + "github.com/spf13/afero" + "github.com/stretchr/testify/require" + "gopkg.in/yaml.v3" +) + +// httpMapping builds a minimal valid UncorsConfig for an HTTP proxy on a free port. +func httpMapping(t *testing.T) (*config.UncorsConfig, int) { + t.Helper() + + port := testutils.GetFreePort(t) + + cfg := &config.UncorsConfig{ + Mappings: config.Mappings{{ + From: hosts.Localhost.HTTPPort(port), + To: hosts.Localhost.HTTP(), + }}, + CacheConfig: config.CacheConfig{ + ExpirationTime: config.DefaultExpirationTime, + MaxSize: config.DefaultMaxSize, + Methods: []string{http.MethodGet}, + }, + } + + return cfg, port +} + +// waitForPort blocks until the TCP address accepts connections or times out. +func waitForPort(t *testing.T, addr string) { + t.Helper() + + const ( + dialTimeout = 100 * time.Millisecond + pollTick = 25 * time.Millisecond + readyWait = 5 * time.Second + ) + + deadline := time.Now().Add(readyWait) + + for time.Now().Before(deadline) { + dialer := &net.Dialer{Timeout: dialTimeout} + + conn, err := dialer.DialContext(context.Background(), "tcp", addr) + if err == nil { + conn.Close() + + return + } + + time.Sleep(pollTick) + } + + t.Fatal("port did not become ready within 5s: " + addr) +} + +// writeConfig marshals cfg to YAML and writes it to path on the real OS filesystem. +func writeConfig(t *testing.T, path string, cfg *config.UncorsConfig) { + t.Helper() + + data, err := yaml.Marshal(cfg) + require.NoError(t, err) + require.NoError(t, os.WriteFile(path, data, 0o600)) +} + +// startProxy starts RunUncors in a goroutine and returns a channel that +// receives the error when it exits. The caller must cancel the context and +// drain the channel to ensure the goroutine has fully stopped. +func startProxy(ctx context.Context, fs afero.Fs, args []string) <-chan error { + errCh := make(chan error, 1) + + go func() { + container := di.NewContainer(di.WithFs(fs), di.WithArgs(args)) + defer func() { + errCh <- container.Close() + }() + + errCh <- cli.RunUncors(ctx, container) + }() + + return errCh +} + +func TestRunUncors(t *testing.T) { + t.Run("returns error when LoadConfiguration fails", func(t *testing.T) { + // No --from/--to flags and no config file → "mappings must not be empty" + container := di.NewContainer(di.WithArgs([]string{})) + defer testutils.Close(t, container) + + err := cli.RunUncors(context.Background(), container) + require.Error(t, err) + }) + + t.Run("returns nil for --version flag", func(t *testing.T) { + container := di.NewContainer(di.WithArgs([]string{"--version"})) + defer testutils.Close(t, container) + + err := cli.RunUncors(context.Background(), container) + require.NoError(t, err) + }) + + t.Run("returns nil for --help flag", func(t *testing.T) { + container := di.NewContainer(di.WithArgs([]string{"--help"})) + defer testutils.Close(t, container) + + err := cli.RunUncors(context.Background(), container) + require.NoError(t, err) + }) + + t.Run("non-interactive: starts server and shuts down on context cancellation", func(t *testing.T) { + cfg, port := httpMapping(t) + fs := afero.NewMemMapFs() + + data, err := yaml.Marshal(cfg) + require.NoError(t, err) + require.NoError(t, afero.WriteFile(fs, "/config.yaml", data, 0o600)) + + ctx, cancel := context.WithCancel(context.Background()) + + errCh := startProxy(ctx, fs, []string{"-c", "/config.yaml", "--interactive=false"}) + + waitForPort(t, net.JoinHostPort("127.0.0.1", strconv.Itoa(port))) + cancel() + + select { + case err := <-errCh: + require.NoError(t, err) + case <-time.After(10 * time.Second): + t.Fatal("RunUncors did not exit after context cancellation") + } + }) + + t.Run("non-interactive: returns error when port is already in use", func(t *testing.T) { + cfg, port := httpMapping(t) + + // Occupy the port so srv.Start fails. + lc := &net.ListenConfig{} + + listener, err := lc.Listen(context.Background(), "tcp4", net.JoinHostPort("127.0.0.1", strconv.Itoa(port))) + require.NoError(t, err) + + defer listener.Close() + + fs := afero.NewMemMapFs() + + data, err := yaml.Marshal(cfg) + require.NoError(t, err) + require.NoError(t, afero.WriteFile(fs, "/config.yaml", data, 0o600)) + + container := di.NewContainer(di.WithFs(fs), di.WithArgs([]string{"-c", "/config.yaml", "--interactive=false"})) + defer testutils.Close(t, container) + + err = cli.RunUncors(context.Background(), container) + require.Error(t, err) + }) + + t.Run("non-interactive: reloads valid config on file change", func(t *testing.T) { + dir := t.TempDir() + configPath := filepath.Join(dir, "config.yaml") + + cfg, port := httpMapping(t) + writeConfig(t, configPath, cfg) + + ctx, cancel := context.WithCancel(context.Background()) + + errCh := startProxy(ctx, afero.NewOsFs(), []string{"-c", configPath, "--interactive=false"}) + + waitForPort(t, net.JoinHostPort("127.0.0.1", strconv.Itoa(port))) + + // Overwrite with same valid config — watcher fires, reloadServer runs. + writeConfig(t, configPath, cfg) + + // Let the debounce + reload settle before stopping. + time.Sleep(200 * time.Millisecond) + + cancel() + + select { + case err := <-errCh: + if err != nil { + t.Logf("RunUncors returned error: %v", err) + } + + require.NoError(t, err) + case <-time.After(10 * time.Second): + t.Fatal("RunUncors did not exit after context cancellation") + } + }) + + t.Run("non-interactive: logs error when config reload produces invalid config", func(t *testing.T) { + dir := t.TempDir() + configPath := filepath.Join(dir, "config.yaml") + + cfg, port := httpMapping(t) + writeConfig(t, configPath, cfg) + + ctx, cancel := context.WithCancel(context.Background()) + + errCh := startProxy(ctx, afero.NewOsFs(), []string{"-c", configPath, "--interactive=false"}) + + waitForPort(t, net.JoinHostPort("127.0.0.1", strconv.Itoa(port))) + + // Write an invalid config (empty mappings) — reloadServer returns an error. + require.NoError(t, os.WriteFile(configPath, []byte("mappings: []\n"), 0o600)) + + // Let the debounce + reload settle before stopping. + time.Sleep(200 * time.Millisecond) + + cancel() + + select { + case err := <-errCh: + require.NoError(t, err) + case <-time.After(10 * time.Second): + t.Fatal("RunUncors did not exit after context cancellation") + } + }) +} diff --git a/internal/commands/generate_certs.go b/internal/commands/generate_certs.go index 6f1ddc6b..01f9cb8b 100644 --- a/internal/commands/generate_certs.go +++ b/internal/commands/generate_certs.go @@ -8,6 +8,7 @@ import ( "github.com/evg4b/uncors/internal/contracts" "github.com/evg4b/uncors/internal/helpers" "github.com/evg4b/uncors/internal/server" + "github.com/evg4b/uncors/internal/tui" "github.com/spf13/afero" "github.com/spf13/pflag" ) @@ -32,7 +33,12 @@ func NewGenerateCertsCommand(options ...Option) *GenerateCertsCommand { } // DefineFlags defines command-line flags for the generate-certs command. -func (c *GenerateCertsCommand) DefineFlags(flags *pflag.FlagSet) { +func (c *GenerateCertsCommand) DefineFlags(flags *pflag.FlagSet, version string) { + flags.Usage = func() { + tui.PrintLogo(flags.Output(), version) + fmt.Fprintln(flags.Output(), "") + fmt.Fprintln(flags.Output(), flags.FlagUsages()) + } flags.IntVar(&c.validityDays, "validity-days", defaultValidityDays, "Certificate validity period in days") flags.BoolVar(&c.force, "force", false, "Force overwrite existing CA certificates") } diff --git a/internal/commands/generate_certs_test.go b/internal/commands/generate_certs_test.go index b7d32773..97e91f25 100644 --- a/internal/commands/generate_certs_test.go +++ b/internal/commands/generate_certs_test.go @@ -18,6 +18,7 @@ const ( configDir = ".config" caCertFile = "ca.crt" caKeyFile = "ca.key" + version = "v0.0.0" ) func TestNewGenerateCertsCommand(t *testing.T) { @@ -40,7 +41,7 @@ func TestGenerateCertsCommand_DefineFlags(t *testing.T) { ) flags := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd.DefineFlags(flags) + cmd.DefineFlags(flags, version) flag := flags.Lookup("validity-days") assert.NotNil(t, flag) @@ -55,7 +56,7 @@ func TestGenerateCertsCommand_DefineFlags(t *testing.T) { ) flags := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd.DefineFlags(flags) + cmd.DefineFlags(flags, version) flag := flags.Lookup("force") assert.NotNil(t, flag) @@ -79,7 +80,7 @@ func TestGenerateCertsCommand_Execute(t *testing.T) { commands.WithOutput(mocks.NoopOutput()), ) flags := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd.DefineFlags(flags) + cmd.DefineFlags(flags, version) err := cmd.Execute() require.NoError(t, err) @@ -112,7 +113,7 @@ func TestGenerateCertsCommand_Execute(t *testing.T) { commands.WithOutput(mocks.NoopOutput()), ) flags := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd.DefineFlags(flags) + cmd.DefineFlags(flags, version) err := flags.Set("validity-days", "730") require.NoError(t, err) @@ -146,7 +147,7 @@ func TestGenerateCertsCommand_Execute(t *testing.T) { commands.WithOutput(mocks.NoopOutput()), ) flags1 := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd1.DefineFlags(flags1) + cmd1.DefineFlags(flags1, version) err := cmd1.Execute() require.NoError(t, err) @@ -155,7 +156,7 @@ func TestGenerateCertsCommand_Execute(t *testing.T) { commands.WithOutput(mocks.NoopOutput()), ) flags2 := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd2.DefineFlags(flags2) + cmd2.DefineFlags(flags2, version) err = cmd2.Execute() require.Error(t, err) }) @@ -173,7 +174,7 @@ func TestGenerateCertsCommand_Execute(t *testing.T) { commands.WithOutput(mocks.NoopOutput()), ) flags1 := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd1.DefineFlags(flags1) + cmd1.DefineFlags(flags1, version) err := cmd1.Execute() require.NoError(t, err) @@ -189,7 +190,7 @@ func TestGenerateCertsCommand_Execute(t *testing.T) { commands.WithOutput(mocks.NoopOutput()), ) flags2 := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd2.DefineFlags(flags2) + cmd2.DefineFlags(flags2, version) err = flags2.Set("force", "true") require.NoError(t, err) @@ -216,7 +217,7 @@ func TestGenerateCertsCommand_Execute(t *testing.T) { commands.WithOutput(mocks.NoopOutput()), ) flags := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd.DefineFlags(flags) + cmd.DefineFlags(flags, version) err := cmd.Execute() require.NoError(t, err) @@ -251,7 +252,7 @@ func TestGenerateCertsCommand_Execute(t *testing.T) { commands.WithOutput(mocks.NoopOutput()), ) flags := pflag.NewFlagSet("test", pflag.ContinueOnError) - cmd.DefineFlags(flags) + cmd.DefineFlags(flags, version) err := cmd.Execute() require.Error(t, err) diff --git a/internal/config/cache_config.go b/internal/config/cache_config.go index f067c437..4d418a6b 100644 --- a/internal/config/cache_config.go +++ b/internal/config/cache_config.go @@ -14,9 +14,9 @@ func (g CacheGlobs) Clone() CacheGlobs { } type CacheConfig struct { - ExpirationTime time.Duration `yaml:"expiration-time"` - MaxSize int64 `yaml:"max-size"` - Methods []string `yaml:"methods"` + ExpirationTime time.Duration `yaml:"expiration-time,omitempty"` + MaxSize int64 `yaml:"max-size,omitempty"` + Methods []string `yaml:"methods,omitempty"` } func (c *CacheConfig) Clone() *CacheConfig { diff --git a/internal/config/config.go b/internal/config/config.go index 2ee03f61..a7b6ece4 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -9,22 +9,37 @@ import ( "gopkg.in/yaml.v3" ) +// ErrVersionRequested is returned when the --version flag is set so that the +// caller can exit cleanly after the version has been printed. +var ErrVersionRequested = errors.New("version requested") + type UncorsConfig struct { Mappings Mappings `yaml:"mappings"` Proxy string `yaml:"proxy"` - Debug bool `yaml:"debug"` CacheConfig CacheConfig `yaml:"cache-config"` Interactive bool `yaml:"-"` } -func LoadConfiguration(fs afero.Fs, args []string) (*UncorsConfig, string, error) { - flags := defineFlags() +func LoadConfiguration(fs afero.Fs, version string, args []string) (*UncorsConfig, string, error) { + flags, err := defineFlags(version) + if err != nil { + return nil, "", err + } - err := flags.Parse(args) + err = flags.Parse(args) if err != nil { return nil, "", fmt.Errorf("failed parsing flags: %w", err) } + printVersion, err := flags.GetBool("version") + if err != nil { + return nil, "", err + } + + if printVersion { + return nil, "", ErrVersionRequested + } + cfg := defaultConfig() configPath, _ := flags.GetString("config") @@ -71,10 +86,6 @@ func applyFlagOverrides(cfg *UncorsConfig, flags *pflag.FlagSet) error { cfg.Proxy, _ = flags.GetString("proxy") } - if flags.Changed("debug") { - cfg.Debug, _ = flags.GetBool("debug") - } - if flags.Changed("interactive") { cfg.Interactive, _ = flags.GetBool("interactive") } diff --git a/internal/config/config_test.go b/internal/config/config_test.go index 9b31b418..6db69077 100644 --- a/internal/config/config_test.go +++ b/internal/config/config_test.go @@ -43,7 +43,6 @@ mappings: Accept-Encoding: deflate raw: demo proxy: http://localhost:8080 -debug: true https-port: 8081 cert-file: /etc/certificates/cert-file.pem key-file: /etc/certificates/key-file.key @@ -82,6 +81,8 @@ func makeTestFs(t *testing.T) afero.Fs { }) } +const version = "v0.0.0" + func TestLoadConfiguration(t *testing.T) { fs := makeTestFs(t) @@ -139,7 +140,6 @@ func TestLoadConfiguration(t *testing.T) { }, }, Proxy: hosts.Localhost.HTTPPort(8080).String(), - Debug: true, CacheConfig: config.CacheConfig{ ExpirationTime: time.Hour, MaxSize: 52428800, @@ -189,11 +189,10 @@ func TestLoadConfiguration(t *testing.T) { }, }, { - name: "CLI proxy and debug flags override config file values", + name: "CLI proxy flag overrides config file value", args: []string{ params.Config, fullConfigPath, "--proxy", "http://newproxy:9999", - "--debug=false", }, expected: &config.UncorsConfig{ Mappings: config.Mappings{ @@ -219,7 +218,6 @@ func TestLoadConfiguration(t *testing.T) { }, }, Proxy: "http://newproxy:9999", - Debug: false, CacheConfig: config.CacheConfig{ ExpirationTime: time.Hour, MaxSize: 52428800, Methods: []string{http.MethodGet, http.MethodPost}, @@ -249,7 +247,7 @@ func TestLoadConfiguration(t *testing.T) { for _, testCase := range tests { t.Run(testCase.name, func(t *testing.T) { - actual, _, err := config.LoadConfiguration(fs, testCase.args) + actual, _, err := config.LoadConfiguration(fs, "", testCase.args) require.NoError(t, err) assert.Equal(t, testCase.expected, actual) @@ -260,13 +258,13 @@ func TestLoadConfiguration(t *testing.T) { t.Run("returns config file path", func(t *testing.T) { t.Run("empty when no config file flag", func(t *testing.T) { args := []string{params.From, hosts.Localhost1.HTTP().String(), params.To, hosts.Github.Host().String()} - _, configPath, err := config.LoadConfiguration(afero.NewMemMapFs(), args) + _, configPath, err := config.LoadConfiguration(afero.NewMemMapFs(), version, args) require.NoError(t, err) assert.Empty(t, configPath) }) t.Run("returns the given config path", func(t *testing.T) { - _, configPath, err := config.LoadConfiguration(fs, []string{params.Config, minimalConfigPath}) + _, configPath, err := config.LoadConfiguration(fs, version, []string{params.Config, minimalConfigPath}) require.NoError(t, err) assert.Equal(t, minimalConfigPath, configPath) }) @@ -331,13 +329,18 @@ func TestLoadConfiguration(t *testing.T) { for _, testCase := range tests { t.Run(testCase.name, func(t *testing.T) { - _, _, err := config.LoadConfiguration(fs, testCase.args) + _, _, err := config.LoadConfiguration(fs, version, testCase.args) assert.EqualError(t, err, testCase.expectedErr) }) } }) } +func TestLoadConfiguration_VersionFlag(t *testing.T) { + _, _, err := config.LoadConfiguration(afero.NewMemMapFs(), "1.2.3", []string{"--version"}) + require.ErrorIs(t, err, config.ErrVersionRequested) +} + func TestUncorsConfigValidator(t *testing.T) { mapFs := testutils.FsFromMap(t, map[string]string{}) diff --git a/internal/config/flags.go b/internal/config/flags.go index 6da8727a..5f4373a8 100644 --- a/internal/config/flags.go +++ b/internal/config/flags.go @@ -1,16 +1,30 @@ package config -import "github.com/spf13/pflag" +import ( + "fmt" -func defineFlags() *pflag.FlagSet { + "github.com/evg4b/uncors/internal/tui" + "github.com/spf13/pflag" +) + +func defineFlags(version string) (*pflag.FlagSet, error) { flags := pflag.NewFlagSet("uncors", pflag.ContinueOnError) - flags.Usage = pflag.Usage + flags.Usage = func() { + tui.PrintLogo(flags.Output(), version) + fmt.Fprintln(flags.Output(), "") + fmt.Fprintln(flags.Output(), flags.FlagUsages()) + } flags.StringSliceP("to", "t", []string{}, "Target host with protocol for the resource to be proxied") flags.StringSliceP("from", "f", []string{}, "Local host with protocol for the resource from which proxying will take place") //nolint: lll flags.String("proxy", "", "HTTP/HTTPS proxy for requests to the real server (uses system proxy by default)") - flags.Bool("debug", false, "Show debug output") flags.StringP("config", "c", "", "Path to the configuration file") - flags.Bool("interactive", true, "") + flags.Bool("interactive", true, "Run application in interactive TUI mode") + flags.BoolP("version", "v", false, "Print the version and exit") + + err := flags.MarkHidden("version") + if err != nil { + return nil, err + } - return flags + return flags, nil } diff --git a/internal/config/rewrite.go b/internal/config/rewrite.go index 2877bfc7..ba0c8804 100644 --- a/internal/config/rewrite.go +++ b/internal/config/rewrite.go @@ -10,7 +10,7 @@ import ( type RewritingOption struct { From string `yaml:"from"` To string `yaml:"to"` - Host urlt.Host `yaml:"host"` + Host urlt.Host `yaml:"host,omitempty"` } func (r RewritingOption) Clone() RewritingOption { diff --git a/internal/config/script.go b/internal/config/script.go index 43bca09b..2450ebeb 100644 --- a/internal/config/script.go +++ b/internal/config/script.go @@ -3,6 +3,7 @@ package config import ( "errors" "fmt" + "strings" "github.com/samber/lo" "github.com/spf13/afero" @@ -14,6 +15,30 @@ type Script struct { File string `yaml:"file"` } +// scriptMarshal is the canonical YAML representation of Script. +// Using a flat struct (no inline) and trimming multi-line script strings avoids +// a gopkg.in/yaml.v3 round-trip bug where strings starting with \n are +// serialized as "|4" block scalars with wrong content indentation. +type scriptMarshal struct { + Path string `yaml:"path,omitempty"` + Method string `yaml:"method,omitempty"` + Queries map[string]string `yaml:"queries,omitempty"` + Headers map[string]string `yaml:"headers,omitempty"` + Script string `yaml:"script,omitempty"` + File string `yaml:"file,omitempty"` +} + +func (s Script) MarshalYAML() (any, error) { + return scriptMarshal{ + Path: s.Matcher.Path, + Method: s.Matcher.Method, + Queries: s.Matcher.Queries, + Headers: s.Matcher.Headers, + Script: strings.TrimSpace(s.Script), + File: s.File, + }, nil +} + func (s *Script) Clone() Script { return Script{ Matcher: s.Matcher.Clone(), diff --git a/internal/config/script_test.go b/internal/config/script_test.go index 4bdb5a7a..f242f284 100644 --- a/internal/config/script_test.go +++ b/internal/config/script_test.go @@ -1,6 +1,7 @@ package config_test import ( + "strings" "testing" "github.com/evg4b/uncors/internal/config" @@ -8,8 +9,55 @@ import ( "github.com/go-http-utils/headers" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gopkg.in/yaml.v3" ) +func TestScript_MarshalYAML(t *testing.T) { + t.Run("inline script with leading newline round-trips without error", func(t *testing.T) { + script := config.Script{ + Matcher: config.RequestMatcher{Path: "/hello"}, + Script: ` +response:set_status(200) +response:set_body("ok") +`, + } + + data, err := yaml.Marshal(script) + require.NoError(t, err) + + var decoded config.Script + + err = yaml.Unmarshal(data, &decoded) + require.NoError(t, err) + + assert.Equal(t, "/hello", decoded.Matcher.Path) + // Leading/trailing whitespace is trimmed on marshal; the script content itself is preserved. + assert.Equal(t, strings.TrimSpace(script.Script), decoded.Script) + }) + + t.Run("matcher fields are preserved in round-trip", func(t *testing.T) { + script := config.Script{ + Matcher: config.RequestMatcher{ + Path: "/api/{id}", + Method: "POST", + }, + File: "/scripts/handler.lua", + } + + data, err := yaml.Marshal(script) + require.NoError(t, err) + + var decoded config.Script + + err = yaml.Unmarshal(data, &decoded) + require.NoError(t, err) + + assert.Equal(t, "/api/{id}", decoded.Matcher.Path) + assert.Equal(t, "POST", decoded.Matcher.Method) + assert.Equal(t, "/scripts/handler.lua", decoded.File) + }) +} + func TestRequestMatcher_Clone(t *testing.T) { original := config.RequestMatcher{ Path: "/api/test", diff --git a/internal/config/watcher.go b/internal/config/watcher.go index 541aee3f..aa4f7e62 100644 --- a/internal/config/watcher.go +++ b/internal/config/watcher.go @@ -34,6 +34,10 @@ func (w *Watcher) Watch(ctx context.Context, onChange func()) error { return errAlreadyWatching } + if w.filePath == "" { + return nil + } + _, err := os.Stat(w.filePath) if err != nil { return fmt.Errorf("failed to watch config file '%s': %w", w.filePath, err) @@ -52,9 +56,10 @@ func (w *Watcher) Watch(ctx context.Context, onChange func()) error { err = fsWatcher.Add(dir) if err != nil { - _ = fsWatcher.Close() - - return fmt.Errorf("failed to watch config directory '%s': %w", dir, err) + return errors.Join( + fsWatcher.Close(), + fmt.Errorf("failed to watch config directory '%s': %w", dir, err), + ) } w.fsWatcher = fsWatcher diff --git a/internal/config/watcher_test.go b/internal/config/watcher_test.go index 12b8f190..330c58a9 100644 --- a/internal/config/watcher_test.go +++ b/internal/config/watcher_test.go @@ -168,6 +168,34 @@ func TestNewConfigWatcher(t *testing.T) { assert.True(t, waitForCall(called, watcherTimeout), "onChange not called after second atomic save") }) + t.Run("returns errAlreadyWatching when Watch called twice", func(t *testing.T) { + tmpDir := t.TempDir() + configFile := filepath.Join(tmpDir, "config.yaml") + require.NoError(t, os.WriteFile(configFile, []byte(""), 0o600)) + + ctx := t.Context() + + watcher := config.NewWatcher(configFile) + err := watcher.Watch(ctx, func() {}) + require.NoError(t, err) + + defer testutils.Close(t, watcher) + + err = watcher.Watch(ctx, func() {}) + require.Error(t, err) + }) + + t.Run("Watch with empty path returns nil immediately", func(t *testing.T) { + watcher := config.NewWatcher("") + err := watcher.Watch(context.Background(), func() {}) + require.NoError(t, err) + }) + + t.Run("Close without Watch returns nil", func(t *testing.T) { + watcher := config.NewWatcher("/some/path.yaml") + require.NoError(t, watcher.Close()) + }) + t.Run("stops watching when context is cancelled", func(t *testing.T) { tmpDir := t.TempDir() configFile := filepath.Join(tmpDir, "config.yaml") diff --git a/internal/di/container.go b/internal/di/container.go index 80cd654b..ca631995 100644 --- a/internal/di/container.go +++ b/internal/di/container.go @@ -3,6 +3,7 @@ package di import ( "errors" "io" + "os" "github.com/evg4b/uncors/internal/commands" "github.com/evg4b/uncors/internal/config" @@ -15,6 +16,7 @@ import ( type Container struct { fs afero.Fs stdout io.Writer + args []string version string cliOutput factory[contracts.Output] @@ -47,12 +49,19 @@ func WithFs(fs afero.Fs) ContainerOption { } } +func WithArgs(args []string) ContainerOption { + return func(c *Container) { + c.args = args + } +} + func NewContainer(options ...ContainerOption) *Container { container := &Container{ fs: afero.NewMemMapFs(), stdout: io.Discard, version: "0.0.0", closers: []io.Closer{}, + args: os.Args, } container = helpers.ApplyOptions(container, options) diff --git a/internal/di/override.go b/internal/di/override.go index 1634e50c..36d0cd02 100644 --- a/internal/di/override.go +++ b/internal/di/override.go @@ -2,13 +2,11 @@ package di import "github.com/evg4b/uncors/internal/contracts" -type OverrideFunc func(c *Container) - -func (c *Container) Override(action OverrideFunc) { +func (c *Container) Override(action ContainerOption) { action(c) } -func OverrideCliOutput(factory func() contracts.Output) OverrideFunc { +func WithCliOutput(factory func() contracts.Output) ContainerOption { return func(c *Container) { c.cliOutput = newFactory(factory) } diff --git a/internal/di/public_api.go b/internal/di/public_api.go index 0571a7c4..a9381563 100644 --- a/internal/di/public_api.go +++ b/internal/di/public_api.go @@ -1,7 +1,10 @@ package di import ( + "errors" "io" + "net" + "strconv" "time" "github.com/evg4b/uncors/internal/commands" @@ -24,6 +27,10 @@ import ( "github.com/spf13/afero" ) +func (c *Container) Args() []string { + return c.args +} + func (c *Container) Fs() afero.Fs { return c.fs } @@ -161,3 +168,26 @@ func (c *Container) Router( return infra.CastToContractsHandler(router), err } + +func (c *Container) Targets(cfg *config.UncorsConfig) ([]server.Target, error) { + groupedMappings := cfg.Mappings.GroupByPort() + targets := make([]server.Target, 0, len(groupedMappings)) + errs := make([]error, 0, len(groupedMappings)) + + for _, group := range groupedMappings { + muxRouter, err := c.Router(group.Mappings, &cfg.CacheConfig, cfg.Proxy) + if err != nil { + errs = append(errs, err) + + continue + } + + targets = append(targets, server.Target{ + Address: net.JoinHostPort("127.0.0.1", strconv.Itoa(group.Port)), + Handler: muxRouter, + EnableTLS: group.Scheme == "https", + }) + } + + return targets, errors.Join(errs...) +} diff --git a/internal/di/public_api_test.go b/internal/di/public_api_test.go index 010070de..f6675d95 100644 --- a/internal/di/public_api_test.go +++ b/internal/di/public_api_test.go @@ -2,6 +2,7 @@ package di_test import ( "bytes" + "net/http" "testing" "time" @@ -252,7 +253,7 @@ func TestContainerOverride(t *testing.T) { overrideApplied := false - container.Override(di.OverrideCliOutput(func() contracts.Output { + container.Override(di.WithCliOutput(func() contracts.Output { overrideApplied = true return customOutput @@ -261,7 +262,7 @@ func TestContainerOverride(t *testing.T) { newContainer := di.NewContainer() defer testutils.Close(t, newContainer) - newContainer.Override(di.OverrideCliOutput(func() contracts.Output { + newContainer.Override(di.WithCliOutput(func() contracts.Output { return customOutput })) @@ -277,7 +278,7 @@ func TestContainerOverride(t *testing.T) { sentinel := container.CliOutput() - container.Override(di.OverrideCliOutput(func() contracts.Output { + container.Override(di.WithCliOutput(func() contracts.Output { return sentinel })) @@ -286,6 +287,92 @@ func TestContainerOverride(t *testing.T) { }) } +func TestContainerTargets(t *testing.T) { + defaultCache := config.CacheConfig{ + ExpirationTime: time.Minute, + MaxSize: 1024, + Methods: []string{http.MethodGet}, + } + + t.Run("returns single HTTP target", func(t *testing.T) { + container := di.NewContainer() + defer testutils.Close(t, container) + + cfg := &config.UncorsConfig{ + Mappings: config.Mappings{{ + From: hosts.Localhost.HTTPPort(18080), + To: hosts.Localhost.HTTP(), + }}, + CacheConfig: defaultCache, + } + + targets, err := container.Targets(cfg) + + require.NoError(t, err) + require.Len(t, targets, 1) + assert.Equal(t, "127.0.0.1:18080", targets[0].Address) + assert.False(t, targets[0].EnableTLS) + assert.NotNil(t, targets[0].Handler) + }) + + t.Run("returns HTTPS target with TLS enabled", func(t *testing.T) { + container := di.NewContainer() + defer testutils.Close(t, container) + + cfg := &config.UncorsConfig{ + Mappings: config.Mappings{{ + From: hosts.Localhost.HTTPSPort(18443), + To: hosts.Localhost.HTTP(), + }}, + CacheConfig: defaultCache, + } + + targets, err := container.Targets(cfg) + + require.NoError(t, err) + require.Len(t, targets, 1) + assert.Equal(t, "127.0.0.1:18443", targets[0].Address) + assert.True(t, targets[0].EnableTLS) + }) + + t.Run("groups two mappings on same port into one target", func(t *testing.T) { + container := di.NewContainer() + defer testutils.Close(t, container) + + cfg := &config.UncorsConfig{ + Mappings: config.Mappings{ + {From: hosts.Localhost1.HTTPPort(19000), To: hosts.Localhost.HTTP()}, + {From: hosts.Localhost2.HTTPPort(19000), To: hosts.Localhost.HTTP()}, + }, + CacheConfig: defaultCache, + } + + targets, err := container.Targets(cfg) + + require.NoError(t, err) + assert.Len(t, targets, 1) + assert.Equal(t, "127.0.0.1:19000", targets[0].Address) + }) + + t.Run("returns two targets for mappings on different ports", func(t *testing.T) { + container := di.NewContainer() + defer testutils.Close(t, container) + + cfg := &config.UncorsConfig{ + Mappings: config.Mappings{ + {From: hosts.Localhost1.HTTPPort(19001), To: hosts.Localhost.HTTP()}, + {From: hosts.Localhost2.HTTPPort(19002), To: hosts.Localhost.HTTP()}, + }, + CacheConfig: defaultCache, + } + + targets, err := container.Targets(cfg) + + require.NoError(t, err) + assert.Len(t, targets, 2) + }) +} + func TestContainerClose(t *testing.T) { t.Run("close with no closers succeeds", func(t *testing.T) { container := di.NewContainer() diff --git a/internal/infra/loggings.go b/internal/infra/loggings.go new file mode 100644 index 00000000..c0531b4d --- /dev/null +++ b/internal/infra/loggings.go @@ -0,0 +1,31 @@ +package infra + +import ( + "io" + "log" + "os" + "path/filepath" +) + +const ( + logFileFlags = os.O_CREATE | os.O_WRONLY | os.O_APPEND + logFilePerm = 0o644 +) + +func SetupLogging() { + path := os.Getenv("UNCORS_LOGGING") + if path == "" { + log.SetOutput(io.Discard) + + return + } + + logFile, err := os.OpenFile(filepath.Clean(path), logFileFlags, logFilePerm) + if err != nil { + log.SetOutput(io.Discard) + + return + } + + log.SetOutput(logFile) +} diff --git a/internal/infra/loggings_test.go b/internal/infra/loggings_test.go new file mode 100644 index 00000000..6a9c1035 --- /dev/null +++ b/internal/infra/loggings_test.go @@ -0,0 +1,54 @@ +package infra_test + +import ( + "io" + "log" + "os" + "path/filepath" + "testing" + + "github.com/evg4b/uncors/internal/infra" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func saveLogWriter(t *testing.T) { + t.Helper() + + orig := log.Writer() + + t.Cleanup(func() { log.SetOutput(orig) }) +} + +func TestSetupLogging(t *testing.T) { + t.Run("discards output when UNCORS_LOGGING is empty", func(t *testing.T) { + saveLogWriter(t) + t.Setenv("UNCORS_LOGGING", "") + + infra.SetupLogging() + + assert.Equal(t, io.Discard, log.Writer()) + }) + + t.Run("writes to file when UNCORS_LOGGING points to a valid path", func(t *testing.T) { + saveLogWriter(t) + logPath := filepath.Join(t.TempDir(), "test.log") + t.Setenv("UNCORS_LOGGING", logPath) + + infra.SetupLogging() + + require.NotEqual(t, io.Discard, log.Writer()) + + _, err := os.Stat(logPath) + assert.NoError(t, err) + }) + + t.Run("discards output when log file cannot be opened", func(t *testing.T) { + saveLogWriter(t) + t.Setenv("UNCORS_LOGGING", "/no-such-dir/test.log") + + infra.SetupLogging() + + assert.Equal(t, io.Discard, log.Writer()) + }) +} diff --git a/internal/server/server.go b/internal/server/server.go index 1ef7b477..47ee3dbc 100644 --- a/internal/server/server.go +++ b/internal/server/server.go @@ -29,6 +29,7 @@ type Target struct { type Server struct { sync.WaitGroup + mu sync.RWMutex listeners []*PortListener manager *HostCertManager tracker IRequestTracker @@ -44,6 +45,7 @@ func New(manager *HostCertManager, tracker IRequestTracker) *Server { } func (s *Server) Start(ctx context.Context, targets []Target) error { + s.mu.Lock() s.listeners = lo.Map(targets, func(target Target, _ int) *PortListener { portCtx, portCtxCancel := context.WithCancel(ctx) @@ -65,6 +67,7 @@ func (s *Server) Start(ctx context.Context, targets []Target) error { return portListener }) + s.mu.Unlock() var launchWaitGroup sync.WaitGroup launchWaitGroup.Add(len(s.listeners)) @@ -111,13 +114,17 @@ func (s *Server) Shutdown(ctx context.Context) error { ctx, cancel := context.WithTimeout(ctx, shutdownTimeout) defer cancel() + s.mu.RLock() + listeners := s.listeners + s.mu.RUnlock() + var ( waitGroup sync.WaitGroup errsMu sync.Mutex errs []error ) - for _, server := range s.listeners { + for _, server := range listeners { waitGroup.Add(1) go func(srv *PortListener) { defer waitGroup.Done() @@ -138,6 +145,9 @@ func (s *Server) Shutdown(ctx context.Context) error { } func (s *Server) Restart(ctx context.Context, targets []Target) error { + s.Add(1) + defer s.Done() + err := s.Shutdown(ctx) if err != nil { return err @@ -151,9 +161,13 @@ func (s *Server) Wait() { } func (s *Server) Close() error { + s.mu.RLock() + listeners := s.listeners + s.mu.RUnlock() + var errs []error - for _, portListener := range s.listeners { + for _, portListener := range listeners { err := portListener.Close() if err != nil { errs = append(errs, err) diff --git a/internal/uncors/app.go b/internal/uncors/app.go deleted file mode 100644 index 8b7b5bdf..00000000 --- a/internal/uncors/app.go +++ /dev/null @@ -1,107 +0,0 @@ -package uncors - -import ( - "context" - "errors" - "net" - "strconv" - - "github.com/evg4b/uncors/internal/contracts" - "github.com/evg4b/uncors/internal/di" - "github.com/evg4b/uncors/internal/server" - "github.com/evg4b/uncors/internal/tui" - - "github.com/evg4b/uncors/internal/config" - "github.com/spf13/afero" -) - -const baseAddress = "127.0.0.1" - -type Uncors struct { - fs afero.Fs - - output contracts.Output - server *server.Server - container *di.Container -} - -func CreateUncors(container *di.Container) *Uncors { - return &Uncors{ - fs: container.Fs(), - output: container.CliOutput(), - container: container, - server: container.Server(), - } -} - -func (app *Uncors) Start(ctx context.Context, uncorsConfig *config.UncorsConfig) error { - tui.PrintLogo(app.output, app.container.Version()) - app.output.Print("") - app.output.WarnBox(tui.DisclaimerMessage) - app.output.Print("") - app.output.InfoBox(uncorsConfig.Mappings.String()) - app.output.Print("") - - targets, err := app.mappingsToTarget(uncorsConfig) - if err != nil { - return err - } - - return app.server.Start(ctx, targets) -} - -func (app *Uncors) Restart(ctx context.Context, uncorsConfig *config.UncorsConfig) error { - app.output.Info("Restarting server....") - - targets, err := app.mappingsToTarget(uncorsConfig) - if err != nil { - return err - } - - err = app.server.Restart(ctx, targets) - if err != nil { - return err - } - - app.output.InfoBox( - "Server restarted", - uncorsConfig.Mappings.String(), - ) - - return nil -} - -func (app *Uncors) Close() error { - return app.server.Close() -} - -func (app *Uncors) Wait() { - app.server.Wait() -} - -func (app *Uncors) Shutdown(ctx context.Context) error { - return app.server.Shutdown(ctx) -} - -func (app *Uncors) mappingsToTarget(uncorsConfig *config.UncorsConfig) ([]server.Target, error) { - groupedMappings := uncorsConfig.Mappings.GroupByPort() - targets := make([]server.Target, 0, len(groupedMappings)) - errs := make([]error, 0, len(groupedMappings)) - - for _, group := range groupedMappings { - muxRouter, err := app.container.Router(group.Mappings, &uncorsConfig.CacheConfig, uncorsConfig.Proxy) - if err != nil { - errs = append(errs, err) - - continue - } - - targets = append(targets, server.Target{ - Address: net.JoinHostPort(baseAddress, strconv.Itoa(group.Port)), - Handler: muxRouter, - EnableTLS: group.Scheme == "https", - }) - } - - return targets, errors.Join(errs...) -} diff --git a/internal/uncors/app_test.go b/internal/uncors/app_test.go deleted file mode 100644 index a7e2b5ec..00000000 --- a/internal/uncors/app_test.go +++ /dev/null @@ -1,575 +0,0 @@ -package uncors_test - -import ( - "context" - "crypto/tls" - "crypto/x509" - "fmt" - "io" - "net/http" - "net/url" - "os" - "path/filepath" - "testing" - "time" - - "github.com/evg4b/uncors/internal/config" - "github.com/evg4b/uncors/internal/di" - "github.com/evg4b/uncors/internal/server" - "github.com/evg4b/uncors/internal/uncors" - "github.com/evg4b/uncors/testing/hosts" - "github.com/evg4b/uncors/testing/testutils" - "github.com/spf13/afero" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -const version = "1.0.0" - -func TestCreateUncors(t *testing.T) { - container := di.NewContainer(di.WithVersion(version)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - assert.NotNil(t, app) -} - -func TestUncorsApp(t *testing.T) { - container := di.NewContainer(di.WithVersion(version)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - fs := container.Fs() - - testResponceHeader := "# Test resrver" - hostFmt := func(host string) string { return fmt.Sprintf("\tHost: %v", host) } - methodFmt := func(method string) string { return fmt.Sprintf("\tMethod: %v", method) } - urlFmt := func(method string) string { return fmt.Sprintf("\tURL: %v", method) } - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - fmt.Fprintln(w, testResponceHeader) - fmt.Fprintln(w, methodFmt(r.Method)) - fmt.Fprintln(w, urlFmt(r.URL.String())) - fmt.Fprintln(w, hostFmt(r.Host)) - })) - defer targetServer.Close() - - homeDir, err := os.UserHomeDir() - require.NoError(t, err) - - certPath, keyPath, err := server.GenerateCA(server.CAConfig{ - Fs: fs, - ValidityDays: 10, - OutputDir: filepath.Join(homeDir, ".config", "uncors"), - }) - require.NoError(t, err) - - caCert, _, err := server.LoadCA(fs, certPath, keyPath) - require.NoError(t, err) - - pool := x509.NewCertPool() - pool.AddCert(caCert) - - client := &http.Client{ - Transport: &http.Transport{ - TLSClientConfig: &tls.Config{ - MinVersion: tls.VersionTLS13, - RootCAs: pool, - ServerName: "127.0.0.1", - }, - }, - } - - port := testutils.GetFreePort(t) - - err = app.Start(t.Context(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPPort(port), To: hosts.Parse(targetServer.URL)}, - }, - }) - require.NoError(t, err) - - defer func() { require.NoError(t, app.Close()) }() - - t.Run("proxy", func(t *testing.T) { - req, err := http.NewRequestWithContext( - t.Context(), - http.MethodGet, - hosts.Loopback.HTTPPort(port).String(), - nil, - ) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - bodyData, err := io.ReadAll(resp.Body) - require.NoError(t, err) - resp.Body.Close() - - assert.Equal(t, http.StatusOK, resp.StatusCode) - - uri, err := url.Parse(targetServer.URL) - require.NoError(t, err) - - assert.Contains(t, string(bodyData), uri.Host) - assert.Contains(t, string(bodyData), methodFmt(http.MethodGet)) - }) -} - -func TestUncorsStart(t *testing.T) { - container := di.NewContainer(di.WithVersion(version)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - fmt.Fprint(w, "OK") - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - - err := app.Start(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPPort(port), To: hosts.Parse(targetServer.URL)}, - }, - }) - require.NoError(t, err) - - defer app.Close() - - req, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - hosts.Loopback.HTTPPort(port).String(), - nil, - ) - require.NoError(t, err) - - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - resp.Body.Close() - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Equal(t, "OK", string(body)) -} - -func TestUncorsRestart(t *testing.T) { - container := di.NewContainer(di.WithVersion(version)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - server1 := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - fmt.Fprint(w, "Server 1") - })) - defer server1.Close() - - server2 := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - fmt.Fprint(w, "Server 2") - })) - defer server2.Close() - - port := testutils.GetFreePort(t) - - err := app.Start(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPPort(port), To: hosts.Parse(server1.URL)}, - }, - }) - require.NoError(t, err) - - defer app.Close() - - req1, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - hosts.Loopback.HTTPPort(port).String(), - nil, - ) - require.NoError(t, err) - resp1, err := http.DefaultClient.Do(req1) - require.NoError(t, err) - body1, err := io.ReadAll(resp1.Body) - require.NoError(t, err) - resp1.Body.Close() - assert.Equal(t, "Server 1", string(body1)) - - err = app.Restart(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPPort(port), To: hosts.Parse(server2.URL)}, - }, - }) - require.NoError(t, err) - time.Sleep(100 * time.Millisecond) - - req2, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - hosts.Loopback.HTTPPort(port).String(), - nil, - ) - require.NoError(t, err) - resp2, err := http.DefaultClient.Do(req2) - require.NoError(t, err) - body2, err := io.ReadAll(resp2.Body) - require.NoError(t, err) - resp2.Body.Close() - assert.Equal(t, "Server 2", string(body2)) -} - -func TestUncorsClose(t *testing.T) { - container := di.NewContainer(di.WithVersion(version)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - err := app.Start(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPPort(port), To: hosts.Parse(targetServer.URL)}, - }, - }) - require.NoError(t, err) - - err = app.Close() - require.NoError(t, err) - - req, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - hosts.Loopback.HTTPPort(port).String(), - nil, - ) - require.NoError(t, err) - _, err = http.DefaultClient.Do(req) - assert.Error(t, err) -} - -func TestUncorsShutdown(t *testing.T) { - container := di.NewContainer(di.WithVersion(version)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - time.Sleep(50 * time.Millisecond) - w.WriteHeader(http.StatusOK) - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - err := app.Start(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPPort(port), To: hosts.Parse(targetServer.URL)}, - }, - }) - require.NoError(t, err) - - ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) - defer cancel() - - err = app.Shutdown(ctx) - assert.NoError(t, err) -} - -func TestUncorsWait(t *testing.T) { - container := di.NewContainer(di.WithVersion(version)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - err := app.Start(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPPort(port), To: hosts.Parse(targetServer.URL)}, - }, - }) - require.NoError(t, err) - - done := make(chan bool) - - go func() { - app.Wait() - - done <- true - }() - go func() { - time.Sleep(100 * time.Millisecond) - app.Close() - }() - - select { - case <-done: - case <-time.After(2 * time.Second): - t.Fatal("Wait() did not return in time") - } -} - -func TestUncorsWithHTTPSMapping(t *testing.T) { - fakeHome := t.TempDir() - t.Setenv("HOME", fakeHome) - - fs := afero.NewOsFs() - require.NoError(t, fs.MkdirAll(fakeHome, 0o755)) - - container := di.NewContainer(di.WithFs(fs)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - fmt.Fprint(w, "HTTPS OK") - })) - defer targetServer.Close() - - caDir := filepath.Join(fakeHome, ".config", "uncors") - certPath, keyPath, err := server.GenerateCA(server.CAConfig{ - Fs: fs, - ValidityDays: 10, - OutputDir: caDir, - }) - require.NoError(t, err) - caCert, _, err := server.LoadCA(fs, certPath, keyPath) - require.NoError(t, err) - - pool := x509.NewCertPool() - pool.AddCert(caCert) - client := &http.Client{ - Transport: &http.Transport{ - TLSClientConfig: &tls.Config{ - MinVersion: tls.VersionTLS13, - RootCAs: pool, - ServerName: "127.0.0.1", - }, - }, - } - - port := testutils.GetFreePort(t) - err = app.Start(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPSPort(port), To: hosts.Parse(targetServer.URL)}, - }, - }) - require.NoError(t, err) - - defer app.Close() - - req, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - hosts.Loopback.HTTPSPort(port).String(), - nil, - ) - require.NoError(t, err) - resp, err := client.Do(req) - require.NoError(t, err) - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - resp.Body.Close() - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Equal(t, "HTTPS OK", string(body)) -} - -func TestUncorsWithMixedHTTPAndHTTPS(t *testing.T) { - fakeHome := t.TempDir() - t.Setenv("HOME", fakeHome) - - fs := afero.NewOsFs() - require.NoError(t, fs.MkdirAll(fakeHome, 0o755)) - - container := di.NewContainer(di.WithFs(fs)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - httpServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - fmt.Fprint(w, "HTTP") - })) - defer httpServer.Close() - - httpsServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - fmt.Fprint(w, "HTTPS") - })) - defer httpsServer.Close() - - caDir := filepath.Join(fakeHome, ".config", "uncors") - certPath, keyPath, err := server.GenerateCA(server.CAConfig{ - Fs: fs, ValidityDays: 10, OutputDir: caDir, - }) - require.NoError(t, err) - caCert, _, err := server.LoadCA(fs, certPath, keyPath) - require.NoError(t, err) - - pool := x509.NewCertPool() - pool.AddCert(caCert) - tlsClient := &http.Client{ - Transport: &http.Transport{ - TLSClientConfig: &tls.Config{ - MinVersion: tls.VersionTLS13, - RootCAs: pool, - ServerName: "127.0.0.1", - }, - }, - } - - httpPort := testutils.GetFreePort(t) - httpsPort := testutils.GetFreePort(t) - - err = app.Start(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - {From: hosts.Loopback.HTTPPort(httpPort), To: hosts.Parse(httpServer.URL)}, - {From: hosts.Loopback.HTTPSPort(httpsPort), To: hosts.Parse(httpsServer.URL)}, - }, - }) - require.NoError(t, err) - - defer app.Close() - - t.Run("HTTP endpoint", func(t *testing.T) { - req, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - hosts.Loopback.HTTPPort(httpPort).String(), - nil, - ) - require.NoError(t, err) - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - resp.Body.Close() - assert.Equal(t, "HTTP", string(body)) - }) - - t.Run("HTTPS endpoint", func(t *testing.T) { - req, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - hosts.Loopback.HTTPSPort(httpsPort).String(), - nil, - ) - require.NoError(t, err) - resp, err := tlsClient.Do(req) - require.NoError(t, err) - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - resp.Body.Close() - assert.Equal(t, "HTTPS", string(body)) - }) -} - -func TestUncorsWithComplexConfiguration(t *testing.T) { - container := di.NewContainer(di.WithVersion(version)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - fs := container.Fs() - - require.NoError(t, fs.MkdirAll("/static", 0o755)) - require.NoError(t, afero.WriteFile(fs, "/static/index.html", []byte("Static"), 0o644)) - require.NoError(t, afero.WriteFile(fs, "/mock.json", []byte(`{"mocked":true}`), 0o644)) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - fmt.Fprint(w, "Proxied") - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - err := app.Start(context.Background(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse(targetServer.URL), - Statics: []config.StaticDirectory{ - {Path: "/static", Dir: "/static", Index: "index.html"}, - }, - Mocks: []config.Mock{ - { - Matcher: config.RequestMatcher{Path: "/api/mock"}, - Response: config.Response{ - Code: 200, File: "/mock.json", - }, - }, - }, - Cache: config.CacheGlobs{"/cache/*"}, - }, - }, - CacheConfig: config.CacheConfig{ - Methods: []string{"GET"}, - ExpirationTime: 1 * time.Minute, - MaxSize: 100 * 1024 * 1024, - }, - }) - require.NoError(t, err) - - defer app.Close() - - t.Run("static content", func(t *testing.T) { - req, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - testutils.JoinPath(hosts.Loopback.HTTPPort(port).String(), "static"), - nil, - ) - require.NoError(t, err) - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - resp.Body.Close() - assert.Contains(t, string(body), "Static") - }) - - t.Run("mock endpoint", func(t *testing.T) { - req, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - testutils.JoinPath(hosts.Loopback.HTTPPort(port).String(), "api", "mock"), - nil, - ) - require.NoError(t, err) - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - resp.Body.Close() - assert.JSONEq(t, `{"mocked":true}`, string(body)) - }) - - t.Run("proxied content", func(t *testing.T) { - req, err := http.NewRequestWithContext( - context.Background(), - http.MethodGet, - testutils.JoinPath(hosts.Loopback.HTTPPort(port).String(), "other"), - nil, - ) - require.NoError(t, err) - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - resp.Body.Close() - assert.Equal(t, "Proxied", string(body)) - }) -} diff --git a/internal/uncors/handler_test.go b/internal/uncors/handler_test.go deleted file mode 100644 index 0b49c115..00000000 --- a/internal/uncors/handler_test.go +++ /dev/null @@ -1,681 +0,0 @@ -package uncors_test - -import ( - "crypto/tls" - "crypto/x509" - "fmt" - "io" - "net/http" - "path/filepath" - "testing" - "time" - - "github.com/evg4b/uncors/internal/config" - "github.com/evg4b/uncors/internal/di" - "github.com/evg4b/uncors/internal/server" - "github.com/evg4b/uncors/internal/uncors" - "github.com/evg4b/uncors/testing/hosts" - "github.com/evg4b/uncors/testing/testutils" - "github.com/spf13/afero" - "github.com/stretchr/testify/assert" - "github.com/stretchr/testify/require" -) - -func TestHandlerWithHTTP(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("X-Test-Header", "test-value") - w.WriteHeader(http.StatusOK) - fmt.Fprintf(w, "Hello from target: %s %s", r.Method, r.URL.Path) //nolint:gosec // G705: test handler - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - - err := app.Start(t.Context(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse(targetServer.URL), - }, - }, - }) - require.NoError(t, err) - - defer app.Close() - - methods := []string{ - http.MethodGet, - http.MethodPost, - http.MethodPut, - http.MethodPatch, - http.MethodDelete, - http.MethodHead, - } - - for _, method := range methods { - t.Run(method, func(t *testing.T) { - url := testutils.JoinPath(hosts.Loopback.HTTPPort(port).String(), "api", method) - req, err := http.NewRequestWithContext(t.Context(), method, url, nil) - require.NoError(t, err) - - req.Header.Set("Content-Type", "application/json") - - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Equal(t, "test-value", resp.Header.Get("X-Test-Header")) - - if method != http.MethodHead { - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - assert.Contains(t, string(body), fmt.Sprintf("Hello from target: %s /api/%s", method, method)) - } - }) - } - - t.Run("OPTIONS request", func(t *testing.T) { - req, err := http.NewRequestWithContext( - t.Context(), - http.MethodOptions, - testutils.JoinPath(hosts.Loopback.HTTPPort(port).String(), "/api/OPTIONS"), - nil, - ) - require.NoError(t, err) - - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Contains(t, resp.Header.Get("Access-Control-Allow-Origin"), "*") - }) -} - -func TestHandlerWithHTTPS(t *testing.T) { - fakeHome := t.TempDir() - t.Setenv("HOME", fakeHome) - - fs := afero.NewOsFs() - require.NoError(t, fs.MkdirAll(fakeHome, 0o755)) - - container := di.NewContainer(di.WithFs(fs)) - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.Header().Set("X-Test-Header", "test-value") - w.WriteHeader(http.StatusOK) - fmt.Fprintf(w, "HTTPS response: %s %s", r.Method, r.URL.Path) //nolint:gosec // G705: test handler - })) - defer targetServer.Close() - - caDir := filepath.Join(fakeHome, ".config", "uncors") - certPath, keyPath, err := server.GenerateCA(server.CAConfig{ - Fs: fs, - ValidityDays: 10, - OutputDir: caDir, - }) - require.NoError(t, err) - - caCert, _, err := server.LoadCA(fs, certPath, keyPath) - require.NoError(t, err) - - pool := x509.NewCertPool() - pool.AddCert(caCert) - - client := &http.Client{ - Transport: &http.Transport{ - TLSClientConfig: &tls.Config{ - MinVersion: tls.VersionTLS13, - RootCAs: pool, - ServerName: "127.0.0.1", - }, - }, - } - - port := testutils.GetFreePort(t) - - err = app.Start(t.Context(), &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPSPort(port), - To: hosts.Parse(targetServer.URL), - }, - }, - }) - require.NoError(t, err) - - defer app.Close() - - methods := []string{ - http.MethodGet, - http.MethodPost, - http.MethodPut, - http.MethodPatch, - http.MethodDelete, - http.MethodHead, - } - - for _, method := range methods { - t.Run(method, func(t *testing.T) { - url := testutils.JoinPath(hosts.Loopback.HTTPSPort(port).String(), "secure", method) - req, err := http.NewRequestWithContext(t.Context(), method, url, nil) - require.NoError(t, err) - - req.Header.Set("Content-Type", "application/json") - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Equal(t, "test-value", resp.Header.Get("X-Test-Header")) - - if method != http.MethodHead { - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - assert.Contains(t, string(body), fmt.Sprintf("HTTPS response: %s /secure/%s", method, method)) - } - }) - } - - t.Run("OPTIONS request", func(t *testing.T) { - req, err := http.NewRequestWithContext( - t.Context(), - http.MethodOptions, - testutils.JoinPath(hosts.Loopback.HTTPSPort(port).String(), "/secure/OPTIONS"), - nil, - ) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Contains(t, resp.Header.Get("Access-Control-Allow-Origin"), "*") - }) -} - -func TestHandlerWithMockMiddleware(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - mockFile := "/mock-response.json" - mockContent := `{"message":"mocked"}` - require.NoError(t, afero.WriteFile(container.Fs(), mockFile, []byte(mockContent), 0o644)) - - port := testutils.GetFreePort(t) - - cfg := &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse("http://example.com"), - Mocks: []config.Mock{ - { - Matcher: config.RequestMatcher{ - Path: "/api/mock", - }, - Response: config.Response{ - Code: http.StatusOK, - File: mockFile, - }, - }, - }, - }, - }, - } - - require.NoError(t, app.Start(t.Context(), cfg)) - defer app.Close() - - url := testutils.JoinPath(hosts.Loopback.HTTPPort(port).String(), "api", "mock") - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.JSONEq(t, mockContent, string(body)) -} - -func TestHandlerWithStaticMiddleware(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - fs := container.Fs() - - staticDir := "/static" - indexFile := filepath.Join(staticDir, "index.html") - textFile := filepath.Join(staticDir, "test.txt") - - require.NoError(t, fs.MkdirAll(staticDir, 0o755)) - require.NoError(t, afero.WriteFile(fs, indexFile, []byte("Static Content"), 0o644)) - require.NoError(t, afero.WriteFile(fs, textFile, []byte("test file content"), 0o644)) - - port := testutils.GetFreePort(t) - - cfg := &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse("http://example.com"), - Statics: []config.StaticDirectory{ - { - Path: staticDir, - Dir: staticDir, - Index: "index.html", - }, - }, - }, - }, - } - - require.NoError(t, app.Start(t.Context(), cfg)) - defer app.Close() - - t.Run("serve index file", func(t *testing.T) { - url := testutils.JoinPath(hosts.Loopback.HTTPPort(port).String(), "static", "/") - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - assert.Contains(t, string(body), "Static Content") - }) - - t.Run("serve specific file", func(t *testing.T) { - url := testutils.JoinPath(hosts.Loopback.HTTPPort(port).String(), "static", "test.txt") - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Equal(t, "test file content", string(body)) - }) -} - -func TestHandlerWithCache(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - callCount := 0 - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - callCount++ - - w.WriteHeader(http.StatusOK) - fmt.Fprintf(w, "Response #%d", callCount) - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - - cfg := &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse(targetServer.URL), - Cache: config.CacheGlobs{ - "/cached/*", - }, - }, - }, - CacheConfig: config.CacheConfig{ - Methods: []string{http.MethodGet}, - ExpirationTime: time.Minute, - MaxSize: 100 * 1024 * 1024, - }, - } - - require.NoError(t, app.Start(t.Context(), cfg)) - defer app.Close() - - client := http.DefaultClient - baseURL := hosts.Loopback.HTTPPort(port).String() - - t.Run("first request", func(t *testing.T) { - url := testutils.JoinPath(baseURL, "cached", "test") - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - assert.Contains(t, string(body), "Response #1") - }) - - t.Run("cached request", func(t *testing.T) { - url := testutils.JoinPath(baseURL, "cached", "test") - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - assert.Contains(t, string(body), "Response #1") - assert.Equal(t, 1, callCount, "should use cached response") - }) - - t.Run("non-cached path", func(t *testing.T) { - url := testutils.JoinPath(baseURL, "other", "path") - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - assert.Contains(t, string(body), "Response #2") - assert.Equal(t, 2, callCount, "should not use cache for different path") - }) -} - -func TestHandlerWithMultipleMappings(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - server1 := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - fmt.Fprint(w, "Server 1") - })) - defer server1.Close() - - server2 := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - fmt.Fprint(w, "Server 2") - })) - defer server2.Close() - - port1 := testutils.GetFreePort(t) - port2 := testutils.GetFreePort(t) - - cfg := &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port1), - To: hosts.Parse(server1.URL), - }, - { - From: hosts.Loopback.HTTPPort(port2), - To: hosts.Parse(server2.URL), - }, - }, - } - - require.NoError(t, app.Start(t.Context(), cfg)) - defer app.Close() - - client := http.DefaultClient - - t.Run("mapping 1", func(t *testing.T) { - url := hosts.Loopback.HTTPPort(port1).String() - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - assert.Equal(t, "Server 1", string(body)) - }) - - t.Run("mapping 2", func(t *testing.T) { - url := hosts.Loopback.HTTPPort(port2).String() - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - assert.Equal(t, "Server 2", string(body)) - }) -} - -func TestHandlerWithRewrite(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - fmt.Fprintf(w, "Path: %s, Host: %s", r.URL.Path, r.Host) //nolint:gosec // G705: test handler - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - - cfg := &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse(targetServer.URL), - Rewrites: []config.RewritingOption{ - { - From: targetServer.URL, - To: hosts.Loopback.HTTPPort(port).String(), - }, - }, - }, - }, - } - - require.NoError(t, app.Start(t.Context(), cfg)) - defer app.Close() - - client := http.DefaultClient - url := hosts.Loopback.HTTPPort(port).String() + "/test" - - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, url, nil) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - assert.Contains(t, string(body), "/test") - assert.Equal(t, http.StatusOK, resp.StatusCode) -} - -func TestHandlerWithRewritePath(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - w.WriteHeader(http.StatusOK) - fmt.Fprintf(w, "Path: %s", r.URL.Path) //nolint:gosec // G705: test handler - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - - cfg := &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse(targetServer.URL), - Rewrites: []config.RewritingOption{ - { - From: "/api/v1", - To: "/api/v2", - }, - }, - }, - }, - } - - require.NoError(t, app.Start(t.Context(), cfg)) - defer app.Close() - - req, err := http.NewRequestWithContext( - t.Context(), - http.MethodGet, - hosts.Loopback.HTTPPort(port).String()+"/api/v1", - nil, - ) - require.NoError(t, err) - - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - body, err := io.ReadAll(resp.Body) - require.NoError(t, err) - - assert.Equal(t, http.StatusOK, resp.StatusCode) - assert.Contains(t, string(body), "/api/v2") -} - -func TestHandlerWithOptions(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - targetServer := testutils.NewServer(t, http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { - w.WriteHeader(http.StatusOK) - })) - defer targetServer.Close() - - port := testutils.GetFreePort(t) - - customHeaders := map[string]string{ - "X-Custom-Header": "custom-value", - } - - cfg := &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse(targetServer.URL), - OptionsHandling: config.OptionsHandling{ - Code: http.StatusNoContent, - Headers: customHeaders, - }, - }, - }, - } - - require.NoError(t, app.Start(t.Context(), cfg)) - defer app.Close() - - url := hosts.Loopback.HTTPPort(port).String() + "/test" - client := &http.Client{} - - req, err := http.NewRequestWithContext(t.Context(), http.MethodOptions, url, nil) - require.NoError(t, err) - - resp, err := client.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - assert.Equal(t, http.StatusNoContent, resp.StatusCode) - assert.Equal(t, "custom-value", resp.Header.Get("X-Custom-Header")) -} - -func TestHandlerWithScript(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - app := uncors.CreateUncors(container) - - port := testutils.GetFreePort(t) - - cfg := &config.UncorsConfig{ - Mappings: []config.Mapping{ - { - From: hosts.Loopback.HTTPPort(port), - To: hosts.Parse("http://example.com"), - Scripts: config.Scripts{ - { - Matcher: config.RequestMatcher{ - Path: "/script", - }, - Script: `response:WriteHeader(201)`, - }, - }, - }, - }, - } - - require.NoError(t, app.Start(t.Context(), cfg)) - defer app.Close() - - reqURL := hosts.Loopback.HTTPPort(port).String() + "/script" - req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, reqURL, nil) - require.NoError(t, err) - - resp, err := http.DefaultClient.Do(req) - require.NoError(t, err) - - defer resp.Body.Close() - - assert.Equal(t, http.StatusCreated, resp.StatusCode) -} diff --git a/internal/uncors_app/app.go b/internal/uncors_app/app.go index ff322c15..8cfe0fcb 100644 --- a/internal/uncors_app/app.go +++ b/internal/uncors_app/app.go @@ -14,7 +14,7 @@ import ( "github.com/evg4b/uncors/internal/di" "github.com/evg4b/uncors/internal/helpers" "github.com/evg4b/uncors/internal/server" - "github.com/evg4b/uncors/internal/uncors" + "github.com/evg4b/uncors/internal/tui" ) const ( @@ -28,7 +28,7 @@ const ( type UncorsApp struct { keys keyMap - app *uncors.Uncors + srv *server.Server output *tuiOutput tracker server.IRequestTracker container *di.Container @@ -78,7 +78,7 @@ func NewUncorsApp( appCtx, cancel := context.WithCancel(context.Background()) - container.Override(di.OverrideCliOutput(func() contracts.Output { + container.Override(di.WithCliOutput(func() contracts.Output { return output })) @@ -88,7 +88,7 @@ func NewUncorsApp( return &UncorsApp{ keys: keys, - app: uncors.CreateUncors(container), + srv: container.Server(), output: output, tracker: container.RequestTracker(), container: container, @@ -278,7 +278,7 @@ func (m *UncorsApp) handleServerStarted() tea.Cmd { newCfg := m.loadConfig() - err := m.app.Restart(m.appContext(), newCfg) + err := m.restart(m.appContext(), newCfg) if err != nil { m.output.Errorf("Failed to restart server: %v", err) } @@ -330,7 +330,19 @@ func (m *UncorsApp) handleShutdown() tea.Cmd { func (m *UncorsApp) startServerCmd() tea.Cmd { return func() tea.Msg { - err := m.app.Start(m.appContext(), m.cfg) + tui.PrintLogo(m.output, m.container.Version()) + m.output.Print("") + m.output.WarnBox(tui.DisclaimerMessage) + m.output.Print("") + m.output.InfoBox(m.cfg.Mappings.String()) + m.output.Print("") + + targets, err := m.container.Targets(m.cfg) + if err != nil { + return serverErrMsg{err: err} + } + + err = m.srv.Start(m.appContext(), targets) if err != nil { return serverErrMsg{err: err} } @@ -376,7 +388,7 @@ func (m *UncorsApp) shutdownCmd() tea.Cmd { ctx, cancel := context.WithTimeout(context.Background(), shutdownTimeout) defer cancel() - _ = m.app.Shutdown(ctx) + _ = m.srv.Shutdown(ctx) return shutdownMsg{} } @@ -390,7 +402,7 @@ func (m *UncorsApp) restartCmd() tea.Cmd { newCfg := m.loadConfig() - err := m.app.Restart(m.appContext(), newCfg) + err := m.restart(m.appContext(), newCfg) if err != nil { m.output.Errorf("Failed to restart: %v", err) } @@ -399,6 +411,24 @@ func (m *UncorsApp) restartCmd() tea.Cmd { } } +func (m *UncorsApp) restart(ctx context.Context, cfg *config.UncorsConfig) error { + m.output.Info("Restarting server....") + + targets, err := m.container.Targets(cfg) + if err != nil { + return err + } + + err = m.srv.Restart(ctx, targets) + if err != nil { + return err + } + + m.output.InfoBox("Server restarted", cfg.Mappings.String()) + + return nil +} + func (m *UncorsApp) versionCheckCmd() tea.Cmd { return func() tea.Msg { time.Sleep(versionCheckDelay) diff --git a/internal/uncors_app/app_internal_test.go b/internal/uncors_app/app_internal_test.go index 37518d9a..a3532c7f 100644 --- a/internal/uncors_app/app_internal_test.go +++ b/internal/uncors_app/app_internal_test.go @@ -52,7 +52,7 @@ func cleanupTestApp(t *testing.T, app *UncorsApp) { t.Helper() app.cancel() - err := app.app.Close() + err := app.srv.Close() require.NoError(t, err) if app.historyWidget != nil && app.historyWidget.hist != nil { @@ -292,7 +292,7 @@ func TestUncorsAppServerErrorRestartShutdownAndFormatting(t *testing.T) { assert.Equal(t, tea.Quit(), cmd()) app.cancel() - err := app.app.Close() + err := app.srv.Close() require.NoError(t, err) }) } @@ -324,7 +324,7 @@ func TestHandleServerStartedWithConfigPath(t *testing.T) { defer func() { app.cancel() - err := app.app.Close() + err := app.srv.Close() require.NoError(t, err) if app.historyWidget != nil && app.historyWidget.hist != nil { @@ -352,7 +352,7 @@ func TestHandleServerStartedWithConfigPath(t *testing.T) { defer func() { app.cancel() - err := app.app.Close() + err := app.srv.Close() require.NoError(t, err) if app.historyWidget != nil && app.historyWidget.hist != nil { @@ -414,7 +414,7 @@ func TestHandleServerStartedCallbackOnFileChange(t *testing.T) { defer func() { // Cancel context first so any in-flight Restart fails fast. - // We deliberately skip app.app.Close() here: closeAll() writes + // We deliberately skip app.srv.Close() here: closeAll() writes // app.closers concurrently with the Restart goroutine's read of // app.closers, which would be a data race. app.cancel() @@ -465,7 +465,7 @@ func TestHandleShutdownWithWatcher(t *testing.T) { assert.Equal(t, tea.Quit(), cmd()) app.cancel() - err = app.app.Close() + err = app.srv.Close() require.NoError(t, err) if app.historyWidget != nil && app.historyWidget.hist != nil { diff --git a/main.go b/main.go index 55a97676..caa42437 100644 --- a/main.go +++ b/main.go @@ -2,227 +2,51 @@ package main import ( "context" - "fmt" - "io" - "log" "os" - "time" - tea "charm.land/bubbletea/v2" - "github.com/evg4b/uncors/internal/config" + "github.com/evg4b/uncors/internal/cli" "github.com/evg4b/uncors/internal/di" - "github.com/evg4b/uncors/internal/helpers" - "github.com/evg4b/uncors/internal/server" + "github.com/evg4b/uncors/internal/infra" "github.com/evg4b/uncors/internal/tui" - "github.com/evg4b/uncors/internal/uncors" - uncorsapp "github.com/evg4b/uncors/internal/uncors_app" "github.com/spf13/afero" - "github.com/spf13/pflag" ) -var Version = "v0.7.0" - -const generateCertsCmd = "generate-certs" +var Version = "v0.0.0" func main() { - exitCode := run() - os.Exit(exitCode) -} - -func run() int { - fs := afero.NewOsFs() + infra.SetupLogging() container := di.NewContainer( - di.WithFs(fs), + di.WithFs(afero.NewOsFs()), di.WithStdout(os.Stdout), di.WithVersion(Version), ) - defer container.Close() - - output := container.CliOutput() - - defer helpers.PanicInterceptor(func(value any) { - output.Error(value) - log.Fatalf("Caught panic: %v", value) - }) - - if len(os.Args) > 1 && os.Args[1] == generateCertsCmd { - return runGenerateCerts(container) - } - - pflag.Usage = func() { - tui.PrintLogo(output, Version) - fmt.Fprintf(output, "Usage of %s:\n", os.Args[0]) - pflag.PrintDefaults() - } - - uncorsConfig, configPath := loadConfiguration(fs) - - if uncorsConfig.Interactive { - return runInteractive(container, configPath, uncorsConfig) - } - - return runNonInteractive(context.Background(), container, configPath, uncorsConfig) -} - -// runGenerateCerts executes the generate-certs sub-command and returns an exit code. -func runGenerateCerts(container *di.Container) int { - cmd := container.GenerateCertsCommand() - output := container.CliOutput() - flags := pflag.NewFlagSet(generateCertsCmd, pflag.ContinueOnError) - cmd.DefineFlags(flags) + defer func() { + handleError(container.Close()) + }() - err := flags.Parse(os.Args[2:]) - if err != nil { - output.Error(err) - log.Printf("Error: %v", err) - - return 1 - } - - err = cmd.Execute() - if err != nil { - output.Error(err) - log.Printf("Error: %v", err) - - return 1 - } + if len(os.Args) >= 2 && os.Args[1] == cli.GenerateCertsCmd { + container.Override(di.WithArgs(os.Args[2:])) - return 0 -} - -// runNonInteractive starts the proxy in non-interactive (headless) mode and -// blocks until the server shuts down. The config file is watched for changes -// when configPath is non-empty. -func runNonInteractive( - ctx context.Context, - container *di.Container, - configPath string, - cfg *config.UncorsConfig, -) int { - output := container.CliOutput() - - app := uncors.CreateUncors(container) - - go server.RequestPrinter(container.RequestTracker(), output) - - startConfigWatcher(ctx, container, configPath, app) - - err := app.Start(ctx, cfg) - if err != nil { - panic(err) - } + err := cli.GenerateCerts(container) + handleError(err) - go startVersionChecker(ctx, container, cfg.Proxy) - - go helpers.GracefulShutdown(ctx, func(shutdownCtx context.Context) error { - log.Println("shutdown signal received") - - return app.Shutdown(shutdownCtx) - }) - - app.Wait() - output.Info("Server was stopped") - - return 0 -} - -// startConfigWatcher begins watching the config file and restarts the proxy on -// every change. The watcher lives for the process lifetime (not closed explicitly). -func startConfigWatcher( - ctx context.Context, - container *di.Container, - configPath string, - app *uncors.Uncors, -) { - if configPath == "" { return } - output := container.CliOutput() - fs := container.Fs() - watcher := config.NewWatcher(configPath) - - err := watcher.Watch(ctx, func() { - defer helpers.PanicInterceptor(func(value any) { - log.Printf("Config reloading error: %v", value) - output.Errorf("Config reloading error: %v", value) - }) - - reloaded, _ := loadConfiguration(fs) - - restartErr := app.Restart(ctx, reloaded) - if restartErr != nil { - log.Printf("Failed to restart server: %v", restartErr) - output.Errorf("Failed to restart server: %v", restartErr) - } - }) - if err != nil { - log.Printf("Failed to start config watcher: %v", err) - output.Errorf("Failed to start config watcher: %v", err) - - return - } + container.Override(di.WithArgs(os.Args[1:])) + err := cli.RunUncors(context.Background(), container) + handleError(err) } -// startVersionChecker waits for a short delay then checks for a newer release. -func startVersionChecker(ctx context.Context, container *di.Container, proxy string) { - const checkDelay = 50 * time.Millisecond - - time.Sleep(checkDelay) +var osExit = os.Exit - container.VersionChecker(proxy). - CheckNewVersion(ctx) -} - -// runInteractive starts the proxy in interactive TUI mode. -func runInteractive(container *di.Container, configPath string, cfg *config.UncorsConfig) int { - app := uncorsapp.NewUncorsApp( - container, - configPath, - cfg, - func() *config.UncorsConfig { - reloaded, _ := loadConfiguration(container.Fs()) - - return reloaded - }, - ) - - _, err := tea.NewProgram(app).Run() - if err != nil { - log.Fatal(err) - } - - return 0 -} - -const ( - logFileName = "uncors.log" - logFileFlags = os.O_CREATE | os.O_WRONLY | os.O_APPEND - logFilePerm = 0o644 -) - -// loadConfiguration loads and validates the configuration from CLI args and the -// config file. It panics on any error so that the PanicInterceptor in run() can -// display a human-readable message and exit cleanly. -func loadConfiguration(fs afero.Fs) (*config.UncorsConfig, string) { - uncorsConfig, configPath, err := config.LoadConfiguration(fs, os.Args) +func handleError(err error) { if err != nil { - panic(err) - } - - if uncorsConfig.Debug { - logFile, err := os.OpenFile(logFileName, logFileFlags, logFilePerm) - if err != nil { - panic(fmt.Sprintf("Failed to open log file: %v", err)) - } + tui.NewCliOutput(os.Stdout). + Error(err) - log.SetOutput(logFile) - log.Print("Enabled debug messages") - } else { - log.SetOutput(io.Discard) + osExit(1) } - - return uncorsConfig, configPath } diff --git a/main_test.go b/main_test.go index 48e7d93a..ebeefff1 100644 --- a/main_test.go +++ b/main_test.go @@ -1,151 +1,126 @@ package main import ( - "context" + "errors" + "io" + "log" "os" "path/filepath" "testing" - "github.com/evg4b/uncors/internal/di" - "github.com/evg4b/uncors/testing/testutils" - "github.com/spf13/afero" + "github.com/evg4b/uncors/internal/cli" + "github.com/evg4b/uncors/internal/infra" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) -// setArgs temporarily overrides os.Args and returns a restore function. -func setArgs(args []string) func() { - old := os.Args - os.Args = args +// saveLogger captures the current log writer and restores it after the test. +func saveLogger(t *testing.T) { + t.Helper() - return func() { os.Args = old } -} + orig := log.Writer() -func TestLoadConfiguration(t *testing.T) { - t.Run("returns config for valid flags", func(t *testing.T) { - defer setArgs([]string{"uncors", "-f", "http://localhost:3000", "-t", "https://api.example.com"})() + t.Cleanup(func() { log.SetOutput(orig) }) +} - cfg, path := loadConfiguration(afero.NewMemMapFs()) +// setArgs temporarily overrides os.Args and restores it via t.Cleanup. +func setArgs(t *testing.T, args []string) { + t.Helper() - require.NotNil(t, cfg) - assert.Empty(t, path) - assert.Len(t, cfg.Mappings, 1) - }) + orig := os.Args + os.Args = args - t.Run("panics when mappings are empty", func(t *testing.T) { - defer setArgs([]string{"uncors"})() + t.Cleanup(func() { os.Args = orig }) +} - assert.Panics(t, func() { - loadConfiguration(afero.NewMemMapFs()) - }) - }) +func TestSetupLogging(t *testing.T) { + t.Run("discards output when UNCORS_LOGGING is empty", func(t *testing.T) { + saveLogger(t) + t.Setenv("UNCORS_LOGGING", "") - t.Run("panics on invalid flags", func(t *testing.T) { - defer setArgs([]string{"uncors", "--no-such-flag"})() + infra.SetupLogging() - assert.Panics(t, func() { - loadConfiguration(afero.NewMemMapFs()) - }) + assert.Equal(t, io.Discard, log.Writer()) }) -} -func TestRunGenerateCerts(t *testing.T) { - t.Run("generates certs and returns 0", func(t *testing.T) { - defer setArgs([]string{"uncors", generateCertsCmd})() + t.Run("writes to file when UNCORS_LOGGING points to a valid path", func(t *testing.T) { + saveLogger(t) + logPath := filepath.Join(t.TempDir(), "test.log") + t.Setenv("UNCORS_LOGGING", logPath) - container := di.NewContainer() - defer testutils.Close(t, container) + infra.SetupLogging() - result := runGenerateCerts(container) + require.NotEqual(t, io.Discard, log.Writer()) - assert.Equal(t, 0, result) + _, err := os.Stat(logPath) + assert.NoError(t, err) }) - t.Run("returns 1 when execute fails", func(t *testing.T) { - defer setArgs([]string{"uncors", generateCertsCmd})() + t.Run("discards output when log file cannot be opened", func(t *testing.T) { + saveLogger(t) + t.Setenv("UNCORS_LOGGING", "/no-such-dir/test.log") - container := di.NewContainer() - defer testutils.Close(t, container) + infra.SetupLogging() - _ = runGenerateCerts(container) - result := runGenerateCerts(container) - - assert.Equal(t, 1, result) + assert.Equal(t, io.Discard, log.Writer()) }) +} - t.Run("returns 1 when flags parse fails", func(t *testing.T) { - defer setArgs([]string{"uncors", generateCertsCmd, "--no-such-flag"})() - - container := di.NewContainer() - defer testutils.Close(t, container) +var errTest = errors.New("something went wrong") - result := runGenerateCerts(container) +func TestHandleError_ExitsOnError(t *testing.T) { + orig := osExit - assert.Equal(t, 1, result) - }) -} + var capturedCode int -func TestLoadConfigurationWithDebug(t *testing.T) { - t.Chdir(t.TempDir()) + osExit = func(code int) { capturedCode = code } - defer setArgs([]string{"uncors", "-f", "http://localhost:3000", "-t", "https://api.example.com", "--debug"})() + t.Cleanup(func() { osExit = orig }) - cfg, _ := loadConfiguration(afero.NewMemMapFs()) + handleError(errTest) - require.NotNil(t, cfg) - assert.True(t, cfg.Debug) + assert.Equal(t, 1, capturedCode) } -func TestLoadConfigurationWithConfigFile(t *testing.T) { - const cfgContent = ` -mappings: - - from: http://localhost:3000 - to: https://api.example.com -` +func TestHandleError_NoopOnNil(t *testing.T) { + orig := osExit + + called := false - defer setArgs([]string{"uncors", "--config", "/config.yaml"})() + osExit = func(_ int) { called = true } - fs := afero.NewMemMapFs() - require.NoError(t, afero.WriteFile(fs, "/config.yaml", []byte(cfgContent), 0o600)) + t.Cleanup(func() { osExit = orig }) - cfg, path := loadConfiguration(fs) + handleError(nil) - require.NotNil(t, cfg) - assert.Equal(t, "/config.yaml", path) - assert.Len(t, cfg.Mappings, 1) + assert.False(t, called) } -func TestStartVersionChecker(t *testing.T) { - t.Run("runs without panic", func(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) +func TestMain_RunUncorsVersionPath(t *testing.T) { + saveLogger(t) + setArgs(t, []string{"uncors", "--version"}) - assert.NotPanics(t, func() { - startVersionChecker(context.Background(), container, "") - }) + assert.NotPanics(t, func() { + main() }) } -func TestStartConfigWatcher(t *testing.T) { - t.Run("logs error for non-existent config path", func(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) +func TestMain_GenerateCertsPath(t *testing.T) { + saveLogger(t) + // Point HOME to a temp dir so CA certificates go there, not ~/.config/uncors. + t.Setenv("HOME", t.TempDir()) + setArgs(t, []string{"uncors", cli.GenerateCertsCmd, "--validity-days=7"}) - assert.NotPanics(t, func() { - startConfigWatcher(context.Background(), container, "/no/such/config.yaml", nil) - }) + assert.NotPanics(t, func() { + main() }) +} - t.Run("creates watcher for existing config file", func(t *testing.T) { - container := di.NewContainer() - defer testutils.Close(t, container) - - tmpDir := t.TempDir() - configFile := filepath.Join(tmpDir, "config.yaml") - require.NoError(t, os.WriteFile(configFile, []byte("proxy: \"\""), 0o600)) +func TestMain_GenerateCertsHelpPath(t *testing.T) { + saveLogger(t) + setArgs(t, []string{"uncors", cli.GenerateCertsCmd, "--help"}) - assert.NotPanics(t, func() { - startConfigWatcher(context.Background(), container, configFile, nil) - }) + assert.NotPanics(t, func() { + main() }) } diff --git a/testing/integration/bin.go b/testing/integration/bin.go new file mode 100644 index 00000000..74a632f1 --- /dev/null +++ b/testing/integration/bin.go @@ -0,0 +1,71 @@ +package integration + +import ( + "context" + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" + "sync" + "testing" +) + +var ( + bin string + compile sync.Once +) + +var ( + repoRoot string + repoRootOnce sync.Once +) + +const UncorsTestVrsion = "v1.2.3" + +func SetupBin(_ *testing.M) { + compile.Do(func() { + //nolint:usetesting // intentional: binary lifetime must span all tests, not one subtest + tmp, err := os.MkdirTemp("", "uncors-test-*") + if err != nil { + panic(err) + } + + bin = filepath.Join(tmp, "uncors") + cmd := exec.CommandContext( + context.Background(), + "go", "build", + "-o", bin, + "-ldflags", fmt.Sprintf("-s -w -X 'main.Version=%s'", UncorsTestVrsion), + repoRootPath(), + ) + + _, err = cmd.CombinedOutput() + if err != nil { + panic(err) + } + }) +} + +func UncorsCommand(t *testing.T, args []string) *exec.Cmd { + return exec.CommandContext(t.Context(), bin, args...) +} + +func repoRootPath() string { + repoRootOnce.Do(func() { + out, err := exec.CommandContext(context.Background(), "go", "list", "-m", "-f", "{{.Dir}}").Output() + if err != nil { + panic(fmt.Sprintf("failed to determine repository root: %v", err)) + } + + repoRoot = strings.TrimSpace(string(out)) + }) + + return repoRoot +} + +func RepoRoot(t *testing.T) string { + t.Helper() + + return repoRootPath() +} diff --git a/testing/integration/proxy.go b/testing/integration/proxy.go index d28a0307..2468cf72 100644 --- a/testing/integration/proxy.go +++ b/testing/integration/proxy.go @@ -3,20 +3,33 @@ package integration import ( + "context" "crypto/x509" + "net" + "strconv" "testing" + "time" + "github.com/evg4b/uncors/internal/cli" "github.com/evg4b/uncors/internal/config" "github.com/evg4b/uncors/internal/di" "github.com/evg4b/uncors/internal/server" - "github.com/evg4b/uncors/internal/uncors" + "github.com/evg4b/uncors/testing/testutils" "github.com/spf13/afero" + "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gopkg.in/yaml.v3" ) // caValidityDays must stay clear of the proxy's expiration warning threshold // (7 days); below it HostCertManager refuses to serve TLS. -const caValidityDays = 30 +const ( + caValidityDays = 30 + configFilePerm = 0o600 + proxyReadyWait = 5 * time.Second + proxyPollTick = 25 * time.Millisecond + proxyDialTimeout = 100 * time.Millisecond +) // bootProxy generates a fresh dev CA, starts uncors in-process with the given // config, and registers shutdown with t.Cleanup. Returns the CA that the client @@ -37,17 +50,57 @@ func bootProxy(t *testing.T, fs afero.Fs, cfg *config.UncorsConfig) *x509.Certif caCert, _, err := server.LoadCA(fs, certPath, keyPath) require.NoError(t, err) - container := di.NewContainer(di.WithFs(fs)) + data, err := yaml.Marshal(cfg) + require.NoError(t, err) - app := uncors.CreateUncors(container) + const configPath = "/uncors-config.yaml" - err = app.Start(t.Context(), cfg) + err = afero.WriteFile(fs, configPath, data, configFilePerm) require.NoError(t, err) - t.Cleanup(func() { - _ = app.Close() - _ = container.Close() - }) + go func() { + // --interactive=false overrides the default (true) so the proxy runs + // in headless mode and actually starts its TCP listeners. + container := di.NewContainer(di.WithFs(fs), di.WithArgs([]string{"-c", configPath, "--interactive=false"})) + defer testutils.Close(t, container) + + err = cli.RunUncors(t.Context(), container) + assert.NoError(t, err) + }() + + waitForMappings(t, cfg) return caCert } + +// waitForMappings polls until every mapped port is accepting TCP connections. +func waitForMappings(t *testing.T, cfg *config.UncorsConfig) { + t.Helper() + + for _, mapping := range cfg.Mappings { + if mapping.From.Port == "" { + continue + } + + addr := net.JoinHostPort("127.0.0.1", mapping.From.Port) + deadline := time.Now().Add(proxyReadyWait) + + for time.Now().Before(deadline) { + dialer := &net.Dialer{Timeout: proxyDialTimeout} + + conn, dialErr := dialer.DialContext(context.Background(), "tcp", addr) + if dialErr == nil { + conn.Close() + + break + } + + time.Sleep(proxyPollTick) + } + + if time.Now().After(deadline) { + port, _ := strconv.Atoi(mapping.From.Port) + t.Fatalf("proxy port %d did not become ready within %s", port, proxyReadyWait) + } + } +} diff --git a/tests/integration/cli/cli_test.go b/tests/integration/cli/cli_test.go new file mode 100644 index 00000000..94b57731 --- /dev/null +++ b/tests/integration/cli/cli_test.go @@ -0,0 +1,28 @@ +package cli_test + +import ( + "fmt" + "testing" + + "github.com/evg4b/uncors/testing/integration" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestVersion(t *testing.T) { + t.Run("short", func(t *testing.T) { + cmd := integration.UncorsCommand(t, []string{"-v"}) + bytes, err := cmd.CombinedOutput() + require.NoError(t, err) + + assert.Equal(t, fmt.Sprintf("%s\n", integration.UncorsTestVrsion), string(bytes)) + }) + + t.Run("full", func(t *testing.T) { + cmd := integration.UncorsCommand(t, []string{"--version"}) + bytes, err := cmd.CombinedOutput() + require.NoError(t, err) + + assert.Equal(t, fmt.Sprintf("%s\n", integration.UncorsTestVrsion), string(bytes)) + }) +} diff --git a/tests/integration/cli/main_test.go b/tests/integration/cli/main_test.go new file mode 100644 index 00000000..e35d4154 --- /dev/null +++ b/tests/integration/cli/main_test.go @@ -0,0 +1,13 @@ +package cli_test + +import ( + "os" + "testing" + + "github.com/evg4b/uncors/testing/integration" +) + +func TestMain(m *testing.M) { + integration.SetupBin(m) + os.Exit(m.Run()) +}