diff --git a/docs/project-specification/01-authoring-projects.md b/docs/project-specification/01-authoring-projects.md index b94ed802..86715e4e 100644 --- a/docs/project-specification/01-authoring-projects.md +++ b/docs/project-specification/01-authoring-projects.md @@ -71,7 +71,6 @@ x-topo: : description: string # Optional required: boolean # Optional - default: string # Optional example: string # Optional ``` diff --git a/docs/project-specification/02-project-configuration.md b/docs/project-specification/02-project-configuration.md index 24268012..e752757b 100644 --- a/docs/project-specification/02-project-configuration.md +++ b/docs/project-specification/02-project-configuration.md @@ -20,8 +20,7 @@ services: platform: linux/arm64 build: context: . - # Optional default: allows running with plain docker compose - # Not used by Implementations that read x-topo.parameters + # Initial value: allows running with plain docker compose args: GREETING: "Hello, World" @@ -65,7 +64,7 @@ x-topo: parameters: MODEL: description: "Model artifact reference" - default: "bartowski/Qwen_Qwen3.5-0.8B-GGUF:SmolLM2-135M-Instruct-Q4_K_M.gguf" + example: "bartowski/Qwen_Qwen3.5-0.8B-GGUF:SmolLM2-135M-Instruct-Q4_K_M.gguf" hints: huggingface.task: text-generation file.format: gguf diff --git a/docs/project-specification/schema/topo-project-specification.json b/docs/project-specification/schema/topo-project-specification.json index 9015ca39..19283f1c 100644 --- a/docs/project-specification/schema/topo-project-specification.json +++ b/docs/project-specification/schema/topo-project-specification.json @@ -44,11 +44,11 @@ }, "required": { "type": "boolean", - "description": "If `true`, Implementations must enforce input or error" + "description": "If `true`, Implementations must enforce that all build args values are set for this parameter when configuring the Project" }, "default": { "type": "string", - "description": "Value used if user skips input (only valid when not required)" + "description": "Deprecated. Topo versions >11.0.1 ignore this property." }, "example": { "type": "string", diff --git a/internal/arguments/interactive_provider.go b/internal/arguments/interactive_provider.go index 3c884b68..eb003f12 100644 --- a/internal/arguments/interactive_provider.go +++ b/internal/arguments/interactive_provider.go @@ -40,19 +40,19 @@ func (p *InteractiveProvider) Provide(args []Arg) ([]ResolvedArg, error) { } } - if arg.Default != "" { - _, err := fmt.Fprintf(p.output, "Default: %s\n", arg.Default) + if len(arg.CurrentValues) > 0 { + _, err := fmt.Fprintf(p.output, "Current: %s\n", formatCurrentValues(arg.CurrentValues)) if err != nil { return nil, err } } - requiredLabel := "" + label := "optional" if arg.Required { - requiredLabel = " (required)" + label = "required" } - _, err = fmt.Fprintf(p.output, "%s%s> ", arg.Name, requiredLabel) + _, err = fmt.Fprintf(p.output, "%s (%s, press Enter to skip)> ", arg.Name, label) if err != nil { return nil, err } diff --git a/internal/arguments/interactive_provider_test.go b/internal/arguments/interactive_provider_test.go index 6203a776..4968c89f 100644 --- a/internal/arguments/interactive_provider_test.go +++ b/internal/arguments/interactive_provider_test.go @@ -40,7 +40,7 @@ func TestInteractiveProvider(t *testing.T) { assert.Equal(t, want, got) assert.Contains(t, output.String(), "The greeting message") assert.Contains(t, output.String(), "Example: Hello") - assert.Contains(t, output.String(), "GREETING (required)>") + assert.Contains(t, output.String(), "GREETING (required, press Enter to skip)>") }) t.Run("skips empty inputs", func(t *testing.T) { @@ -48,13 +48,25 @@ func TestInteractiveProvider(t *testing.T) { output := &bytes.Buffer{} provider := arguments.NewInteractiveProvider(input, output) - args := []arguments.Arg{ - {Name: "OPTIONAL", Required: false}, - } + got, err := provider.Provide([]arguments.Arg{{Name: "OPTIONAL"}}) + + require.NoError(t, err) + assert.Empty(t, got) + }) + + t.Run("shows current values", func(t *testing.T) { + input := strings.NewReader("\n") + output := &bytes.Buffer{} + provider := arguments.NewInteractiveProvider(input, output) + args := []arguments.Arg{{ + Name: "GREETING", + CurrentValues: []string{"Hello", ""}, + }} got, err := provider.Provide(args) require.NoError(t, err) assert.Empty(t, got) + assert.Contains(t, output.String(), `Current: ["Hello",""]`) }) } diff --git a/internal/arguments/provider.go b/internal/arguments/provider.go index 4c457a52..ac6b4f38 100644 --- a/internal/arguments/provider.go +++ b/internal/arguments/provider.go @@ -1,11 +1,16 @@ package arguments +import ( + "encoding/json" + "fmt" +) + type Arg struct { - Name string - Description string - Required bool - Example string - Default string + Name string + Description string + Required bool + Example string + CurrentValues []string } type ResolvedArg struct { @@ -16,3 +21,11 @@ type ResolvedArg struct { type Provider interface { Provide(args []Arg) ([]ResolvedArg, error) } + +func formatCurrentValues(values []string) string { + formatted, err := json.Marshal(values) + if err != nil { + return fmt.Sprintf("%q", values) + } + return string(formatted) +} diff --git a/internal/arguments/strict_provider_chain.go b/internal/arguments/strict_provider_chain.go index efd438c4..0f3ce7b3 100644 --- a/internal/arguments/strict_provider_chain.go +++ b/internal/arguments/strict_provider_chain.go @@ -1,8 +1,10 @@ package arguments -import "strings" - -import "fmt" +import ( + "fmt" + "slices" + "strings" +) // StrictProviderChain chains multiple providers and ensures all required arguments are resolved. // It stops early once all required arguments are satisfied. @@ -34,16 +36,12 @@ func (p *StrictProviderChain) Provide(args []Arg) ([]ResolvedArg, error) { remaining = filterProvided(remaining, provided) - if allRequiredProvided(args, provided) { + if allRequiredResolved(args, provided) { break } } - if len(remaining) > 0 { - defaultNonProvided(remaining, provided) - } - - if err := validateRequiredProvided(args, provided); err != nil { + if err := validateRequiredResolved(args, provided); err != nil { return nil, err } @@ -61,13 +59,16 @@ type MissingArgsError []Arg func (e MissingArgsError) Error() string { var msg strings.Builder - msg.WriteString("missing required parameters:\n") + msg.WriteString("missing value(s) for required parameters:\n") for _, arg := range e { fmt.Fprintf(&msg, " %s:\n", arg.Name) fmt.Fprintf(&msg, " description: %s\n", arg.Description) if arg.Example != "" { fmt.Fprintf(&msg, " example: %s\n", arg.Example) } + if len(arg.CurrentValues) > 0 { + fmt.Fprintf(&msg, " # current: %s\n", formatCurrentValues(arg.CurrentValues)) + } } return msg.String() } @@ -82,32 +83,30 @@ func filterProvided(args []Arg, provided map[string]string) []Arg { return remaining } -func allRequiredProvided(args []Arg, provided map[string]string) bool { +func allRequiredResolved(args []Arg, provided map[string]string) bool { for _, arg := range args { - if arg.Required { - if value, exists := provided[arg.Name]; !exists || value == "" { - return false - } + if arg.Required && !isResolved(arg, provided) { + return false } } return true } -func defaultNonProvided(remaining []Arg, provided map[string]string) { - for _, arg := range remaining { - if arg.Default != "" { - provided[arg.Name] = arg.Default - } +func isResolved(arg Arg, provided map[string]string) bool { + if value, exists := provided[arg.Name]; exists { + return value != "" + } + if len(arg.CurrentValues) == 0 { + return false } + return !slices.Contains(arg.CurrentValues, "") } -func validateRequiredProvided(args []Arg, provided map[string]string) error { +func validateRequiredResolved(args []Arg, provided map[string]string) error { var missing []Arg for _, arg := range args { - if arg.Required { - if value, exists := provided[arg.Name]; !exists || value == "" { - missing = append(missing, arg) - } + if arg.Required && !isResolved(arg, provided) { + missing = append(missing, arg) } } diff --git a/internal/arguments/strict_provider_chain_test.go b/internal/arguments/strict_provider_chain_test.go index 4a0c427f..6449adc5 100644 --- a/internal/arguments/strict_provider_chain_test.go +++ b/internal/arguments/strict_provider_chain_test.go @@ -160,27 +160,36 @@ func TestStrictMultiProvider(t *testing.T) { assert.Equal(t, want, got) }) - t.Run("provides resolved args when default provided", func(t *testing.T) { - provider1 := arguments.NewStaticProvider() - multi := arguments.NewStrictProviderChain(provider1) - args := []arguments.Arg{ - { - Name: "CINNAMON", - Required: true, - Default: "filled", - }, - } + t.Run("allows required arguments with non-empty current values", func(t *testing.T) { + provider := arguments.NewStaticProvider() + multi := arguments.NewStrictProviderChain(provider) + args := []arguments.Arg{{ + Name: "CINNAMON", + Required: true, + CurrentValues: []string{"current", "${CINNAMON}"}, + }} got, err := multi.Provide(args) require.NoError(t, err) - want := []arguments.ResolvedArg{ - {Name: "CINNAMON", Value: "filled"}, + assert.Empty(t, got) + }) + + t.Run("errors when any current value is empty", func(t *testing.T) { + provider := arguments.NewStaticProvider() + multi := arguments.NewStrictProviderChain(provider) + arg := arguments.Arg{ + Name: "CINNAMON", + Required: true, + CurrentValues: []string{"current", ""}, } - assert.Equal(t, want, got) + + _, err := multi.Provide([]arguments.Arg{arg}) + + assert.Equal(t, arguments.MissingArgsError{arg}, err) }) - t.Run("does not provide resolved args when no default provided", func(t *testing.T) { + t.Run("does not resolve omitted optional arguments", func(t *testing.T) { provider1 := arguments.NewStaticProvider() multi := arguments.NewStrictProviderChain(provider1) args := []arguments.Arg{ @@ -207,19 +216,21 @@ func TestMissingArgsError(t *testing.T) { Example: "Hello", }, { - Name: "PORT", - Description: "Port number", + Name: "PORT", + Description: "Port number", + CurrentValues: []string{"8080", ""}, }, } got := err.Error() - want := `missing required parameters: + want := `missing value(s) for required parameters: GREETING: description: The greeting message example: Hello PORT: description: Port number + # current: ["8080",""] ` assert.Equal(t, want, got) }) diff --git a/internal/project/definition.go b/internal/project/definition.go index 8392c999..edb55901 100644 --- a/internal/project/definition.go +++ b/internal/project/definition.go @@ -3,6 +3,7 @@ package project import ( "fmt" "io" + "strings" "github.com/arm/topo/internal/output/logger" "gopkg.in/yaml.v3" @@ -11,7 +12,8 @@ import ( const ComposeFilename = "compose.yaml" type Project struct { - Metadata Metadata + Metadata Metadata + currentParameterValues map[string][]string } type Metadata struct { @@ -26,7 +28,6 @@ type Parameter struct { Description string Required bool Example string - Default string } func FromContent(reader io.Reader) (Project, error) { @@ -34,14 +35,23 @@ func FromContent(reader io.Reader) (Project, error) { XTopo Metadata `yaml:"x-topo"` } - var parsed composeFile + var document yaml.Node decoder := yaml.NewDecoder(reader) - if err := decoder.Decode(&parsed); err != nil { + if err := decoder.Decode(&document); err != nil { + return Project{}, fmt.Errorf("failed to decode project: %w", err) + } + if len(document.Content) == 0 { + return Project{}, fmt.Errorf("failed to decode project: compose file is empty") + } + + var parsed composeFile + if err := document.Decode(&parsed); err != nil { return Project{}, fmt.Errorf("failed to decode project: %w", err) } return Project{ - Metadata: parsed.XTopo, + Metadata: parsed.XTopo, + currentParameterValues: parseCurrentParameterValues(document.Content[0]), }, nil } @@ -57,7 +67,6 @@ type rawParameter struct { Description string `yaml:"description"` Required bool `yaml:"required"` Example string `yaml:"example,omitempty"` - Default string `yaml:"default,omitempty"` } func (t *Metadata) UnmarshalYAML(node *yaml.Node) error { @@ -69,11 +78,11 @@ func (t *Metadata) UnmarshalYAML(node *yaml.Node) error { t.Name = raw.Name t.Description = raw.Description t.Features = raw.Features - parametersNode := findMetadataNode(node, "parameters") + parametersNode := findMappingValue(node, "parameters") parameters := raw.Parameters if len(parameters) == 0 && len(raw.Args) > 0 { logger.Warn("x-topo.args is deprecated; use x-topo.parameters instead") - parametersNode = findMetadataNode(node, "args") + parametersNode = findMappingValue(node, "args") parameters = raw.Args } t.Parameters = parseParametersInOrder(parametersNode, parameters) @@ -81,19 +90,54 @@ func (t *Metadata) UnmarshalYAML(node *yaml.Node) error { return nil } -func findMetadataNode(node *yaml.Node, key string) *yaml.Node { +func findMappingValue(node *yaml.Node, key string) *yaml.Node { + node = resolveAlias(node) + if node == nil || node.Kind != yaml.MappingNode { + return nil + } for i := 0; i < len(node.Content); i += 2 { if node.Content[i].Value == key { - return node.Content[i+1] + return resolveAlias(node.Content[i+1]) } } return nil } +func parseCurrentParameterValues(root *yaml.Node) map[string][]string { + values := make(map[string][]string) + services := findMappingValue(root, "services") + if services == nil { + return values + } + + for i := 0; i < len(services.Content); i += 2 { + build := findMappingValue(services.Content[i+1], "build") + args := findMappingValue(build, "args") + if args == nil { + continue + } + + switch args.Kind { + case yaml.MappingNode: + for j := 0; j < len(args.Content); j += 2 { + name := args.Content[j].Value + value := resolveAlias(args.Content[j+1]).Value + values[name] = append(values[name], value) + } + case yaml.SequenceNode: + for _, node := range args.Content { + name, value, _ := strings.Cut(resolveAlias(node).Value, "=") + values[name] = append(values[name], value) + } + } + } + + return values +} + func parseParametersInOrder(parametersNode *yaml.Node, parametersMap map[string]rawParameter) []Parameter { var result []Parameter - parametersNode = resolveAlias(parametersNode) if parametersNode == nil { return result } @@ -106,7 +150,6 @@ func parseParametersInOrder(parametersNode *yaml.Node, parametersMap map[string] Description: metadata.Description, Required: metadata.Required, Example: metadata.Example, - Default: metadata.Default, }) } } diff --git a/internal/project/project_test.go b/internal/project/project_test.go index 6dbb736d..f042854c 100644 --- a/internal/project/project_test.go +++ b/internal/project/project_test.go @@ -5,6 +5,7 @@ import ( "fmt" "os" "path/filepath" + "strings" "testing" "github.com/arm/topo/internal/arguments" @@ -66,6 +67,33 @@ services: assert.FileExists(t, composeFilePath) }) + t.Run("preserves current build arg values", func(t *testing.T) { + dir := t.TempDir() + destDir := filepath.Join(dir, "demo") + composeFileContents := `services: + app: + build: + args: + GREETING: ${GREETING} + app-2: + build: + args: + GREETING: "goodbye!" +x-topo: + parameters: + GREETING: + required: true +` + mockSource := mockSourceWithContent(t, composeFileContents) + provider := arguments.NewInteractiveProvider(strings.NewReader("\n"), &bytes.Buffer{}) + + err := project.Clone(destDir, mockSource, arguments.NewStrictProviderChain(provider)) + + require.NoError(t, err) + composeFilePath := filepath.Join(destDir, project.ComposeFilename) + assert.Equal(t, composeFileContents, testutil.RequireReadFile(t, composeFilePath)) + }) + t.Run("removes destination directory when parameter resolution fails", func(t *testing.T) { dir := t.TempDir() destDir := filepath.Join(dir, "demo") @@ -74,7 +102,7 @@ services: app: build: args: - GREETING: ${GREETING} + GREETING: "" x-topo: parameters: GREETING: @@ -159,4 +187,53 @@ x-topo: assert.YAMLEq(t, want, got) }) + + t.Run("rejects empty input for required parameters when any current value is empty", func(t *testing.T) { + dir := t.TempDir() + composeFileContents := `services: + configured: + build: + args: + FOO: current + empty: + build: + args: + FOO: "" +x-topo: + parameters: + FOO: + required: true + default: default +` + composeFilePath := filepath.Join(dir, project.ComposeFilename) + testutil.RequireWriteFile(t, composeFilePath, composeFileContents) + argProvider := arguments.NewStrictProviderChain(arguments.NewStaticProvider()) + + err := project.ResolveAndApplyArgs(composeFilePath, argProvider) + + require.ErrorContains(t, err, "missing value(s) for required parameters") + assert.Equal(t, composeFileContents, testutil.RequireReadFile(t, composeFilePath)) + }) + + t.Run("recognizes current values in sequence build args", func(t *testing.T) { + dir := t.TempDir() + composeFileContents := `services: + app: + build: + args: ["FOO=current"] +x-topo: + parameters: + FOO: + required: true + default: default +` + composeFilePath := filepath.Join(dir, project.ComposeFilename) + testutil.RequireWriteFile(t, composeFilePath, composeFileContents) + argProvider := arguments.NewStrictProviderChain(arguments.NewStaticProvider()) + + err := project.ResolveAndApplyArgs(composeFilePath, argProvider) + + require.NoError(t, err) + assert.Equal(t, composeFileContents, testutil.RequireReadFile(t, composeFilePath)) + }) } diff --git a/internal/project/resolution.go b/internal/project/resolution.go index eb1c045f..82d913e0 100644 --- a/internal/project/resolution.go +++ b/internal/project/resolution.go @@ -5,17 +5,23 @@ import ( ) func Resolve(p Project, argProvider arguments.Provider) ([]arguments.ResolvedArg, error) { - resolvedArgs, err := argProvider.Provide(castParameters(p.Metadata.Parameters)) + resolvedArgs, err := argProvider.Provide(castParameters(p.Metadata.Parameters, p.currentParameterValues)) if err != nil { return nil, err } return resolvedArgs, nil } -func castParameters(toCast []Parameter) []arguments.Arg { - casted := make([]arguments.Arg, len(toCast)) - for i, parameter := range toCast { - casted[i] = arguments.Arg(parameter) +func castParameters(parameters []Parameter, currentValues map[string][]string) []arguments.Arg { + casted := make([]arguments.Arg, len(parameters)) + for i, parameter := range parameters { + casted[i] = arguments.Arg{ + Name: parameter.Name, + Description: parameter.Description, + Required: parameter.Required, + Example: parameter.Example, + CurrentValues: currentValues[parameter.Name], + } } return casted }