diff --git a/commands/docker.go b/commands/docker.go index 6e1b7349..622a8f6b 100644 --- a/commands/docker.go +++ b/commands/docker.go @@ -58,6 +58,10 @@ func (d *KoolDocker) Execute(args []string) (err error) { d.dockerRun.AppendArgs("--env", "ASUSER="+asuser) } + for _, envVar := range environment.AgentEnvVars(d.envStorage) { + d.dockerRun.AppendArgs("--env", envVar) + } + if len(d.Flags.EnvVariables) > 0 { for _, envVar := range d.Flags.EnvVariables { d.dockerRun.AppendArgs("--env", envVar) diff --git a/commands/docker_test.go b/commands/docker_test.go index 172af55a..9b7aa14f 100644 --- a/commands/docker_test.go +++ b/commands/docker_test.go @@ -169,6 +169,23 @@ func TestEnvFlagNewDockerCommand(t *testing.T) { } } +func TestAgentEnvNewDockerCommand(t *testing.T) { + f := newFakeKoolDocker() + f.shell.(*shell.FakeShell).MockIsTerminal = false + f.envStorage.(*environment.FakeEnvStorage).Envs["CLAUDECODE"] = "1" + cmd := NewDockerCommand(f) + cmd.SetArgs([]string{"image"}) + + if err := cmd.Execute(); err != nil { + t.Errorf("unexpected error executing docker command; error: %v", err) + } + + argsAppend := f.dockerRun.(*builder.FakeCommand).ArgsAppend + if len(argsAppend) != 4 || argsAppend[0] != "--env" || argsAppend[1] != "CLAUDECODE=1" { + t.Errorf("bad arguments to KoolDocker.dockerRun Command with detected agent: %v", argsAppend) + } +} + func TestVolumesFlagNewDockerCommand(t *testing.T) { f := newFakeKoolDocker() f.shell.(*shell.FakeShell).MockIsTerminal = false diff --git a/commands/exec.go b/commands/exec.go index 087aabbd..e5464dc2 100644 --- a/commands/exec.go +++ b/commands/exec.go @@ -98,6 +98,10 @@ func (e *KoolExec) Execute(args []string) (err error) { e.checkUser(args[0]) + for _, envVar := range environment.AgentEnvVars(e.env) { + e.composeExec.AppendArgs("--env", envVar) + } + if len(e.Flags.EnvVariables) > 0 { for _, envVar := range e.Flags.EnvVariables { e.composeExec.AppendArgs("--env", envVar) diff --git a/commands/exec_test.go b/commands/exec_test.go index 9d823a42..6010272e 100644 --- a/commands/exec_test.go +++ b/commands/exec_test.go @@ -149,6 +149,22 @@ func TestEnvFlagNewExecCommand(t *testing.T) { } } +func TestAgentEnvNewExecCommand(t *testing.T) { + f := newFakeKoolExec() + f.env.(*environment.FakeEnvStorage).Envs["OPENCODE"] = "1" + cmd := NewExecCommand(f) + cmd.SetArgs([]string{"service", "command"}) + + if err := cmd.Execute(); err != nil { + t.Errorf("unexpected error executing exec command; error: %v", err) + } + + argsAppend := f.composeExec.(*builder.FakeCommand).ArgsAppend + if len(argsAppend) != 2 || argsAppend[0] != "--env" || argsAppend[1] != "OPENCODE=1" { + t.Errorf("bad arguments to KoolExec.composeExec Command with detected agent: %v", argsAppend) + } +} + func TestDetachFlagNewExecCommand(t *testing.T) { f := newFakeKoolExec() cmd := NewExecCommand(f) diff --git a/commands/proxy.go b/commands/proxy.go new file mode 100644 index 00000000..90ce6807 --- /dev/null +++ b/commands/proxy.go @@ -0,0 +1,30 @@ +package commands + +import ( + "kool-dev/kool/core/environment" + "kool-dev/kool/core/shell" + "kool-dev/kool/services/proxy" + + "github.com/spf13/cobra" +) + +// AddKoolProxy adds commands for managing Kool's local proxy. +func AddKoolProxy(root *cobra.Command) { + proxyCommand := &cobra.Command{ + Use: "proxy", + Short: "Manage Kool's local proxy", + } + proxyCommand.AddCommand(&cobra.Command{ + Use: "trust", + Short: "Trust the local proxy certificate authority", + Args: cobra.NoArgs, + RunE: func(cmd *cobra.Command, args []string) error { + sh := shell.NewShell() + sh.SetInStream(cmd.InOrStdin()) + sh.SetOutStream(cmd.OutOrStdout()) + sh.SetErrStream(cmd.ErrOrStderr()) + return proxy.NewManager(sh, environment.NewEnvStorage()).Trust() + }, + }) + root.AddCommand(proxyCommand) +} diff --git a/commands/root.go b/commands/root.go index 7cbd52a2..e379884c 100644 --- a/commands/root.go +++ b/commands/root.go @@ -1,6 +1,7 @@ package commands import ( + "errors" "fmt" "io" "kool-dev/kool/core/environment" @@ -32,6 +33,7 @@ var AddCommands AddCommandsFN = func(root *cobra.Command) { AddKoolInfo(root) AddKoolLogs(root) AddKoolPreset(root) + AddKoolProxy(root) AddKoolRestart(root) AddKoolRun(root) AddKoolSelfUpdate(root) @@ -47,7 +49,7 @@ const DEV_VERSION = "0.0.0-dev" var version string = DEV_VERSION -var rootCmd = NewRootCmd(environment.NewEnvStorage()) +var rootCmd = newRootCmd(environment.NewEnvStorage(), true) var originalWorkingDir = "" @@ -60,6 +62,11 @@ func init() { // NewRootCmd creates the root command func NewRootCmd(env environment.EnvStorage) (cmd *cobra.Command) { + return newRootCmd(env, false) +} + +func newRootCmd(env environment.EnvStorage, initializeEnvironment bool) (cmd *cobra.Command) { + environmentInitialized := false cmd = &cobra.Command{ Args: cobra.ArbitraryArgs, Use: "kool", @@ -75,20 +82,17 @@ Complete documentation is available at https://kool.dev/docs`, DisableFlagsInUseLine: true, PersistentPreRunE: func(cmd *cobra.Command, args []string) (err error) { cmd.SilenceUsage = true - - if verbose := cmd.Flags().Lookup("verbose"); verbose != nil && verbose.Value.String() == "true" { - env.Set("KOOL_VERBOSE", verbose.Value.String()) - } - - if !hasWarnedDevelopmentVersion && version == DEV_VERSION && shell.NewTerminalChecker().IsTerminal(cmd.OutOrStdout()) { - shell.NewShell().Warning("Warning: you are executing a development version of kool.") - hasWarnedDevelopmentVersion = true - } - workDirFlag := cmd.Flags().Lookup("working_dir") if workDirFlag != nil && workDirFlag.Value.String() != "" { workDir := workDirFlag.Value.String() + currentWorkingDir, getwdErr := os.Getwd() + if getwdErr != nil { + return getwdErr + } + if originalWorkingDir == "" { + originalWorkingDir = currentWorkingDir + } if originalWorkingDir != "" { // having an original working dir set means we have // already changed the working dir before and we are in @@ -108,17 +112,27 @@ Complete documentation is available at https://kool.dev/docs`, if err = os.Chdir(workDir); err != nil { return } - - if originalWorkingDir == "" { - // we only set the original working dir if it is not set - // yet. This is to avoid overriding the original working - // dir in recursive calls. - if originalWorkingDir, err = os.Getwd(); err != nil { - return - } + if workDir, err = os.Getwd(); err != nil { + return } + env.Set("PWD", workDir) + + } + if initializeEnvironment && !environmentInitialized { + environment.InitEnvironmentVariables(env) + environmentInitialized = true + } + if workspaceError := env.Get("KOOL_WORKSPACE_ERROR"); workspaceError != "" { + return errors.New(workspaceError) + } + + if verbose := cmd.Flags().Lookup("verbose"); verbose != nil && verbose.Value.String() == "true" { + env.Set("KOOL_VERBOSE", verbose.Value.String()) + } - environment.NewEnvStorage().Set("PWD", workDir) + if !hasWarnedDevelopmentVersion && version == DEV_VERSION && shell.NewTerminalChecker().IsTerminal(cmd.OutOrStdout()) { + shell.NewShell().Warning("Warning: you are executing a development version of kool.") + hasWarnedDevelopmentVersion = true } return @@ -165,7 +179,27 @@ func Execute() error { func setRecursiveCall(root *cobra.Command) { shell.RecursiveCall = func(args []string, in io.Reader, out, err io.Writer) error { - childRoot := NewRootCmd(environment.NewEnvStorage()) + currentDirectory, getwdErr := os.Getwd() + if getwdErr != nil { + return getwdErr + } + currentPWD := os.Getenv("PWD") + currentOriginalWorkingDir := originalWorkingDir + defer func() { + _ = os.Chdir(currentDirectory) + _ = os.Setenv("PWD", currentPWD) + originalWorkingDir = currentOriginalWorkingDir + }() + + initializeEnvironment := hasWorkingDirArg(args) + if initializeEnvironment { + restoreEnvironment := clearDirectoryEnvironment(currentDirectory) + defer restoreEnvironment() + restoreWorkspace := environment.IsolateWorkspace() + defer restoreWorkspace() + originalWorkingDir = currentDirectory + } + childRoot := newRootCmd(environment.NewEnvStorage(), initializeEnvironment) childRoot.SetArgs(args) @@ -179,6 +213,46 @@ func setRecursiveCall(root *cobra.Command) { } } +func hasWorkingDirArg(args []string) bool { + for _, arg := range args { + if arg == "-w" || arg == "--working_dir" || strings.HasPrefix(arg, "--working_dir=") || strings.HasPrefix(arg, "-w=") || strings.HasPrefix(arg, "-w") && len(arg) > 2 { + return true + } + } + return false +} + +func clearDirectoryEnvironment(directory string) func() { + previous := make(map[string]string) + for _, entry := range os.Environ() { + parts := strings.SplitN(entry, "=", 2) + previous[parts[0]] = parts[1] + } + keys := []string{ + "PWD", "KOOL_NAME", "KOOL_WORKSPACES_ENABLED", "KOOL_PROXY_ENABLED", + "KOOL_WORKSPACE", "KOOL_WORKSPACE_PROVIDER", "KOOL_WORKSPACE_ERROR", + "KOOL_WORKSPACE_SOURCE", "KOOL_WORKSPACE_SOURCE_PROJECT", "KOOL_WORKSPACE_NAME", + "KOOL_WORKSPACE_PATH", "KOOL_WORKSPACE_PROJECT", "KOOL_WORKSPACE_SERVICES", + "KOOL_PROXY_DOMAIN", "KOOL_PROXY_HOST", "COMPOSE_PROJECT_NAME", "COMPOSE_FILE", + "KOOL_GLOBAL_NETWORK", + } + keys = append(keys, environment.LoadedEnvKeys(directory)...) + for _, key := range keys { + _ = os.Unsetenv(key) + } + return func() { + for _, entry := range os.Environ() { + key := strings.SplitN(entry, "=", 2)[0] + if _, exists := previous[key]; !exists { + _ = os.Unsetenv(key) + } + } + for key, value := range previous { + _ = os.Setenv(key, value) + } + } +} + // RootCmd exposes the root command func RootCmd() *cobra.Command { return rootCmd diff --git a/commands/root_test.go b/commands/root_test.go index 607f2fec..9857b7db 100644 --- a/commands/root_test.go +++ b/commands/root_test.go @@ -7,6 +7,7 @@ import ( "kool-dev/kool/core/environment" "kool-dev/kool/core/shell" "os" + "path/filepath" "strings" "testing" @@ -191,6 +192,113 @@ func TestVerboseFlagRootCommand(t *testing.T) { } } +func TestWorkingDirectoryInitializesTargetEnvironment(t *testing.T) { + originalDirectory, err := os.Getwd() + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + _ = os.Chdir(originalDirectory) + originalWorkingDir = "" + environment.CleanupWorkspace() + }) + originalWorkingDir = "" + target := t.TempDir() + if err = os.WriteFile(filepath.Join(target, ".env"), []byte("TARGET_ENV=loaded\n"), 0644); err != nil { + t.Fatal(err) + } + + env := environment.NewFakeEnvStorage() + root := newRootCmd(env, true) + root.AddCommand(&cobra.Command{Use: "target", Run: func(*cobra.Command, []string) {}}) + root.SetArgs([]string{"-w", target, "target"}) + if err = root.Execute(); err != nil { + t.Fatal(err) + } + + expectedTarget, err := filepath.EvalSymlinks(target) + if err != nil { + t.Fatal(err) + } + if got := env.Get("PWD"); got != expectedTarget { + t.Errorf("expected target PWD %q, got %q", expectedTarget, got) + } + if got := env.Get("TARGET_ENV"); got != "loaded" { + t.Errorf("expected target .env to be loaded, got %q", got) + } +} + +func TestRecursiveWorkingDirectoryReplacesParentEnvironment(t *testing.T) { + originalDirectory, err := os.Getwd() + if err != nil { + t.Fatal(err) + } + parent := t.TempDir() + target := t.TempDir() + if err = os.WriteFile(filepath.Join(parent, ".env"), []byte("SHARED_ENV=parent\nPARENT_ONLY=present\n"), 0644); err != nil { + t.Fatal(err) + } + if err = os.WriteFile(filepath.Join(target, ".env"), []byte("SHARED_ENV=target\n"), 0644); err != nil { + t.Fatal(err) + } + if err = os.Chdir(parent); err != nil { + t.Fatal(err) + } + if err = os.Unsetenv("SHARED_ENV"); err != nil { + t.Fatal(err) + } + if err = os.Unsetenv("PARENT_ONLY"); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + _ = os.Chdir(originalDirectory) + _ = os.Unsetenv("SHARED_ENV") + _ = os.Unsetenv("PARENT_ONLY") + originalWorkingDir = "" + }) + parentEnv := environment.NewEnvStorage() + environment.InitEnvironmentVariables(parentEnv) + + restore := clearDirectoryEnvironment(parent) + defer restore() + if err = os.Chdir(target); err != nil { + t.Fatal(err) + } + env := environment.NewEnvStorage() + environment.InitEnvironmentVariables(env) + + if got := env.Get("SHARED_ENV"); got != "target" { + t.Errorf("expected target environment value, got %q", got) + } + if got := env.Get("PARENT_ONLY"); got != "" { + t.Errorf("expected parent-only environment value to be cleared, got %q", got) + } +} + +func TestRecursiveWorkingDirectoryReplacesGlobalNetwork(t *testing.T) { + originalDirectory, err := os.Getwd() + if err != nil { + t.Fatal(err) + } + parent := t.TempDir() + target := t.TempDir() + if err = os.WriteFile(filepath.Join(target, ".env"), []byte("KOOL_GLOBAL_NETWORK=target_network\n"), 0644); err != nil { + t.Fatal(err) + } + t.Setenv("KOOL_GLOBAL_NETWORK", "kool_global") + restore := clearDirectoryEnvironment(parent) + defer restore() + if err := os.Chdir(target); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.Chdir(originalDirectory) }) + env := environment.NewEnvStorage() + environment.InitEnvironmentVariables(env) + if got := env.Get("KOOL_GLOBAL_NETWORK"); got != "target_network" { + t.Fatalf("expected target global network, got %q", got) + } +} + func TestRecursiveCall(t *testing.T) { recursive := &cobra.Command{ Use: "recursive", @@ -233,6 +341,14 @@ func TestMultipleRecursiveCall(t *testing.T) { } } +func TestHasWorkingDirArgAcceptsAttachedShorthand(t *testing.T) { + for _, arg := range []string{"-w=/tmp/project", "-w/tmp/project"} { + if !hasWorkingDirArg([]string{"status", arg}) { + t.Errorf("expected %q to initialize the target environment", arg) + } + } +} + func TestAddCommands(t *testing.T) { root := NewRootCmd(environment.NewFakeEnvStorage()) @@ -247,6 +363,7 @@ func TestAddCommands(t *testing.T) { "info": false, "logs": false, "preset": false, + "proxy": false, "restart": false, "run": false, "self-update": false, diff --git a/commands/share.go b/commands/share.go index 7c57ce97..8e1d7d48 100644 --- a/commands/share.go +++ b/commands/share.go @@ -65,7 +65,7 @@ func (s *KoolShare) validSubdomain(subdomain string) bool { func (s *KoolShare) Execute(args []string) (err error) { var isRunning bool - if isRunning, _, _, err = s.status.getServiceInfo(s.Flags.Service); err != nil { + if isRunning, _, _, err = s.status.getCurrentServiceInfo(s.Flags.Service); err != nil { return } diff --git a/commands/start.go b/commands/start.go index 337084e4..57af1bc1 100644 --- a/commands/start.go +++ b/commands/start.go @@ -6,6 +6,7 @@ import ( "kool-dev/kool/core/environment" "kool-dev/kool/core/network" "kool-dev/kool/services/checker" + "kool-dev/kool/services/proxy" "kool-dev/kool/services/updater" "strings" @@ -84,18 +85,22 @@ func AddKoolStart(root *cobra.Command) { // Execute runs the rebuild logic func (r *KoolRebuild) Execute(args []string) (err error) { - if err = r.Shell().Interactive(r.pull); err != nil { + if err = r.Shell().Interactive(r.pull, args...); err != nil { return } - err = r.Shell().Interactive(r.build) + err = r.Shell().Interactive(r.build, args...) return } // Execute runs the start logic with incoming arguments func (s *KoolStart) Execute(args []string) (err error) { + if args, err = selectWorkspaceServices(s.envStorage, args); err != nil { + return + } + if s.Flags.Rebuild { - if err = s.rebuild(); err != nil { + if err = s.rebuild(args); err != nil { return } } @@ -107,6 +112,9 @@ func (s *KoolStart) Execute(args []string) (err error) { if !s.Flags.Foreground { s.start.AppendArgs("-d") } + if isWorkspace(s.envStorage) { + s.start.AppendArgs("--no-deps") + } if err = s.checkDependencies(); err != nil { if strings.HasPrefix(err.Error(), "no configuration file provided: not found") { @@ -114,12 +122,27 @@ func (s *KoolStart) Execute(args []string) (err error) { } return } + proxyManager := proxy.NewManager(s.Shell(), s.envStorage) + finishProxy := func(bool) error { return nil } + if proxyEnabled(s.envStorage) { + var proxyErr error + finishProxy, proxyErr = proxyManager.Prepare(args) + if proxyErr != nil { + return proxyErr + } + } + defer func() { + persistRoutes := err == nil && !s.Flags.Foreground + if finishErr := finishProxy(persistRoutes); err == nil { + err = finishErr + } + }() err = s.Shell().Interactive(s.start, args...) return } -func (s *KoolStart) rebuild() (err error) { +func (s *KoolStart) rebuild(args []string) (err error) { var task = NewKoolTask("Updating service's images", s.rebuilder) task.SetFrameOutput(false) @@ -128,7 +151,7 @@ func (s *KoolStart) rebuild() (err error) { task.Shell().SetOutStream(s.Shell().OutStream()) task.Shell().SetErrStream(s.Shell().ErrStream()) - err = task.Run(nil) + err = task.Run(args) return } diff --git a/commands/start_test.go b/commands/start_test.go index 7fea168d..62d4da5a 100644 --- a/commands/start_test.go +++ b/commands/start_test.go @@ -140,6 +140,77 @@ func TestStartServicesCommand(t *testing.T) { } } +func TestStartWorkspaceServicesCommand(t *testing.T) { + koolStart := newFakeKoolStart() + koolStart.envStorage.Set("KOOL_WORKSPACE", "true") + koolStart.envStorage.Set("KOOL_WORKSPACE_SERVICES", "app,node") + + if err := koolStart.Execute(nil); err != nil { + t.Fatal(err) + } + + interactiveArgs := koolStart.shell.(*shell.FakeShell).ArgsInteractive["start"] + expected := []string{"app", "node"} + if !startedServicesAreEqual(interactiveArgs, expected) { + t.Errorf("expected workspace services %v, got %v", expected, interactiveArgs) + } + args := koolStart.start.(*builder.FakeCommand).ArgsAppend + if !containsArg(args, "--no-deps") { + t.Errorf("expected --no-deps for workspace start, got %v", args) + } +} + +func TestRebuildWorkspaceServicesCommand(t *testing.T) { + koolStart := newFakeKoolStart() + koolStart.envStorage.Set("KOOL_WORKSPACE", "true") + koolStart.envStorage.Set("KOOL_WORKSPACE_SERVICES", "app,node") + koolStart.Flags.Rebuild = true + rebuilder := koolStart.rebuilder.(*KoolRebuild) + rebuilder.shell.(*shell.FakeShell).MockOutStream = io.Discard + + if err := koolStart.Execute(nil); err != nil { + t.Fatal(err) + } + + for _, command := range []string{"pull", "build"} { + args := rebuilder.shell.(*shell.FakeShell).ArgsInteractive[command] + if !startedServicesAreEqual(args, []string{"app", "node"}) { + t.Errorf("expected %s to target workspace services, got %v", command, args) + } + } +} + +func TestStartWorkspaceRequiresServices(t *testing.T) { + koolStart := newFakeKoolStart() + koolStart.envStorage.Set("KOOL_WORKSPACE", "true") + + if err := koolStart.Execute(nil); err == nil || !strings.Contains(err.Error(), "workspaces in kool.yml") { + t.Fatalf("expected workspace configuration error, got %v", err) + } +} + +func TestStartRejectsSharedServiceInWorkspace(t *testing.T) { + koolStart := newFakeKoolStart() + koolStart.envStorage.Set("KOOL_WORKSPACE", "true") + koolStart.envStorage.Set("KOOL_WORKSPACE_SERVICES", "app,node") + + if err := koolStart.Execute([]string{"database"}); err == nil || !strings.Contains(err.Error(), "not enabled for workspaces") { + t.Fatalf("expected shared service error, got %v", err) + } +} + +func TestForegroundStartDoesNotPersistProxyRoutes(t *testing.T) { + koolStart := newFakeKoolStart() + koolStart.Flags.Foreground = true + + if err := koolStart.Execute(nil); err != nil { + t.Fatal(err) + } + if containsArg(koolStart.start.(*builder.FakeCommand).ArgsAppend, "-d") { + t.Error("did not expect detached mode for foreground start") + } +} + func TestFailedDependenciesStartCommand(t *testing.T) { koolStart := newFakeKoolStart() koolStart.check.(*checker.FakeChecker).MockError = errors.New("dependencies") @@ -195,3 +266,12 @@ func startedServicesAreEqual(a, b []string) bool { } return true } + +func containsArg(args []string, expected string) bool { + for _, arg := range args { + if arg == expected { + return true + } + } + return false +} diff --git a/commands/status.go b/commands/status.go index ca2acf91..e086fb85 100644 --- a/commands/status.go +++ b/commands/status.go @@ -6,6 +6,7 @@ import ( "kool-dev/kool/core/network" "kool-dev/kool/core/shell" "kool-dev/kool/services/checker" + "sort" "strings" "sync" @@ -21,16 +22,18 @@ type KoolStatus struct { env environment.EnvStorage getServicesCmd builder.Command + getProjectsCmd builder.Command getServiceIDCmd builder.Command + getProjectServiceIDCmd builder.Command getServiceStatusPortCmd builder.Command table shell.TableWriter } type statusService struct { - service, state, ports string - running string - err error + project, service, state, ports string + running string + err error } func AddKoolStatus(root *cobra.Command) { @@ -51,7 +54,9 @@ func NewKoolStatus() *KoolStatus { network.NewHandler(defaultKoolService.shell), environment.NewEnvStorage(), builder.NewCommand("docker", "compose", "config", "--services"), + builder.NewCommand("docker", "ps", "--all"), builder.NewCommand("docker", "compose", "ps", "--all", "--quiet"), + builder.NewCommand("docker", "ps", "--all", "--quiet"), builder.NewCommand("docker", "ps", "--all", "--format", "{{.Status}}|{{.Ports}}"), shell.NewTableWriter(), } @@ -59,49 +64,132 @@ func NewKoolStatus() *KoolStatus { // Execute runs the status logic with incoming arguments. func (s *KoolStatus) Execute(args []string) (err error) { - var services []string + if !workspacesEnabled(s.env) { + return s.executeLegacy() + } + var projects []statusProject if err = s.checkDependencies(); err != nil { return } - if services, err = s.getServices(); err != nil { + if projects, err = s.getProjects(); err != nil { return - } else if len(services) == 0 { + } else if len(projects) == 0 { + s.Shell().Warning("No services found.") + return + } + serviceCount := 0 + for _, project := range projects { + serviceCount += len(project.services) + } + if serviceCount == 0 { s.Shell().Warning("No services found.") return } - chStatus := make(chan *statusService, len(services)) + chStatus := make(chan *statusService, serviceCount) s.table.SetWriter(s.Shell().OutStream()) - s.table.AppendHeader("Service", "Running", "Ports", "State") + s.table.AppendHeader("Project", "Service", "Running", "Ports", "State") go func() { var wg sync.WaitGroup defer close(chStatus) - for _, service := range services { - wg.Add(1) - go s.fetchServiceInfo(service, chStatus, &wg) + for _, project := range projects { + for _, service := range project.services { + wg.Add(1) + go s.fetchServiceInfo(project.name, service, chStatus, &wg) + } } wg.Wait() }() + var statuses []*statusService for ss := range chStatus { if ss.err != nil { err = ss.err return } + statuses = append(statuses, ss) + } + sort.Slice(statuses, func(i, j int) bool { + if statuses[i].project == statuses[j].project { + return statuses[i].service < statuses[j].service + } + return statuses[i].project < statuses[j].project + }) + for _, ss := range statuses { + s.table.AppendRow(ss.project, ss.service, ss.running, ss.ports, ss.state) + } + s.table.Render() + return +} - s.table.AppendRow(ss.service, ss.running, ss.ports, ss.state) +func (s *KoolStatus) executeLegacy() (err error) { + if err = s.checkDependencies(); err != nil { + return + } + services, err := s.getServices() + if err != nil { + return err + } + if len(services) == 0 { + s.Shell().Warning("No services found.") + return nil } + chStatus := make(chan *statusService, len(services)) + s.table.SetWriter(s.Shell().OutStream()) + s.table.AppendHeader("Service", "Running", "Ports", "State") + go func() { + var wg sync.WaitGroup + defer close(chStatus) + for _, service := range services { + wg.Add(1) + go s.fetchLegacyServiceInfo(service, chStatus, &wg) + } + wg.Wait() + }() + for status := range chStatus { + if status.err != nil { + return status.err + } + s.table.AppendRow(status.service, status.running, status.ports, status.state) + } s.table.SortBy(1) s.table.Render() - return + return nil +} + +type statusProject struct { + name string + services []string +} + +func (s *KoolStatus) getProjects() ([]statusProject, error) { + services, err := s.getServices() + if err != nil { + return nil, err + } + mainProject := sourceProject(s.env) + projects := []statusProject{{name: mainProject, services: services}} + workspaceProjectNames := []string{} + if isWorkspace(s.env) { + workspaceProjectNames = append(workspaceProjectNames, currentProject(s.env)) + } else if workspaceProjectNames, err = activeWorkspaceProjects(s.Shell(), s.getProjectsCmd, s.env); err != nil { + return nil, err + } + workspaceProjectServices := configuredWorkspaceServices(s.env) + for _, project := range workspaceProjectNames { + if project != "" && project != mainProject { + projects = append(projects, statusProject{name: project, services: workspaceProjectServices}) + } + } + return projects, nil } func (s *KoolStatus) checkDependencies() (err error) { @@ -154,17 +242,16 @@ func (s *KoolStatus) getServices() (services []string, err error) { services = append(services, s) } } - return } -func (s *KoolStatus) fetchServiceInfo(service string, chStatus chan *statusService, wg *sync.WaitGroup) { +func (s *KoolStatus) fetchServiceInfo(project, service string, chStatus chan *statusService, wg *sync.WaitGroup) { var isRunning bool defer wg.Done() - ss := &statusService{service: service, running: "Not running"} - isRunning, ss.state, ss.ports, ss.err = s.getServiceInfo(service) + ss := &statusService{project: project, service: service, running: "Not running"} + isRunning, ss.state, ss.ports, ss.err = s.getServiceInfo(project, service) if isRunning { ss.running = "Running" } @@ -172,9 +259,23 @@ func (s *KoolStatus) fetchServiceInfo(service string, chStatus chan *statusServi chStatus <- ss } -func (s *KoolStatus) getServiceInfo(service string) (isRunning bool, status, port string, err error) { +func (s *KoolStatus) fetchLegacyServiceInfo(service string, chStatus chan *statusService, wg *sync.WaitGroup) { + defer wg.Done() + status := &statusService{service: service, running: "Not running"} var serviceID string - if serviceID, err = s.Shell().Exec(s.getServiceIDCmd, service); err == nil && serviceID != "" { + if serviceID, status.err = s.Shell().Exec(s.getServiceIDCmd, service); status.err == nil && serviceID != "" { + status.state, status.ports = s.getStatusPort(serviceID) + if strings.HasPrefix(status.state, "Up") { + status.running = "Running" + } + } + chStatus <- status +} + +func (s *KoolStatus) getServiceInfo(project, service string) (isRunning bool, status, port string, err error) { + var serviceID string + if serviceID, err = s.Shell().Exec(s.getProjectServiceIDCmd, "--filter", "label=com.docker.compose.project="+project, "--filter", "label=com.docker.compose.service="+service, "--filter", "label=com.docker.compose.oneoff=False"); err == nil && serviceID != "" { + serviceID = strings.Fields(serviceID)[0] status, port = s.getStatusPort(serviceID) if strings.HasPrefix(status, "Up") { isRunning = true @@ -183,6 +284,18 @@ func (s *KoolStatus) getServiceInfo(service string) (isRunning bool, status, por return } +func (s *KoolStatus) getCurrentServiceInfo(service string) (bool, string, string, error) { + if workspacesEnabled(s.env) { + return s.getServiceInfo(currentProject(s.env), service) + } + serviceID, err := s.Shell().Exec(s.getServiceIDCmd, service) + if err != nil || serviceID == "" { + return false, "", "", err + } + status, port := s.getStatusPort(serviceID) + return strings.HasPrefix(status, "Up"), status, port, nil +} + func (s *KoolStatus) getStatusPort(serviceID string) (status string, port string) { var output string diff --git a/commands/status_test.go b/commands/status_test.go index af0d7351..e2a57cca 100644 --- a/commands/status_test.go +++ b/commands/status_test.go @@ -17,6 +17,21 @@ type FakeRaceShell struct { shell.FakeShell } +type statusRecordingShell struct { + shell.FakeShell + args [][]string +} + +func (f *statusRecordingShell) Exec(command builder.Command, extraArgs ...string) (string, error) { + if len(extraArgs) > 0 { + f.args = append(f.args, append([]string(nil), extraArgs...)) + } + if fake, ok := command.(*builder.FakeCommand); ok { + return fake.MockExecOut, fake.MockExecError + } + return "", nil +} + func (f *FakeRaceShell) Exec(command builder.Command, extraArgs ...string) (string, error) { output := command.(*builder.FakeCommand).MockExecOut return output, nil @@ -31,11 +46,14 @@ func newFakeKoolStatus() *KoolStatus { &builder.FakeCommand{}, &builder.FakeCommand{}, &builder.FakeCommand{}, + &builder.FakeCommand{}, + &builder.FakeCommand{}, &shell.FakeTableWriter{}, } fs.shell.(*shell.FakeShell).MockErrStream = io.Discard fs.shell.(*shell.FakeShell).MockOutStream = io.Discard + fs.env.Set("KOOL_NAME", "example") return fs } @@ -157,6 +175,86 @@ func TestNoServicesStatusCommand(t *testing.T) { } } +func TestStatusFiltersWorkspaceServices(t *testing.T) { + f := newFakeKoolStatus() + f.shell = &FakeRaceShell{FakeShell: shell.FakeShell{MockErrStream: io.Discard, MockOutStream: io.Discard}} + f.env.Set("KOOL_WORKSPACE", "true") + f.env.Set("KOOL_WORKSPACES_ENABLED", "true") + f.env.Set("KOOL_WORKSPACE_SERVICES", "app,node") + f.env.Set("KOOL_WORKSPACE_SOURCE_PROJECT", "example") + f.env.Set("KOOL_WORKSPACE_PROJECT", "example-workspace-task-a") + f.getServicesCmd.(*builder.FakeCommand).MockExecOut = "app\ndatabase\nnode" + f.getProjectServiceIDCmd.(*builder.FakeCommand).MockExecOut = "100" + + cmd := NewStatusCommand(f) + if err := cmd.Execute(); err != nil { + t.Fatal(err) + } + + output := f.table.(*shell.FakeTableWriter).TableOut + for _, expected := range []string{"example | database", "example-workspace-task-a | app", "example-workspace-task-a | node"} { + if !strings.Contains(output, expected) { + t.Errorf("expected status to contain %q, got %q", expected, output) + } + } + if strings.Contains(output, "example-workspace-task-a | database") { + t.Errorf("did not expect shared database in workspace project, got %q", output) + } +} + +func TestStatusShowsMainAndAllActiveWorkspaces(t *testing.T) { + f := newFakeKoolStatus() + f.shell = &FakeRaceShell{FakeShell: shell.FakeShell{MockErrStream: io.Discard, MockOutStream: io.Discard}} + f.env.Set("KOOL_NAME", "example") + f.env.Set("KOOL_WORKSPACES_ENABLED", "true") + f.env.Set("KOOL_WORKSPACE_SERVICES", "app") + f.getServicesCmd.(*builder.FakeCommand).MockExecOut = "app\ndatabase" + f.getProjectsCmd.(*builder.FakeCommand).MockExecOut = "example-workspace-task-b\nexample-workspace-task-a\nexample-workspace-task-b" + f.getProjectServiceIDCmd.(*builder.FakeCommand).MockExecOut = "100" + + if err := NewStatusCommand(f).Execute(); err != nil { + t.Fatal(err) + } + + output := f.table.(*shell.FakeTableWriter).TableOut + for _, expected := range []string{ + "example | app", + "example | database", + "example-workspace-task-a | app", + "example-workspace-task-b | app", + } { + if !strings.Contains(output, expected) { + t.Errorf("expected status to contain %q, got %q", expected, output) + } + } + if strings.Contains(output, "example-workspace-task-a | database") || strings.Contains(output, "example-workspace-task-b | database") { + t.Errorf("did not expect shared database in workspace projects, got %q", output) + } +} + +func TestWorkspaceServiceInfoExcludesOneOffContainers(t *testing.T) { + f := newFakeKoolStatus() + recorder := &statusRecordingShell{} + f.shell = recorder + f.getProjectServiceIDCmd.(*builder.FakeCommand).MockExecOut = "regular-id\noneoff-id\n" + f.getServiceStatusPortCmd.(*builder.FakeCommand).MockExecOut = "Up 1 minute|80/tcp" + + running, _, _, err := f.getServiceInfo("example-workspace-task", "app") + if err != nil { + t.Fatal(err) + } + if !running { + t.Error("expected canonical service container to be running") + } + var recorded []string + for _, args := range recorder.args { + recorded = append(recorded, args...) + } + if !strings.Contains(strings.Join(recorded, " "), "label=com.docker.compose.oneoff=False") { + t.Fatalf("expected one-off container filter, got %v", recorder.args) + } +} + func TestFailedGetServicesStatusCommand(t *testing.T) { f := newFakeKoolStatus() @@ -227,6 +325,8 @@ func TestServicesOrderStatusCommand(t *testing.T) { &builder.FakeCommand{}, &builder.FakeCommand{}, &builder.FakeCommand{}, + &builder.FakeCommand{}, + &builder.FakeCommand{}, &shell.FakeTableWriter{}, } @@ -238,6 +338,7 @@ func TestServicesOrderStatusCommand(t *testing.T) { } f.getServicesCmd.(*builder.FakeCommand).MockExecOut = `cache app` + f.env.Set("KOOL_NAME", "example") f.getServiceIDCmd.(*builder.FakeCommand).MockExecOut = "output" f.getServiceStatusPortCmd.(*builder.FakeCommand).MockExecOut = "output" diff --git a/commands/stop.go b/commands/stop.go index 3829f7ec..10f9fc7e 100644 --- a/commands/stop.go +++ b/commands/stop.go @@ -2,7 +2,9 @@ package commands import ( "kool-dev/kool/core/builder" + "kool-dev/kool/core/environment" "kool-dev/kool/services/checker" + "kool-dev/kool/services/proxy" "time" "github.com/spf13/cobra" @@ -18,9 +20,11 @@ type KoolStop struct { DefaultKoolService Flags *KoolStopFlags - check checker.Checker - down builder.Command - rm builder.Command + check checker.Checker + env environment.EnvStorage + getProjects builder.Command + down builder.Command + rm builder.Command } func AddKoolStop(root *cobra.Command) { @@ -39,6 +43,8 @@ func NewKoolStop() *KoolStop { *defaultKoolService, &KoolStopFlags{false}, checker.NewChecker(defaultKoolService.shell), + environment.NewEnvStorage(), + builder.NewCommand("docker", "ps", "--all"), builder.NewCommand("docker", "compose", "down"), builder.NewCommand("docker", "compose", "rm"), } @@ -51,6 +57,26 @@ func (s *KoolStop) Execute(args []string) (err error) { if err = s.check.Check(); err != nil { return } + if len(args) == 0 && workspacesEnabled(s.env) && !isWorkspace(s.env) { + var projects []string + if projects, err = activeWorkspaceProjects(s.Shell(), s.getProjects, s.env); err != nil { + return + } + for _, project := range projects { + workspaceDown := builder.NewCommand("docker", "compose", "--project-name", project, "down", "--remove-orphans") + if s.Flags.Purge { + workspaceDown.AppendArgs("--volumes") + } + if err = s.Shell().Interactive(workspaceDown); err != nil { + return + } + if proxyEnabled(s.env) { + if proxyErr := proxy.NewManager(s.Shell(), s.env).RemoveProject(project, nil); proxyErr != nil { + s.Shell().Warning("Could not remove proxy routes:", proxyErr) + } + } + } + } if len(args) == 0 { s.down.AppendArgs("--remove-orphans") @@ -76,6 +102,14 @@ func (s *KoolStop) Execute(args []string) (err error) { } err = s.Shell().Interactive(stopCommand) + if err != nil { + return + } + if proxyEnabled(s.env) { + if proxyErr := proxy.NewManager(s.Shell(), s.env).Remove(args); proxyErr != nil { + s.Shell().Warning("Could not remove proxy routes:", proxyErr) + } + } time.Sleep(time.Second * 2) return } diff --git a/commands/stop_test.go b/commands/stop_test.go index d531328c..cdb5866e 100644 --- a/commands/stop_test.go +++ b/commands/stop_test.go @@ -4,16 +4,40 @@ import ( "errors" "io" "kool-dev/kool/core/builder" + "kool-dev/kool/core/environment" "kool-dev/kool/core/shell" "kool-dev/kool/services/checker" + "strings" "testing" ) +type workspaceStopShell struct { + shell.FakeShell + commands []string +} + +func (s *workspaceStopShell) Exec(command builder.Command, extraArgs ...string) (string, error) { + if strings.HasPrefix(command.String(), "docker ps --all") { + return "example-workspace-task-b\nexample-workspace-task-a\n", nil + } + if strings.HasPrefix(command.String(), "docker inspect kool-proxy") { + return "", errors.New("proxy is not running") + } + return "", nil +} + +func (s *workspaceStopShell) Interactive(command builder.Command, extraArgs ...string) error { + s.commands = append(s.commands, strings.TrimSpace(command.String()+" "+strings.Join(extraArgs, " "))) + return nil +} + func newFakeKoolStop() *KoolStop { fs := &KoolStop{ *(newDefaultKoolService().Fake()), &KoolStopFlags{false}, &checker.FakeChecker{}, + environment.NewFakeEnvStorage(), + &builder.FakeCommand{}, &builder.FakeCommand{}, &builder.FakeCommand{}, } @@ -126,6 +150,28 @@ func TestNewStopPurgeCommandWithServices(t *testing.T) { } } +func TestStopMainStopsWorkspacesFirst(t *testing.T) { + f := newFakeKoolStop() + recorder := &workspaceStopShell{} + f.shell = recorder + f.env.Set("KOOL_NAME", "example") + f.env.Set("KOOL_WORKSPACES_ENABLED", "true") + f.getProjects = builder.NewCommand("docker", "ps", "--all") + f.down = builder.NewCommand("docker", "compose", "down") + + if err := f.Execute(nil); err != nil { + t.Fatal(err) + } + expected := []string{ + "docker compose --project-name example-workspace-task-a down --remove-orphans", + "docker compose --project-name example-workspace-task-b down --remove-orphans", + "docker compose down --remove-orphans", + } + if strings.Join(recorder.commands, "\n") != strings.Join(expected, "\n") { + t.Errorf("expected workspace projects to stop before main:\n%v\ngot:\n%v", expected, recorder.commands) + } +} + func TestNewFailingDependenciesCheckStopCommand(t *testing.T) { f := newFakeKoolStop() diff --git a/commands/workspace.go b/commands/workspace.go new file mode 100644 index 00000000..9b2b86f4 --- /dev/null +++ b/commands/workspace.go @@ -0,0 +1,108 @@ +package commands + +import ( + "fmt" + "kool-dev/kool/core/builder" + "kool-dev/kool/core/environment" + "kool-dev/kool/core/parser" + "kool-dev/kool/core/shell" + "sort" + "strings" +) + +func isWorkspace(env environment.EnvStorage) bool { + return env.IsTrue("KOOL_WORKSPACE") +} + +func workspacesEnabled(env environment.EnvStorage) bool { + return env.IsTrue("KOOL_WORKSPACES_ENABLED") +} + +func proxyEnabled(env environment.EnvStorage) bool { + return env.IsTrue("KOOL_PROXY_ENABLED") +} + +func workspaceServices(env environment.EnvStorage) []string { + configured := env.Get("KOOL_WORKSPACE_SERVICES") + if configured == "" { + return nil + } + return strings.Split(configured, ",") +} + +func configuredWorkspaceServices(env environment.EnvStorage) []string { + if services := workspaceServices(env); len(services) > 0 { + return services + } + config, err := parser.LoadKoolYaml(env.Get("PWD")) + if err != nil { + return nil + } + services := append([]string(nil), config.Workspaces...) + sort.Strings(services) + return services +} + +func sourceProject(env environment.EnvStorage) string { + if project := env.Get("KOOL_WORKSPACE_SOURCE_PROJECT"); project != "" { + return project + } + if project := env.Get("COMPOSE_PROJECT_NAME"); project != "" && !isWorkspace(env) { + return project + } + return env.Get("KOOL_NAME") +} + +func currentProject(env environment.EnvStorage) string { + if isWorkspace(env) { + return env.Get("KOOL_WORKSPACE_PROJECT") + } + return sourceProject(env) +} + +func activeWorkspaceProjects(sh shell.Shell, command builder.Command, env environment.EnvStorage) ([]string, error) { + output, err := sh.Exec(command, + "--filter", "label="+environment.WorkspaceSourceLabel+"="+sourceProject(env), + "--format", `{{.Label "com.docker.compose.project"}}`, + ) + if err != nil { + return nil, err + } + prefix := sourceProject(env) + "-workspace-" + seen := make(map[string]bool) + var projects []string + for _, project := range strings.Split(strings.ReplaceAll(output, "\r\n", "\n"), "\n") { + project = strings.TrimSpace(project) + if strings.HasPrefix(project, prefix) && !seen[project] { + projects = append(projects, project) + seen[project] = true + } + } + sort.Strings(projects) + return projects, nil +} + +func selectWorkspaceServices(env environment.EnvStorage, requested []string) ([]string, error) { + if !isWorkspace(env) { + return requested, nil + } + + configured := workspaceServices(env) + if len(configured) == 0 { + return nil, fmt.Errorf("no services are configured under workspaces in kool.yml") + } + if len(requested) == 0 { + return configured, nil + } + + allowed := make(map[string]bool, len(configured)) + for _, service := range configured { + allowed[service] = true + } + for _, service := range requested { + if !allowed[service] { + return nil, fmt.Errorf("service %s is not enabled for workspaces", service) + } + } + return requested, nil +} diff --git a/core/environment/agent.go b/core/environment/agent.go new file mode 100644 index 00000000..cbf1e94b --- /dev/null +++ b/core/environment/agent.go @@ -0,0 +1,46 @@ +package environment + +import "strings" + +var agentEnvVars = []string{ + "AI_AGENT", + "CURSOR_AGENT", + "GEMINI_CLI", + "CODEX_SANDBOX", + "CODEX_CI", + "CODEX_THREAD_ID", + "AUGMENT_AGENT", + "OPENCODE_CLIENT", + "OPENCODE", + "AMP_CURRENT_THREAD_ID", + "CLAUDECODE", + "CLAUDE_CODE", + "CLAUDE_CODE_IS_COWORK", + "REPL_ID", + "COPILOT_MODEL", + "COPILOT_ALLOW_ALL", + "COPILOT_CLI", + "ANTIGRAVITY_AGENT", + "PI_CODING_AGENT", + "KIRO_AGENT_PATH", +} + +// AgentEnvVars returns agent environment variables that are safe to forward. +func AgentEnvVars(env EnvStorage) []string { + values := make(map[string]string) + for _, entry := range env.All() { + parts := strings.SplitN(entry, "=", 2) + if len(parts) == 2 { + values[parts[0]] = parts[1] + } + } + + var forwarded []string + for _, name := range agentEnvVars { + if value, exists := values[name]; exists { + forwarded = append(forwarded, name+"="+value) + } + } + + return forwarded +} diff --git a/core/environment/agent_test.go b/core/environment/agent_test.go new file mode 100644 index 00000000..ab2bdd46 --- /dev/null +++ b/core/environment/agent_test.go @@ -0,0 +1,22 @@ +package environment + +import ( + "reflect" + "testing" +) + +func TestAgentEnvVars(t *testing.T) { + env := NewFakeEnvStorage() + env.Envs = map[string]string{ + "AI_AGENT": "custom", + "OPENCODE": "1", + "CLAUDECODE": "", + "COPILOT_GITHUB_TOKEN": "secret", + "UNRELATED": "value", + } + + expected := []string{"AI_AGENT=custom", "OPENCODE=1", "CLAUDECODE="} + if actual := AgentEnvVars(env); !reflect.DeepEqual(actual, expected) { + t.Fatalf("expected %v, got %v", expected, actual) + } +} diff --git a/core/environment/env.go b/core/environment/env.go index dd8a2462..7f6798db 100644 --- a/core/environment/env.go +++ b/core/environment/env.go @@ -1,12 +1,15 @@ package environment import ( + "kool-dev/kool/core/parser" "log" "os" + "path/filepath" "strings" ) var envFiles = []string{".env.local", ".env"} +var loadedEnvKeys = make(map[string][]string) // InitEnvironmentVariables handles the reading of .env files and // setting up important environment variables necessary for kool @@ -35,16 +38,37 @@ func InitEnvironmentVariables(envStorage EnvStorage) { log.Fatal("Could not evaluate working directory - ", err) } envStorage.Set("PWD", workDir) + envDirectory := canonicalDirectory(workDir) + loadedEnvKeys[envDirectory] = nil for _, envFile := range envFiles { if _, err = os.Stat(envFile); os.IsNotExist(err) { continue } + before := envKeySet(envStorage.All()) err = envStorage.Load(envFile) if err != nil { log.Fatal("Failure loading environment file ", envFile, " error: '", err, "'") } + for key := range envKeySet(envStorage.All()) { + if !before[key] { + loadedEnvKeys[envDirectory] = append(loadedEnvKeys[envDirectory], key) + } + } + } + + config := loadKoolConfig(workDir) + if config != nil && len(config.Workspaces) > 0 { + envStorage.Set("KOOL_WORKSPACES_ENABLED", "true") + initRift(envStorage, workDir) + initGitWorktree(envStorage, workDir) + initSourceProject(envStorage, workDir) + } + if config != nil && config.Proxy != nil { + envStorage.Set("KOOL_PROXY_ENABLED", "true") + initSourceProject(envStorage, workDir) + initProxy(envStorage, config) } // Now that we loaded up the files, we will check for @@ -60,3 +84,53 @@ func InitEnvironmentVariables(envStorage EnvStorage) { initAsuser(envStorage) } + +// LoadedEnvKeys returns variables introduced from a directory's environment files. +func LoadedEnvKeys(directory string) []string { + return append([]string(nil), loadedEnvKeys[canonicalDirectory(directory)]...) +} + +func canonicalDirectory(directory string) string { + canonical, err := filepath.EvalSymlinks(directory) + if err == nil { + return canonical + } + canonical, err = filepath.Abs(directory) + if err == nil { + return canonical + } + return directory +} + +func envKeySet(entries []string) map[string]bool { + keys := make(map[string]bool, len(entries)) + for _, entry := range entries { + keys[strings.SplitN(entry, "=", 2)[0]] = true + } + return keys +} + +// loadKoolConfig decodes the kool config file for the given directory, or +// returns nil when there is none - or when it cannot be decoded, since +// environment setup runs before we have any means of reporting the failure. +func loadKoolConfig(workDir string) *parser.KoolYaml { + config, err := parser.LoadKoolYaml(workDir) + if err != nil { + return nil + } + return config +} + +func initSourceProject(envStorage EnvStorage, workDir string) { + if envStorage.Get("KOOL_WORKSPACE_SOURCE_PROJECT") != "" { + return + } + project := envStorage.Get("COMPOSE_PROJECT_NAME") + if project == "" { + project = composeSourceProject(envStorage, workDir) + } + if project == "" { + project = composeProjectName(filepath.Base(workDir)) + } + envStorage.Set("KOOL_WORKSPACE_SOURCE_PROJECT", project) +} diff --git a/core/environment/env_test.go b/core/environment/env_test.go index 366256fc..3b010784 100644 --- a/core/environment/env_test.go +++ b/core/environment/env_test.go @@ -79,3 +79,47 @@ func TestInitEnvironmentVariablesOverridesStalePWD(t *testing.T) { t.Errorf("expecting $PWD to be overridden to '%s', got '%s'", workDir, envWorkDir) } } + +func TestInitSourceProjectUsesComposeName(t *testing.T) { + workDir := t.TempDir() + if err := os.WriteFile(filepath.Join(workDir, "compose.yml"), []byte("name: ${PROJECT_NAME:-custom-project}\nservices: {}\n"), 0644); err != nil { + t.Fatal(err) + } + env := NewFakeEnvStorage() + + initSourceProject(env, workDir) + + if got := env.Get("KOOL_WORKSPACE_SOURCE_PROJECT"); got != "custom-project" { + t.Errorf("expected Compose name custom-project, got %q", got) + } +} + +func TestInitEnvironmentSkipsWorkspaceDetectionWithoutConfiguration(t *testing.T) { + originalDirectory, err := os.Getwd() + if err != nil { + t.Fatal(err) + } + workDir := t.TempDir() + if err = os.Chdir(workDir); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.Chdir(originalDirectory) }) + + originalGitWorktreeOutput := gitWorktreeOutput + gitCalled := false + gitWorktreeOutput = func(string, ...string) ([]byte, error) { + gitCalled = true + return nil, nil + } + t.Cleanup(func() { gitWorktreeOutput = originalGitWorktreeOutput }) + + env := NewFakeEnvStorage() + InitEnvironmentVariables(env) + + if gitCalled { + t.Error("expected Git workspace detection to be skipped") + } + if env.IsTrue("KOOL_WORKSPACES_ENABLED") || env.IsTrue("KOOL_PROXY_ENABLED") { + t.Error("expected workspace and proxy features to remain disabled") + } +} diff --git a/core/environment/git_worktree.go b/core/environment/git_worktree.go new file mode 100644 index 00000000..0aba1165 --- /dev/null +++ b/core/environment/git_worktree.go @@ -0,0 +1,94 @@ +package environment + +import ( + "fmt" + "os/exec" + "path/filepath" + "strings" +) + +var gitWorktreeOutput = func(workDir string, args ...string) ([]byte, error) { + arguments := append([]string{"-C", workDir}, args...) + return exec.Command("git", arguments...).Output() +} + +func initGitWorktree(envStorage EnvStorage, workDir string) { + if envStorage.Get("KOOL_WORKSPACE_PROVIDER") != "" { + return + } + + rootOutput, err := gitWorktreeOutput(workDir, "rev-parse", "--show-toplevel") + if err != nil { + return + } + root := strings.TrimSpace(string(rootOutput)) + + gitDirOutput, err := gitWorktreeOutput(workDir, "rev-parse", "--git-dir") + if err != nil { + setGitWorktreeError(envStorage, err) + return + } + commonDirOutput, err := gitWorktreeOutput(workDir, "rev-parse", "--git-common-dir") + if err != nil { + setGitWorktreeError(envStorage, err) + return + } + + gitDir := absoluteGitPath(workDir, strings.TrimSpace(string(gitDirOutput))) + commonDir := absoluteGitPath(workDir, strings.TrimSpace(string(commonDirOutput))) + if gitDir == commonDir { + return + } + + listOutput, err := gitWorktreeOutput(workDir, "worktree", "list", "--porcelain", "-z") + if err != nil { + setGitWorktreeError(envStorage, err) + return + } + source := firstGitWorktree(listOutput) + if source == "" { + envStorage.Set("KOOL_WORKSPACE_ERROR", "could not resolve Git worktree source") + return + } + + initWorkspaceContext(envStorage, workDir, source, root, "worktree", true, hasDuplicateWorktreeBasename(listOutput, root)) +} + +func absoluteGitPath(workDir, path string) string { + if !filepath.IsAbs(path) { + path = filepath.Join(workDir, path) + } + absolute, err := filepath.Abs(path) + if err != nil { + return filepath.Clean(path) + } + return filepath.Clean(absolute) +} + +func firstGitWorktree(output []byte) string { + for _, field := range strings.Split(string(output), "\x00") { + if strings.HasPrefix(field, "worktree ") { + return strings.TrimPrefix(field, "worktree ") + } + } + return "" +} + +func hasDuplicateWorktreeBasename(output []byte, workspace string) bool { + name := filepath.Base(workspace) + workspace = canonicalDirectory(workspace) + for _, field := range strings.Split(string(output), "\x00") { + if !strings.HasPrefix(field, "worktree ") { + continue + } + candidate := strings.TrimPrefix(field, "worktree ") + if filepath.Base(candidate) == name && canonicalDirectory(candidate) != workspace { + return true + } + } + return false +} + +func setGitWorktreeError(envStorage EnvStorage, err error) { + envStorage.Set("KOOL_WORKSPACE_ERROR", fmt.Sprintf("could not resolve Git worktree: %v", err)) +} diff --git a/core/environment/git_worktree_test.go b/core/environment/git_worktree_test.go new file mode 100644 index 00000000..f6ccf824 --- /dev/null +++ b/core/environment/git_worktree_test.go @@ -0,0 +1,114 @@ +package environment + +import ( + "fmt" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestInitGitWorktree(t *testing.T) { + root := t.TempDir() + source := filepath.Join(root, "My App") + workspace := filepath.Join(root, "trees", "task-a") + workDir := filepath.Join(workspace, "packages", "api") + gitDir := filepath.Join(source, ".git", "worktrees", "task-a") + commonDir := filepath.Join(source, ".git") + + originalGitWorktreeOutput := gitWorktreeOutput + gitWorktreeOutput = fakeGitWorktreeOutput(map[string]string{ + "rev-parse --show-toplevel": workspace, + "rev-parse --git-dir": gitDir, + "rev-parse --git-common-dir": commonDir, + "worktree list --porcelain -z": "worktree " + source + "\x00HEAD abc\x00\x00worktree " + workspace + "\x00HEAD def\x00", + }) + defer func() { gitWorktreeOutput = originalGitWorktreeOutput }() + + env := NewFakeEnvStorage() + initGitWorktree(env, workDir) + + if !env.IsTrue("KOOL_WORKSPACE") { + t.Error("expected linked Git worktree to be a workspace") + } + if got := env.Get("KOOL_WORKSPACE_PROVIDER"); got != "worktree" { + t.Errorf("expected worktree provider, got %q", got) + } + if got := env.Get("KOOL_WORKSPACE_SOURCE"); got != source { + t.Errorf("expected source %q, got %q", source, got) + } + if got := env.Get("KOOL_WORKSPACE_PATH"); got != workspace { + t.Errorf("expected workspace path %q, got %q", workspace, got) + } + expectedProject := "myapp-workspace-" + composeProjectName(workspaceIdentity(workspace)) + if got := env.Get("COMPOSE_PROJECT_NAME"); got != expectedProject { + t.Errorf("expected worktree Compose project %q, got %q", expectedProject, got) + } +} + +func TestInitGitWorktreeIgnoresMainWorktree(t *testing.T) { + root := t.TempDir() + originalGitWorktreeOutput := gitWorktreeOutput + gitWorktreeOutput = fakeGitWorktreeOutput(map[string]string{ + "rev-parse --show-toplevel": root, + "rev-parse --git-dir": ".git", + "rev-parse --git-common-dir": ".git", + }) + defer func() { gitWorktreeOutput = originalGitWorktreeOutput }() + + env := NewFakeEnvStorage() + initGitWorktree(env, root) + + if env.IsTrue("KOOL_WORKSPACE") { + t.Error("expected main Git worktree to remain the source project") + } + if got := env.Get("KOOL_WORKSPACE_PROVIDER"); got != "" { + t.Errorf("expected no workspace provider, got %q", got) + } +} + +func TestInitGitWorktreePreservesRiftContext(t *testing.T) { + env := NewFakeEnvStorage() + env.Set("KOOL_WORKSPACE_PROVIDER", "rift") + + originalGitWorktreeOutput := gitWorktreeOutput + gitWorktreeOutput = func(string, ...string) ([]byte, error) { + return nil, fmt.Errorf("Git should not be called") + } + defer func() { gitWorktreeOutput = originalGitWorktreeOutput }() + + initGitWorktree(env, t.TempDir()) + + if got := env.Get("KOOL_WORKSPACE_PROVIDER"); got != "rift" { + t.Errorf("expected Rift context to be preserved, got %q", got) + } +} + +func TestDuplicateGitWorktreeBasenamesUseUniqueHosts(t *testing.T) { + root := t.TempDir() + first := filepath.Join(root, "one", "task") + second := filepath.Join(root, "two", "task") + for _, workspace := range []string{first, second} { + if err := os.MkdirAll(workspace, 0755); err != nil { + t.Fatal(err) + } + } + output := []byte("worktree " + first + "\x00HEAD abc\x00\x00worktree " + second + "\x00HEAD def\x00") + if !hasDuplicateWorktreeBasename(output, first) || !hasDuplicateWorktreeBasename(output, second) { + t.Fatal("expected duplicate worktree basenames to require unique hosts") + } + if hasDuplicateWorktreeBasename(output, filepath.Join(root, "three", "unique")) { + t.Fatal("did not expect a unique worktree basename to require a hash") + } +} + +func fakeGitWorktreeOutput(outputs map[string]string) func(string, ...string) ([]byte, error) { + return func(_ string, args ...string) ([]byte, error) { + key := strings.Join(args, " ") + output, found := outputs[key] + if !found { + return nil, fmt.Errorf("unexpected Git command: %s", key) + } + return []byte(output), nil + } +} diff --git a/core/environment/proxy.go b/core/environment/proxy.go new file mode 100644 index 00000000..3c15e9a2 --- /dev/null +++ b/core/environment/proxy.go @@ -0,0 +1,31 @@ +package environment + +import ( + "kool-dev/kool/core/parser" + "os" + "strings" + + "github.com/compose-spec/compose-go/template" +) + +func initProxy(envStorage EnvStorage, config *parser.KoolYaml) { + if config == nil || config.Proxy == nil { + return + } + + domain, err := template.Substitute(config.Proxy.Domain, os.LookupEnv) + if err != nil { + return + } + domain = strings.TrimSuffix(strings.TrimSpace(domain), ".") + if domain == "" { + return + } + + envStorage.Set("KOOL_PROXY_DOMAIN", domain) + host := domain + if envStorage.IsTrue("KOOL_WORKSPACE") { + host = envStorage.Get("KOOL_WORKSPACE_NAME") + ".workspace." + domain + } + envStorage.Set("KOOL_PROXY_HOST", host) +} diff --git a/core/environment/proxy_test.go b/core/environment/proxy_test.go new file mode 100644 index 00000000..e6f9003b --- /dev/null +++ b/core/environment/proxy_test.go @@ -0,0 +1,57 @@ +package environment + +import ( + "os" + "path/filepath" + "testing" +) + +func TestInitProxySourceHost(t *testing.T) { + workDir := t.TempDir() + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte("proxy:\n domain: exlink.localhost\n routes: {}\n"), 0644); err != nil { + t.Fatal(err) + } + env := NewFakeEnvStorage() + + initProxy(env, loadKoolConfig(workDir)) + + if got := env.Get("KOOL_PROXY_DOMAIN"); got != "exlink.localhost" { + t.Errorf("expected proxy domain, got %q", got) + } + if got := env.Get("KOOL_PROXY_HOST"); got != "exlink.localhost" { + t.Errorf("expected source proxy host, got %q", got) + } +} + +func TestInitProxyWorkspaceHost(t *testing.T) { + workDir := t.TempDir() + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte("proxy:\n domain: exlink.localhost\n routes: {}\n"), 0644); err != nil { + t.Fatal(err) + } + env := NewFakeEnvStorage() + env.Set("KOOL_WORKSPACE", "true") + env.Set("KOOL_WORKSPACE_NAME", "vite-smoke") + + initProxy(env, loadKoolConfig(workDir)) + + if got := env.Get("KOOL_PROXY_HOST"); got != "vite-smoke.workspace.exlink.localhost" { + t.Errorf("expected workspace proxy host, got %q", got) + } +} + +func TestInitProxyGitWorktreeHost(t *testing.T) { + workDir := t.TempDir() + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte("proxy:\n domain: exlink.localhost\n routes: {}\n"), 0644); err != nil { + t.Fatal(err) + } + env := NewFakeEnvStorage() + env.Set("KOOL_WORKSPACE", "true") + env.Set("KOOL_WORKSPACE_NAME", "vite-smoke") + env.Set("KOOL_WORKSPACE_PROVIDER", "worktree") + + initProxy(env, loadKoolConfig(workDir)) + + if got := env.Get("KOOL_PROXY_HOST"); got != "vite-smoke.workspace.exlink.localhost" { + t.Errorf("expected Git worktree proxy host, got %q", got) + } +} diff --git a/core/environment/rift.go b/core/environment/rift.go new file mode 100644 index 00000000..d8cce0e9 --- /dev/null +++ b/core/environment/rift.go @@ -0,0 +1,68 @@ +package environment + +import ( + "fmt" + "os" + "os/exec" + "path/filepath" + "strings" +) + +var riftAncestors = func(workDir string) ([]byte, error) { + return exec.Command("rift", "ancestors", workDir).Output() +} + +func initRift(envStorage EnvStorage, workDir string) { + workspace, found := findRiftWorkspace(workDir) + if !found { + return + } + + output, err := riftAncestors(workDir) + if err != nil { + envStorage.Set("KOOL_WORKSPACE_ERROR", fmt.Sprintf("could not resolve Rift workspace ancestry: %v", err)) + return + } + + var ancestors []string + for _, ancestor := range strings.Split(strings.TrimSpace(string(output)), "\n") { + if ancestor != "" { + ancestors = append(ancestors, ancestor) + } + } + source := workspace + if len(ancestors) > 0 { + source = ancestors[len(ancestors)-1] + } + + initWorkspaceContext(envStorage, workDir, source, workspace, "rift", len(ancestors) > 0, false) +} + +func findRiftWorkspace(workDir string) (string, bool) { + for directory := filepath.Clean(workDir); ; directory = filepath.Dir(directory) { + if _, err := os.Stat(filepath.Join(directory, ".rift")); err == nil { + return directory, true + } + + parent := filepath.Dir(directory) + if parent == directory { + return "", false + } + } +} + +func composeProjectName(name string) string { + name = strings.ToLower(name) + name = strings.Map(func(character rune) rune { + if character >= 'a' && character <= 'z' || character >= '0' && character <= '9' || character == '-' || character == '_' { + return character + } + return -1 + }, name) + + name = strings.TrimLeft(name, "-_") + if name == "" { + return "kool" + } + return name +} diff --git a/core/environment/rift_test.go b/core/environment/rift_test.go new file mode 100644 index 00000000..fab1130c --- /dev/null +++ b/core/environment/rift_test.go @@ -0,0 +1,193 @@ +package environment + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestInitRiftWorkspace(t *testing.T) { + root := t.TempDir() + source := filepath.Join(root, "My App") + workspace := filepath.Join(root, ".rifts", "My App", "task-a") + workDir := filepath.Join(workspace, "packages", "api") + + for _, directory := range []string{source, workDir} { + if err := os.MkdirAll(directory, 0755); err != nil { + t.Fatal(err) + } + } + if err := os.WriteFile(filepath.Join(workspace, ".rift"), []byte("workspace"), 0644); err != nil { + t.Fatal(err) + } + + originalRiftAncestors := riftAncestors + riftAncestors = func(string) ([]byte, error) { + return []byte(source + "\n"), nil + } + defer func() { riftAncestors = originalRiftAncestors }() + + env := NewFakeEnvStorage() + initRift(env, workDir) + + if !env.IsTrue("KOOL_WORKSPACE") { + t.Error("expected created Rift to be a workspace") + } + if got := env.Get("KOOL_WORKSPACE_PROVIDER"); got != "rift" { + t.Errorf("expected Rift provider, got %q", got) + } + if got := env.Get("KOOL_WORKSPACE_SOURCE"); got != source { + t.Errorf("expected KOOL_WORKSPACE_SOURCE %q, got %q", source, got) + } + if got := env.Get("KOOL_WORKSPACE_PATH"); got != workspace { + t.Errorf("expected KOOL_WORKSPACE_PATH %q, got %q", workspace, got) + } + expectedProject := "myapp-workspace-" + composeProjectName(workspaceIdentity(workspace)) + if got := env.Get("COMPOSE_PROJECT_NAME"); got != expectedProject { + t.Errorf("expected COMPOSE_PROJECT_NAME %q, got %q", expectedProject, got) + } +} + +func TestWorkspaceHostNameIsDNSLabel(t *testing.T) { + for input, expected := range map[string]string{ + "Feature Login": "feature-login", + "TASK_api": "task-api", + "---": "workspace", + } { + if actual := workspaceHostName(input); actual != expected { + t.Errorf("expected workspace host %q for %q, got %q", expected, input, actual) + } + } + if actual := workspaceHostName(strings.Repeat("a", 70)); len(actual) != 63 { + t.Errorf("expected workspace host to be limited to 63 characters, got %d", len(actual)) + } +} + +func TestInitRiftUsesConfiguredSourceProjectName(t *testing.T) { + workspace := t.TempDir() + if err := os.WriteFile(filepath.Join(workspace, ".rift"), []byte("workspace"), 0644); err != nil { + t.Fatal(err) + } + + originalRiftAncestors := riftAncestors + riftAncestors = func(string) ([]byte, error) { + return []byte("/projects/app\n"), nil + } + defer func() { riftAncestors = originalRiftAncestors }() + + env := NewFakeEnvStorage() + env.Set("COMPOSE_PROJECT_NAME", "custom-project") + initRift(env, workspace) + + expected := "custom-project-workspace-" + composeProjectName(workspaceIdentity(workspace)) + if got := env.Get("COMPOSE_PROJECT_NAME"); got != expected { + t.Errorf("expected COMPOSE_PROJECT_NAME %q, got %q", expected, got) + } +} + +func TestInitRiftWorkspaceComposeProject(t *testing.T) { + t.Cleanup(CleanupWorkspace) + root := t.TempDir() + source := filepath.Join(root, "app") + workspace := filepath.Join(root, ".rifts", "app", "task-a") + if err := os.MkdirAll(workspace, 0755); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(workspace, ".rift"), []byte("workspace"), 0644); err != nil { + t.Fatal(err) + } + compose := `name: custom-app +services: + database: + image: mysql + node: + image: node + app: + image: php + ports: + - "80:80" +` + if err := os.WriteFile(filepath.Join(workspace, "compose.yml"), []byte(compose), 0644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(workspace, "kool.yml"), []byte("workspaces: [app, node]\n"), 0644); err != nil { + t.Fatal(err) + } + + originalRiftAncestors := riftAncestors + riftAncestors = func(string) ([]byte, error) { return []byte(source + "\n"), nil } + defer func() { riftAncestors = originalRiftAncestors }() + + env := NewFakeEnvStorage() + initRift(env, workspace) + + if !env.IsTrue("KOOL_WORKSPACE") { + t.Error("expected KOOL_WORKSPACE to be enabled") + } + expectedProject := "custom-app-workspace-" + composeProjectName(workspaceIdentity(workspace)) + if got := env.Get("COMPOSE_PROJECT_NAME"); got != expectedProject { + t.Errorf("expected workspace Compose project %q, got %q", expectedProject, got) + } + if got := env.Get("KOOL_WORKSPACE_SERVICES"); got != "app,node" { + t.Errorf("expected sorted workspace services, got %q", got) + } + files := strings.Split(env.Get("COMPOSE_FILE"), string(os.PathListSeparator)) + if len(files) != 2 { + t.Fatalf("expected Compose file and generated override, got %v", files) + } + override, err := os.ReadFile(files[1]) + if err != nil { + t.Fatal(err) + } + for _, expected := range []string{"ports: !reset []", "container_name: !reset null"} { + if !strings.Contains(string(override), expected) { + t.Errorf("expected generated override to contain %q, got:\n%s", expected, override) + } + } +} + +func TestInitRiftPreservesComposeFileProjectName(t *testing.T) { + workspace := t.TempDir() + if err := os.WriteFile(filepath.Join(workspace, ".rift"), []byte("workspace"), 0644); err != nil { + t.Fatal(err) + } + if err := os.WriteFile(filepath.Join(workspace, "compose.yml"), []byte("name: custom-project\nservices: {}\n"), 0644); err != nil { + t.Fatal(err) + } + + originalRiftAncestors := riftAncestors + riftAncestors = func(string) ([]byte, error) { + return []byte("/projects/app\n"), nil + } + defer func() { riftAncestors = originalRiftAncestors }() + + env := NewFakeEnvStorage() + initRift(env, workspace) + + expected := "custom-project-workspace-" + composeProjectName(workspaceIdentity(workspace)) + if got := env.Get("COMPOSE_PROJECT_NAME"); got != expected { + t.Errorf("expected top-level Compose project name %q, got %q", expected, got) + } +} + +func TestInitRiftOriginalWorkspaceSetsSourceContext(t *testing.T) { + workspace := t.TempDir() + if err := os.WriteFile(filepath.Join(workspace, ".rift"), []byte("workspace"), 0644); err != nil { + t.Fatal(err) + } + + originalRiftAncestors := riftAncestors + riftAncestors = func(string) ([]byte, error) { return nil, nil } + defer func() { riftAncestors = originalRiftAncestors }() + + env := NewFakeEnvStorage() + initRift(env, workspace) + + if env.IsTrue("KOOL_WORKSPACE") { + t.Error("expected original workspace not to be a created workspace") + } + if got := env.Get("KOOL_WORKSPACE_SOURCE"); got != workspace { + t.Errorf("expected KOOL_WORKSPACE_SOURCE %q, got %q", workspace, got) + } +} diff --git a/core/environment/workspace.go b/core/environment/workspace.go new file mode 100644 index 00000000..f134a085 --- /dev/null +++ b/core/environment/workspace.go @@ -0,0 +1,229 @@ +package environment + +import ( + "crypto/sha256" + "fmt" + "os" + "path/filepath" + "sort" + "strings" + + "github.com/compose-spec/compose-go/template" + "gopkg.in/yaml.v2" +) + +var workspaceOverrideFile string + +// WorkspaceSourceLabel identifies containers owned by a source project's workspaces. +const WorkspaceSourceLabel = "dev.kool.workspace.source" + +func initWorkspaceContext(envStorage EnvStorage, workDir, source, workspace, provider string, active, uniqueHost bool) { + envStorage.Set("KOOL_WORKSPACE_PROVIDER", provider) + envStorage.Set("KOOL_WORKSPACE_SOURCE", source) + if envStorage.Get("KOOL_NAME") == "" { + envStorage.Set("KOOL_NAME", filepath.Base(source)) + } + + sourceProject := envStorage.Get("COMPOSE_PROJECT_NAME") + if sourceProject == "" { + sourceProject = composeSourceProject(envStorage, workDir) + if sourceProject == "" { + sourceProject = composeProjectName(filepath.Base(source)) + } + } + envStorage.Set("KOOL_WORKSPACE_SOURCE_PROJECT", sourceProject) + if !active { + return + } + + workspaceName := filepath.Base(workspace) + workspaceProject := sourceProject + "-workspace-" + composeProjectName(workspaceIdentity(workspace)) + workspaceHost := workspaceName + if uniqueHost { + workspaceHost = workspaceIdentity(workspace) + } + envStorage.Set("KOOL_WORKSPACE", "true") + envStorage.Set("KOOL_WORKSPACE_NAME", workspaceHostName(workspaceHost)) + envStorage.Set("KOOL_WORKSPACE_PATH", workspace) + envStorage.Set("KOOL_WORKSPACE_PROJECT", workspaceProject) + envStorage.Set("COMPOSE_PROJECT_NAME", workspaceProject) + initWorkspaceCompose(envStorage, workDir) +} + +func workspaceIdentity(workspace string) string { + canonical, err := filepath.EvalSymlinks(workspace) + if err != nil { + canonical, err = filepath.Abs(workspace) + if err != nil { + canonical = workspace + } + } + digest := sha256.Sum256([]byte(canonical)) + suffix := fmt.Sprintf("-%x", digest[:4]) + name := filepath.Base(workspace) + if len(name) > 63-len(suffix) { + name = strings.TrimRight(name[:63-len(suffix)], "-_") + } + return name + suffix +} + +func workspaceHostName(name string) string { + name = strings.ToLower(name) + var host strings.Builder + separator := false + for _, character := range name { + if character >= 'a' && character <= 'z' || character >= '0' && character <= '9' { + host.WriteRune(character) + separator = false + } else if host.Len() > 0 && !separator { + host.WriteByte('-') + separator = true + } + } + result := strings.Trim(host.String(), "-") + if result == "" { + return "workspace" + } + if len(result) > 63 { + result = strings.TrimRight(result[:63], "-") + } + return result +} + +type workspaceComposeConfig struct { + Name string `yaml:"name"` +} + +func initWorkspaceCompose(envStorage EnvStorage, workDir string) []string { + koolConfig := loadKoolConfig(workDir) + if koolConfig == nil || len(koolConfig.Workspaces) == 0 { + return nil + } + unique := make(map[string]bool, len(koolConfig.Workspaces)) + var selected []string + for _, service := range koolConfig.Workspaces { + if service != "" && !unique[service] { + selected = append(selected, service) + unique[service] = true + } + } + sort.Strings(selected) + if len(selected) == 0 { + return nil + } + + files := composeFiles(envStorage, workDir) + if len(files) == 0 { + return nil + } + + for _, file := range files { + content, err := os.ReadFile(file) + if err != nil { + return nil + } + + config := workspaceComposeConfig{} + if yaml.Unmarshal(content, &config) != nil { + return nil + } + } + + CleanupWorkspace() + override, err := os.CreateTemp("", "kool-workspace-*.yml") + if err != nil { + return nil + } + + var content strings.Builder + content.WriteString("services:\n") + for _, name := range selected { + fmt.Fprintf(&content, " %q:\n", name) + content.WriteString(" ports: !reset []\n") + content.WriteString(" container_name: !reset null\n") + content.WriteString(" labels:\n") + fmt.Fprintf(&content, " %q: %q\n", WorkspaceSourceLabel, envStorage.Get("KOOL_WORKSPACE_SOURCE_PROJECT")) + } + + if _, err = override.WriteString(content.String()); err != nil { + _ = override.Close() + _ = os.Remove(override.Name()) + return nil + } + if err = override.Close(); err != nil { + _ = os.Remove(override.Name()) + return nil + } + workspaceOverrideFile = override.Name() + + files = append(files, workspaceOverrideFile) + separator := envStorage.Get("COMPOSE_PATH_SEPARATOR") + if separator == "" { + separator = string(os.PathListSeparator) + } + envStorage.Set("COMPOSE_FILE", strings.Join(files, separator)) + envStorage.Set("KOOL_WORKSPACE_SERVICES", strings.Join(selected, ",")) + return selected +} + +// CleanupWorkspace removes the generated Compose override for this process. +func CleanupWorkspace() { + if workspaceOverrideFile != "" { + _ = os.Remove(workspaceOverrideFile) + workspaceOverrideFile = "" + } +} + +// IsolateWorkspace gives a recursive command its own generated Compose override. +func IsolateWorkspace() func() { + parentOverride := workspaceOverrideFile + workspaceOverrideFile = "" + return func() { + CleanupWorkspace() + workspaceOverrideFile = parentOverride + } +} + +func composeSourceProject(envStorage EnvStorage, workDir string) string { + project := "" + for _, file := range composeFiles(envStorage, workDir) { + content, err := os.ReadFile(file) + if err != nil { + continue + } + + config := workspaceComposeConfig{} + if yaml.Unmarshal(content, &config) == nil && config.Name != "" { + resolved, err := template.Substitute(config.Name, os.LookupEnv) + if err == nil { + project = composeProjectName(resolved) + } + } + } + return project +} + +func composeFiles(envStorage EnvStorage, workDir string) []string { + if configured := envStorage.Get("COMPOSE_FILE"); configured != "" { + separator := envStorage.Get("COMPOSE_PATH_SEPARATOR") + if separator == "" { + separator = string(os.PathListSeparator) + } + + files := strings.Split(configured, separator) + for index, file := range files { + if !filepath.IsAbs(file) { + files[index] = filepath.Join(workDir, file) + } + } + return files + } + + for _, name := range []string{"compose.yaml", "compose.yml", "docker-compose.yml", "docker-compose.yaml"} { + file := filepath.Join(workDir, name) + if _, err := os.Stat(file); err == nil { + return []string{file} + } + } + return nil +} diff --git a/core/environment/workspace_test.go b/core/environment/workspace_test.go new file mode 100644 index 00000000..a094bc03 --- /dev/null +++ b/core/environment/workspace_test.go @@ -0,0 +1,94 @@ +package environment + +import ( + "os" + "path/filepath" + "strings" + "testing" +) + +func TestWorkspaceIdentityDistinguishesDuplicateBasenames(t *testing.T) { + root := t.TempDir() + first := filepath.Join(root, "one", "task") + second := filepath.Join(root, "two", "task") + for _, workspace := range []string{first, second} { + if err := os.MkdirAll(workspace, 0755); err != nil { + t.Fatal(err) + } + } + if workspaceIdentity(first) == workspaceIdentity(second) { + t.Fatal("expected duplicate workspace basenames to have distinct identities") + } +} + +func TestWorkspaceIdentityPreservesHashForLongBasenames(t *testing.T) { + root := t.TempDir() + name := strings.Repeat("a", 70) + first := filepath.Join(root, "one", name) + second := filepath.Join(root, "two", name) + for _, workspace := range []string{first, second} { + if err := os.MkdirAll(workspace, 0755); err != nil { + t.Fatal(err) + } + } + firstIdentity := workspaceIdentity(first) + secondIdentity := workspaceIdentity(second) + if len(firstIdentity) > 63 || len(secondIdentity) > 63 { + t.Fatalf("expected identities to fit a DNS label: %q %q", firstIdentity, secondIdentity) + } + if firstIdentity == secondIdentity { + t.Fatal("expected long duplicate basenames to preserve distinct hash suffixes") + } +} + +func TestWorkspaceHostDoesNotExposeIdentityHash(t *testing.T) { + workspace := filepath.Join(t.TempDir(), "quiet-yarrow") + if err := os.MkdirAll(workspace, 0755); err != nil { + t.Fatal(err) + } + env := NewFakeEnvStorage() + initWorkspaceContext(env, workspace, filepath.Dir(workspace), workspace, "worktree", true, false) + + if got := env.Get("KOOL_WORKSPACE_NAME"); got != "quiet-yarrow" { + t.Fatalf("expected clean workspace host name, got %q", got) + } + if project := env.Get("KOOL_WORKSPACE_PROJECT"); !strings.Contains(project, workspaceIdentity(workspace)) { + t.Fatalf("expected internal project identity to remain unique, got %q", project) + } +} + +func TestIsolateWorkspacePreservesParentOverride(t *testing.T) { + parent, err := os.CreateTemp("", "kool-workspace-parent-*.yml") + if err != nil { + t.Fatal(err) + } + if err = parent.Close(); err != nil { + t.Fatal(err) + } + t.Cleanup(func() { + _ = os.Remove(parent.Name()) + CleanupWorkspace() + }) + workspaceOverrideFile = parent.Name() + + restore := IsolateWorkspace() + child, err := os.CreateTemp("", "kool-workspace-child-*.yml") + if err != nil { + t.Fatal(err) + } + if err = child.Close(); err != nil { + t.Fatal(err) + } + workspaceOverrideFile = child.Name() + restore() + + if _, err = os.Stat(parent.Name()); err != nil { + t.Fatalf("expected parent override to remain: %v", err) + } + if _, err = os.Stat(child.Name()); !os.IsNotExist(err) { + t.Fatalf("expected child override to be removed, got %v", err) + } + if workspaceOverrideFile != parent.Name() { + t.Fatalf("expected parent override to be restored, got %q", workspaceOverrideFile) + } +} diff --git a/core/parser/discover.go b/core/parser/discover.go new file mode 100644 index 00000000..dff0ca44 --- /dev/null +++ b/core/parser/discover.go @@ -0,0 +1,37 @@ +package parser + +import ( + "os" + "path/filepath" +) + +// koolYamlFileNames lists the accepted kool config file names, in lookup order. +var koolYamlFileNames = []string{"kool.yml", "kool.yaml"} + +// FindKoolYaml returns the path of the kool config file within the given +// directory, preferring kool.yml over kool.yaml. It returns ErrKoolYmlNotFound +// when neither file exists. +func FindKoolYaml(dir string) (file string, err error) { + for _, name := range koolYamlFileNames { + candidate := filepath.Join(dir, name) + if _, err = os.Stat(candidate); err == nil { + return candidate, nil + } + } + + return "", ErrKoolYmlNotFound +} + +// LoadKoolYaml finds and decodes the kool config file within the given +// directory. It returns ErrKoolYmlNotFound when no config file exists; any +// other error means a file was found but could not be decoded, so callers can +// tell a malformed config apart from a missing one. +func LoadKoolYaml(dir string) (parsed *KoolYaml, err error) { + var file string + + if file, err = FindKoolYaml(dir); err != nil { + return + } + + return ParseKoolYaml(file) +} diff --git a/core/parser/discover_test.go b/core/parser/discover_test.go new file mode 100644 index 00000000..d627ffd6 --- /dev/null +++ b/core/parser/discover_test.go @@ -0,0 +1,88 @@ +package parser + +import ( + "errors" + "os" + "path/filepath" + "testing" +) + +func TestFindKoolYamlPrefersYml(t *testing.T) { + dir := t.TempDir() + for _, name := range []string{"kool.yml", "kool.yaml"} { + if err := os.WriteFile(filepath.Join(dir, name), []byte("scripts: {}\n"), 0644); err != nil { + t.Fatal(err) + } + } + + file, err := FindKoolYaml(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if expected := filepath.Join(dir, "kool.yml"); file != expected { + t.Errorf("expected %q, got %q", expected, file) + } +} + +func TestFindKoolYamlFallsBackToYaml(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "kool.yaml"), []byte("scripts: {}\n"), 0644); err != nil { + t.Fatal(err) + } + + file, err := FindKoolYaml(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if expected := filepath.Join(dir, "kool.yaml"); file != expected { + t.Errorf("expected %q, got %q", expected, file) + } +} + +func TestFindKoolYamlNotFound(t *testing.T) { + if _, err := FindKoolYaml(t.TempDir()); !errors.Is(err, ErrKoolYmlNotFound) { + t.Errorf("expected ErrKoolYmlNotFound, got %v", err) + } +} + +func TestLoadKoolYaml(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "kool.yml"), []byte("workspaces:\n - app\n"), 0644); err != nil { + t.Fatal(err) + } + + parsed, err := LoadKoolYaml(dir) + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + if len(parsed.Workspaces) != 1 || parsed.Workspaces[0] != "app" { + t.Errorf("expected workspaces [app], got %v", parsed.Workspaces) + } +} + +func TestLoadKoolYamlMissing(t *testing.T) { + if _, err := LoadKoolYaml(t.TempDir()); !errors.Is(err, ErrKoolYmlNotFound) { + t.Errorf("expected ErrKoolYmlNotFound, got %v", err) + } +} + +// A malformed config must be distinguishable from a missing one, so callers +// can choose to report it rather than silently behaving as if unconfigured. +func TestLoadKoolYamlMalformed(t *testing.T) { + dir := t.TempDir() + if err := os.WriteFile(filepath.Join(dir, "kool.yml"), []byte("scripts: [oops\n"), 0644); err != nil { + t.Fatal(err) + } + + _, err := LoadKoolYaml(dir) + if err == nil { + t.Fatal("expected a decoding error, got none") + } + + if errors.Is(err, ErrKoolYmlNotFound) { + t.Error("malformed config reported as not found") + } +} diff --git a/core/parser/parser.go b/core/parser/parser.go index c35662f7..9e5d8aec 100644 --- a/core/parser/parser.go +++ b/core/parser/parser.go @@ -2,8 +2,6 @@ package parser import ( "errors" - "os" - "path" "sort" "strings" @@ -37,25 +35,16 @@ func (p *DefaultParser) AddLookupPath(rootPath string) (err error) { p.lookedUp = make(map[string]bool) } - ymlPath := path.Join(rootPath, "kool.yml") - yamlPath := path.Join(rootPath, "kool.yaml") - - if _, err = os.Stat(ymlPath); err == nil { - koolFile = ymlPath - } else if _, err = os.Stat(yamlPath); err == nil { - koolFile = yamlPath + if koolFile, err = FindKoolYaml(rootPath); err != nil { + return } - if koolFile == "" { - err = ErrKoolYmlNotFound - } else { - if !p.lookedUp[koolFile] { - p.targetFiles = append(p.targetFiles, koolFile) - } - - p.lookedUp[koolFile] = true + if !p.lookedUp[koolFile] { + p.targetFiles = append(p.targetFiles, koolFile) } + p.lookedUp[koolFile] = true + return } diff --git a/core/parser/yml.go b/core/parser/yml.go index 69d2cb27..15d5915e 100644 --- a/core/parser/yml.go +++ b/core/parser/yml.go @@ -21,6 +21,22 @@ type yamlMarshalFnType func(interface{}) ([]byte, error) type KoolYaml struct { Scripts map[string]interface{} `yaml:"scripts"` ScriptDetails map[string]ScriptDetail `yaml:"-"` + Proxy *ProxyConfig `yaml:"proxy,omitempty"` + Workspaces []string `yaml:"workspaces,omitempty"` +} + +// ProxyConfig describes routes managed by Kool's local proxy. +type ProxyConfig struct { + Domain string `yaml:"domain"` + HTTPS bool `yaml:"https,omitempty"` + Network string `yaml:"network,omitempty"` + Routes map[string]ProxyRouteConfig `yaml:"routes"` +} + +// ProxyRouteConfig describes host and port mappings for a Compose service. +type ProxyRouteConfig struct { + Ports []string `yaml:"ports"` + Hosts []string `yaml:"hosts"` } // ScriptDetail describes a kool.yml script with context @@ -107,6 +123,8 @@ func (y *KoolYaml) Parse(filePath string) (err error) { y.Scripts = parsed.Scripts y.ScriptDetails = parsed.ScriptDetails + y.Proxy = parsed.Proxy + y.Workspaces = parsed.Workspaces return } diff --git a/core/parser/yml_test.go b/core/parser/yml_test.go index c1cb149e..584e8a23 100644 --- a/core/parser/yml_test.go +++ b/core/parser/yml_test.go @@ -17,6 +17,33 @@ const KoolYmlOK = `scripts: - line 2 ` +func TestParseWorkspaceAndProxyConfig(t *testing.T) { + file := path.Join(t.TempDir(), "kool.yml") + content := `workspaces: [app, node] +proxy: + domain: app.localhost + https: true + routes: + app: + ports: ["80:8080"] + hosts: ["@", "*"] +` + if err := os.WriteFile(file, []byte(content), os.ModePerm); err != nil { + t.Fatal(err) + } + + parsed, err := ParseKoolYaml(file) + if err != nil { + t.Fatal(err) + } + if len(parsed.Workspaces) != 2 || parsed.Workspaces[0] != "app" || parsed.Workspaces[1] != "node" { + t.Errorf("unexpected workspaces: %v", parsed.Workspaces) + } + if parsed.Proxy == nil || parsed.Proxy.Domain != "app.localhost" || !parsed.Proxy.HTTPS || parsed.Proxy.Routes["app"].Ports[0] != "80:8080" || parsed.Proxy.Routes["app"].Hosts[1] != "*" { + t.Errorf("unexpected proxy config: %#v", parsed.Proxy) + } +} + func TestParseKoolYaml(t *testing.T) { var ( err error diff --git a/docs/15-Snippets/Workspaces.md b/docs/15-Snippets/Workspaces.md new file mode 100644 index 00000000..2e3b3757 --- /dev/null +++ b/docs/15-Snippets/Workspaces.md @@ -0,0 +1,188 @@ +# Workspaces and local proxy + +Kool detects workspaces managed by [Rift](https://github.com/anomalyco/rift) or linked [Git worktrees](https://git-scm.com/docs/git-worktree). Each workspace runs selected application services in a separate Compose project while databases, caches, and other infrastructure remain in the original project. + +## Configuration + +Configure workspace services and proxy routes once in `kool.yml`: + +```yaml +workspaces: + - app + - node + +proxy: + domain: "${APP_DOMAIN:-app.localhost}" + routes: + app: + ports: ["80:80"] + hosts: ["@", "*"] + node: + ports: + - "3001:3001" + +scripts: + # Existing scripts remain here. +``` + +Both features are opt-in. Without a non-empty `workspaces` list, Rift and Git worktree detection is skipped and existing Kool commands retain their legacy behavior. Without `proxy`, Kool does not inspect or manage the global Caddy proxy. + +A port uses `listen:target`: + +```text +"80:8080" + host port : service container port +``` + +Every port is mapped across every host in the service route. Host values use these conventions: + +```text +@ proxy.domain +* wildcard before proxy.domain +api api.proxy.domain +admin.example.test complete hostname +``` + +`hosts` is optional and defaults to `["@", "*"]`, routing both the base domain and wildcard subdomains. Set it explicitly to `["@"]` when wildcard routing is not wanted. + +Relative hosts use the workspace base automatically. For example, `api` becomes `api..workspace.`. A complete hostname remains unchanged. + +`proxy.network` optionally changes the shared Docker network. It defaults to `KOOL_GLOBAL_NETWORK`, which defaults to `kool_global`. + +Kool injects `KOOL_PROXY_DOMAIN` and the context-specific `KOOL_PROXY_HOST` before Compose starts services. They can be passed into application or Vite configuration: + +```yaml +environment: + VITE_PUBLIC_ORIGIN: "http://${KOOL_PROXY_HOST}:3001" +``` + +No Kool or proxy labels are required in `docker-compose.yml`. + +## Compose services + +Workspace services should mount the current project normally and join the external shared network: + +```yaml +services: + app: + image: example/app + volumes: + - .:/app:delegated + networks: + - kool_global + + node: + image: node:22 + working_dir: /app + volumes: + - .:/app:delegated + networks: + - kool_global + + database: + image: mysql:8 + networks: + - kool_global + +networks: + kool_global: + external: true + name: "${KOOL_GLOBAL_NETWORK:-kool_global}" +``` + +Kool removes published ports from proxied services because Caddy owns the listen ports. It also removes fixed `container_name` values from workspace services. Shared infrastructure must already be running from the original workspace. + +## Managed proxy + +When proxy routes are configured, Kool manages a global `kool-proxy` container using the pinned `caddy:2.10-alpine` image. + +Caddy: + +- Joins the shared Docker network. +- Publishes each configured listen port. +- Exposes its Admin API only at `127.0.0.1:2019`. +- Does not mount the Docker socket. +- Handles HTTP streaming and WebSocket upgrades, including Vite HMR. + +Kool gives proxied services deterministic network aliases such as: + +```text +example-app +example-node +example-workspace-task-a-app +example-workspace-task-a-node +``` + +`kool start` registers routes through Caddy's Admin API. `kool stop` removes only the current source or workspace project's routes. + +For `domain: app.localhost`, routes are generated as follows: + +```text +Source: +app.localhost -> app +*.app.localhost -> app + +Workspace task-a: +task-a.workspace.app.localhost -> task-a app +*.task-a.workspace.app.localhost -> task-a app +``` + +Routes using different listen ports can use the same hostname. For example, `80:80` can target the application while `3001:3001` targets Vite. + +If an existing `kool-proxy` container does not publish a newly configured listen port, remove it and run `kool start` again: + +```bash +docker rm -f kool-proxy +kool start +``` + +## Local HTTPS + +Enable Caddy's internal certificate authority with `proxy.https`: + +```yaml +proxy: + domain: "${APP_DOMAIN:-app.localhost}" + https: true + routes: + app: + ports: ["443:80"] + node: + ports: ["3001:3001"] + hosts: ["@"] +``` + +All configured listeners use TLS when `https` is enabled. Caddy issues and renews certificates for source and workspace hosts, including wildcards. The `kool_proxy` Docker volume stores autosaved routes under `config/` and CA and certificate data under `data/`. + +After starting the proxy, trust its local root CA on macOS or Linux: + +```bash +kool start +kool proxy trust +``` + +Trust installation requires administrator privileges. On WSL, `kool proxy trust` updates the Linux trust store; browsers running on Windows also require importing the root certificate into the Windows certificate store. + +HTTPS does not currently add an automatic HTTP-to-HTTPS redirect. Configure only the TLS listen ports clients should use. + +## Commands + +From the original workspace, `kool start` starts the project normally. From a Rift or linked Git worktree it starts only services listed under `workspaces`, using an internal project name such as `example-workspace-task-a`. + +Existing commands remain transparent: + +```bash +kool start +kool stop +kool restart +kool exec app php artisan test +kool run artisan test +kool logs app +kool status +``` + +From the original project, `kool status` shows the main project and all active workspaces. Inside a workspace, it shows the main project and the current workspace. + +`kool stop` inside a workspace removes only that workspace's containers and proxy routes. The original project and other workspaces remain running. Running `kool stop` without service arguments from the original project stops every active workspace before stopping the main project. + +This mode requires Docker Compose support for the `!reset` merge tag. diff --git a/kool.yml b/kool.yml index 19a4d091..0f5c784b 100644 --- a/kool.yml +++ b/kool.yml @@ -13,7 +13,10 @@ scripts: compile: - kool run fmt - kool run go build -buildvcs=false -o kool-cli - install: mv ./kool-cli /usr/local/bin/kool + install: + - sudo install -m 0755 ./kool-cli /usr/local/bin/.kool.new + - sudo mv /usr/local/bin/.kool.new /usr/local/bin/kool + - rm -f $HOME/.local/bin/kool fmt: kool run go:linux fmt ./... lint: kool docker --volume=kool_gopath:/go golangci/golangci-lint:v2.11.4 golangci-lint run -v # Reachability-aware vulnerability scan of our own code and Go dependencies. diff --git a/main.go b/main.go index f6a7ba40..6b53a112 100644 --- a/main.go +++ b/main.go @@ -11,7 +11,6 @@ import ( func main() { log.SetFlags(log.Ldate | log.Ltime | log.Lmicroseconds) - environment.InitEnvironmentVariables(environment.NewEnvStorage()) if err := commands.Execute(); err != nil { shell.NewShell().Println(err) @@ -19,8 +18,10 @@ func main() { if ex, ok := err.(shell.ErrExitable); ok { code = ex.Code } + environment.CleanupWorkspace() os.Exit(code) } + environment.CleanupWorkspace() os.Exit(0) } diff --git a/services/proxy/manager.go b/services/proxy/manager.go new file mode 100644 index 00000000..99e85b80 --- /dev/null +++ b/services/proxy/manager.go @@ -0,0 +1,1571 @@ +package proxy + +import ( + "bytes" + "crypto/sha256" + "encoding/json" + "errors" + "fmt" + "io" + "kool-dev/kool/core/builder" + "kool-dev/kool/core/environment" + "kool-dev/kool/core/parser" + "kool-dev/kool/core/shell" + "net/http" + "os" + "path/filepath" + "reflect" + "runtime" + "sort" + "strconv" + "strings" + "sync/atomic" + "time" + + "github.com/compose-spec/compose-go/template" + "golang.org/x/sys/unix" +) + +const ( + caddyContainer = "kool-proxy" + caddyImage = "caddy:2.10-alpine" + caddyVolume = "kool_proxy" + caddyAdminURL = "http://127.0.0.1:2019" + caddyAdminHost = "kool-proxy-admin" + caddyAdminNet = "kool_proxy_admin" + caddyStartCmd = "exec caddy run --config /etc/caddy/caddy.json" + defaultNetwork = "kool_global" +) + +// Manager controls Kool's local proxy routes. +type Manager interface { + Prepare([]string) (func(bool) error, error) + Remove([]string) error + RemoveProject(string, []string) error + Trust() error +} + +type route struct { + Service string + Listen int + Target int + Hosts []string + HTTPS bool +} + +type config struct { + Domain string + HTTPS bool + Network string + Routes []route +} + +// DefaultManager manages a global Caddy container through its Admin API. +type DefaultManager struct { + shell shell.Shell + env environment.EnvStorage + http *http.Client + adminURL string + projectOverride string + generation string +} + +var prepareGeneration atomic.Uint64 + +// NewManager creates a proxy manager for the current project. +func NewManager(sh shell.Shell, env environment.EnvStorage) Manager { + return &DefaultManager{shell: sh, env: env, http: &http.Client{Timeout: 2 * time.Second}, adminURL: caddyAdminURL} +} + +// Prepare ensures Caddy is running, applies network aliases, and registers routes. +func (m *DefaultManager) Prepare(services []string) (finish func(bool) error, err error) { + finish = func(bool) error { return nil } + var cfg *config + if cfg, err = m.loadConfig(); err != nil || cfg == nil { + return + } + m.generation = fmt.Sprintf("%d-%d", os.Getpid(), prepareGeneration.Add(1)) + + requested := make(map[string]bool, len(services)) + for _, service := range services { + requested[service] = true + } + var routes []route + for _, route := range cfg.Routes { + if len(requested) == 0 || requested[route.Service] { + routes = append(routes, route) + } + } + previousComposeFiles := m.env.Get("COMPOSE_FILE") + var override string + if len(routes) > 0 { + if override, err = m.createAliasOverride(cfg.Network, routes); err != nil { + return + } + } + cleanupOverride := func() { + m.env.Set("COMPOSE_FILE", previousComposeFiles) + if override != "" { + _ = os.Remove(override) + } + } + + var unlock func() + if unlock, err = m.acquireConfigLock(); err != nil { + cleanupOverride() + return func(bool) error { return nil }, err + } + lockHeld := true + defer func() { + if lockHeld { + unlock() + } + }() + + if len(routes) == 0 { + running, inspectErr := m.shell.Exec(builder.NewCommand("docker", "inspect", "--format", "{{.State.Running}}", caddyContainer)) + if inspectErr != nil { + unlock() + lockHeld = false + cleanupOverride() + return func(bool) error { return nil }, nil + } + if running != "true" { + if err = m.shell.Interactive(builder.NewCommand("docker", "start"), caddyContainer); err != nil { + cleanupOverride() + return func(bool) error { return nil }, err + } + } + if err = m.waitForCaddy(); err != nil { + cleanupOverride() + return func(bool) error { return nil }, err + } + } else if err = m.ensureCaddy(cfg.Network, cfg.Routes); err != nil { + cleanupOverride() + return func(bool) error { return nil }, err + } + + var snapshot []byte + var snapshotExists bool + if snapshot, snapshotExists, err = m.snapshotApps(); err != nil { + cleanupOverride() + return func(bool) error { return nil }, err + } + rollback := func() error { return m.restoreApps(snapshot, snapshotExists) } + if err = m.registerTLSUnlocked(cfg.Routes); err != nil { + err = errors.Join(err, rollback()) + cleanupOverride() + return func(bool) error { return nil }, err + } + for _, route := range routes { + if _, err = m.registerUnlocked(route); err != nil { + err = errors.Join(err, rollback()) + cleanupOverride() + return func(bool) error { return nil }, err + } + } + if err = m.reconcileRoutesUnlocked(cfg.Routes); err != nil { + err = errors.Join(err, rollback()) + cleanupOverride() + return func(bool) error { return nil }, err + } + committed, committedExists, snapshotErr := m.snapshotApps() + if snapshotErr != nil { + snapshotErr = errors.Join(snapshotErr, rollback()) + cleanupOverride() + return func(bool) error { return nil }, snapshotErr + } + unlock() + lockHeld = false + finished := false + finish = func(success bool) error { + if finished { + return nil + } + finished = true + cleanupOverride() + if success { + return m.withConfigLock(func() error { + finalSnapshot, finalExists, snapshotErr := m.snapshotApps() + if snapshotErr != nil { + return snapshotErr + } + apply := func() error { + if err := m.registerTLSUnlocked(cfg.Routes); err != nil { + return err + } + for _, route := range routes { + if _, err := m.registerUnlocked(route); err != nil { + return err + } + } + return m.reconcileRoutesUnlocked(cfg.Routes) + } + if applyErr := apply(); applyErr != nil { + return errors.Join(applyErr, m.restoreApps(finalSnapshot, finalExists)) + } + return nil + }) + } + return m.withConfigLock(func() error { + return m.rollbackGeneration(committed, committedExists, snapshot, snapshotExists) + }) + } + return +} + +// Remove deletes all proxy routes belonging to the current project. +func (m *DefaultManager) Remove(services []string) error { + _, err := m.loadConfig() + if err != nil { + return err + } + if _, err = m.shell.Exec(builder.NewCommand("docker", "inspect", caddyContainer)); err != nil { + return nil + } + return m.withConfigLock(func() error { + if err = m.removeProjectRoutesUnlocked(services); err != nil { + return err + } + if len(services) == 0 { + if err = m.deleteRoute(m.tlsID()); err != nil { + return err + } + } + return m.stopIfUnusedUnlocked() + }) +} + +// RemoveProject removes proxy routes for an explicit Compose project. +func (m *DefaultManager) RemoveProject(project string, services []string) error { + manager := *m + manager.projectOverride = project + if _, err := manager.shell.Exec(builder.NewCommand("docker", "inspect", caddyContainer)); err != nil { + return nil + } + return manager.withConfigLock(func() error { + if err := manager.removeProjectRoutesUnlocked(services); err != nil { + return err + } + if len(services) == 0 { + if err := manager.deleteRoute(manager.tlsID()); err != nil { + return err + } + } + return manager.stopIfUnusedUnlocked() + }) +} + +// Trust installs Caddy's local root CA in the host trust store. +func (m *DefaultManager) Trust() error { + if _, err := m.shell.Exec(builder.NewCommand("docker", "inspect", caddyContainer)); err != nil { + return errors.New("kool proxy is not running; run kool start first") + } + certificate, err := os.CreateTemp("", "kool-proxy-root-*.crt") + if err != nil { + return err + } + path := certificate.Name() + if err = certificate.Close(); err != nil { + return err + } + defer func() { _ = os.Remove(path) }() + if err = m.shell.Interactive(builder.NewCommand("docker", "cp"), caddyContainer+":/var/lib/caddy/data/caddy/pki/authorities/local/root.crt", path); err != nil { + return fmt.Errorf("could not export proxy root certificate; start an HTTPS proxy route first: %w", err) + } + + switch runtime.GOOS { + case "darwin": + return m.shell.Interactive(builder.NewCommand("sudo", "security", "add-trusted-cert", "-d", "-r", "trustRoot", "-k", "/Library/Keychains/System.keychain"), path) + case "linux": + destination := "/usr/local/share/ca-certificates/kool-proxy.crt" + if err = m.shell.Interactive(builder.NewCommand("sudo", "cp"), path, destination); err != nil { + return err + } + return m.shell.Interactive(builder.NewCommand("sudo", "update-ca-certificates")) + default: + return fmt.Errorf("automatic proxy trust is not supported on %s", runtime.GOOS) + } +} + +func (m *DefaultManager) loadConfig() (*config, error) { + parsed, err := parser.LoadKoolYaml(m.env.Get("PWD")) + if errors.Is(err, parser.ErrKoolYmlNotFound) { + return nil, nil + } + if err != nil || parsed.Proxy == nil { + return nil, err + } + domain, err := template.Substitute(parsed.Proxy.Domain, os.LookupEnv) + domain = strings.TrimSuffix(strings.TrimSpace(domain), ".") + if err != nil || domain == "" { + if err == nil { + err = errors.New("proxy.domain cannot be empty") + } + return nil, err + } + if strings.Contains(domain, "://") { + return nil, errors.New("proxy.domain must be a hostname without a URL scheme") + } + network := parsed.Proxy.Network + if network == "" { + network = m.env.Get("KOOL_GLOBAL_NETWORK") + } + if network == "" { + network = defaultNetwork + } + if network, err = template.Substitute(network, os.LookupEnv); err != nil { + return nil, err + } + if network == caddyAdminNet { + return nil, fmt.Errorf("proxy.network %q is reserved for proxy administration", network) + } + + cfg := &config{Domain: domain, HTTPS: parsed.Proxy.HTTPS, Network: network} + for service, routeConfig := range parsed.Proxy.Routes { + if len(routeConfig.Ports) == 0 { + return nil, fmt.Errorf("proxy route %s requires at least one port", service) + } + if len(routeConfig.Hosts) == 0 { + routeConfig.Hosts = []string{"@", "*"} + } + var hosts []string + for _, host := range routeConfig.Hosts { + var resolvedHost string + if resolvedHost, err = template.Substitute(host, os.LookupEnv); err != nil { + return nil, err + } + resolvedHost = strings.TrimSpace(resolvedHost) + if resolvedHost == "" { + return nil, fmt.Errorf("proxy route %s contains an empty host", service) + } + hosts = append(hosts, resolvedHost) + } + listenPorts := make(map[int]bool) + for _, portMapping := range routeConfig.Ports { + var resolved string + if resolved, err = template.Substitute(portMapping, os.LookupEnv); err != nil { + return nil, err + } + parts := strings.SplitN(resolved, ":", 2) + if len(parts) != 2 { + return nil, fmt.Errorf("proxy route %s must use listen:target", service) + } + listen, listenErr := strconv.Atoi(strings.TrimSpace(parts[0])) + target, targetErr := strconv.Atoi(strings.TrimSpace(parts[1])) + if listenErr != nil || targetErr != nil || listen < 1 || listen > 65535 || target < 1 || target > 65535 { + return nil, fmt.Errorf("proxy route %s has invalid ports %q", service, resolved) + } + if listen == 2019 { + return nil, fmt.Errorf("proxy route %s cannot listen on reserved admin port 2019", service) + } + if listenPorts[listen] { + return nil, fmt.Errorf("proxy route %s maps listen port %d more than once", service, listen) + } + listenPorts[listen] = true + cfg.Routes = append(cfg.Routes, route{Service: service, Listen: listen, Target: target, Hosts: hosts, HTTPS: cfg.HTTPS}) + } + } + sort.Slice(cfg.Routes, func(i, j int) bool { + if cfg.Routes[i].Service == cfg.Routes[j].Service { + if cfg.Routes[i].Listen == cfg.Routes[j].Listen { + return cfg.Routes[i].Target < cfg.Routes[j].Target + } + return cfg.Routes[i].Listen < cfg.Routes[j].Listen + } + return cfg.Routes[i].Service < cfg.Routes[j].Service + }) + bindings := make(map[string]string) + for _, route := range cfg.Routes { + for _, host := range m.routeHosts(route) { + host = strings.ToLower(strings.TrimSuffix(host, ".")) + binding := fmt.Sprintf("%d:%s", route.Listen, host) + if service := bindings[binding]; service != "" && service != route.Service { + return nil, fmt.Errorf("proxy routes %s and %s both use %s", service, route.Service, binding) + } + bindings[binding] = route.Service + } + } + return cfg, nil +} + +func (m *DefaultManager) createAliasOverride(network string, routes []route) (string, error) { + file, err := os.CreateTemp("", "kool-proxy-*.yml") + if err != nil { + return "", err + } + + var content strings.Builder + content.WriteString("services:\n") + written := make(map[string]bool) + for _, route := range routes { + if written[route.Service] { + continue + } + fmt.Fprintf(&content, " %q:\n ports: !reset []\n networks:\n %q:\n aliases:\n - %q\n", route.Service, defaultNetwork, m.alias(route.Service)) + written[route.Service] = true + } + fmt.Fprintf(&content, "networks:\n %q:\n name: %q\n external: true\n", defaultNetwork, network) + if _, err = file.WriteString(content.String()); err == nil { + err = file.Close() + } else { + _ = file.Close() + } + if err != nil { + _ = os.Remove(file.Name()) + return "", err + } + + separator := m.env.Get("COMPOSE_PATH_SEPARATOR") + if separator == "" { + separator = string(os.PathListSeparator) + } + composeFiles := m.env.Get("COMPOSE_FILE") + if composeFiles == "" { + for _, name := range []string{"compose.yaml", "compose.yml", "docker-compose.yml", "docker-compose.yaml"} { + if _, statErr := os.Stat(filepath.Join(m.env.Get("PWD"), name)); statErr == nil { + composeFiles = name + break + } + } + } + m.env.Set("COMPOSE_FILE", composeFiles+separator+file.Name()) + return file.Name(), nil +} + +func (m *DefaultManager) ensureCaddy(network string, routes []route) (err error) { + inspect := builder.NewCommand("docker", "inspect", "--format", "{{.State.Running}}", caddyContainer) + running, err := m.shell.Exec(inspect) + bindings := make(map[int][]string) + var preservedApps []byte + var preservedAppsExist bool + var preservedNetworks []string + var originalBindings map[int][]string + expanding := false + defer func() { + if err == nil || !expanding { + return + } + _ = m.shell.Interactive(builder.NewCommand("docker", "rm", "--force"), caddyContainer) + rollbackErr := m.createCaddy(originalBindings) + if rollbackErr == nil { + rollbackErr = m.connectCaddyNetworks(preservedNetworks) + } + if rollbackErr == nil { + rollbackErr = m.waitForCaddy() + } + if rollbackErr == nil && preservedAppsExist { + rollbackErr = m.restoreApps(preservedApps, true) + } + err = errors.Join(err, rollbackErr) + }() + if err == nil { + if running != "true" { + if err = m.shell.Interactive(builder.NewCommand("docker", "start"), caddyContainer); err != nil { + return err + } + if err = m.waitForCaddy(); err != nil { + return err + } + running = "true" + } + compatibility, inspectErr := m.shell.Exec(builder.NewCommand("docker", "inspect", "--format", "{{json .NetworkSettings.Networks}}|{{json .Config.Entrypoint}}|{{json .Config.Cmd}}", caddyContainer)) + if inspectErr != nil { + return inspectErr + } + expectedCommand := `["/bin/sh"]|["-c","` + strings.ReplaceAll(caddyStartCmd, `"`, `\"`) + `"]` + if !strings.Contains(compatibility, `"`+caddyAdminNet+`"`) || !strings.HasSuffix(compatibility, expectedCommand) { + if preservedApps, preservedAppsExist, err = m.snapshotApps(); err != nil { + return err + } + bindingsJSON, inspectErr := m.shell.Exec(builder.NewCommand("docker", "inspect", "--format", "{{json .HostConfig.PortBindings}}", caddyContainer)) + if inspectErr != nil { + return inspectErr + } + bindings = parsePortBindings(bindingsJSON) + originalBindings = copyPortBindings(bindings) + delete(bindings, 2019) + networks, inspectErr := m.shell.Exec(builder.NewCommand("docker", "inspect", "--format", "{{json .NetworkSettings.Networks}}", caddyContainer)) + if inspectErr != nil { + return inspectErr + } + preservedNetworks = parseDockerObjectKeys(networks) + if err = m.shell.Interactive(builder.NewCommand("docker", "rm", "--force"), caddyContainer); err != nil { + return err + } + expanding = true + err = errors.New("legacy proxy container removed") + } else { + for _, route := range routes { + if _, portErr := m.shell.Exec(builder.NewCommand("docker", "port", caddyContainer), fmt.Sprintf("%d/tcp", route.Listen)); portErr == nil { + if len(bindings[route.Listen]) == 0 { + bindings[route.Listen] = []string{fmt.Sprintf("%d:%d", route.Listen, route.Listen)} + } + continue + } + if preservedApps, preservedAppsExist, err = m.snapshotApps(); err != nil { + return err + } + bindingsJSON, inspectErr := m.shell.Exec(builder.NewCommand("docker", "inspect", "--format", "{{json .HostConfig.PortBindings}}", caddyContainer)) + if inspectErr != nil { + return inspectErr + } + bindings = parsePortBindings(bindingsJSON) + originalBindings = copyPortBindings(bindings) + delete(bindings, 2019) + networks, inspectErr := m.shell.Exec(builder.NewCommand("docker", "inspect", "--format", "{{json .NetworkSettings.Networks}}", caddyContainer)) + if inspectErr != nil { + return inspectErr + } + preservedNetworks = parseDockerObjectKeys(networks) + if err = m.shell.Interactive(builder.NewCommand("docker", "rm", "--force"), caddyContainer); err != nil { + return err + } + expanding = true + err = errors.New("proxy listener expansion required") + break + } + } + } + if err != nil { + for _, route := range routes { + if len(bindings[route.Listen]) == 0 { + bindings[route.Listen] = []string{fmt.Sprintf("%d:%d", route.Listen, route.Listen)} + } + } + if err = m.createCaddy(bindings); err != nil { + return err + } + } else if running != "true" { + if err = m.shell.Interactive(builder.NewCommand("docker", "start"), caddyContainer); err != nil { + return err + } + } + networks, err := m.shell.Exec(builder.NewCommand("docker", "inspect", "--format", "{{json .NetworkSettings.Networks}}", caddyContainer)) + if err != nil { + return err + } + if !strings.Contains(networks, `"`+network+`"`) { + if err = m.shell.Interactive(builder.NewCommand("docker", "network", "connect", network), caddyContainer); err != nil { + return err + } + } + for _, previousNetwork := range preservedNetworks { + if previousNetwork == caddyAdminNet || previousNetwork == network { + continue + } + if err = m.shell.Interactive(builder.NewCommand("docker", "network", "connect", previousNetwork), caddyContainer); err != nil { + return err + } + } + + for _, route := range routes { + if _, err = m.shell.Exec(builder.NewCommand("docker", "port", caddyContainer), fmt.Sprintf("%d/tcp", route.Listen)); err != nil { + return fmt.Errorf("kool proxy does not publish port %d; remove %s and retry", route.Listen, caddyContainer) + } + } + if err = m.waitForCaddy(); err != nil { + return err + } + if preservedAppsExist { + if err = m.restoreApps(preservedApps, true); err != nil { + return err + } + } + expanding = false + return nil +} + +func (m *DefaultManager) createCaddy(bindings map[int][]string) error { + configPath, err := m.ensureBaseConfig() + if err != nil { + return err + } + if _, err = m.shell.Exec(builder.NewCommand("docker", "network", "inspect", caddyAdminNet)); err != nil { + if err = m.shell.Interactive(builder.NewCommand("docker", "network", "create"), caddyAdminNet); err != nil { + return err + } + } + args := []string{"run", "-d", "--name", caddyContainer, "--restart", "unless-stopped", "--network", caddyAdminNet, "--network-alias", caddyAdminHost, "-p", "127.0.0.1:2019:2019"} + var sortedPorts []int + for port := range bindings { + if port == 2019 { + continue + } + sortedPorts = append(sortedPorts, port) + } + sort.Ints(sortedPorts) + for _, port := range sortedPorts { + for _, binding := range bindings[port] { + args = append(args, "-p", binding) + } + } + args = append(args, "-v", configPath+":/etc/caddy/caddy.json:ro", "-v", caddyVolume+":/var/lib/caddy", "-e", "XDG_CONFIG_HOME=/var/lib/caddy/config", "-e", "XDG_DATA_HOME=/var/lib/caddy/data", "--entrypoint", "/bin/sh", caddyImage, "-c", caddyStartCmd) + return m.shell.Interactive(builder.NewCommand("docker"), args...) +} + +func (m *DefaultManager) connectCaddyNetworks(networks []string) error { + for _, network := range networks { + if network == caddyAdminNet { + continue + } + if err := m.shell.Interactive(builder.NewCommand("docker", "network", "connect", network), caddyContainer); err != nil { + return err + } + } + return nil +} + +func (m *DefaultManager) waitForCaddy() error { + for attempt := 0; attempt < 30; attempt++ { + response, err := m.request(http.MethodGet, m.adminURL+"/config/", nil) + if err == nil { + _ = response.Body.Close() + if response.StatusCode < 500 { + return nil + } + } + time.Sleep(100 * time.Millisecond) + } + return errors.New("kool proxy did not become ready") +} + +func copyPortBindings(bindings map[int][]string) map[int][]string { + copy := make(map[int][]string, len(bindings)) + for port, specs := range bindings { + copy[port] = append([]string(nil), specs...) + } + return copy +} + +func parsePortBindings(raw string) map[int][]string { + result := make(map[int][]string) + var bindings map[string][]struct { + HostIP string `json:"HostIp"` + HostPort string `json:"HostPort"` + } + if json.Unmarshal([]byte(raw), &bindings) != nil { + return result + } + for key, published := range bindings { + port, err := strconv.Atoi(strings.TrimSuffix(key, "/tcp")) + if err != nil { + continue + } + for _, binding := range published { + spec := binding.HostPort + ":" + strconv.Itoa(port) + if binding.HostIP != "" { + hostIP := binding.HostIP + if strings.Contains(hostIP, ":") && !strings.HasPrefix(hostIP, "[") { + hostIP = "[" + hostIP + "]" + } + spec = hostIP + ":" + spec + } + result[port] = append(result[port], spec) + } + } + return result +} + +func parseDockerObjectKeys(raw string) []string { + var object map[string]interface{} + if json.Unmarshal([]byte(raw), &object) != nil { + return nil + } + keys := make([]string, 0, len(object)) + for key := range object { + keys = append(keys, key) + } + sort.Strings(keys) + return keys +} + +func (m *DefaultManager) ensureBaseConfig() (string, error) { + home, err := os.UserHomeDir() + if err != nil { + return "", err + } + directory := filepath.Join(home, ".kool", "proxy") + if err = os.MkdirAll(directory, 0755); err != nil { + return "", err + } + path := filepath.Join(directory, "caddy.json") + content := []byte(`{"admin":{"listen":"0.0.0.0:2019"},"apps":{"http":{"servers":{}},"tls":{"automation":{"policies":[]}}}}`) + if err = os.WriteFile(path, content, 0644); err != nil { + return "", err + } + return path, nil +} + +func (m *DefaultManager) register(route route) error { + _, err := m.registerWithResult(route) + return err +} + +func (m *DefaultManager) registerWithResult(route route) (created bool, err error) { + err = m.withConfigLock(func() error { + created, err = m.registerUnlocked(route) + return err + }) + return +} + +func (m *DefaultManager) registerUnlocked(route route) (bool, error) { + serverID := "kool-" + strconv.Itoa(route.Listen) + server := caddyServerConfig(route) + serverBody, _ := json.Marshal(server) + serverURL := m.adminURL + "/config/apps/http/servers/" + serverID + response, err := m.request(http.MethodGet, serverURL, nil) + if err != nil { + return false, err + } + status := response.StatusCode + serverMissing := status == http.StatusNotFound + if status >= 400 && !serverMissing { + defer func() { _ = response.Body.Close() }() + return false, responseError(response) + } + if serverMissing { + if err = response.Body.Close(); err != nil { + return false, err + } + } else { + var responseBody []byte + if responseBody, err = io.ReadAll(response.Body); err != nil { + _ = response.Body.Close() + return false, err + } + if err = response.Body.Close(); err != nil { + return false, err + } + serverMissing = bytes.Equal(bytes.TrimSpace(responseBody), []byte("null")) + if !serverMissing { + if err = m.reconfigureListenerMode(serverURL, responseBody, route); err != nil { + return false, err + } + } + } + if serverMissing { + response, err = m.request(http.MethodPost, serverURL, serverBody) + if err != nil { + return false, err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 { + return false, responseError(response) + } + } + if route.HTTPS { + if err = m.setTLSPolicy(serverURL); err != nil { + return false, err + } + } else if err = m.disableAutomaticHTTPS(serverURL); err != nil { + return false, err + } + + routeConfig := map[string]interface{}{ + "@id": m.routeID(route), + "match": []interface{}{map[string]interface{}{"host": m.routeHosts(route)}}, + "handle": []interface{}{map[string]interface{}{ + "@id": m.projectMarker() + "-" + m.routeID(route) + "-" + m.generation, + "handler": "reverse_proxy", + "upstreams": []interface{}{map[string]string{"dial": m.alias(route.Service) + ":" + strconv.Itoa(route.Target)}}, + }}, + "terminal": true, + } + return m.upsertRouteUnlocked(serverURL, routeConfig) +} + +func (m *DefaultManager) upsertRouteUnlocked(serverURL string, routeConfig map[string]interface{}) (bool, error) { + routesURL := serverURL + "/routes" + response, err := m.request(http.MethodGet, routesURL, nil) + if err != nil { + return false, err + } + if response.StatusCode >= 400 && response.StatusCode != http.StatusNotFound { + defer func() { _ = response.Body.Close() }() + return false, responseError(response) + } + + var routes []json.RawMessage + if response.StatusCode != http.StatusNotFound { + if err = json.NewDecoder(response.Body).Decode(&routes); err != nil && err != io.EOF { + _ = response.Body.Close() + return false, err + } + } + if err = response.Body.Close(); err != nil { + return false, err + } + + routeBody, _ := json.Marshal(routeConfig) + routeID, _ := routeConfig["@id"].(string) + replaced := false + for index, existing := range routes { + var metadata caddyRouteMetadata + if json.Unmarshal(existing, &metadata) == nil && metadata.ID == routeID { + routes[index] = routeBody + replaced = true + break + } + } + if !replaced { + var incoming caddyRouteMetadata + _ = json.Unmarshal(routeBody, &incoming) + for _, existing := range routes { + var metadata caddyRouteMetadata + if json.Unmarshal(existing, &metadata) != nil || m.routeBelongsToProject(metadata) { + continue + } + if hostSetsOverlap(routeMetadataHosts(metadata), routeMetadataHosts(incoming)) { + return false, fmt.Errorf("proxy listener host conflict with route %s", metadata.ID) + } + } + routes = append(routes, routeBody) + } + routes = orderCaddyRoutes(routes) + body, _ := json.Marshal(routes) + response, err = m.request(http.MethodPatch, routesURL, body) + if err != nil { + return false, err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 { + return false, responseError(response) + } + return !replaced, nil +} + +func hostSetsOverlap(first, second []string) bool { + normalized := make(map[string]bool, len(first)) + for _, host := range first { + normalized[strings.ToLower(strings.TrimSuffix(host, "."))] = true + } + for _, host := range second { + host = strings.ToLower(strings.TrimSuffix(host, ".")) + if normalized[host] { + return true + } + } + return false +} + +type caddyRouteMetadata struct { + ID string `json:"@id"` + Match []struct { + Hosts []string `json:"host"` + } `json:"match"` + Handle []struct { + ID string `json:"@id"` + Upstreams []struct { + Dial string `json:"dial"` + } `json:"upstreams"` + } `json:"handle"` +} + +func (m *DefaultManager) reconcileRoutes(desiredRoutes []route) error { + return m.withConfigLock(func() error { + return m.reconcileRoutesUnlocked(desiredRoutes) + }) +} + +func (m *DefaultManager) reconcileRoutesUnlocked(desiredRoutes []route) error { + desired := make(map[string]bool, len(desiredRoutes)) + for _, route := range desiredRoutes { + desired[m.routeID(route)] = true + } + return m.filterProjectRoutesUnlocked(func(route caddyRouteMetadata) bool { + return !desired[route.ID] + }) +} + +func (m *DefaultManager) removeProjectRoutes(services []string) error { + return m.withConfigLock(func() error { + return m.removeProjectRoutesUnlocked(services) + }) +} + +func (m *DefaultManager) removeProjectRoutesUnlocked(services []string) error { + aliases := make(map[string]bool, len(services)) + for _, service := range services { + aliases[m.alias(service)] = true + } + return m.filterProjectRoutesUnlocked(func(route caddyRouteMetadata) bool { + if len(aliases) == 0 { + return true + } + for _, alias := range routeMetadataAliases(route) { + if aliases[alias] { + return true + } + } + return false + }) +} + +func (m *DefaultManager) filterProjectRoutesUnlocked(shouldRemove func(caddyRouteMetadata) bool) error { + serversURL := m.adminURL + "/config/apps/http/servers" + response, err := m.request(http.MethodGet, serversURL, nil) + if err != nil { + return err + } + if response.StatusCode == http.StatusNotFound { + _ = response.Body.Close() + return nil + } + if response.StatusCode >= 400 { + defer func() { _ = response.Body.Close() }() + return responseError(response) + } + var servers map[string]struct { + Routes []json.RawMessage `json:"routes"` + } + if err = json.NewDecoder(response.Body).Decode(&servers); err != nil { + _ = response.Body.Close() + return err + } + if err = response.Body.Close(); err != nil { + return err + } + + for serverID, server := range servers { + routes := server.Routes[:0] + changed := false + for _, rawRoute := range server.Routes { + var metadata caddyRouteMetadata + if json.Unmarshal(rawRoute, &metadata) == nil && m.routeBelongsToProject(metadata) && shouldRemove(metadata) { + changed = true + continue + } + routes = append(routes, rawRoute) + } + if !changed { + continue + } + body, _ := json.Marshal(orderCaddyRoutes(routes)) + routesURL := serversURL + "/" + serverID + "/routes" + response, err = m.request(http.MethodPatch, routesURL, body) + if err != nil { + return err + } + if response.StatusCode >= 400 { + defer func() { _ = response.Body.Close() }() + return responseError(response) + } + if err = response.Body.Close(); err != nil { + return err + } + } + return nil +} + +func (m *DefaultManager) stopIfUnusedUnlocked() error { + serversURL := m.adminURL + "/config/apps/http/servers" + response, err := m.request(http.MethodGet, serversURL, nil) + if err != nil { + return err + } + if response.StatusCode == http.StatusNotFound { + _ = response.Body.Close() + return m.shell.Interactive(builder.NewCommand("docker", "stop"), caddyContainer) + } + if response.StatusCode >= 400 { + defer func() { _ = response.Body.Close() }() + return responseError(response) + } + var servers map[string]struct { + Routes []json.RawMessage `json:"routes"` + } + if err = json.NewDecoder(response.Body).Decode(&servers); err != nil { + _ = response.Body.Close() + return err + } + if err = response.Body.Close(); err != nil { + return err + } + for _, server := range servers { + for _, rawRoute := range server.Routes { + var metadata caddyRouteMetadata + if json.Unmarshal(rawRoute, &metadata) != nil { + continue + } + for _, handle := range metadata.Handle { + if strings.HasPrefix(handle.ID, "kool-project-") { + return nil + } + } + } + } + return m.shell.Interactive(builder.NewCommand("docker", "stop"), caddyContainer) +} + +func (m *DefaultManager) withConfigLock(action func() error) error { + unlock, err := m.acquireConfigLock() + if err != nil { + return err + } + defer unlock() + return action() +} + +func (m *DefaultManager) acquireConfigLock() (func(), error) { + home, err := os.UserHomeDir() + if err != nil { + return nil, err + } + directory := filepath.Join(home, ".kool", "proxy") + if err = os.MkdirAll(directory, 0755); err != nil { + return nil, err + } + lock, err := os.OpenFile(filepath.Join(directory, "config.lock"), os.O_CREATE|os.O_RDWR, 0600) + if err != nil { + return nil, err + } + if err = unix.Flock(int(lock.Fd()), unix.LOCK_EX); err != nil { + _ = lock.Close() + return nil, err + } + return func() { + _ = unix.Flock(int(lock.Fd()), unix.LOCK_UN) + _ = lock.Close() + }, nil +} + +func (m *DefaultManager) routeBelongsToProject(route caddyRouteMetadata) bool { + prefix := m.projectMarker() + "-" + for _, handle := range route.Handle { + if strings.HasPrefix(handle.ID, prefix) { + return true + } + } + return false +} + +func (m *DefaultManager) projectMarker() string { + digest := sha256.Sum256([]byte(m.project())) + return fmt.Sprintf("kool-project-%x", digest[:8]) +} + +func routeMetadataAliases(route caddyRouteMetadata) []string { + var aliases []string + for _, handle := range route.Handle { + for _, upstream := range handle.Upstreams { + alias := upstream.Dial + if separator := strings.LastIndex(alias, ":"); separator >= 0 { + alias = alias[:separator] + } + aliases = append(aliases, alias) + } + } + return aliases +} + +func orderCaddyRoutes(routes []json.RawMessage) []json.RawMessage { + ordered := append([]json.RawMessage(nil), routes...) + for broad := 0; broad < len(ordered); broad++ { + var broadMetadata caddyRouteMetadata + if json.Unmarshal(ordered[broad], &broadMetadata) != nil || !strings.HasPrefix(broadMetadata.ID, "kool-") { + continue + } + for specific := broad + 1; specific < len(ordered); specific++ { + var specificMetadata caddyRouteMetadata + if json.Unmarshal(ordered[specific], &specificMetadata) != nil || !strings.HasPrefix(specificMetadata.ID, "kool-") { + continue + } + if routeHostsAreMoreSpecific(routeMetadataHosts(specificMetadata), routeMetadataHosts(broadMetadata)) { + route := ordered[specific] + copy(ordered[broad+1:specific+1], ordered[broad:specific]) + ordered[broad] = route + broad-- + break + } + } + } + return ordered +} + +func routeMetadataHosts(route caddyRouteMetadata) []string { + var hosts []string + for _, match := range route.Match { + hosts = append(hosts, match.Hosts...) + } + return hosts +} + +func routeHostsAreMoreSpecific(hosts, other []string) bool { + return len(hosts) > 0 && hostPatternsCover(other, hosts) && !hostPatternsCover(hosts, other) +} + +func hostPatternsCover(patterns, hosts []string) bool { + for _, host := range hosts { + covered := false + for _, pattern := range patterns { + if hostPatternCovers(pattern, host) { + covered = true + break + } + } + if !covered { + return false + } + } + return true +} + +func hostPatternCovers(pattern, host string) bool { + pattern = strings.ToLower(strings.TrimSuffix(pattern, ".")) + host = strings.ToLower(strings.TrimSuffix(host, ".")) + if !strings.HasPrefix(pattern, "*.") { + return pattern == host + } + pattern = strings.TrimPrefix(pattern, "*.") + host = strings.TrimPrefix(host, "*.") + return host == pattern || strings.HasSuffix(host, "."+pattern) +} + +func caddyServerConfig(route route) map[string]interface{} { + server := map[string]interface{}{"listen": []string{":" + strconv.Itoa(route.Listen)}, "routes": []interface{}{}} + if route.HTTPS { + server["tls_connection_policies"] = []interface{}{map[string]interface{}{}} + } else { + server["automatic_https"] = map[string]interface{}{"disable": true} + } + return server +} + +func (m *DefaultManager) reconfigureListenerMode(serverURL string, server []byte, route route) error { + var config struct { + AutomaticHTTPS struct { + Disable bool `json:"disable"` + } `json:"automatic_https"` + TLSConnectionPolicies json.RawMessage `json:"tls_connection_policies"` + Routes []json.RawMessage `json:"routes"` + } + if err := json.Unmarshal(server, &config); err != nil { + return fmt.Errorf("could not inspect proxy listener %d: %w", route.Listen, err) + } + hasTLS := len(config.TLSConnectionPolicies) > 0 && string(config.TLSConnectionPolicies) != "null" + modeDiffers := route.HTTPS && config.AutomaticHTTPS.Disable || !route.HTTPS && hasTLS + if !modeDiffers { + return nil + } + for _, rawRoute := range config.Routes { + var metadata caddyRouteMetadata + if json.Unmarshal(rawRoute, &metadata) != nil || !m.routeBelongsToProject(metadata) { + return fmt.Errorf("proxy listener %d cannot change protocol while another project uses it", route.Listen) + } + } + + var replacement map[string]interface{} + if err := json.Unmarshal(server, &replacement); err != nil { + return fmt.Errorf("could not inspect proxy listener %d: %w", route.Listen, err) + } + delete(replacement, "automatic_https") + delete(replacement, "tls_connection_policies") + for key, value := range caddyServerConfig(route) { + if key != "routes" { + replacement[key] = value + } + } + body, _ := json.Marshal(replacement) + response, err := m.request(http.MethodPatch, serverURL, body) + if err != nil { + return err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 { + return responseError(response) + } + return nil +} + +func (m *DefaultManager) disableAutomaticHTTPS(serverURL string) error { + body, _ := json.Marshal(map[string]interface{}{"disable": true}) + response, err := m.request(http.MethodPost, serverURL+"/automatic_https", body) + if err != nil { + return err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 { + return responseError(response) + } + return nil +} + +func (m *DefaultManager) setTLSPolicy(serverURL string) error { + url := serverURL + "/tls_connection_policies" + response, err := m.request(http.MethodGet, url, nil) + if err != nil { + return err + } + method := http.MethodPatch + if response.StatusCode == http.StatusNotFound { + method = http.MethodPost + } else if response.StatusCode >= 400 { + defer func() { _ = response.Body.Close() }() + return responseError(response) + } + _ = response.Body.Close() + + policies, _ := json.Marshal([]interface{}{map[string]interface{}{}}) + response, err = m.request(method, url, policies) + if err != nil { + return err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 { + return responseError(response) + } + return nil +} + +func (m *DefaultManager) registerTLS(routes []route) error { + return m.withConfigLock(func() error { + return m.registerTLSUnlocked(routes) + }) +} + +func (m *DefaultManager) registerTLSUnlocked(routes []route) error { + seen := make(map[string]bool) + var subjects []string + for _, route := range routes { + if !route.HTTPS { + continue + } + for _, host := range m.routeHosts(route) { + if !seen[host] { + subjects = append(subjects, host) + seen[host] = true + } + } + } + if len(subjects) == 0 { + return m.deleteRoute(m.tlsID()) + } + sort.Strings(subjects) + + tlsURL := m.adminURL + "/config/apps/tls" + response, err := m.request(http.MethodGet, tlsURL, nil) + if err != nil { + return err + } + status := response.StatusCode + _ = response.Body.Close() + if status == http.StatusNotFound { + body := []byte(`{"automation":{"policies":[]}}`) + if response, err = m.request(http.MethodPost, tlsURL, body); err != nil { + return err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 { + return responseError(response) + } + } else if status >= 400 { + return fmt.Errorf("caddy Admin API returned %s", response.Status) + } + + if err = m.deleteRoute(m.tlsID()); err != nil { + return err + } + policy := map[string]interface{}{ + "@id": m.tlsID(), + "subjects": subjects, + "issuers": []interface{}{map[string]string{"module": "internal"}}, + } + body, _ := json.Marshal(policy) + response, err = m.request(http.MethodPost, tlsURL+"/automation/policies", body) + if err != nil { + return err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 { + return responseError(response) + } + return nil +} + +func (m *DefaultManager) deleteRoute(id string) error { + response, err := m.request(http.MethodDelete, m.adminURL+"/id/"+id, nil) + if err != nil { + return err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 && response.StatusCode != http.StatusNotFound { + return responseError(response) + } + return nil +} + +func (m *DefaultManager) snapshotApps() ([]byte, bool, error) { + response, err := m.request(http.MethodGet, m.adminURL+"/config/apps", nil) + if err != nil { + return nil, false, err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode == http.StatusNotFound { + return nil, false, nil + } + if response.StatusCode >= 400 { + return nil, false, responseError(response) + } + body, err := io.ReadAll(response.Body) + return body, true, err +} + +func (m *DefaultManager) restoreApps(snapshot []byte, existed bool) error { + method := http.MethodPatch + if !existed { + method = http.MethodDelete + } + response, err := m.request(method, m.adminURL+"/config/apps", snapshot) + if err != nil { + return err + } + defer func() { _ = response.Body.Close() }() + if response.StatusCode >= 400 && response.StatusCode != http.StatusNotFound { + return responseError(response) + } + return nil +} + +func (m *DefaultManager) restoreAppsIfUnchanged(committed []byte, committedExists bool, snapshot []byte, snapshotExists bool) error { + current, currentExists, err := m.snapshotApps() + if err != nil || currentExists != committedExists || !bytes.Equal(bytes.TrimSpace(current), bytes.TrimSpace(committed)) { + return err + } + return m.restoreApps(snapshot, snapshotExists) +} + +func (m *DefaultManager) rollbackGeneration(committed []byte, committedExists bool, snapshot []byte, snapshotExists bool) error { + current, currentExists, err := m.snapshotApps() + if err != nil { + return err + } + if currentExists == committedExists && bytes.Equal(bytes.TrimSpace(current), bytes.TrimSpace(committed)) { + return m.restoreApps(snapshot, snapshotExists) + } + merged, mergeErr := m.mergeGenerationRollback(current, committed, snapshot) + if mergeErr != nil { + return mergeErr + } + return m.restoreApps(merged, true) +} + +func (m *DefaultManager) mergeGenerationRollback(current, committed, snapshot []byte) ([]byte, error) { + states := make(map[string]map[string]interface{}, 3) + for name, data := range map[string][]byte{"current": current, "committed": committed, "snapshot": snapshot} { + state := make(map[string]interface{}) + states[name] = state + if len(bytes.TrimSpace(data)) > 0 { + if err := json.Unmarshal(data, &state); err != nil { + return nil, fmt.Errorf("could not merge %s proxy state: %w", name, err) + } + } + } + currentServers := caddyServers(states["current"]) + committedServers := caddyServers(states["committed"]) + snapshotServers := caddyServers(states["snapshot"]) + previousRoutes := make(map[string]interface{}) + committedRoutes := make(map[string]interface{}) + for _, server := range snapshotServers { + for _, route := range caddyRoutes(server) { + previousRoutes[caddyObjectID(route)] = route + } + } + for _, server := range committedServers { + for _, route := range caddyRoutes(server) { + committedRoutes[caddyObjectID(route)] = route + } + } + failedOwnsProjectRoute := false + newerProjectRouteExists := false + for serverID, server := range currentServers { + serverConfig, _ := server.(map[string]interface{}) + var routes []interface{} + generationOwned := false + nonGenerationRouteSurvives := false + for _, route := range caddyRoutes(server) { + owned := routeHasGeneration(route, m.generation) + if owned { + generationOwned = true + failedOwnsProjectRoute = true + if previous := previousRoutes[caddyObjectID(route)]; previous != nil { + routes = append(routes, previous) + } + continue + } + nonGenerationRouteSurvives = true + if routeBelongsToMarker(route, m.projectMarker()) { + committedRoute, existed := committedRoutes[caddyObjectID(route)] + newerProjectRouteExists = newerProjectRouteExists || !existed || !reflect.DeepEqual(route, committedRoute) + } + routes = append(routes, route) + } + serverConfig["routes"] = routes + if generationOwned && !nonGenerationRouteSurvives { + if snapshotServer, existed := snapshotServers[serverID]; existed { + currentServers[serverID] = snapshotServer + } else { + delete(currentServers, serverID) + } + } + } + currentPolicies, currentAutomation := caddyPolicies(states["current"]) + committedPolicies, _ := caddyPolicies(states["committed"]) + snapshotPolicies, _ := caddyPolicies(states["snapshot"]) + currentPolicy, currentPresent := objectByID(currentPolicies, m.tlsID()) + committedPolicy, committedPresent := objectByID(committedPolicies, m.tlsID()) + if failedOwnsProjectRoute && !newerProjectRouteExists && currentPresent == committedPresent && reflect.DeepEqual(currentPolicy, committedPolicy) { + replacement, replacementPresent := objectByID(snapshotPolicies, m.tlsID()) + currentAutomation["policies"] = replaceObjectByID(currentPolicies, m.tlsID(), replacement, replacementPresent) + } + return json.Marshal(states["current"]) +} + +func nestedMap(parent map[string]interface{}, keys ...string) map[string]interface{} { + current := parent + for _, key := range keys { + next, _ := current[key].(map[string]interface{}) + if next == nil { + next = make(map[string]interface{}) + current[key] = next + } + current = next + } + return current +} + +func caddyServers(apps map[string]interface{}) map[string]interface{} { + return nestedMap(apps, "http", "servers") +} + +func caddyRoutes(server interface{}) []interface{} { + config, _ := server.(map[string]interface{}) + routes, _ := config["routes"].([]interface{}) + return routes +} + +func caddyObjectID(object interface{}) string { + config, _ := object.(map[string]interface{}) + id, _ := config["@id"].(string) + return id +} + +func routeHasGeneration(route interface{}, generation string) bool { + config, _ := route.(map[string]interface{}) + handles, _ := config["handle"].([]interface{}) + for _, handle := range handles { + if strings.HasSuffix(caddyObjectID(handle), "-"+generation) { + return true + } + } + return false +} + +func routeBelongsToMarker(route interface{}, marker string) bool { + config, _ := route.(map[string]interface{}) + handles, _ := config["handle"].([]interface{}) + for _, handle := range handles { + if strings.HasPrefix(caddyObjectID(handle), marker+"-") { + return true + } + } + return false +} + +func caddyPolicies(apps map[string]interface{}) ([]interface{}, map[string]interface{}) { + automation := nestedMap(apps, "tls", "automation") + policies, _ := automation["policies"].([]interface{}) + return policies, automation +} + +func objectByID(objects []interface{}, id string) (interface{}, bool) { + for _, object := range objects { + if caddyObjectID(object) == id { + return object, true + } + } + return nil, false +} + +func replaceObjectByID(objects []interface{}, id string, replacement interface{}, replacementPresent bool) []interface{} { + result := make([]interface{}, 0, len(objects)+1) + replaced := false + for _, object := range objects { + if caddyObjectID(object) == id { + if replacementPresent { + result = append(result, replacement) + } + replaced = true + continue + } + result = append(result, object) + } + if !replaced && replacementPresent { + result = append(result, replacement) + } + return result +} + +func (m *DefaultManager) request(method, url string, body []byte) (*http.Response, error) { + request, err := http.NewRequest(method, url, bytes.NewReader(body)) + if err != nil { + return nil, err + } + request.Header.Set("Content-Type", "application/json") + return m.http.Do(request) +} + +func (m *DefaultManager) alias(service string) string { + return m.project() + "-" + service +} + +func (m *DefaultManager) routeID(route route) string { + return fmt.Sprintf("kool-%s-%s-%d-%d", m.project(), route.Service, route.Listen, route.Target) +} + +func (m *DefaultManager) tlsID() string { + return "kool-" + m.project() + "-tls" +} + +func (m *DefaultManager) routeHosts(route route) []string { + base := m.env.Get("KOOL_PROXY_HOST") + seen := make(map[string]bool) + var hosts []string + for _, host := range route.Hosts { + switch { + case host == "@": + host = base + case host == "*": + host = "*." + base + case !strings.Contains(host, "."): + host += "." + base + } + if !seen[host] { + hosts = append(hosts, host) + seen[host] = true + } + } + return hosts +} + +func (m *DefaultManager) project() string { + if m.projectOverride != "" { + return m.projectOverride + } + if project := m.env.Get("KOOL_WORKSPACE_PROJECT"); project != "" && m.env.IsTrue("KOOL_WORKSPACE") { + return project + } + if project := m.env.Get("KOOL_WORKSPACE_SOURCE_PROJECT"); project != "" { + return project + } + if project := m.env.Get("COMPOSE_PROJECT_NAME"); project != "" { + return project + } + return m.env.Get("KOOL_NAME") +} + +func responseError(response *http.Response) error { + body, _ := io.ReadAll(response.Body) + return fmt.Errorf("caddy Admin API returned %s: %s", response.Status, strings.TrimSpace(string(body))) +} diff --git a/services/proxy/manager_test.go b/services/proxy/manager_test.go new file mode 100644 index 00000000..fe676412 --- /dev/null +++ b/services/proxy/manager_test.go @@ -0,0 +1,1113 @@ +package proxy + +import ( + "encoding/json" + "io" + "kool-dev/kool/core/environment" + "kool-dev/kool/core/shell" + "net/http" + "net/http/httptest" + "os" + "path/filepath" + "strings" + "testing" +) + +func TestLoadConfig(t *testing.T) { + workDir := t.TempDir() + content := `proxy: + domain: app.localhost + network: shared + routes: + app: + ports: ["80:8080", "8080:8081"] + hosts: ["@", "*"] + node: + ports: ["3001:3001"] +` + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte(content), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + config, err := manager.loadConfig() + if err != nil { + t.Fatal(err) + } + if config.Domain != "app.localhost" || config.Network != "shared" || len(config.Routes) != 3 { + t.Fatalf("unexpected proxy config: %#v", config) + } + if config.Routes[0].Service != "app" || config.Routes[0].Listen != 80 || config.Routes[0].Target != 8080 || len(config.Routes[0].Hosts) != 2 { + t.Errorf("unexpected app route: %#v", config.Routes[0]) + } + if config.Routes[2].Service != "node" || len(config.Routes[2].Hosts) != 2 || config.Routes[2].Hosts[0] != "@" || config.Routes[2].Hosts[1] != "*" { + t.Errorf("expected node route to default to base and wildcard hosts, got %#v", config.Routes[2]) + } +} + +func TestRouteHosts(t *testing.T) { + env := environment.NewFakeEnvStorage() + env.Set("KOOL_PROXY_HOST", "task-a.workspace.app.localhost") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + hosts := manager.routeHosts(route{Hosts: []string{"@", "*", "api", "admin.example.test"}}) + expected := []string{ + "task-a.workspace.app.localhost", + "*.task-a.workspace.app.localhost", + "api.task-a.workspace.app.localhost", + "admin.example.test", + } + if strings.Join(hosts, ",") != strings.Join(expected, ",") { + t.Errorf("expected hosts %v, got %v", expected, hosts) + } +} + +func TestLoadConfigRejectsInvalidRoute(t *testing.T) { + workDir := t.TempDir() + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte("proxy:\n domain: app.localhost\n routes:\n app:\n ports: [invalid]\n hosts: ['@']\n"), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + if _, err := manager.loadConfig(); err == nil || !strings.Contains(err.Error(), "listen:target") { + t.Fatalf("expected route format error, got %v", err) + } +} + +func TestLoadConfigRejectsEquivalentResolvedHosts(t *testing.T) { + workDir := t.TempDir() + content := "proxy:\n domain: app.localhost\n routes:\n app:\n ports: ['80:80']\n hosts: ['@']\n node:\n ports: ['80:3001']\n hosts: ['app.localhost']\n" + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte(content), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + env.Set("KOOL_PROXY_HOST", "app.localhost") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + if _, err := manager.loadConfig(); err == nil || !strings.Contains(err.Error(), "both use 80:app.localhost") { + t.Fatalf("expected equivalent host conflict, got %v", err) + } +} + +func TestLoadConfigRejectsReservedAdminNetwork(t *testing.T) { + workDir := t.TempDir() + content := "proxy:\n domain: app.localhost\n network: kool_proxy_admin\n routes:\n app:\n ports: ['80:80']\n" + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte(content), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + if _, err := manager.loadConfig(); err == nil || !strings.Contains(err.Error(), "reserved for proxy administration") { + t.Fatalf("expected reserved network error, got %v", err) + } +} + +func TestLoadConfigRejectsReservedAdminPort(t *testing.T) { + workDir := t.TempDir() + content := "proxy:\n domain: app.localhost\n routes:\n app:\n ports: ['2019:8080']\n" + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte(content), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + if _, err := manager.loadConfig(); err == nil || !strings.Contains(err.Error(), "reserved admin port 2019") { + t.Fatalf("expected reserved port error, got %v", err) + } +} + +func TestCreateAliasOverride(t *testing.T) { + workDir := t.TempDir() + if err := os.WriteFile(filepath.Join(workDir, "compose.yml"), []byte("services: {}\n"), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + env.Set("KOOL_NAME", "example") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + file, err := manager.createAliasOverride("kool_global", []route{{Service: "app", Listen: 80, Target: 80}, {Service: "app", Listen: 8080, Target: 3000}}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.Remove(file) }) + content, err := os.ReadFile(file) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(content), "example-app") || !strings.Contains(string(content), "kool_global") { + t.Errorf("unexpected alias override:\n%s", content) + } + if strings.Count(string(content), `"app":`) != 1 { + t.Errorf("expected one service override for multiple routes, got:\n%s", content) + } + if !strings.Contains(env.Get("COMPOSE_FILE"), file) { + t.Errorf("expected COMPOSE_FILE to include %s, got %s", file, env.Get("COMPOSE_FILE")) + } +} + +func TestCreateAliasOverrideMapsCustomExternalNetwork(t *testing.T) { + workDir := t.TempDir() + if err := os.WriteFile(filepath.Join(workDir, "compose.yml"), []byte("services: {}\nnetworks:\n kool_global:\n external: true\n"), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + env.Set("KOOL_NAME", "example") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + file, err := manager.createAliasOverride("my_network", []route{{Service: "app", Listen: 80, Target: 80}}) + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = os.Remove(file) }) + content, err := os.ReadFile(file) + if err != nil { + t.Fatal(err) + } + for _, expected := range []string{`"kool_global":`, `name: "my_network"`, "external: true"} { + if !strings.Contains(string(content), expected) { + t.Errorf("expected override to contain %q, got:\n%s", expected, content) + } + } + if strings.Contains(string(content), `"my_network":`) { + t.Errorf("custom network name must not be used as an undeclared Compose key:\n%s", content) + } +} + +func TestPrepareRestoresComposeFile(t *testing.T) { + workDir := t.TempDir() + if err := os.WriteFile(filepath.Join(workDir, "kool.yml"), []byte("proxy:\n domain: app.localhost\n routes:\n app:\n ports: ['80:80']\n"), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + env.Set("COMPOSE_FILE", "compose.yml:compose.dev.yml") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + + cleanup, err := manager.Prepare([]string{"app"}) + if err == nil { + t.Fatal("expected fake shell to fail while ensuring Caddy") + } + _ = cleanup(false) + if value := env.Get("COMPOSE_FILE"); value != "compose.yml:compose.dev.yml" { + t.Fatalf("expected COMPOSE_FILE to be restored, got %q", value) + } +} + +func TestPrepareReconcilesWhenProxyRoutesBecomeEmpty(t *testing.T) { + state, server := newCaddyRouteState(t) + workDir := t.TempDir() + configPath := filepath.Join(workDir, "kool.yml") + if err := os.WriteFile(configPath, []byte("proxy:\n domain: app.localhost\n routes: {}\n"), 0644); err != nil { + t.Fatal(err) + } + env := environment.NewFakeEnvStorage() + env.Set("PWD", workDir) + env.Set("COMPOSE_PROJECT_NAME", "example") + env.Set("KOOL_PROXY_HOST", "app.localhost") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = server.URL + mustRegister(t, manager, route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@", "*"}}) + + cleanup, err := manager.Prepare(nil) + if err != nil { + t.Fatal(err) + } + if err = cleanup(true); err != nil { + t.Fatal(err) + } + state.requireRoutes(t, "kool-80", nil) +} + +func TestRegisterWorkspaceRoute(t *testing.T) { + serverCreated := false + var events []string + var registered map[string]interface{} + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/routes"): + _, _ = response.Write([]byte(`[]`)) + case request.Method == http.MethodGet: + if !serverCreated { + _, _ = response.Write([]byte("null")) + return + } + _, _ = response.Write([]byte(`{"listen":[":80"],"routes":[]}`)) + case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/kool-80"): + serverCreated = true + events = append(events, "server") + return + case request.Method == http.MethodDelete: + response.WriteHeader(http.StatusNotFound) + return + case request.Method == http.MethodPatch && strings.HasSuffix(request.URL.Path, "/routes"): + if !serverCreated { + http.Error(response, "server was not created", http.StatusInternalServerError) + return + } + events = append(events, "route") + body, _ := io.ReadAll(request.Body) + var routes []map[string]interface{} + _ = json.Unmarshal(body, &routes) + registered = routes[0] + return + } + response.WriteHeader(http.StatusOK) + })) + defer server.Close() + + env := environment.NewFakeEnvStorage() + env.Set("KOOL_WORKSPACE", "true") + env.Set("KOOL_WORKSPACE_NAME", "task-a") + env.Set("KOOL_WORKSPACE_PROJECT", "example-workspace-task-a") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = server.URL + + env.Set("KOOL_PROXY_HOST", "task-a.workspace.app.localhost") + if err := manager.register(route{Service: "app", Listen: 80, Target: 8080, Hosts: []string{"@", "*"}}); err != nil { + t.Fatal(err) + } + if strings.Join(events, ",") != "server,route" { + t.Errorf("expected server creation before route registration, got %v", events) + } + encoded, _ := json.Marshal(registered) + result := string(encoded) + for _, expected := range []string{"task-a.workspace.app.localhost", "*.task-a.workspace.app.localhost", "example-workspace-task-a-app:8080"} { + if !strings.Contains(result, expected) { + t.Errorf("expected registered route to contain %q, got %s", expected, result) + } + } +} + +func TestCaddyHTTPServerDisablesAutomaticHTTPS(t *testing.T) { + server := caddyServerConfig(route{Listen: 3001}) + automaticHTTPS, ok := server["automatic_https"].(map[string]interface{}) + if !ok || automaticHTTPS["disable"] != true { + t.Errorf("expected HTTP server to disable automatic HTTPS, got %#v", server) + } + if _, exists := server["tls_connection_policies"]; exists { + t.Errorf("did not expect HTTP server to configure TLS, got %#v", server) + } +} + +func TestCaddyHTTPSServerPreservesTLSConfiguration(t *testing.T) { + server := caddyServerConfig(route{Listen: 3001, HTTPS: true}) + if _, exists := server["tls_connection_policies"]; !exists { + t.Errorf("expected HTTPS server to configure TLS, got %#v", server) + } + if _, exists := server["automatic_https"]; exists { + t.Errorf("did not expect HTTPS server to disable automatic HTTPS, got %#v", server) + } +} + +func TestRegisterRouteWithExistingServer(t *testing.T) { + serverPostCalled := false + routePostCalled := false + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/routes"): + _, _ = response.Write([]byte(`[]`)) + case request.Method == http.MethodGet: + _, _ = response.Write([]byte(`{"listen":[":80"],"routes":[]}`)) + case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/kool-80"): + serverPostCalled = true + case request.Method == http.MethodPatch && strings.HasSuffix(request.URL.Path, "/routes"): + routePostCalled = true + } + })) + defer server.Close() + + env := environment.NewFakeEnvStorage() + env.Set("KOOL_NAME", "example") + env.Set("KOOL_PROXY_HOST", "app.localhost") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = server.URL + + if err := manager.register(route{Service: "app", Listen: 80, Target: 8080, Hosts: []string{"@"}}); err != nil { + t.Fatal(err) + } + if serverPostCalled { + t.Error("did not expect an existing server to be recreated") + } + if !routePostCalled { + t.Error("expected route to be registered on the existing server") + } +} + +func TestRegisterRouteReportsWhetherRouteWasCreated(t *testing.T) { + state, server := newCaddyRouteState(t) + manager := testRouteManager(server.URL, "example", "app.localhost") + proxyRoute := route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@"}} + + created, err := manager.registerWithResult(proxyRoute) + if err != nil { + t.Fatal(err) + } + if !created { + t.Error("expected first registration to report a newly created route") + } + created, err = manager.registerWithResult(proxyRoute) + if err != nil { + t.Fatal(err) + } + if created { + t.Error("expected repeated registration to report an existing route") + } + state.requireRoutes(t, "kool-80", []string{"kool-example-app-80-80"}) +} + +func TestRouteGenerationDistinguishesConcurrentPreparations(t *testing.T) { + state, server := newCaddyRouteState(t) + first := testRouteManager(server.URL, "example", "app.localhost") + second := testRouteManager(server.URL, "example", "app.localhost") + first.generation = "first" + second.generation = "second" + proxyRoute := route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@"}} + mustRegister(t, first, proxyRoute) + mustRegister(t, second, proxyRoute) + + raw := state.routes["kool-80"][0] + if !strings.Contains(string(raw), "-second") || strings.Contains(string(raw), "-first") { + t.Fatalf("expected latest preparation generation in route marker, got %s", raw) + } +} + +func TestRollbackGenerationRemovesOnlyFailedConcurrentRoutes(t *testing.T) { + state, server := newCaddyRouteState(t) + failed := testRouteManager(server.URL, "failed", "failed.localhost") + failed.generation = "failed-generation" + succeeded := testRouteManager(server.URL, "succeeded", "succeeded.localhost") + succeeded.generation = "succeeded-generation" + proxyRoute := route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@"}} + mustRegister(t, failed, proxyRoute) + mustRegister(t, succeeded, proxyRoute) + + if err := failed.rollbackGeneration([]byte(`{}`), true, []byte(`{}`), true); err != nil { + t.Fatal(err) + } + state.requireRoutes(t, "kool-80", []string{"kool-succeeded-app-80-80"}) +} + +func TestRollbackGenerationRestoresPreviousRouteAfterConcurrentUpdate(t *testing.T) { + state, server := newCaddyRouteState(t) + failed := testRouteManager(server.URL, "app", "app.localhost") + failed.generation = "previous" + proxyRoute := route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@"}} + mustRegister(t, failed, proxyRoute) + snapshot, _, err := failed.snapshotApps() + if err != nil { + t.Fatal(err) + } + + failed.generation = "failed" + mustRegister(t, failed, proxyRoute) + committed, _, err := failed.snapshotApps() + if err != nil { + t.Fatal(err) + } + other := testRouteManager(server.URL, "other", "other.localhost") + other.generation = "other" + mustRegister(t, other, proxyRoute) + + if err = failed.rollbackGeneration(committed, true, snapshot, true); err != nil { + t.Fatal(err) + } + state.requireRoutes(t, "kool-80", []string{"kool-app-app-80-80", "kool-other-app-80-80"}) + if !strings.Contains(string(state.routes["kool-80"][0]), "-previous") { + t.Fatalf("expected previous route generation to be restored, got %s", state.routes["kool-80"][0]) + } +} + +func TestMergeGenerationRollbackPreservesFieldsAndRestoresProtocolAndTLS(t *testing.T) { + manager := testRouteManager("", "app", "app.localhost") + manager.generation = "failed" + current := []byte(`{ + "http":{"servers":{"kool-80":{"listen":[":80"],"tls_connection_policies":[{}],"routes":[{"@id":"route","handle":[{"@id":"marker-failed"}]}]}}}, + "tls":{"automation":{"policies":[]}}, + "unrelated":{"keep":true} +}`) + committed := append([]byte(nil), current...) + snapshot := []byte(`{ + "http":{"servers":{"kool-80":{"listen":[":80"],"automatic_https":{"disable":true},"routes":[{"@id":"route","handle":[{"@id":"marker-previous"}]}]}}}, + "tls":{"automation":{"policies":[{"@id":"kool-app-tls","subjects":["app.localhost"]}]}}, + "unrelated":{"keep":true} +}`) + + merged, err := manager.mergeGenerationRollback(current, committed, snapshot) + if err != nil { + t.Fatal(err) + } + result := string(merged) + for _, expected := range []string{`"listen":[":80"]`, `"automatic_https":{"disable":true}`, `"marker-previous"`, `"kool-app-tls"`, `"unrelated":{"keep":true}`} { + if !strings.Contains(result, expected) { + t.Errorf("expected merged state to preserve or restore %s, got %s", expected, result) + } + } + for _, unexpected := range []string{`"tls_connection_policies"`, `"marker-failed"`} { + if strings.Contains(result, unexpected) { + t.Errorf("did not expect merged state to contain %s, got %s", unexpected, result) + } + } +} + +func TestMergeGenerationRollbackPreservesSharedListenerAndNewerTLS(t *testing.T) { + manager := testRouteManager("", "app", "app.localhost") + manager.generation = "failed" + marker := manager.projectMarker() + current := []byte(strings.ReplaceAll(`{ + "http":{"servers":{"kool-443":{"listen":[":443"],"tls_connection_policies":[{}],"routes":[ + {"@id":"app-route","handle":[{"@id":"PROJECT_MARKER-app-route-newer"}]}, + {"@id":"other-route","handle":[{"@id":"other-project-route"}]} + ]}}}, + "tls":{"automation":{"policies":[{"@id":"kool-app-tls","subjects":["new.localhost"]}]}} +}`, "PROJECT_MARKER", marker)) + committed := []byte(strings.ReplaceAll(`{ + "http":{"servers":{"kool-443":{"listen":[":443"],"tls_connection_policies":[{}],"routes":[ + {"@id":"app-route","handle":[{"@id":"PROJECT_MARKER-app-route-failed"}]} + ]}}}, + "tls":{"automation":{"policies":[{"@id":"kool-app-tls","subjects":["new.localhost"]}]}} +}`, "PROJECT_MARKER", marker)) + snapshot := []byte(`{ + "http":{"servers":{"kool-443":{"listen":[":443"],"automatic_https":{"disable":true},"routes":[]}}}, + "tls":{"automation":{"policies":[{"@id":"kool-app-tls","subjects":["old.localhost"]}]}} +}`) + + merged, err := manager.mergeGenerationRollback(current, committed, snapshot) + if err != nil { + t.Fatal(err) + } + result := string(merged) + for _, expected := range []string{`"tls_connection_policies":[{}]`, `"new.localhost"`, marker + `-app-route-newer`, `"other-project-route"`} { + if !strings.Contains(result, expected) { + t.Errorf("expected concurrent state %s to be preserved, got %s", expected, result) + } + } + if strings.Contains(result, "old.localhost") || strings.Contains(result, "automatic_https") { + t.Errorf("did not expect stale protocol or TLS state to be restored, got %s", result) + } +} + +func TestMergeGenerationRollbackRestoresTLSWithUnchangedProjectRoute(t *testing.T) { + manager := testRouteManager("", "app", "app.localhost") + manager.generation = "failed" + marker := manager.projectMarker() + unchanged := `{"@id":"existing","handle":[{"@id":"` + marker + `-existing-old"}]}` + failed := `{"@id":"failed","handle":[{"@id":"` + marker + `-failed-failed"}]}` + current := []byte(`{"http":{"servers":{"kool-443":{"routes":[` + unchanged + `,` + failed + `,{"@id":"other","handle":[{"@id":"other-new"}]}]}}},"tls":{"automation":{"policies":[{"@id":"kool-app-tls","subjects":["failed.localhost"]}]}}}`) + committed := []byte(`{"http":{"servers":{"kool-443":{"routes":[` + unchanged + `,` + failed + `]}}},"tls":{"automation":{"policies":[{"@id":"kool-app-tls","subjects":["failed.localhost"]}]}}}`) + snapshot := []byte(`{"http":{"servers":{"kool-443":{"routes":[` + unchanged + `]}}},"tls":{"automation":{"policies":[{"@id":"kool-app-tls","subjects":["previous.localhost"]}]}}}`) + + merged, err := manager.mergeGenerationRollback(current, committed, snapshot) + if err != nil { + t.Fatal(err) + } + result := string(merged) + if !strings.Contains(result, "previous.localhost") || strings.Contains(result, "failed.localhost") { + t.Fatalf("expected prior TLS policy to be restored despite unchanged project route, got %s", result) + } +} + +func TestRegisterRouteRejectsHostClaimedByAnotherProject(t *testing.T) { + _, server := newCaddyRouteState(t) + first := testRouteManager(server.URL, "first", "app.localhost") + second := testRouteManager(server.URL, "second", "app.localhost") + proxyRoute := route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@"}} + mustRegister(t, first, proxyRoute) + + err := second.register(proxyRoute) + if err == nil || !strings.Contains(err.Error(), "host conflict") { + t.Fatalf("expected cross-project host conflict, got %v", err) + } +} + +func TestRegisterRouteRejectsPartiallyOverlappingHosts(t *testing.T) { + _, server := newCaddyRouteState(t) + first := testRouteManager(server.URL, "first", "app.localhost") + second := testRouteManager(server.URL, "second", "app.localhost") + mustRegister(t, first, route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@", "*"}}) + + err := second.register(route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@"}}) + if err == nil || !strings.Contains(err.Error(), "host conflict") { + t.Fatalf("expected shared exact host pattern conflict, got %v", err) + } +} + +func TestRegisterRouteChangesListenerMode(t *testing.T) { + tests := []struct { + name string + server string + https bool + }{ + {name: "HTTPS on HTTP listener", server: `{"listen":[":80"],"automatic_https":{"disable":true},"routes":[]}`, https: true}, + {name: "HTTP on HTTPS listener", server: `{"listen":[":80"],"tls_connection_policies":[{}],"routes":[]}`}, + } + for _, test := range tests { + t.Run(test.name, func(t *testing.T) { + listenerPatched := false + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/routes"): + _, _ = response.Write([]byte(`[]`)) + case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/kool-80"): + _, _ = response.Write([]byte(test.server)) + case request.Method == http.MethodPatch && strings.HasSuffix(request.URL.Path, "/kool-80"): + listenerPatched = true + default: + response.WriteHeader(http.StatusOK) + } + })) + defer server.Close() + + env := environment.NewFakeEnvStorage() + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = server.URL + if err := manager.register(route{Service: "app", Listen: 80, Target: 8080, HTTPS: test.https}); err != nil { + t.Fatal(err) + } + if !listenerPatched { + t.Error("expected existing listener protocol to be replaced") + } + }) + } +} + +func TestRegisterRouteRejectsListenerModeChangeUsedByAnotherProject(t *testing.T) { + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + _, _ = response.Write([]byte(`{"listen":[":80"],"automatic_https":{"disable":true},"routes":[{"@id":"kool-other-app-80-80","handle":[{"@id":"kool-project-other-kool-other-app-80-80"}]}]}`)) + })) + defer server.Close() + + env := environment.NewFakeEnvStorage() + env.Set("KOOL_NAME", "example") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = server.URL + err := manager.register(route{Service: "app", Listen: 80, Target: 8080, HTTPS: true}) + if err == nil || !strings.Contains(err.Error(), "another project uses it") { + t.Fatalf("expected listener ownership conflict, got %v", err) + } +} + +func TestRegisterRouteCreatesServerAfterNotFound(t *testing.T) { + serverCreated := false + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodGet && strings.HasSuffix(request.URL.Path, "/routes") && serverCreated: + _, _ = response.Write([]byte(`[]`)) + case request.Method == http.MethodGet: + response.WriteHeader(http.StatusNotFound) + case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/kool-80"): + serverCreated = true + } + })) + defer server.Close() + + env := environment.NewFakeEnvStorage() + env.Set("KOOL_NAME", "example") + env.Set("KOOL_PROXY_HOST", "app.localhost") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = server.URL + + if err := manager.register(route{Service: "app", Listen: 80, Target: 8080, Hosts: []string{"@"}}); err != nil { + t.Fatal(err) + } + if !serverCreated { + t.Error("expected a missing server to be created after a 404") + } +} + +func TestRegisterTLS(t *testing.T) { + var policy map[string]interface{} + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + switch { + case request.Method == http.MethodGet: + _, _ = response.Write([]byte(`{"automation":{"policies":[]}}`)) + case request.Method == http.MethodDelete: + response.WriteHeader(http.StatusNotFound) + case request.Method == http.MethodPost && strings.HasSuffix(request.URL.Path, "/policies"): + body, _ := io.ReadAll(request.Body) + _ = json.Unmarshal(body, &policy) + default: + response.WriteHeader(http.StatusOK) + } + })) + defer server.Close() + + env := environment.NewFakeEnvStorage() + env.Set("KOOL_NAME", "example") + env.Set("KOOL_PROXY_HOST", "app.localhost") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = server.URL + + if err := manager.registerTLS([]route{ + {Service: "app", Listen: 443, Target: 80, Hosts: []string{"@", "*"}, HTTPS: true}, + {Service: "node", Listen: 3001, Target: 3001, Hosts: []string{"@"}, HTTPS: true}, + }); err != nil { + t.Fatal(err) + } + encoded, _ := json.Marshal(policy) + result := string(encoded) + for _, expected := range []string{"kool-example-tls", "app.localhost", "*.app.localhost", "internal"} { + if !strings.Contains(result, expected) { + t.Errorf("expected TLS policy to contain %q, got %s", expected, result) + } + } +} + +func TestRegisterTLSRemovesPolicyWhenHTTPSIsDisabled(t *testing.T) { + deleted := "" + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if request.Method == http.MethodDelete { + deleted = request.URL.Path + } + response.WriteHeader(http.StatusOK) + })) + defer server.Close() + + env := environment.NewFakeEnvStorage() + env.Set("KOOL_NAME", "example") + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = server.URL + if err := manager.registerTLS([]route{{Service: "app", Listen: 80, Target: 80}}); err != nil { + t.Fatal(err) + } + if deleted != "/id/kool-example-tls" { + t.Fatalf("expected stale project TLS policy to be deleted, got %q", deleted) + } +} + +func TestStopIfUnusedStopsOnlyAfterLastManagedRoute(t *testing.T) { + state, server := newCaddyRouteState(t) + manager := testRouteManager(server.URL, "example", "app.localhost") + fakeShell := &shell.FakeShell{} + manager.shell = fakeShell + manager.generation = "test" + proxyRoute := route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@"}} + mustRegister(t, manager, proxyRoute) + + if err := manager.stopIfUnusedUnlocked(); err != nil { + t.Fatal(err) + } + if fakeShell.CalledInteractive["docker"] { + t.Fatal("did not expect proxy to stop while a managed route remains") + } + state.routes["kool-80"] = nil + if err := manager.stopIfUnusedUnlocked(); err != nil { + t.Fatal(err) + } + if !fakeShell.CalledInteractive["docker"] || strings.Join(fakeShell.ArgsInteractive["docker"], " ") != "kool-proxy" { + t.Fatalf("expected final route removal to stop proxy, got %v", fakeShell.ArgsInteractive) + } +} + +func TestRestoreAppsReplacesPreviousConfiguration(t *testing.T) { + var method string + var restored []byte + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + method = request.Method + restored, _ = io.ReadAll(request.Body) + response.WriteHeader(http.StatusOK) + })) + defer server.Close() + + manager := NewManager(&shell.FakeShell{}, environment.NewFakeEnvStorage()).(*DefaultManager) + manager.adminURL = server.URL + snapshot := []byte(`{"http":{"servers":{"kool-80":{"listen":[":80"]}}}}`) + if err := manager.restoreApps(snapshot, true); err != nil { + t.Fatal(err) + } + if method != http.MethodPatch || string(restored) != string(snapshot) { + t.Fatalf("expected previous apps configuration to be restored, got %s %s", method, restored) + } +} + +func TestRestoreAppsDeletesConfigurationWhenPreviouslyMissing(t *testing.T) { + method := "" + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + method = request.Method + response.WriteHeader(http.StatusNotFound) + })) + defer server.Close() + + manager := NewManager(&shell.FakeShell{}, environment.NewFakeEnvStorage()).(*DefaultManager) + manager.adminURL = server.URL + if err := manager.restoreApps(nil, false); err != nil { + t.Fatal(err) + } + if method != http.MethodDelete { + t.Fatalf("expected newly created apps configuration to be deleted, got %s", method) + } +} + +func TestBaseConfigUsesReachableAdminListener(t *testing.T) { + home := t.TempDir() + t.Setenv("HOME", home) + manager := NewManager(&shell.FakeShell{}, environment.NewFakeEnvStorage()).(*DefaultManager) + path, err := manager.ensureBaseConfig() + if err != nil { + t.Fatal(err) + } + content, err := os.ReadFile(path) + if err != nil { + t.Fatal(err) + } + if !strings.Contains(string(content), `"listen":"0.0.0.0:2019"`) { + t.Fatalf("expected Docker-published Admin listener, got %s", content) + } +} + +func TestProxyStartupDoesNotResumeUntrustedAutosave(t *testing.T) { + if strings.Contains(caddyStartCmd, "--resume") || strings.Contains(caddyStartCmd, "autosave") { + t.Fatalf("proxy startup must not resume an unvalidated autosave: %s", caddyStartCmd) + } +} + +func TestProxyCompatibilityRequiresSecureStoredCommand(t *testing.T) { + expectedCommand := `["/bin/sh"]|["-c","` + strings.ReplaceAll(caddyStartCmd, `"`, `\"`) + `"]` + secure := `{"kool_proxy_admin":{}}|` + expectedCommand + legacy := `{"kool_proxy_admin":{}}|["/bin/sh"]|["-c","caddy run --resume"]` + if !strings.HasSuffix(secure, expectedCommand) { + t.Fatal("expected secure stored command to be accepted") + } + if strings.HasSuffix(legacy, expectedCommand) { + t.Fatal("expected legacy resume command to require container recreation") + } +} + +func TestParseExistingProxyPortsAndNetworks(t *testing.T) { + bindings := parsePortBindings(`{"80/tcp":[{"HostIp":"127.0.0.1","HostPort":"8080"}],"3001/tcp":[{"HostIp":"0.0.0.0","HostPort":"3001"}],"443/tcp":[{"HostIp":"::1","HostPort":"8443"}]}`) + if strings.Join(bindings[80], ",") != "127.0.0.1:8080:80" || strings.Join(bindings[3001], ",") != "0.0.0.0:3001:3001" || strings.Join(bindings[443], ",") != "[::1]:8443:443" { + t.Fatalf("expected exact existing proxy bindings to be preserved, got %v", bindings) + } + networks := parseDockerObjectKeys(`{"project_b":{},"kool_proxy_admin":{},"project_a":{}}`) + expected := []string{"kool_proxy_admin", "project_a", "project_b"} + if strings.Join(networks, ",") != strings.Join(expected, ",") { + t.Fatalf("expected sorted existing networks %v, got %v", expected, networks) + } +} + +func TestCopyPortSetPreservesOriginalExpansionPorts(t *testing.T) { + original := map[int][]string{80: {"127.0.0.1:8080:80"}} + expanded := copyPortBindings(original) + expanded[3001] = []string{"3001:3001"} + if strings.Join(original[80], ",") != "127.0.0.1:8080:80" || original[3001] != nil || expanded[3001] == nil { + t.Fatalf("expected rollback ports to remain independent, original=%v expanded=%v", original, expanded) + } +} + +func TestAdminBindingIsExcludedFromRecreatedListenerBindings(t *testing.T) { + bindings := parsePortBindings(`{"2019/tcp":[{"HostIp":"0.0.0.0","HostPort":"2019"}],"80/tcp":[{"HostIp":"0.0.0.0","HostPort":"80"}]}`) + replacement := copyPortBindings(bindings) + delete(replacement, 2019) + if bindings[2019] == nil || replacement[2019] != nil || replacement[80] == nil { + t.Fatalf("expected admin binding excluded only from replacement, original=%v replacement=%v", bindings, replacement) + } +} + +func TestRestoreAppsIfUnchangedPreservesConcurrentUpdate(t *testing.T) { + current := []byte(`{"http":{"servers":{"current":{}}}}`) + restored := false + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + if request.Method == http.MethodGet { + _, _ = response.Write(current) + return + } + restored = true + })) + defer server.Close() + + manager := NewManager(&shell.FakeShell{}, environment.NewFakeEnvStorage()).(*DefaultManager) + manager.adminURL = server.URL + committed := []byte(`{"http":{"servers":{"committed":{}}}}`) + if err := manager.restoreAppsIfUnchanged(committed, true, []byte(`{}`), true); err != nil { + t.Fatal(err) + } + if restored { + t.Fatal("did not expect rollback to overwrite concurrent Caddy changes") + } +} + +func TestWorkspaceRoutePrecedence(t *testing.T) { + proxyRoute := route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@", "*"}} + viteRoute := route{Service: "node", Listen: 3001, Target: 3001, Hosts: []string{"@", "*"}} + + t.Run("source started before workspace", func(t *testing.T) { + state, server := newCaddyRouteState(t) + source := testRouteManager(server.URL, "exlink", "exlink.localhost") + workspace := testRouteManager(server.URL, "exlink-workspace-task-a", "task-a.workspace.exlink.localhost") + + mustRegister(t, source, proxyRoute) + mustRegister(t, workspace, proxyRoute) + + state.requireRoutes(t, "kool-80", []string{ + "kool-exlink-workspace-task-a-app-80-80", + "kool-exlink-app-80-80", + }) + state.requireUpstream(t, "kool-80", "kool-exlink-workspace-task-a-app-80-80", "exlink-workspace-task-a-app:80") + }) + + t.Run("workspace started before and routes re-registered", func(t *testing.T) { + state, server := newCaddyRouteState(t) + source := testRouteManager(server.URL, "exlink", "exlink.localhost") + workspace := testRouteManager(server.URL, "exlink-workspace-task-a", "task-a.workspace.exlink.localhost") + + mustRegister(t, workspace, proxyRoute) + mustRegister(t, source, proxyRoute) + mustRegister(t, source, proxyRoute) + mustRegister(t, workspace, proxyRoute) + + state.requireRoutes(t, "kool-80", []string{ + "kool-exlink-workspace-task-a-app-80-80", + "kool-exlink-app-80-80", + }) + }) + + t.Run("multiple workspaces and unrelated project retain order", func(t *testing.T) { + state, server := newCaddyRouteState(t) + source := testRouteManager(server.URL, "exlink", "exlink.localhost") + other := testRouteManager(server.URL, "exlink-api", "other.localhost") + workspaceA := testRouteManager(server.URL, "exlink-workspace-task-a", "task-a.workspace.exlink.localhost") + workspaceB := testRouteManager(server.URL, "exlink-workspace-task-b", "task-b.workspace.exlink.localhost") + + mustRegister(t, source, proxyRoute) + mustRegister(t, other, proxyRoute) + mustRegister(t, workspaceA, proxyRoute) + mustRegister(t, workspaceB, proxyRoute) + mustRegister(t, workspaceA, proxyRoute) + + state.requireRoutes(t, "kool-80", []string{ + "kool-exlink-workspace-task-a-app-80-80", + "kool-exlink-workspace-task-b-app-80-80", + "kool-exlink-app-80-80", + "kool-exlink-api-app-80-80", + }) + }) + + t.Run("app and Vite listeners use the same precedence", func(t *testing.T) { + state, server := newCaddyRouteState(t) + source := testRouteManager(server.URL, "exlink", "exlink.localhost") + workspace := testRouteManager(server.URL, "exlink-workspace-task-a", "task-a.workspace.exlink.localhost") + + for _, manager := range []*DefaultManager{source, workspace} { + mustRegister(t, manager, proxyRoute) + mustRegister(t, manager, viteRoute) + } + + state.requireRoutes(t, "kool-80", []string{ + "kool-exlink-workspace-task-a-app-80-80", + "kool-exlink-app-80-80", + }) + state.requireRoutes(t, "kool-3001", []string{ + "kool-exlink-workspace-task-a-node-3001-3001", + "kool-exlink-node-3001-3001", + }) + }) + + t.Run("stopping one workspace removes only its routes", func(t *testing.T) { + state, server := newCaddyRouteState(t) + source := testRouteManager(server.URL, "exlink", "exlink.localhost") + workspaceA := testRouteManager(server.URL, "exlink-workspace-task-a", "task-a.workspace.exlink.localhost") + workspaceB := testRouteManager(server.URL, "exlink-workspace-task-b", "task-b.workspace.exlink.localhost") + + for _, manager := range []*DefaultManager{source, workspaceA, workspaceB} { + mustRegister(t, manager, proxyRoute) + mustRegister(t, manager, viteRoute) + } + if err := workspaceA.removeProjectRoutes(nil); err != nil { + t.Fatal(err) + } + + state.requireRoutes(t, "kool-80", []string{ + "kool-exlink-workspace-task-b-app-80-80", + "kool-exlink-app-80-80", + }) + state.requireRoutes(t, "kool-3001", []string{ + "kool-exlink-workspace-task-b-node-3001-3001", + "kool-exlink-node-3001-3001", + }) + if err := source.removeProjectRoutes(nil); err != nil { + t.Fatal(err) + } + state.requireRoutes(t, "kool-80", []string{"kool-exlink-workspace-task-b-app-80-80"}) + state.requireRoutes(t, "kool-3001", []string{"kool-exlink-workspace-task-b-node-3001-3001"}) + }) + + t.Run("changed configuration removes stale routes on every listener", func(t *testing.T) { + state, server := newCaddyRouteState(t) + source := testRouteManager(server.URL, "exlink", "exlink.localhost") + oldApp := route{Service: "app", Listen: 80, Target: 80, Hosts: []string{"@", "*"}} + oldVite := route{Service: "node", Listen: 3001, Target: 3001, Hosts: []string{"@", "*"}} + newApp := route{Service: "app", Listen: 8080, Target: 8080, Hosts: []string{"@", "*"}} + + mustRegister(t, source, oldApp) + mustRegister(t, source, oldVite) + if err := source.reconcileRoutes([]route{newApp}); err != nil { + t.Fatal(err) + } + + state.requireRoutes(t, "kool-80", nil) + state.requireRoutes(t, "kool-3001", nil) + }) +} + +type caddyRouteState struct { + routes map[string][]json.RawMessage +} + +func newCaddyRouteState(t *testing.T) (*caddyRouteState, *httptest.Server) { + t.Helper() + state := &caddyRouteState{routes: make(map[string][]json.RawMessage)} + server := httptest.NewServer(http.HandlerFunc(func(response http.ResponseWriter, request *http.Request) { + const serverPrefix = "/config/apps/http/servers/" + if request.URL.Path == "/config/apps" { + switch request.Method { + case http.MethodGet: + servers := make(map[string]map[string][]json.RawMessage, len(state.routes)) + for serverID, routes := range state.routes { + servers[serverID] = map[string][]json.RawMessage{"routes": routes} + } + _ = json.NewEncoder(response).Encode(map[string]interface{}{"http": map[string]interface{}{"servers": servers}, "tls": map[string]interface{}{"automation": map[string]interface{}{"policies": []interface{}{}}}}) + case http.MethodPatch: + var apps map[string]interface{} + if err := json.NewDecoder(request.Body).Decode(&apps); err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + servers := caddyServers(apps) + state.routes = make(map[string][]json.RawMessage, len(servers)) + for serverID, server := range servers { + for _, route := range caddyRoutes(server) { + raw, _ := json.Marshal(route) + state.routes[serverID] = append(state.routes[serverID], raw) + } + } + } + return + } + if request.Method == http.MethodGet && request.URL.Path == strings.TrimSuffix(serverPrefix, "/") { + servers := make(map[string]map[string][]json.RawMessage, len(state.routes)) + for serverID, routes := range state.routes { + servers[serverID] = map[string][]json.RawMessage{"routes": routes} + } + _ = json.NewEncoder(response).Encode(servers) + return + } + switch { + case request.Method == http.MethodDelete && strings.HasPrefix(request.URL.Path, "/id/"): + id := strings.TrimPrefix(request.URL.Path, "/id/") + for serverID, routes := range state.routes { + for index, route := range routes { + if caddyRouteID(route) == id { + state.routes[serverID] = append(routes[:index], routes[index+1:]...) + response.WriteHeader(http.StatusOK) + return + } + } + } + response.WriteHeader(http.StatusNotFound) + return + case !strings.HasPrefix(request.URL.Path, serverPrefix): + response.WriteHeader(http.StatusOK) + return + } + + path := strings.TrimPrefix(request.URL.Path, serverPrefix) + parts := strings.Split(path, "/") + serverID := parts[0] + isRoutes := len(parts) == 2 && parts[1] == "routes" + switch { + case request.Method == http.MethodGet && isRoutes: + routes, exists := state.routes[serverID] + if !exists { + response.WriteHeader(http.StatusNotFound) + return + } + _ = json.NewEncoder(response).Encode(routes) + case request.Method == http.MethodGet: + routes, exists := state.routes[serverID] + if !exists { + _, _ = response.Write([]byte("null")) + return + } + _ = json.NewEncoder(response).Encode(map[string]interface{}{"routes": routes}) + case request.Method == http.MethodPost && len(parts) == 1: + state.routes[serverID] = []json.RawMessage{} + case request.Method == http.MethodPatch && isRoutes: + var routes []json.RawMessage + if err := json.NewDecoder(request.Body).Decode(&routes); err != nil { + http.Error(response, err.Error(), http.StatusBadRequest) + return + } + state.routes[serverID] = routes + default: + response.WriteHeader(http.StatusOK) + } + })) + t.Cleanup(server.Close) + return state, server +} + +func testRouteManager(adminURL, project, host string) *DefaultManager { + env := environment.NewFakeEnvStorage() + env.Set("KOOL_PROXY_HOST", host) + env.Set("COMPOSE_PROJECT_NAME", project) + manager := NewManager(&shell.FakeShell{}, env).(*DefaultManager) + manager.adminURL = adminURL + return manager +} + +func mustRegister(t *testing.T, manager *DefaultManager, route route) { + t.Helper() + if err := manager.register(route); err != nil { + t.Fatal(err) + } +} + +func (state *caddyRouteState) requireRoutes(t *testing.T, serverID string, expected []string) { + t.Helper() + var actual []string + for _, route := range state.routes[serverID] { + actual = append(actual, caddyRouteID(route)) + } + if strings.Join(actual, ",") != strings.Join(expected, ",") { + t.Fatalf("expected %s routes %v, got %v", serverID, expected, actual) + } +} + +func (state *caddyRouteState) requireUpstream(t *testing.T, serverID, routeID, expected string) { + t.Helper() + for _, rawRoute := range state.routes[serverID] { + if caddyRouteID(rawRoute) != routeID { + continue + } + var routeConfig struct { + Handle []struct { + Upstreams []struct { + Dial string `json:"dial"` + } `json:"upstreams"` + } `json:"handle"` + } + if json.Unmarshal(rawRoute, &routeConfig) != nil || len(routeConfig.Handle) == 0 || len(routeConfig.Handle[0].Upstreams) == 0 { + t.Fatalf("route %s has no upstream", routeID) + } + if actual := routeConfig.Handle[0].Upstreams[0].Dial; actual != expected { + t.Fatalf("expected route %s upstream %s, got %s", routeID, expected, actual) + } + return + } + t.Fatalf("route %s was not found", routeID) +} + +func caddyRouteID(route json.RawMessage) string { + var metadata caddyRouteMetadata + _ = json.Unmarshal(route, &metadata) + return metadata.ID +}