diff --git a/coderd/templatebuilder/compose.go b/coderd/templatebuilder/compose.go new file mode 100644 index 0000000000..381bb8772f --- /dev/null +++ b/coderd/templatebuilder/compose.go @@ -0,0 +1,257 @@ +package templatebuilder + +import ( + "archive/tar" + "bytes" + "encoding/json" + "time" + + "golang.org/x/xerrors" +) + +// ComposeRequest describes which base template and modules to render. +type ComposeRequest struct { + BaseTemplateID string + // RegistryURL is the module registry base URL from the deployment + // config (CODER_TEMPLATE_BUILDER_REGISTRY_URL). + RegistryURL string + Modules []ComposeModule +} + +// ComposeModule identifies a module to include and the variable values +// to render into its module block. +type ComposeModule struct { + ID string + // Variables maps variable names to HCL literal values for + // non-sensitive, non-computed variables. + Variables map[string]string +} + +// ComposeResult holds the rendered Terraform files ready for bundling. +type ComposeResult struct { + // MainTF is the rendered base template. + MainTF []byte + // ModulesTF is the concatenated rendered module blocks. Empty when + // no modules are selected. + ModulesTF []byte +} + +// Compose renders a base template and selected modules into Terraform +// source files. It extracts the coder_agent resource name from the +// rendered base HCL and wires it into each module block. +func Compose(req ComposeRequest) (*ComposeResult, error) { + mainTF, err := renderBase(req.BaseTemplateID) + if err != nil { + return nil, err + } + + if len(req.Modules) == 0 { + return &ComposeResult{MainTF: mainTF}, nil + } + + agentName, err := ExtractAgentResourceName(mainTF) + if err != nil { + return nil, xerrors.Errorf("extract agent name: %w", err) + } + + catalog, err := loadCatalogMap() + if err != nil { + return nil, err + } + + baseOS := BaseTemplateOS(req.BaseTemplateID) + if err := validateModules(req.Modules, catalog, baseOS); err != nil { + return nil, err + } + + modulesTF, err := renderModules(req.Modules, catalog, req.RegistryURL, agentName) + if err != nil { + return nil, err + } + + return &ComposeResult{ + MainTF: mainTF, + ModulesTF: modulesTF, + }, nil +} + +// renderBase renders the base template for the given example ID. +func renderBase(baseTemplateID string) ([]byte, error) { + renderCtx := DefaultBaseRenderContext(baseTemplateID) + mainTF, err := RenderBaseTemplate(baseTemplateID, "main.tf.tmpl", renderCtx) + if err != nil { + return nil, xerrors.Errorf("render base template: %w", err) + } + return mainTF, nil +} + +// loadCatalogMap loads the module catalog and returns it as a map keyed +// by module ID. +func loadCatalogMap() (map[string]ModuleManifest, error) { + modules, err := LoadModules() + if err != nil { + return nil, xerrors.Errorf("load module catalog: %w", err) + } + catalog := make(map[string]ModuleManifest, len(modules)) + for _, m := range modules { + catalog[m.ID] = m + } + return catalog, nil +} + +// validateModules checks that all requested modules exist, are +// OS-compatible, have no duplicates, and have no conflicts. +func validateModules(requested []ComposeModule, catalog map[string]ModuleManifest, baseOS BaseOS) error { + seen := make(map[string]bool, len(requested)) + for _, cm := range requested { + if seen[cm.ID] { + return xerrors.Errorf("duplicate module %q", cm.ID) + } + seen[cm.ID] = true + + manifest, ok := catalog[cm.ID] + if !ok { + return xerrors.Errorf("unknown module %q", cm.ID) + } + if !manifest.CompatibleWithOS(string(baseOS)) { + return xerrors.Errorf("module %q is not compatible with OS %q", cm.ID, baseOS) + } + } + + // Check conflicts bidirectionally so that order does not matter. + for _, cm := range requested { + manifest := catalog[cm.ID] + for _, conflict := range manifest.ConflictsWith { + if seen[conflict] { + return xerrors.Errorf("module %q conflicts with %q", cm.ID, conflict) + } + } + } + + return nil +} + +// renderModules renders each module template and concatenates the +// results with newline separators. +func renderModules( + requested []ComposeModule, + catalog map[string]ModuleManifest, + registryURL, agentName string, +) ([]byte, error) { + var buf bytes.Buffer + for _, cm := range requested { + manifest := catalog[cm.ID] + + modFS, err := ModuleTemplateFS(cm.ID) + if err != nil { + return nil, xerrors.Errorf("module template FS for %q: %w", cm.ID, err) + } + + vars := mergeModuleVariables(manifest, cm.Variables) + modCtx := ModuleRenderContext{ + RegistryBase: registryURL, + PinnedVersion: manifest.PinnedVersion, + AgentResourceName: agentName, + Variables: vars, + } + + rendered, err := RenderModuleTemplate(modFS, cm.ID+".tf.tmpl", modCtx) + if err != nil { + return nil, xerrors.Errorf("render module %q: %w", cm.ID, err) + } + + if buf.Len() > 0 { + _ = buf.WriteByte('\n') + } + _, _ = buf.Write(rendered) + } + return buf.Bytes(), nil +} + +// mergeModuleVariables builds the final Variables map for a module template. +// It starts with manifest defaults for all non-computed, non-sensitive +// variables, then overlays caller-supplied values. This ensures every +// variable referenced in the template has a value. +func mergeModuleVariables(manifest ModuleManifest, callerVars map[string]string) map[string]string { + merged := make(map[string]string, len(manifest.Variables)) + for _, v := range manifest.Variables { + if v.Computed || v.Sensitive { + continue + } + if len(v.Default) > 0 && isSimpleJSONValue(v.Default) { + // json.RawMessage values for simple types (e.g. `""`, + // `false`, `13337`) are valid HCL literals. + merged[v.Name] = string(v.Default) + } else if !v.Required { + // Non-required variables without an explicit default use + // null, which tells Terraform to apply the module's own + // default. + merged[v.Name] = "null" + } + // Required variables without defaults are left out so that + // missingkey=error surfaces the omission at render time. + } + for k, val := range callerVars { + merged[k] = val + } + return merged +} + +// isSimpleJSONValue returns true if raw is a valid JSON string, number, +// bool, or null. Arrays and objects are rejected; the template builder +// only supports simple variable types. +func isSimpleJSONValue(raw json.RawMessage) bool { + var v interface{} + if err := json.Unmarshal(raw, &v); err != nil { + return false + } + switch v.(type) { + case string, float64, bool, nil: + return true + default: + return false + } +} + +// BundleTar packages the compose result into a tar archive suitable for +// the Coder file store. +func BundleTar(result *ComposeResult) ([]byte, error) { + if result == nil { + return nil, xerrors.New("nil ComposeResult") + } + + var buf bytes.Buffer + tw := tar.NewWriter(&buf) + + if err := writeTarFile(tw, "main.tf", result.MainTF); err != nil { + return nil, xerrors.Errorf("write main.tf to tar: %w", err) + } + + if len(result.ModulesTF) > 0 { + if err := writeTarFile(tw, "modules.tf", result.ModulesTF); err != nil { + return nil, xerrors.Errorf("write modules.tf to tar: %w", err) + } + } + + if err := tw.Close(); err != nil { + return nil, xerrors.Errorf("close tar writer: %w", err) + } + + return buf.Bytes(), nil +} + +// writeTarFile adds a single file entry to a tar writer. It uses a zero +// timestamp for reproducible archives. +func writeTarFile(tw *tar.Writer, name string, data []byte) error { + hdr := &tar.Header{ + Name: name, + Mode: 0o644, + Size: int64(len(data)), + ModTime: time.Unix(0, 0), + } + if err := tw.WriteHeader(hdr); err != nil { + return err + } + _, err := tw.Write(data) + return err +} diff --git a/coderd/templatebuilder/compose_internal_test.go b/coderd/templatebuilder/compose_internal_test.go new file mode 100644 index 0000000000..53edf70957 --- /dev/null +++ b/coderd/templatebuilder/compose_internal_test.go @@ -0,0 +1,98 @@ +package templatebuilder + +import ( + "encoding/json" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestIsSimpleJSONValue(t *testing.T) { + t.Parallel() + + tests := []struct { + name string + raw json.RawMessage + want bool + }{ + {"String", json.RawMessage(`"hello"`), true}, + {"EmptyString", json.RawMessage(`""`), true}, + {"True", json.RawMessage(`true`), true}, + {"False", json.RawMessage(`false`), true}, + {"Null", json.RawMessage(`null`), true}, + {"PositiveInt", json.RawMessage(`42`), true}, + {"NegativeInt", json.RawMessage(`-1`), true}, + {"Float", json.RawMessage(`3.14`), true}, + {"Array", json.RawMessage(`[1,2]`), false}, + {"Object", json.RawMessage(`{"a":1}`), false}, + {"Empty", json.RawMessage(``), false}, + {"Nil", nil, false}, + {"MalformedString", json.RawMessage(`"unclosed`), false}, + {"MalformedBool", json.RawMessage(`truesomething`), false}, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + got := isSimpleJSONValue(tc.raw) + require.Equal(t, tc.want, got) + }) + } +} + +func TestMergeModuleVariables(t *testing.T) { + t.Parallel() + + manifest := ModuleManifest{ + Variables: []ModuleVariable{ + {Name: "agent_id", Type: "string", Computed: true}, + {Name: "api_key", Type: "string", Sensitive: true}, + {Name: "port", Type: "number", Default: json.RawMessage(`13337`)}, + {Name: "enabled", Type: "bool", Default: json.RawMessage(`false`)}, + {Name: "optional_no_default", Type: "string", Required: false}, + {Name: "required_no_default", Type: "string", Required: true}, + }, + } + + t.Run("DefaultsApplied", func(t *testing.T) { + t.Parallel() + merged := mergeModuleVariables(manifest, nil) + require.Equal(t, "13337", merged["port"]) + require.Equal(t, "false", merged["enabled"]) + }) + + t.Run("ComputedAndSensitiveSkipped", func(t *testing.T) { + t.Parallel() + merged := mergeModuleVariables(manifest, nil) + require.NotContains(t, merged, "agent_id") + require.NotContains(t, merged, "api_key") + }) + + t.Run("NonRequiredWithoutDefaultGetsNull", func(t *testing.T) { + t.Parallel() + merged := mergeModuleVariables(manifest, nil) + require.Equal(t, "null", merged["optional_no_default"]) + }) + + t.Run("RequiredWithoutDefaultOmitted", func(t *testing.T) { + t.Parallel() + merged := mergeModuleVariables(manifest, nil) + require.NotContains(t, merged, "required_no_default") + }) + + t.Run("CallerOverridesDefault", func(t *testing.T) { + t.Parallel() + merged := mergeModuleVariables(manifest, map[string]string{ + "port": "9999", + }) + require.Equal(t, "9999", merged["port"]) + }) + + t.Run("CallerProvidesRequired", func(t *testing.T) { + t.Parallel() + merged := mergeModuleVariables(manifest, map[string]string{ + "required_no_default": `"value"`, + }) + require.Equal(t, `"value"`, merged["required_no_default"]) + }) +} diff --git a/coderd/templatebuilder/compose_test.go b/coderd/templatebuilder/compose_test.go new file mode 100644 index 0000000000..175c0e78cb --- /dev/null +++ b/coderd/templatebuilder/compose_test.go @@ -0,0 +1,277 @@ +package templatebuilder_test + +import ( + "archive/tar" + "bytes" + "errors" + "io" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/coder/coder/v2/coderd/templatebuilder" +) + +func TestCompose(t *testing.T) { + t.Parallel() + + t.Run("BaseOnly", func(t *testing.T) { + t.Parallel() + result, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + }) + require.NoError(t, err) + require.NotEmpty(t, result.MainTF) + require.Contains(t, string(result.MainTF), `resource "coder_agent" "main"`) + require.Empty(t, result.ModulesTF) + }) + + t.Run("BaseWithModuleAndVariableOverride", func(t *testing.T) { + t.Parallel() + result, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + { + ID: "code-server", + Variables: map[string]string{ + "port": "9999", + }, + }, + }, + }) + require.NoError(t, err) + require.NotEmpty(t, result.MainTF) + require.NotEmpty(t, result.ModulesTF) + + modules := string(result.ModulesTF) + require.Contains(t, modules, `module "code-server"`) + require.Contains(t, modules, `coder_agent.main.id`) + require.Contains(t, modules, `registry.coder.com`) + require.Contains(t, modules, `port = 9999`) + }) + + t.Run("AWSLinuxAgentName", func(t *testing.T) { + t.Parallel() + result, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "aws-linux", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + {ID: "git-commit-signing"}, + }, + }) + require.NoError(t, err) + require.Contains(t, string(result.ModulesTF), `coder_agent.dev.id`) + }) + + t.Run("SensitiveVariable", func(t *testing.T) { + t.Parallel() + result, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + {ID: "claude-code"}, + }, + }) + require.NoError(t, err) + modules := string(result.ModulesTF) + // claude-code has a sensitive variable (claude_code_oauth_token) + // that renders as a top-level variable block + var. reference. + require.Contains(t, modules, `variable "claude_code_oauth_token"`) + require.Contains(t, modules, `sensitive = true`) + require.Contains(t, modules, `var.claude_code_oauth_token`) + }) + + t.Run("MultipleModulesWithRequiredVariable", func(t *testing.T) { + t.Parallel() + result, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + {ID: "code-server"}, + { + ID: "git-clone", + Variables: map[string]string{ + "url": `"https://github.com/coder/coder"`, + }, + }, + }, + }) + require.NoError(t, err) + modules := string(result.ModulesTF) + require.Contains(t, modules, `module "code-server"`) + require.Contains(t, modules, `module "git-clone"`) + require.Contains(t, modules, `"https://github.com/coder/coder"`) + }) + + t.Run("CustomRegistryURL", func(t *testing.T) { + t.Parallel() + result, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.internal.corp", + Modules: []templatebuilder.ComposeModule{ + {ID: "code-server"}, + }, + }) + require.NoError(t, err) + require.Contains(t, string(result.ModulesTF), `registry.internal.corp`) + }) + + t.Run("DuplicateModuleError", func(t *testing.T) { + t.Parallel() + _, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + {ID: "code-server"}, + {ID: "code-server"}, + }, + }) + require.Error(t, err) + require.Contains(t, err.Error(), `duplicate module "code-server"`) + }) + + t.Run("ConflictingModuleError", func(t *testing.T) { + t.Parallel() + _, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + {ID: "code-server"}, + {ID: "vscode-web"}, + }, + }) + require.Error(t, err) + require.Contains(t, err.Error(), "conflicts with") + }) + + t.Run("UnknownBase", func(t *testing.T) { + t.Parallel() + _, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "nonexistent", + RegistryURL: "https://registry.coder.com", + }) + require.Error(t, err) + require.Contains(t, err.Error(), "unknown base template") + }) + + t.Run("UnknownModule", func(t *testing.T) { + t.Parallel() + _, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + {ID: "nonexistent-module"}, + }, + }) + require.Error(t, err) + require.Contains(t, err.Error(), `unknown module "nonexistent-module"`) + }) + + t.Run("MissingRequiredVariable", func(t *testing.T) { + t.Parallel() + // git-clone has a required "url" variable with no default. + // Omitting it should cause a render error from missingkey=error. + _, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + {ID: "git-clone"}, + }, + }) + require.Error(t, err) + require.Contains(t, err.Error(), "render module") + }) +} + +func TestBundleTar(t *testing.T) { + t.Parallel() + + t.Run("NilResult", func(t *testing.T) { + t.Parallel() + _, err := templatebuilder.BundleTar(nil) + require.Error(t, err) + require.Contains(t, err.Error(), "nil") + }) + + t.Run("MainOnly", func(t *testing.T) { + t.Parallel() + result := &templatebuilder.ComposeResult{ + MainTF: []byte("resource {}"), + } + data, err := templatebuilder.BundleTar(result) + require.NoError(t, err) + + files := extractTar(t, data) + require.Contains(t, files, "main.tf") + require.NotContains(t, files, "modules.tf") + require.Equal(t, "resource {}", files["main.tf"]) + }) + + t.Run("MainAndModules", func(t *testing.T) { + t.Parallel() + result := &templatebuilder.ComposeResult{ + MainTF: []byte("resource {}"), + ModulesTF: []byte("module {}"), + } + data, err := templatebuilder.BundleTar(result) + require.NoError(t, err) + + files := extractTar(t, data) + require.Contains(t, files, "main.tf") + require.Contains(t, files, "modules.tf") + require.Equal(t, "resource {}", files["main.tf"]) + require.Equal(t, "module {}", files["modules.tf"]) + }) + + t.Run("RoundTrip", func(t *testing.T) { + t.Parallel() + result, err := templatebuilder.Compose(templatebuilder.ComposeRequest{ + BaseTemplateID: "docker", + RegistryURL: "https://registry.coder.com", + Modules: []templatebuilder.ComposeModule{ + {ID: "code-server"}, + }, + }) + require.NoError(t, err) + + data, err := templatebuilder.BundleTar(result) + require.NoError(t, err) + + files := extractTar(t, data) + require.Equal(t, string(result.MainTF), files["main.tf"]) + require.Equal(t, string(result.ModulesTF), files["modules.tf"]) + }) + + t.Run("ReproducibleArchive", func(t *testing.T) { + t.Parallel() + result := &templatebuilder.ComposeResult{ + MainTF: []byte("resource {}"), + ModulesTF: []byte("module {}"), + } + data1, err := templatebuilder.BundleTar(result) + require.NoError(t, err) + data2, err := templatebuilder.BundleTar(result) + require.NoError(t, err) + require.Equal(t, data1, data2, "identical inputs should produce identical archives") + }) +} + +// extractTar reads a tar archive and returns a map of filename to content. +func extractTar(t *testing.T, data []byte) map[string]string { + t.Helper() + tr := tar.NewReader(bytes.NewReader(data)) + files := make(map[string]string) + for { + hdr, err := tr.Next() + if errors.Is(err, io.EOF) { + break + } + require.NoError(t, err) + body, err := io.ReadAll(tr) + require.NoError(t, err) + files[hdr.Name] = string(body) + } + return files +}