diff --git a/docs/coverage/aws/ssm.md b/docs/coverage/aws/ssm.md index e09c0a4c1..bd5853e49 100644 --- a/docs/coverage/aws/ssm.md +++ b/docs/coverage/aws/ssm.md @@ -41,13 +41,24 @@ Documents is an OPTIONAL capability, discovered by type assertion. It covers | `UpdateDocument` | | | `UpdateDocumentDefaultVersion` | | +### ManagedNodes + +ManagedNodes is an OPTIONAL capability, discovered by type assertion. Every + +| Operation | Description | +| --- | --- | +| `DescribeInstanceInformation` | | + ### RunCommand RunCommand is an OPTIONAL capability, discovered by type assertion. | Operation | Description | | --- | --- | -| `GetCommandInvocation` | | +| `CancelCommand` | CancelCommand cancels the command on the given instances, or on all of | +| `GetCommandInvocation` | GetCommandInvocation reports one instance's run. pluginName selects a | +| `ListCommandInvocations` | ListCommandInvocations returns the matching invocations, newest command | +| `ListCommands` | ListCommands returns the matching commands, newest first. | | `SendCommand` | | ### ServiceSettings diff --git a/docs/coverage/coverage.json b/docs/coverage/coverage.json index 877eb82e7..3c3b4d761 100644 --- a/docs/coverage/coverage.json +++ b/docs/coverage/coverage.json @@ -12512,12 +12512,34 @@ } ] }, + { + "name": "ManagedNodes", + "doc": "ManagedNodes is an OPTIONAL capability, discovered by type assertion. Every", + "operations": [ + { + "name": "DescribeInstanceInformation" + } + ] + }, { "name": "RunCommand", "doc": "RunCommand is an OPTIONAL capability, discovered by type assertion.", "operations": [ { - "name": "GetCommandInvocation" + "name": "CancelCommand", + "doc": "CancelCommand cancels the command on the given instances, or on all of" + }, + { + "name": "GetCommandInvocation", + "doc": "GetCommandInvocation reports one instance's run. pluginName selects a" + }, + { + "name": "ListCommandInvocations", + "doc": "ListCommandInvocations returns the matching invocations, newest command" + }, + { + "name": "ListCommands", + "doc": "ListCommands returns the matching commands, newest first." }, { "name": "SendCommand" diff --git a/internal/settle/settle.go b/internal/settle/settle.go index d2c95f671..5a496f4ff 100644 --- a/internal/settle/settle.go +++ b/internal/settle/settle.go @@ -34,6 +34,9 @@ const ( DefaultExperimentInitiateSettle = 1 * time.Second // FIS experiment initiating->running + DefaultCommandDeliverySettle = 1 * time.Second // SSM Run Command invocation Pending->InProgress + DefaultCommandRunSettle = 5 * time.Second // SSM Run Command invocation InProgress->Success + DefaultCacheSettle = 2 * time.Second // ElastiCache/Redis/Memorystore creating->available DefaultCacheModifySettle = 1 * time.Second // cache modifying->available DefaultClusterSettle = 3 * time.Second // Redshift/MemoryDB/Bigtable creating->available diff --git a/providers/aws/aws.go b/providers/aws/aws.go index 6f4568be2..0eff2f8a7 100644 --- a/providers/aws/aws.go +++ b/providers/aws/aws.go @@ -447,6 +447,8 @@ func newProvider(o *config.Options, shared *GlobalServices) *Provider { p.SecretsManager.SetKMSCrypto(kmsCrypto) p.SSM.SetKMSCrypto(kmsCrypto) p.SSM.SetInstanceResolver(p.EC2) + // Run Command writes invocation output to the command's OutputS3BucketName. + p.SSM.SetOutputStore(p.S3) // ECS-registered container instances surface as managed EC2 instances, so // #159 (ECS) composes with #300 (EC2 managed-resource visibility). p.ECS.SetManagedInstanceLauncher(p.EC2) diff --git a/providers/aws/ssm/command_list.go b/providers/aws/ssm/command_list.go new file mode 100644 index 000000000..6c19ddcc6 --- /dev/null +++ b/providers/aws/ssm/command_list.go @@ -0,0 +1,436 @@ +package ssm + +import ( + "context" + "slices" + "sort" + "time" + + "github.com/stackshy/cloudemu/v2/errors" + ssmdriver "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" +) + +// Command filter keys and limits. +const ( + filterInvokedAfter = "InvokedAfter" + filterInvokedBefore = "InvokedBefore" + filterStatus = "Status" + filterExecutionStage = "ExecutionStage" + filterDocumentName = "DocumentName" + + stageExecuting = "Executing" + stageComplete = "Complete" + + maxCommandFilters = 5 + commandIDLength = 36 +) + +// commandStatusValues are the Status filter values ListCommands accepts. +func commandStatusValues() []string { + return []string{ + "Pending", "InProgress", "Success", "Cancelled", "Failed", "TimedOut", "AccessDenied", //nolint:misspell // SSM literal. + "DeliveryTimedOut", "ExecutionTimedOut", "Incomplete", "NoInstancesInTag", "LimitExceeded", + } +} + +// invocationStatusValues are the Status filter values ListCommandInvocations +// accepts. +func invocationStatusValues() []string { + return []string{ + "Pending", "InProgress", "Delayed", "Success", "Cancelled", "Failed", "TimedOut", "AccessDenied", //nolint:misspell // SSM literal. + "DeliveryTimedOut", "ExecutionTimedOut", "Undeliverable", "InvalidPlatform", "Terminated", + } +} + +// commandMatcher is a validated CommandQuery. +type commandMatcher struct { + q ssmdriver.CommandQuery + after *time.Time + before *time.Time + status string + stage string + doc string +} + +// newCommandMatcher validates q. forCommands selects the ListCommands rules; +// ListCommandInvocations does not take ExecutionStage. +func newCommandMatcher(q ssmdriver.CommandQuery, forCommands bool) (*commandMatcher, error) { + if q.CommandID != "" && len(q.CommandID) != commandIDLength { + return nil, ssmErrf(excValidation, errors.InvalidArgument, + "1 validation error detected: Value '%s' at 'commandId' failed to satisfy constraint: "+ + "Member must have length greater than or equal to 36", q.CommandID) + } + + if q.InstanceID != "" && !instanceIDPattern.MatchString(q.InstanceID) { + return nil, patternErr(q.InstanceID, "instanceId", `(^i-(\w{8}|\w{17})$)|(^mi-\w{17}$)`) + } + + if len(q.Filters) > maxCommandFilters { + return nil, ssmErrf(excValidation, errors.InvalidArgument, + "1 validation error detected: Value at 'filters' failed to satisfy constraint: "+ + "Member must have length less than or equal to %d", maxCommandFilters) + } + + mt := &commandMatcher{q: q} + + for _, f := range q.Filters { + if err := mt.add(f, forCommands); err != nil { + return nil, err + } + } + + return mt, nil +} + +func (mt *commandMatcher) add(f ssmdriver.CommandFilter, forCommands bool) error { + statuses := invocationStatusValues() + if forCommands { + statuses = commandStatusValues() + } + + switch f.Key { + case filterInvokedAfter: + return parseFilterTime(f.Value, &mt.after) + case filterInvokedBefore: + return parseFilterTime(f.Value, &mt.before) + case filterStatus: + if !slices.Contains(statuses, f.Value) { + return ssmErrf(excValidation, errors.InvalidArgument, "The filter value %s is not a valid Status.", f.Value) + } + + mt.status = f.Value + case filterExecutionStage: + return mt.setStage(f.Value, forCommands) + case filterDocumentName: + mt.doc = f.Value + default: + return ssmErrf(excInvalidFilterKey, errors.InvalidArgument, "The filter key %s is not valid.", f.Key) + } + + return nil +} + +// setStage applies an ExecutionStage filter, which only ListCommands takes. +func (mt *commandMatcher) setStage(v string, forCommands bool) error { + if !forCommands { + return ssmErrf(excInvalidFilterKey, errors.InvalidArgument, + "The ExecutionStage filter can't be used with ListCommandInvocations.") + } + + if v != stageExecuting && v != stageComplete { + return ssmErrf(excValidation, errors.InvalidArgument, + "The filter value %s is not a valid ExecutionStage. Use Executing or Complete.", v) + } + + mt.stage = v + + return nil +} + +// parseFilterTime reads an InvokedAfter/InvokedBefore value into *dst. +func parseFilterTime(v string, dst **time.Time) error { + t, err := time.Parse(time.RFC3339, v) + if err != nil { + return ssmErrf(excValidation, errors.InvalidArgument, + "The filter value %s is not a valid timestamp. Use the format 2024-07-07T00:00:00Z.", v) + } + + *dst = &t + + return nil +} + +// matchCommand checks the command-level criteria (everything but InstanceID). +func (mt *commandMatcher) matchCommand(cmd *ssmdriver.Command) bool { + switch { + case mt.q.CommandID != "" && cmd.CommandID != mt.q.CommandID: + return false + case mt.after != nil && cmd.RequestedDateTime.Before(*mt.after): + return false + case mt.before != nil && cmd.RequestedDateTime.After(*mt.before): + return false + case mt.doc != "" && cmd.DocumentName != mt.doc: + return false + } + + return true +} + +// matchStatus checks the Status and ExecutionStage filters. +func (mt *commandMatcher) matchStatus(status, details string, terminal bool) bool { + if mt.status != "" && mt.status != status && mt.status != details { + return false + } + + switch mt.stage { + case stageExecuting: + return !terminal + case stageComplete: + return terminal + default: + return true + } +} + +// sortedRecords returns every command record, newest first. The caller holds +// cmdMu. +func (m *Mock) sortedRecords() []*commandRecord { + all := m.commands.All() + out := make([]*commandRecord, 0, len(all)) + + for _, r := range all { + out = append(out, r) + } + + sort.Slice(out, func(i, j int) bool { + a, b := out[i].Command, out[j].Command + if !a.RequestedDateTime.Equal(b.RequestedDateTime) { + return a.RequestedDateTime.After(b.RequestedDateTime) + } + + return a.CommandID < b.CommandID + }) + + return out +} + +func (r *commandRecord) hasInstance(id string) bool { + for _, inv := range r.Invocations { + if inv.InstanceID == id { + return true + } + } + + return false +} + +// ListCommands returns the matching commands, newest first. +func (m *Mock) ListCommands(ctx context.Context, q ssmdriver.CommandQuery) ([]ssmdriver.Command, error) { + mt, err := newCommandMatcher(q, true) + if err != nil { + return nil, err + } + + now := m.opts.Clock.Now() + m.flushOutputs(ctx, now) + + m.cmdMu.RLock() + defer m.cmdMu.RUnlock() + + out := make([]ssmdriver.Command, 0) + + for _, r := range m.sortedRecords() { + if !mt.matchCommand(&r.Command) || (q.InstanceID != "" && !r.hasInstance(q.InstanceID)) { + continue + } + + cmd, _ := r.observe(now) + if mt.matchStatus(cmd.Status, cmd.StatusDetails, int(cmd.CompletedCount) == int(cmd.TargetCount)) { + out = append(out, cmd) + } + } + + return out, nil +} + +// ListCommandInvocations returns the matching invocations, newest command +// first. +func (m *Mock) ListCommandInvocations( + ctx context.Context, q ssmdriver.CommandQuery, details bool, +) ([]ssmdriver.CommandInvocation, error) { + mt, err := newCommandMatcher(q, false) + if err != nil { + return nil, err + } + + now := m.opts.Clock.Now() + m.flushOutputs(ctx, now) + + m.cmdMu.RLock() + defer m.cmdMu.RUnlock() + + out := make([]ssmdriver.CommandInvocation, 0) + + for _, r := range m.sortedRecords() { + if !mt.matchCommand(&r.Command) { + continue + } + + _, states := r.observe(now) + + for i, inv := range r.Invocations { + s := states[i] + if (q.InstanceID != "" && inv.InstanceID != q.InstanceID) || !mt.matchStatus(s.status, s.details, s.terminal()) { + continue + } + + ci := r.render(i, &s) + if !details { + ci.Plugins = nil + } + + out = append(out, ci) + } + } + + return out, nil +} + +// GetCommandInvocation returns one instance's run, narrowed to pluginName +// when it is set. +// +// An unknown pair is InvocationDoesNotExist rather than a fabricated success: +// a caller polling a command it never sent has a real bug, and answering +// "Success" would bury it. +func (m *Mock) GetCommandInvocation( + ctx context.Context, commandID, instanceID, pluginName string, +) (*ssmdriver.CommandInvocation, error) { + if len(commandID) != commandIDLength { + return nil, ssmErrf(excValidation, errors.InvalidArgument, + "1 validation error detected: Value '%s' at 'commandId' failed to satisfy constraint: "+ + "Member must have length greater than or equal to 36", commandID) + } + + now := m.opts.Clock.Now() + m.flushOutputs(ctx, now) + + m.cmdMu.RLock() + defer m.cmdMu.RUnlock() + + r, ok := m.commands.Get(commandID) + if !ok { + return nil, invocationMissing(commandID, instanceID) + } + + for i, inv := range r.Invocations { + if inv.InstanceID != instanceID { + continue + } + + _, states := r.observe(now) + ci := r.render(i, &states[i]) + + if err := narrowToPlugin(&ci, pluginName); err != nil { + return nil, err + } + + return &ci, nil + } + + return nil, invocationMissing(commandID, instanceID) +} + +func invocationMissing(commandID, instanceID string) error { + return ssmErrf(excInvocationDoesNotExist, errors.NotFound, + "No invocation of command %s on instance %s.", commandID, instanceID) +} + +// narrowToPlugin fills the invocation's plugin fields from the named plugin, +// or from the first one when none is named. +func narrowToPlugin(ci *ssmdriver.CommandInvocation, name string) error { + if len(ci.Plugins) == 0 { + if name != "" { + return ssmErrf(excInvalidPluginName, errors.InvalidArgument, "The plugin name %s is not valid.", name) + } + + return nil + } + + p := &ci.Plugins[0] + + if name != "" { + idx := slices.IndexFunc(ci.Plugins, func(c ssmdriver.CommandPlugin) bool { return c.Name == name }) + if idx < 0 { + return ssmErrf(excInvalidPluginName, errors.InvalidArgument, "The plugin name %s is not valid.", name) + } + + p = &ci.Plugins[idx] + ci.Status, ci.StatusDetails = p.Status, p.StatusDetails + } + + ci.PluginName = p.Name + ci.ResponseCode = p.ResponseCode + ci.ExecutionStartTime, ci.ExecutionEndTime = p.ResponseStartDateTime, p.ResponseFinishDateTime + ci.Stdout = p.Output + ci.StandardOutputURL, ci.StandardErrorURL = p.StandardOutputURL, p.StandardErrorURL + + return nil +} + +// render builds invocation i from its observed state. +func (r *commandRecord) render(i int, s *invocationState) ssmdriver.CommandInvocation { + cmd := &r.Command + inv := r.Invocations[i] + + ci := ssmdriver.CommandInvocation{ + CommandID: cmd.CommandID, InstanceID: inv.InstanceID, Comment: cmd.Comment, + DocumentName: cmd.DocumentName, DocumentVersion: cmd.DocumentVersion, RequestedDateTime: cmd.RequestedDateTime, + Status: s.status, StatusDetails: s.details, ResponseCode: s.code, + ExecutionStartTime: s.start, ExecutionEndTime: s.end, ServiceRole: cmd.ServiceRole, + Notification: cmd.Notification, CloudWatchOutput: cmd.CloudWatchOutput, + Plugins: make([]ssmdriver.CommandPlugin, 0, len(r.Steps)), + } + + pluginStatus := s.status + if pluginStatus == ssmdriver.CommandDelayed { + pluginStatus = ssmdriver.CommandPending + } + + for _, step := range r.Steps { + p := ssmdriver.CommandPlugin{ + Name: step.Name, Status: pluginStatus, StatusDetails: s.details, ResponseCode: s.code, + ResponseStartDateTime: s.start, ResponseFinishDateTime: s.end, + } + + if cmd.OutputS3BucketName != "" { + p.OutputS3Region, p.OutputS3BucketName = cmd.OutputS3Region, cmd.OutputS3BucketName + p.OutputS3KeyPrefix = r.outputFolder(inv.InstanceID, step) + p.StandardOutputURL = r.outputURL(p.OutputS3KeyPrefix + "/stdout") + p.StandardErrorURL = r.outputURL(p.OutputS3KeyPrefix + "/stderr") + } + + ci.Plugins = append(ci.Plugins, p) + } + + // An invocation carries the output URLs only when the document has one + // plugin. + if len(ci.Plugins) == 1 { + ci.StandardOutputURL, ci.StandardErrorURL = ci.Plugins[0].StandardOutputURL, ci.Plugins[0].StandardErrorURL + } + + return ci +} + +// CancelCommand cancels the unfinished invocations of a command, on the given +// instances or on all of them. Finished invocations keep their result. +func (m *Mock) CancelCommand(_ context.Context, commandID string, instanceIDs []string) error { + m.cmdMu.Lock() + defer m.cmdMu.Unlock() + + r, ok := m.commands.Get(commandID) + if !ok { + return ssmErrf(excInvalidCommandID, errors.NotFound, "The command ID %s is not valid.", commandID) + } + + for _, id := range instanceIDs { + if !r.hasInstance(id) { + return ssmErrf(excInvalidInstanceID, errors.NotFound, + "Instance %s is not a target of command %s.", id, commandID) + } + } + + now := m.opts.Clock.Now() + _, states := r.observe(now) + + for i, inv := range r.Invocations { + if len(instanceIDs) > 0 && !slices.Contains(instanceIDs, inv.InstanceID) { + continue + } + + if !states[i].terminal() { + inv.CancelledAt = now + } + } + + return nil +} diff --git a/providers/aws/ssm/command_output.go b/providers/aws/ssm/command_output.go new file mode 100644 index 000000000..347694fc5 --- /dev/null +++ b/providers/aws/ssm/command_output.go @@ -0,0 +1,97 @@ +package ssm + +import ( + "context" + "path" + "strings" + "time" + + ssmdriver "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" +) + +// OutputStore is the slice of the S3 mock Run Command writes output through. +type OutputStore interface { + PutObject(ctx context.Context, bucket, key string, data []byte, contentType string, metadata map[string]string) error +} + +// SetOutputStore wires S3 in so a command that names OutputS3BucketName gets +// its output written there. Without it the bucket is only echoed back. +func (m *Mock) SetOutputStore(s OutputStore) { + m.outputStore = s +} + +// outputFolder is the S3 folder of one plugin's output, as the SSM agent +// lays it out: ////. +func (r *commandRecord) outputFolder(instanceID string, step commandStep) string { + plugin := strings.ReplaceAll(step.Action, ":", "") + + stepDir := step.Name + if step.Legacy { + stepDir = "0." + plugin + } + + parts := []string{r.Command.CommandID, instanceID, plugin, stepDir} + if prefix := strings.Trim(r.Command.OutputS3KeyPrefix, "/"); prefix != "" { + parts = append([]string{prefix}, parts...) + } + + return path.Join(parts...) +} + +// outputURL is the virtual S3 URL SSM reports for an output object. +func (r *commandRecord) outputURL(key string) string { + return "https://s3." + r.Command.OutputS3Region + ".amazonaws.com/" + r.Command.OutputS3BucketName + "/" + key +} + +// outputWrite is one object waiting to be written. +type outputWrite struct { + bucket, key string +} + +// flushOutputs writes the output of every invocation that finished by now +// and has not been written yet. The objects are written after cmdMu is +// released, so an S3 notification target cannot deadlock on this mock. A +// failed write (the bucket is missing, say) is dropped: real SSM reports the +// command result regardless. +func (m *Mock) flushOutputs(ctx context.Context, now time.Time) bool { + if m.outputStore == nil { + return false + } + + var writes []outputWrite + + m.cmdMu.Lock() + + for _, r := range m.commands.All() { + if r.Command.OutputS3BucketName == "" { + continue + } + + for i, inv := range r.Invocations { + if inv.OutputWritten { + continue + } + + s := r.invocation(i, now) + if s.status != ssmdriver.CommandSuccess && s.details != detailsExecutionTimedOut { + continue + } + + inv.OutputWritten = true + + for _, step := range r.Steps { + writes = append(writes, outputWrite{ + bucket: r.Command.OutputS3BucketName, key: r.outputFolder(inv.InstanceID, step) + "/stdout", + }) + } + } + } + + m.cmdMu.Unlock() + + for _, w := range writes { + _ = m.outputStore.PutObject(ctx, w.bucket, w.key, []byte{}, "text/plain", nil) + } + + return len(writes) > 0 +} diff --git a/providers/aws/ssm/command_params.go b/providers/aws/ssm/command_params.go new file mode 100644 index 000000000..eaa641775 --- /dev/null +++ b/providers/aws/ssm/command_params.go @@ -0,0 +1,296 @@ +package ssm + +import ( + "encoding/json" + "regexp" + "slices" + "sort" + "strconv" + "strings" + + "github.com/stackshy/cloudemu/v2/errors" +) + +// Document parameter types a Command document can declare. +const ( + paramStringList = "StringList" + paramInteger = "Integer" + paramBoolean = "Boolean" + paramStringMap = "StringMap" + paramMapList = "MapList" +) + +// excInvalidParameters is what SendCommand returns for parameters that do not +// fit the document. +const excInvalidParameters = "InvalidParameters" + +// paramSpec is one declared parameter with the constraints SendCommand +// checks. A negative bound is unset. +type paramSpec struct { + typ string + defaults []string + hasDefault bool + allowedValues []string + allowedPattern string + minItems int + maxItems int + minChars int + maxChars int +} + +// commandStep is one plugin a Command document runs. Name is the step name +// (schema 2.x) or the plugin name (schema 1.2). +type commandStep struct { + Name string `json:"name"` + Action string `json:"action"` + // Legacy is set for a schema 1.2 runtimeConfig plugin, whose output sits + // under a "0." folder. + Legacy bool `json:"legacy,omitempty"` +} + +// parameterSpecs reads the declared parameters and their constraints. +func parameterSpecs(v any) map[string]paramSpec { + params, ok := v.(map[string]any) + if !ok { + return nil + } + + out := make(map[string]paramSpec, len(params)) + + for name, raw := range params { + spec, _ := raw.(map[string]any) + p := paramSpec{ + typ: scalarString(spec["type"]), + allowedPattern: scalarString(spec["allowedPattern"]), + minItems: intField(spec, "minItems"), + maxItems: intField(spec, "maxItems"), + minChars: intField(spec, "minChars"), + maxChars: intField(spec, "maxChars"), + } + + if def, ok := spec["default"]; ok { + p.hasDefault = true + p.defaults = paramValues(def) + } + + if allowed, ok := spec["allowedValues"].([]any); ok { + for _, a := range allowed { + p.allowedValues = append(p.allowedValues, scalarString(a)) + } + } + + out[name] = p + } + + return out +} + +// intField reads a numeric constraint, or -1 when it is absent. +func intField(spec map[string]any, key string) int { + n, err := strconv.Atoi(scalarString(spec[key])) + if err != nil { + return -1 + } + + return n +} + +// paramValues renders a default as the value list SendCommand would carry. +func paramValues(v any) []string { + switch t := v.(type) { + case nil: + return nil + case []any: + out := make([]string, 0, len(t)) + for _, e := range t { + out = append(out, jsonScalar(e)) + } + + return out + default: + return []string{jsonScalar(t)} + } +} + +// jsonScalar renders a scalar as text and anything else as JSON. +func jsonScalar(v any) string { + if s := scalarString(v); s != "" || v == "" { + return s + } + + b, _ := json.Marshal(v) + + return string(b) +} + +// documentSteps lists the plugins of a Command document: the runtimeConfig +// plugins of schema 1.2 or the mainSteps of schema 2.x. +func documentSteps(root map[string]any) []commandStep { + var out []commandStep + + if rc, ok := root["runtimeConfig"].(map[string]any); ok { + names := make([]string, 0, len(rc)) + for name := range rc { + names = append(names, name) + } + + sort.Strings(names) + + for _, name := range names { + out = append(out, commandStep{Name: name, Action: name, Legacy: true}) + } + } + + steps, _ := root["mainSteps"].([]any) + for _, s := range steps { + if step, ok := s.(map[string]any); ok { + out = append(out, commandStep{Name: scalarString(step["name"]), Action: scalarString(step["action"])}) + } + } + + return out +} + +// validateParameters checks SendCommand parameters against the document's +// declared parameters and returns them with the defaults filled in. +func validateParameters(doc string, specs map[string]paramSpec, given map[string][]string) (map[string][]string, error) { + undeclared := make([]string, 0) + + for name := range given { + if _, ok := specs[name]; !ok { + undeclared = append(undeclared, name) + } + } + + if len(undeclared) > 0 { + sort.Strings(undeclared) + + return nil, ssmErrf(excInvalidParameters, errors.InvalidArgument, + "Parameters %v are not defined in document %s.", undeclared, doc) + } + + names := make([]string, 0, len(specs)) + for name := range specs { + names = append(names, name) + } + + sort.Strings(names) + + resolved := make(map[string][]string, len(specs)) + + for _, name := range names { + spec := specs[name] + + values, ok := given[name] + if !ok { + if !spec.hasDefault { + return nil, ssmErrf(excInvalidParameters, errors.InvalidArgument, + "Parameter \"%s\" is required by document %s but was not supplied.", name, doc) + } + + resolved[name] = spec.defaults + + continue + } + + if err := spec.check(name, values); err != nil { + return nil, err + } + + resolved[name] = values + } + + return resolved, nil +} + +// check validates the values supplied for one parameter. +func (p *paramSpec) check(name string, values []string) error { + if err := p.checkShape(name, values); err != nil { + return err + } + + for _, v := range values { + if err := p.checkValue(name, v); err != nil { + return err + } + } + + return nil +} + +// checkShape checks the value count and each value's type. +func (p *paramSpec) checkShape(name string, values []string) error { + if p.typ == paramStringList { + if p.minItems >= 0 && len(values) < p.minItems { + return invalidParam(name, "needs at least %d values", p.minItems) + } + + if p.maxItems >= 0 && len(values) > p.maxItems { + return invalidParam(name, "accepts at most %d values", p.maxItems) + } + + return nil + } + + if len(values) != 1 { + return invalidParam(name, "of type %s takes exactly one value", p.typ) + } + + return p.checkScalar(name, values[0]) +} + +// checkScalar checks that a single value parses as the parameter's type. +func (p *paramSpec) checkScalar(name, v string) error { + switch p.typ { + case paramInteger: + if _, err := strconv.Atoi(v); err != nil { + return invalidParam(name, "must be an Integer, got %q", v) + } + case paramBoolean: + if v != "true" && v != "false" { + return invalidParam(name, "must be a Boolean, got %q", v) + } + case paramStringMap: + var m map[string]any + if json.Unmarshal([]byte(v), &m) != nil { + return invalidParam(name, "must be a StringMap (a JSON object)") + } + case paramMapList: + var l []map[string]any + if json.Unmarshal([]byte(v), &l) != nil { + return invalidParam(name, "must be a MapList (a JSON list of objects)") + } + } + + return nil +} + +// checkValue checks one value against allowedValues, allowedPattern and the +// character limits. +func (p *paramSpec) checkValue(name, v string) error { + if len(p.allowedValues) > 0 && !slices.Contains(p.allowedValues, v) { + return invalidParam(name, "value %q is not one of the allowed values [%s]", v, strings.Join(p.allowedValues, ", ")) + } + + if p.allowedPattern != "" { + // The pattern must match the whole value. One Go cannot compile is + // not checked rather than rejecting every send. + if re, err := regexp.Compile(`^(?:` + p.allowedPattern + `)$`); err == nil && !re.MatchString(v) { + return invalidParam(name, "value %q does not match the allowed pattern %s", v, p.allowedPattern) + } + } + + if p.minChars >= 0 && len(v) < p.minChars { + return invalidParam(name, "value must have at least %d characters", p.minChars) + } + + if p.maxChars >= 0 && len(v) > p.maxChars { + return invalidParam(name, "value must have at most %d characters", p.maxChars) + } + + return nil +} + +func invalidParam(name, format string, args ...any) error { + return ssmErrf(excInvalidParameters, errors.InvalidArgument, "Parameter \"%s\" "+format+".", append([]any{name}, args...)...) +} diff --git a/providers/aws/ssm/command_status.go b/providers/aws/ssm/command_status.go new file mode 100644 index 000000000..b98c6b75c --- /dev/null +++ b/providers/aws/ssm/command_status.go @@ -0,0 +1,255 @@ +package ssm + +import ( + "strconv" + "strings" + "time" + + ssmdriver "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" +) + +// Command lifecycle. +// +// A send stores a commandRecord and nothing else changes on its own: the +// status of every invocation is worked out from the clock at read time, so +// reads are stable for a given clock and survive snapshot/restore. Only +// CancelCommand writes a status (the cancel time). +// +// Invocations run in batches of MaxConcurrency. Each batch spends the delivery +// window Pending and the run window InProgress before it reports Success, and +// the next batch starts when it ends. Both windows are zero unless async settle +// is on, so by default a command is finished when SendCommand returns. +// +// An instance a tag Target selected while it was not running stays Delayed +// until TimeoutSeconds passes and then times out (DeliveryTimedOut). An +// executionTimeout parameter shorter than the run window makes the invocation +// time out (ExecutionTimedOut). + +// StatusDetails values beyond the plain statuses. +const ( + detailsDeliveryTimedOut = "DeliveryTimedOut" + detailsExecutionTimedOut = "ExecutionTimedOut" + detailsIncomplete = "Incomplete" + detailsNoInstancesInTag = "NoInstancesInTag" +) + +// noResponse is the ResponseCode of an invocation that has not finished. +const noResponse = -1 + +// commandRecord is one stored send. Fields are exported for the snapshot and +// guarded by Mock.cmdMu. +type commandRecord struct { + // Command holds what the send recorded. Status and counts are filled in + // on each read. + Command ssmdriver.Command `json:"command"` + Steps []commandStep `json:"steps,omitempty"` + // ExecutionTimeout is the executionTimeout parameter in seconds, or 0. + ExecutionTimeout int `json:"executionTimeout,omitempty"` + // Delivery and Run are the settle windows in force when it was sent. + Delivery time.Duration `json:"delivery,omitempty"` + Run time.Duration `json:"run,omitempty"` + Invocations []*invocationRecord `json:"invocations"` +} + +// invocationRecord is the stored part of one invocation. +type invocationRecord struct { + InstanceID string `json:"instanceId"` + // Offline is set when the instance was not running at send time. + Offline bool `json:"offline,omitempty"` + CancelledAt time.Time `json:"cancelledAt,omitzero"` //nolint:misspell // SSM spells the status Cancelled. + // OutputWritten is set once the output has gone to S3. + OutputWritten bool `json:"outputWritten,omitempty"` +} + +// invocationState is an invocation observed at one instant. +type invocationState struct { + status string + details string + code int32 + start time.Time + end time.Time +} + +// terminal reports whether the invocation can no longer change. +func (s *invocationState) terminal() bool { + switch s.status { + case ssmdriver.CommandPending, ssmdriver.CommandInProgress, ssmdriver.CommandDelayed: + return false + default: + return true + } +} + +// batchSize is how many invocations run at once under MaxConcurrency. +func (r *commandRecord) batchSize() int { + n := len(r.Invocations) + + return max(1, countOrPercent(r.Command.MaxConcurrency, n, true)) +} + +// maxErrors is the error count the command tolerates before it fails. +func (r *commandRecord) maxErrors() int { + return countOrPercent(r.Command.MaxErrors, len(r.Invocations), false) +} + +// countOrPercent reads "10" or "10%" of total. A percentage of concurrency +// rounds up, one of errors rounds down. +func countOrPercent(v string, total int, roundUp bool) int { + if p, ok := strings.CutSuffix(v, "%"); ok { + n, _ := strconv.Atoi(p) + if roundUp { + return (n*total + percent - 1) / percent + } + + return n * total / percent + } + + n, _ := strconv.Atoi(v) + + return n +} + +func stateOf(status, details string) invocationState { + return invocationState{status: status, details: details, code: noResponse} +} + +// invocation observes invocation i at now. +func (r *commandRecord) invocation(i int, now time.Time) invocationState { + inv := r.Invocations[i] + if inv.Offline { + return r.offlineInvocation(inv, now) + } + + batch := time.Duration(i / r.batchSize()) + start := r.Command.RequestedDateTime.Add(r.Delivery + batch*r.Run) + finish := start.Add(r.Run) + done := invocationState{status: ssmdriver.CommandSuccess, details: ssmdriver.CommandSuccess, start: start, end: finish} + + if limit := time.Duration(r.ExecutionTimeout) * time.Second; r.ExecutionTimeout > 0 && limit < r.Run { + finish = start.Add(limit) + done = stateOf(ssmdriver.CommandTimedOut, detailsExecutionTimedOut) + done.start, done.end = start, finish + } + + switch { + case !inv.CancelledAt.IsZero() && inv.CancelledAt.Before(finish): + s := stateOf(ssmdriver.CommandCancelled, ssmdriver.CommandCancelled) + if !inv.CancelledAt.Before(start) { + s.start, s.end = start, inv.CancelledAt + } + + return s + case now.Before(start): + return stateOf(ssmdriver.CommandPending, ssmdriver.CommandPending) + case now.Before(finish): + s := stateOf(ssmdriver.CommandInProgress, ssmdriver.CommandInProgress) + s.start = start + + return s + default: + return done + } +} + +// offlineInvocation observes an invocation whose instance was not running: +// it waits for delivery until TimeoutSeconds passes. +func (r *commandRecord) offlineInvocation(inv *invocationRecord, now time.Time) invocationState { + expires := r.Command.RequestedDateTime.Add(time.Duration(r.Command.TimeoutSeconds) * time.Second) + + switch { + case !inv.CancelledAt.IsZero() && inv.CancelledAt.Before(expires): + return stateOf(ssmdriver.CommandCancelled, ssmdriver.CommandCancelled) + case now.Before(expires): + return stateOf(ssmdriver.CommandDelayed, ssmdriver.CommandDelayed) + default: + return stateOf(ssmdriver.CommandTimedOut, detailsDeliveryTimedOut) + } +} + +// tally counts the invocation outcomes a command status is built from. +type tally struct { + total, completed, started int + succeeded, canceled, errs, deliveryTimeouts, executionTimeouts int +} + +func (t *tally) add(s *invocationState) { + t.total++ + + switch { + case s.status == ssmdriver.CommandSuccess: + t.succeeded++ + case s.status == ssmdriver.CommandCancelled: + t.canceled++ + case s.details == detailsDeliveryTimedOut: + t.deliveryTimeouts++ + case s.details == detailsExecutionTimedOut: + t.executionTimeouts++ + t.errs++ + case s.status == ssmdriver.CommandFailed: + t.errs++ + } + + if s.terminal() { + t.completed++ + } + + if s.status != ssmdriver.CommandPending && s.status != ssmdriver.CommandDelayed { + t.started++ + } +} + +// status is the command Status and StatusDetails for the tallied outcomes. +func (t *tally) status(maxErrors int) (status, details string) { + switch { + case t.total == 0: + return ssmdriver.CommandSuccess, detailsNoInstancesInTag + case t.completed < t.total && t.started > 0: + return ssmdriver.CommandInProgress, ssmdriver.CommandInProgress + case t.completed < t.total: + return ssmdriver.CommandPending, ssmdriver.CommandPending + default: + return t.finalStatus(maxErrors) + } +} + +// finalStatus is the status of a command whose invocations have all finished. +func (t *tally) finalStatus(maxErrors int) (status, details string) { + switch { + case t.canceled > 0: + return ssmdriver.CommandCancelled, ssmdriver.CommandCancelled + case t.succeeded == t.total: + return ssmdriver.CommandSuccess, ssmdriver.CommandSuccess + case t.errs > maxErrors && t.executionTimeouts == t.errs: + return ssmdriver.CommandTimedOut, detailsExecutionTimedOut + case t.errs > maxErrors: + return ssmdriver.CommandFailed, ssmdriver.CommandFailed + case t.succeeded == 0 && t.deliveryTimeouts == t.total: + return ssmdriver.CommandTimedOut, detailsDeliveryTimedOut + default: + return ssmdriver.CommandFailed, detailsIncomplete + } +} + +// observe returns the command with its status and counts at now, and every +// invocation's state. +func (r *commandRecord) observe(now time.Time) (ssmdriver.Command, []invocationState) { + cmd := r.Command + states := make([]invocationState, len(r.Invocations)) + + var t tally + + for i := range r.Invocations { + states[i] = r.invocation(i, now) + t.add(&states[i]) + } + + // Counts are bounded by the target count: at most 50 explicit ids plus + // the tag matches. + cmd.TargetCount = int32(t.total) //nolint:gosec // bounded, see above. + cmd.CompletedCount = int32(t.completed) //nolint:gosec // bounded, see above. + cmd.ErrorCount = int32(t.errs) //nolint:gosec // bounded, see above. + cmd.DeliveryTimedOutCount = int32(t.deliveryTimeouts) //nolint:gosec // bounded, see above. + cmd.Status, cmd.StatusDetails = t.status(r.maxErrors()) + + return cmd, states +} diff --git a/providers/aws/ssm/document_content.go b/providers/aws/ssm/document_content.go index beba61825..9597ca51e 100644 --- a/providers/aws/ssm/document_content.go +++ b/providers/aws/ssm/document_content.go @@ -34,6 +34,11 @@ type contentMeta struct { parameters []ssmdriver.DocumentParameter platformTypes []string hash string + // paramSpecs holds the declared parameters with the constraints + // SendCommand checks, keyed by name. + paramSpecs map[string]paramSpec + // steps lists the plugins a Command document runs, in order. + steps []commandStep } // supportedSchemas lists the schemaVersion values each document type accepts. @@ -80,6 +85,8 @@ func parseContent(content, format, docType string) (*contentMeta, error) { meta.description, _ = root["description"].(string) meta.parameters = documentParameters(root["parameters"]) + meta.paramSpecs = parameterSpecs(root["parameters"]) + meta.steps = documentSteps(root) meta.platformTypes = platformTypes(root, docType) return meta, nil diff --git a/providers/aws/ssm/document_validation_test.go b/providers/aws/ssm/document_validation_test.go index ee6769531..c9adc4064 100644 --- a/providers/aws/ssm/document_validation_test.go +++ b/providers/aws/ssm/document_validation_test.go @@ -2,6 +2,7 @@ package ssm_test import ( "context" + stderrors "errors" "strings" "testing" @@ -17,16 +18,29 @@ func TestSendCommandAcceptsAWSOwnedNames(t *testing.T) { "arn:aws:ssm:us-east-1::document/AWS-RunInspecChecks", } + // A catalog document checks its required parameters; a name the catalog + // lacks is accepted as it stands. Either way the name resolves. for _, doc := range names { - if _, err := m.SendCommand(context.Background(), ssmdriver.CommandConfig{ - InstanceIDs: []string{"i-0123"}, DocumentName: doc, - }); err != nil { + _, err := m.SendCommand(context.Background(), ssmdriver.CommandConfig{ + InstanceIDs: []string{"i-0123456789abcdef0"}, DocumentName: doc, + }) + + if err == nil { + continue + } + + var ex interface{ SSMException() (string, int) } + if !stderrors.As(err, &ex) { + t.Fatalf("SendCommand %s: %v", doc, err) + } + + if name, _ := ex.SSMException(); name != "InvalidParameters" { t.Errorf("SendCommand %s: %v", doc, err) } } _, err := m.SendCommand(context.Background(), ssmdriver.CommandConfig{ - InstanceIDs: []string{"i-0123"}, DocumentName: "my-missing-doc", + InstanceIDs: []string{"i-0123456789abcdef0"}, DocumentName: "my-missing-doc", }) wantException(t, err, "InvalidDocument") diff --git a/providers/aws/ssm/documents_test.go b/providers/aws/ssm/documents_test.go index 74cc9f8c9..dc2cb0332 100644 --- a/providers/aws/ssm/documents_test.go +++ b/providers/aws/ssm/documents_test.go @@ -419,14 +419,19 @@ func TestSendCommandResolvesDocument(t *testing.T) { ctx := context.Background() m := newMock() send := func(doc string) error { - _, err := m.SendCommand(ctx, ssmdriver.CommandConfig{InstanceIDs: []string{"i-0123"}, DocumentName: doc}) + _, err := m.SendCommand(ctx, ssmdriver.CommandConfig{InstanceIDs: []string{"i-0123456789abcdef0"}, DocumentName: doc}) return err } wantException(t, send("No-Such-Doc"), "InvalidDocument") wantException(t, send("AWS-StopEC2Instance"), "InvalidDocument") - if err := send("AWS-RunShellScript"); err != nil { + wantException(t, send("AWS-RunShellScript"), "InvalidParameters") + + if _, err := m.SendCommand(ctx, ssmdriver.CommandConfig{ + InstanceIDs: []string{"i-0123456789abcdef0"}, DocumentName: "AWS-RunShellScript", + Parameters: map[string][]string{"commands": {"uptime"}}, + }); err != nil { t.Fatalf("send AWS-RunShellScript: %v", err) } diff --git a/providers/aws/ssm/driver/driver.go b/providers/aws/ssm/driver/driver.go index 94ee164b1..0a314d36c 100644 --- a/providers/aws/ssm/driver/driver.go +++ b/providers/aws/ssm/driver/driver.go @@ -1,7 +1,7 @@ // Package driver defines the AWS-native Systems Manager families that have no -// portable counterpart: Run Command and Documents. Parameter Store stays in -// services/parameterstore/driver because Azure App Configuration and GCP -// Secret Manager share its shape. +// portable counterpart: Run Command, managed nodes and Documents. Parameter +// Store stays in services/parameterstore/driver because Azure App +// Configuration and GCP Secret Manager share its shape. // // Each family is an optional capability the SSM wire handler discovers by type // assertion on the configured Parameter Store driver. @@ -12,50 +12,6 @@ import ( "time" ) -// CommandInvocation is the result of a Run Command execution on one instance. -type CommandInvocation struct { - CommandID string - InstanceID string - DocumentName string - Status string - ResponseCode int32 - Stdout string - Stderr string -} - -// CommandTarget identifies managed nodes by a Key/Values criterion, e.g. -// {Key: "tag:Name", Values: ["web"]}. It mirrors the SSM Target shape and is an -// alternative to listing InstanceIDs explicitly. -type CommandTarget struct { - Key string - Values []string -} - -// CommandConfig describes a Run Command send. Either InstanceIDs or Targets -// (or both) must be supplied; Targets select managed nodes by tag/attribute. -type CommandConfig struct { - InstanceIDs []string - Targets []CommandTarget - DocumentName string - Comment string - Parameters map[string][]string -} - -// RunCommand is an OPTIONAL capability, discovered by type assertion. -// -// Targets are validated: sending to an instance that does not exist is -// InvalidInstanceId, as it is against the real service. -// -// IMPORTANT: an emulated instance has no guest operating system, so nothing -// executes. Invocations report success and empty output. This exercises a -// caller's send/poll orchestration (that it waits for a terminal status, reads -// the response code, and handles failure) but it does NOT validate the script -// itself. A caller whose bootstrap script is wrong will still see success here. -type RunCommand interface { - SendCommand(ctx context.Context, cfg CommandConfig) (string, error) - GetCommandInvocation(ctx context.Context, commandID, instanceID string) (*CommandInvocation, error) -} - // Document formats, statuses, owners and hash types, as the SSM API spells them. const ( FormatJSON = "JSON" diff --git a/providers/aws/ssm/driver/run_command.go b/providers/aws/ssm/driver/run_command.go new file mode 100644 index 000000000..b97c11b55 --- /dev/null +++ b/providers/aws/ssm/driver/run_command.go @@ -0,0 +1,209 @@ +package driver + +import ( + "context" + "time" +) + +// Command, invocation and plugin statuses, as the SSM API spells them. +const ( + CommandPending = "Pending" + CommandInProgress = "InProgress" + CommandDelayed = "Delayed" + CommandSuccess = "Success" + CommandCancelled = "Cancelled" //nolint:misspell // SSM API literal. + CommandFailed = "Failed" + CommandTimedOut = "TimedOut" +) + +// CommandTarget identifies managed nodes by a Key/Values criterion, e.g. +// {Key: "tag:Name", Values: ["web"]}. It mirrors the SSM Target shape and is an +// alternative to listing InstanceIDs explicitly. +type CommandTarget struct { + Key string + Values []string +} + +// NotificationConfig is where Run Command sends status notifications. It is +// recorded and echoed back. +type NotificationConfig struct { + NotificationArn string + NotificationEvents []string + NotificationType string +} + +// CloudWatchOutputConfig names the log group command output goes to. It is +// recorded and echoed back. +type CloudWatchOutputConfig struct { + LogGroupName string + OutputEnabled bool +} + +// CommandConfig describes a Run Command send. Either InstanceIDs or Targets +// (or both) must be supplied; Targets select managed nodes by tag/attribute. +// Empty MaxConcurrency, MaxErrors and DocumentVersion take the service +// defaults (50, 0 and $DEFAULT), and a zero TimeoutSeconds is 3600. +type CommandConfig struct { + InstanceIDs []string + Targets []CommandTarget + DocumentName string + DocumentVersion string + Comment string + Parameters map[string][]string + TimeoutSeconds int32 + MaxConcurrency string + MaxErrors string + OutputS3Region string + OutputS3BucketName string + OutputS3KeyPrefix string + ServiceRoleArn string + Notification *NotificationConfig + CloudWatchOutput *CloudWatchOutputConfig +} + +// Command is a sent command as SendCommand and ListCommands report it. Status +// and the counts reflect the time it is read. +type Command struct { + CommandID string + DocumentName string + DocumentVersion string + Comment string + ExpiresAfter time.Time + Parameters map[string][]string + InstanceIDs []string + Targets []CommandTarget + RequestedDateTime time.Time + Status string + StatusDetails string + OutputS3Region string + OutputS3BucketName string + OutputS3KeyPrefix string + MaxConcurrency string + MaxErrors string + TargetCount int32 + CompletedCount int32 + ErrorCount int32 + DeliveryTimedOutCount int32 + ServiceRole string + Notification *NotificationConfig + CloudWatchOutput *CloudWatchOutputConfig + TimeoutSeconds int32 +} + +// CommandPlugin is the result of one document step on one instance. +type CommandPlugin struct { + Name string + Status string + StatusDetails string + ResponseCode int32 + ResponseStartDateTime time.Time + ResponseFinishDateTime time.Time + Output string + StandardOutputURL string + StandardErrorURL string + OutputS3Region string + OutputS3BucketName string + OutputS3KeyPrefix string +} + +// CommandInvocation is a command's run on one instance. From +// GetCommandInvocation, PluginName, the timing, the response code and the +// output are those of the selected plugin (the first one when none is named). +type CommandInvocation struct { + CommandID string + InstanceID string + InstanceName string + Comment string + DocumentName string + DocumentVersion string + RequestedDateTime time.Time + Status string + StatusDetails string + PluginName string + ResponseCode int32 + ExecutionStartTime time.Time + ExecutionEndTime time.Time + Stdout string + Stderr string + StandardOutputURL string + StandardErrorURL string + ServiceRole string + Notification *NotificationConfig + CloudWatchOutput *CloudWatchOutputConfig + Plugins []CommandPlugin +} + +// CommandFilter is one ListCommands or ListCommandInvocations filter. Keys are +// InvokedAfter, InvokedBefore, Status, ExecutionStage (ListCommands only) and +// DocumentName. +type CommandFilter struct { + Key string + Value string +} + +// CommandQuery selects commands or invocations. Every set field must match. +type CommandQuery struct { + CommandID string + InstanceID string + Filters []CommandFilter +} + +// RunCommand is an OPTIONAL capability, discovered by type assertion. +// +// Targets are validated: sending to an instance that does not exist, or that +// is not running, is InvalidInstanceId, as it is against the real service. +// Parameters are checked against the declared parameters of the document +// version that is sent. +// +// IMPORTANT: an emulated instance has no guest operating system, so nothing +// executes. Invocations report success and empty output. This exercises a +// caller's send/poll orchestration (that it waits for a terminal status, reads +// the response code, and handles failure) but it does NOT validate the script +// itself. A caller whose bootstrap script is wrong will still see success here. +type RunCommand interface { + SendCommand(ctx context.Context, cfg CommandConfig) (*Command, error) + // GetCommandInvocation reports one instance's run. pluginName selects a + // step and may be empty. + GetCommandInvocation(ctx context.Context, commandID, instanceID, pluginName string) (*CommandInvocation, error) + // ListCommands returns the matching commands, newest first. + ListCommands(ctx context.Context, q CommandQuery) ([]Command, error) + // ListCommandInvocations returns the matching invocations, newest command + // first. Plugins are filled only when details is set. + ListCommandInvocations(ctx context.Context, q CommandQuery, details bool) ([]CommandInvocation, error) + // CancelCommand cancels the command on the given instances, or on all of + // them when instanceIDs is empty. Invocations that already finished keep + // their result. + CancelCommand(ctx context.Context, commandID string, instanceIDs []string) error +} + +// InstanceInformation describes one managed node. +type InstanceInformation struct { + InstanceID string + PingStatus string + LastPingDateTime time.Time + AgentVersion string + IsLatestVersion bool + PlatformType string + PlatformName string + PlatformVersion string + ResourceType string + IPAddress string + ComputerName string + SourceID string + SourceType string +} + +// InstanceInformationFilter is one DescribeInstanceInformation filter. Keys +// are InstanceIds, PingStatus, PlatformType, ResourceType, AgentVersion, +// SourceIds, SourceTypes, tag-key and tag:. +type InstanceInformationFilter struct { + Key string + Values []string +} + +// ManagedNodes is an OPTIONAL capability, discovered by type assertion. Every +// EC2 instance that is not terminated counts as a managed node: running ones +// are Online and the rest ConnectionLost. +type ManagedNodes interface { + DescribeInstanceInformation(ctx context.Context, filters []InstanceInformationFilter) ([]InstanceInformation, error) +} diff --git a/providers/aws/ssm/errors.go b/providers/aws/ssm/errors.go index ff252aa5a..6920b5fae 100644 --- a/providers/aws/ssm/errors.go +++ b/providers/aws/ssm/errors.go @@ -21,6 +21,10 @@ const ( excInvalidDocumentVersion = "InvalidDocumentVersion" excInvalidFilterKey = "InvalidFilterKey" excInvalidInstanceID = "InvalidInstanceId" + excInvalidCommandID = "InvalidCommandId" + excInvocationDoesNotExist = "InvocationDoesNotExist" + excInvalidPluginName = "InvalidPluginName" + excInvalidInstanceInfoFilter = "InvalidInstanceInformationFilterValue" excInvalidPermissionType = "InvalidPermissionType" excMaxDocumentSizeExceeded = "MaxDocumentSizeExceeded" excValidation = "ValidationException" diff --git a/providers/aws/ssm/instance_information.go b/providers/aws/ssm/instance_information.go new file mode 100644 index 000000000..56cb5dc1e --- /dev/null +++ b/providers/aws/ssm/instance_information.go @@ -0,0 +1,158 @@ +package ssm + +import ( + "context" + "slices" + "sort" + "strings" + + "github.com/stackshy/cloudemu/v2/errors" + ssmdriver "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" + computedriver "github.com/stackshy/cloudemu/v2/services/compute/driver" +) + +var _ ssmdriver.ManagedNodes = (*Mock)(nil) + +// Managed-node values the emulator reports for every EC2 instance. +const ( + pingOnline = "Online" + pingConnectionLost = "ConnectionLost" + resourceEC2 = "EC2Instance" + sourceEC2 = "AWS::EC2::Instance" + agentVersion = "3.3.1611.0" + + nodeKeyPingStatus = "PingStatus" + nodeKeyPlatformType = "PlatformType" + nodeKeyPlatformTypes = "PlatformTypes" + nodeKeyResourceType = "ResourceType" + nodeKeyAgentVersion = "AgentVersion" + nodeKeyInstanceIDs = "InstanceIds" + nodeKeySourceIDs = "SourceIds" + nodeKeySourceTypes = "SourceTypes" + nodeKeyTagKey = "tag-key" +) + +// DescribeInstanceInformation lists the managed nodes: every EC2 instance that +// is not terminated. Filters are AND-combined and the values of one filter +// OR-combined. +func (m *Mock) DescribeInstanceInformation( + ctx context.Context, filters []ssmdriver.InstanceInformationFilter, +) ([]ssmdriver.InstanceInformation, error) { + for _, f := range filters { + if err := validNodeFilter(f); err != nil { + return nil, err + } + } + + out := make([]ssmdriver.InstanceInformation, 0) + if m.instanceResolver == nil { + return out, nil + } + + found, err := m.instanceResolver.DescribeInstances(ctx, nil, nil, + computedriver.DescribeInstancesOptions{IncludeManagedResources: true}) + if err != nil { + return nil, err + } + + now := m.opts.Clock.Now() + + for i := range found { + inst := &found[i] + if !managedState(inst.State) { + continue + } + + info := nodeInformation(inst) + info.LastPingDateTime = now + + if slices.IndexFunc(filters, func(f ssmdriver.InstanceInformationFilter) bool { return !nodeMatches(&info, inst, f) }) < 0 { + out = append(out, info) + } + } + + sort.Slice(out, func(i, j int) bool { return out[i].InstanceID < out[j].InstanceID }) + + return out, nil +} + +// nodeInformation describes one instance as a managed node. +func nodeInformation(inst *computedriver.Instance) ssmdriver.InstanceInformation { + info := ssmdriver.InstanceInformation{ + InstanceID: inst.ID, PingStatus: pingConnectionLost, AgentVersion: agentVersion, IsLatestVersion: true, + PlatformType: platformLinux, PlatformName: "Amazon Linux", PlatformVersion: "2023", + ResourceType: resourceEC2, IPAddress: inst.PrivateIP, SourceID: inst.ID, SourceType: sourceEC2, + } + + if inst.State == stateRunning { + info.PingStatus = pingOnline + } + + if strings.EqualFold(inst.OSType, platformWindows) { + info.PlatformType, info.PlatformName, info.PlatformVersion = platformWindows, "Microsoft Windows Server 2022 Datacenter", "10.0.20348" + } + + if inst.PrivateIP != "" { + info.ComputerName = "ip-" + strings.ReplaceAll(inst.PrivateIP, ".", "-") + ".ec2.internal" + } + + return info +} + +// validNodeFilter rejects an unknown key or an out-of-range value. +func validNodeFilter(f ssmdriver.InstanceInformationFilter) error { + var allowed []string + + switch f.Key { + case nodeKeyPingStatus: + allowed = []string{pingOnline, pingConnectionLost, "Inactive"} + case nodeKeyPlatformType, nodeKeyPlatformTypes: + allowed = []string{platformWindows, platformLinux, platformMacOS} + case nodeKeyResourceType: + allowed = []string{resourceEC2, "ManagedInstance"} + case nodeKeyInstanceIDs, nodeKeyAgentVersion, "ActivationIds", "IamRole", "AssociationStatus", + nodeKeySourceIDs, nodeKeySourceTypes, nodeKeyTagKey: + default: + if !strings.HasPrefix(f.Key, "tag:") { + return ssmErrf(excInvalidFilterKey, errors.InvalidArgument, "The filter key %s is not valid.", f.Key) + } + } + + for _, v := range f.Values { + if allowed != nil && !slices.Contains(allowed, v) { + return ssmErrf(excInvalidInstanceInfoFilter, errors.InvalidArgument, + "The filter value %s is not valid for %s.", v, f.Key) + } + } + + return nil +} + +// nodeMatches applies one filter. Activations, IAM roles and associations +// do not exist for an emulated EC2 node, so those filters match nothing. +func nodeMatches(info *ssmdriver.InstanceInformation, inst *computedriver.Instance, f ssmdriver.InstanceInformationFilter) bool { + switch f.Key { + case nodeKeyInstanceIDs, nodeKeySourceIDs: + return slices.Contains(f.Values, info.InstanceID) + case nodeKeyPingStatus: + return slices.Contains(f.Values, info.PingStatus) + case nodeKeyPlatformType, nodeKeyPlatformTypes: + return slices.Contains(f.Values, info.PlatformType) + case nodeKeyResourceType: + return slices.Contains(f.Values, info.ResourceType) + case nodeKeySourceTypes: + return slices.Contains(f.Values, info.SourceType) + case nodeKeyAgentVersion: + return slices.Contains(f.Values, info.AgentVersion) + case nodeKeyTagKey: + return slices.ContainsFunc(f.Values, func(k string) bool { _, ok := inst.Tags[k]; return ok }) + default: + if key, ok := strings.CutPrefix(f.Key, "tag:"); ok { + v, has := inst.Tags[key] + + return has && slices.Contains(f.Values, v) + } + + return false + } +} diff --git a/providers/aws/ssm/policy_eval.go b/providers/aws/ssm/policy_eval.go index 46c29462a..090c096ec 100644 --- a/providers/aws/ssm/policy_eval.go +++ b/providers/aws/ssm/policy_eval.go @@ -30,8 +30,11 @@ type policyNotice struct { // deleted and each action publishes a "Parameter Store Policy Action" event. // It reports whether anything changed. It has the Tickable signature of the // shared scheduler. Reads also evaluate policies, so nothing has to call it. +// It also writes the S3 output of Run Command invocations that finished. func (m *Mock) Tick(now time.Time) bool { - return m.evaluateAllPolicies(context.Background(), now) + changed := m.evaluateAllPolicies(context.Background(), now) + + return m.flushOutputs(context.Background(), now) || changed } // evaluateAllPolicies runs the due policies of every parameter. diff --git a/providers/aws/ssm/run_command.go b/providers/aws/ssm/run_command.go index 283cd6ded..7d7f750d0 100644 --- a/providers/aws/ssm/run_command.go +++ b/providers/aws/ssm/run_command.go @@ -2,14 +2,46 @@ package ssm import ( "context" + "regexp" + "strconv" "strings" + "time" "github.com/stackshy/cloudemu/v2/errors" "github.com/stackshy/cloudemu/v2/internal/idgen" + "github.com/stackshy/cloudemu/v2/internal/settle" ssmdriver "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" computedriver "github.com/stackshy/cloudemu/v2/services/compute/driver" ) +var _ ssmdriver.RunCommand = (*Mock)(nil) + +// SendCommand request limits and defaults from the SSM API reference. +const ( + defaultCommandTimeout = 3600 + minCommandTimeout = 30 + maxCommandTimeout = 2592000 + maxCommandComment = 100 + maxCommandInstances = 50 + maxCommandTargets = 5 + defaultMaxConcurrency = "50" + defaultMaxErrors = "0" + percent = 100 +) + +// EC2 instance states Run Command cares about. +const ( + stateRunning = "running" + stateTerminated = "terminated" + stateShuttingDown = "shutting-down" +) + +var ( + maxConcurrencyPattern = regexp.MustCompile(`^([1-9]\d*|[1-9]\d%|[1-9]%|100%)$`) + maxErrorsPattern = regexp.MustCompile(`^([1-9]\d*|0|[1-9]\d%|\d%|100%)$`) + instanceIDPattern = regexp.MustCompile(`^(i-(\w{8}|\w{17})|mi-\w{17})$`) +) + // InstanceResolver is the slice of the compute mock this package needs to // check that a Run Command target exists. type InstanceResolver interface { @@ -23,61 +55,210 @@ func (m *Mock) SetInstanceResolver(r InstanceResolver) { m.instanceResolver = r } -// SendCommand records a Run Command send and returns its command id. +// commandDoc is the document version a send runs. +type commandDoc struct { + name string + specs map[string]paramSpec + steps []commandStep + strict bool // false for an AWS-owned name the catalog does not hold +} + +// SendCommand validates a Run Command send, records it and returns the +// command as it stands right after the send. // -// Nothing executes: an emulated instance has no guest operating system. The -// invocation is recorded as successful so a caller's send/poll loop runs to -// completion, but the script itself is never validated. See driver.RunCommand. +// Nothing executes: an emulated instance has no guest operating system. Each +// invocation moves Pending -> InProgress -> Success over the settle windows +// (immediately when async settle is off), so a caller's send/poll loop runs to +// completion. See driver.RunCommand. // //nolint:gocritic // hugeParam: interface method signature cannot be changed. -func (m *Mock) SendCommand(ctx context.Context, cfg ssmdriver.CommandConfig) (string, error) { - // Real SSM accepts EITHER explicit InstanceIds OR tag/attribute Targets; - // supplying neither is a ValidationException. - if len(cfg.InstanceIDs) == 0 && len(cfg.Targets) == 0 { - return "", errors.New(errors.InvalidArgument, "either instance IDs or targets must be specified") +func (m *Mock) SendCommand(ctx context.Context, cfg ssmdriver.CommandConfig) (*ssmdriver.Command, error) { + if err := validateCommandConfig(&cfg); err != nil { + return nil, err } - if cfg.DocumentName == "" { - return "", errors.New(errors.InvalidArgument, "DocumentName is required") + doc, err := m.commandDocument(cfg.DocumentName, cfg.DocumentVersion) + if err != nil { + return nil, err } - docName, err := m.commandDocument(cfg.DocumentName) - if err != nil { - return "", err + resolved := cfg.Parameters + + if doc.strict { + if resolved, err = validateParameters(doc.name, doc.specs, cfg.Parameters); err != nil { + return nil, err + } } // Real SSM answers InvalidInstanceId when an explicitly listed target is not // a managed instance, and that is the single most common Run Command failure // during bring-up. Accepting any id hides it until the caller runs for real. if err := m.checkTargets(ctx, cfg.InstanceIDs); err != nil { - return "", err + return nil, err } // Resolve tag/attribute Targets to concrete instance ids. Unlike an explicit // id, a Target that matches nothing is not an error. Real SSM accepts the // command with a TargetCount of zero. - resolved := m.resolveTargets(ctx, cfg.Targets) - instanceIDs := dedupeStrings(append(append([]string{}, cfg.InstanceIDs...), resolved...)) - - commandID := idgen.UUID() - - for _, instanceID := range instanceIDs { - m.commands.Set(commandKey(commandID, instanceID), ssmdriver.CommandInvocation{ - CommandID: commandID, - InstanceID: instanceID, - DocumentName: docName, - Status: "Success", - ResponseCode: 0, - }) + invocations := make([]*invocationRecord, 0, len(cfg.InstanceIDs)) + seen := map[string]bool{} + + for _, id := range cfg.InstanceIDs { + if !seen[id] { + seen[id] = true + + invocations = append(invocations, &invocationRecord{InstanceID: id}) + } + } + + for _, t := range m.resolveTargets(ctx, cfg.Targets) { + if !seen[t.ID] { + seen[t.ID] = true + + invocations = append(invocations, &invocationRecord{InstanceID: t.ID, Offline: t.State != stateRunning}) + } + } + + rec := m.newCommandRecord(&cfg, doc, resolved, invocations) + + // The response reports the command as real SSM does at acceptance: + // Pending with nothing completed, even when it finishes at once here. The + // caller learns the outcome by polling. + cmd := rec.Command + cmd.Status, cmd.StatusDetails = ssmdriver.CommandPending, ssmdriver.CommandPending + cmd.TargetCount = int32(len(invocations)) //nolint:gosec // at most 50 explicit ids plus tag matches. + + m.cmdMu.Lock() + m.commands.Set(rec.Command.CommandID, rec) + m.cmdMu.Unlock() + + m.flushOutputs(ctx, m.opts.Clock.Now()) + + return &cmd, nil +} + +// newCommandRecord builds the stored record of a send. +func (m *Mock) newCommandRecord(cfg *ssmdriver.CommandConfig, doc commandDoc, + resolved map[string][]string, invocations []*invocationRecord, +) *commandRecord { + now := m.opts.Clock.Now() + timeout := cfg.TimeoutSeconds + + if timeout == 0 { + timeout = defaultCommandTimeout + } + + version := cfg.DocumentVersion + if version == "" { + version = versionDefault + } + + rec := &commandRecord{ + Command: ssmdriver.Command{ + CommandID: idgen.UUID(), DocumentName: doc.name, DocumentVersion: version, Comment: cfg.Comment, + Parameters: cfg.Parameters, InstanceIDs: cfg.InstanceIDs, Targets: cfg.Targets, + RequestedDateTime: now, OutputS3Region: cfg.OutputS3Region, OutputS3BucketName: cfg.OutputS3BucketName, + OutputS3KeyPrefix: cfg.OutputS3KeyPrefix, MaxConcurrency: orDefault(cfg.MaxConcurrency, defaultMaxConcurrency), + MaxErrors: orDefault(cfg.MaxErrors, defaultMaxErrors), ServiceRole: cfg.ServiceRoleArn, + Notification: cfg.Notification, CloudWatchOutput: cfg.CloudWatchOutput, TimeoutSeconds: timeout, + }, + Steps: doc.steps, + Delivery: m.opts.SettleDuration(settle.DefaultCommandDeliverySettle), + Run: m.opts.SettleDuration(settle.DefaultCommandRunSettle), + Invocations: invocations, + } + + if rec.Command.OutputS3BucketName != "" && rec.Command.OutputS3Region == "" { + rec.Command.OutputS3Region = m.opts.Region + } + + if len(resolved["executionTimeout"]) == 1 { + rec.ExecutionTimeout, _ = strconv.Atoi(resolved["executionTimeout"][0]) + } + + total := time.Duration(timeout) * time.Second + if rec.ExecutionTimeout > 0 { + total += time.Duration(rec.ExecutionTimeout) * time.Second + } else { + total += defaultCommandTimeout * time.Second + } + + rec.Command.ExpiresAfter = now.Add(total) + + return rec +} + +func orDefault(v, def string) string { + if v == "" { + return def } - return commandID, nil + return v } -// commandDocument resolves a SendCommand DocumentName (a name or an ARN) -// through the customer documents and the AWS-owned catalog. Only Command -// documents can be sent. -func (m *Mock) commandDocument(ref string) (string, error) { +// validateCommandConfig applies the SendCommand request constraints real SSM +// checks before it looks at the document. +// +//nolint:gocyclo // one flat check per request field. +func validateCommandConfig(cfg *ssmdriver.CommandConfig) error { + // Real SSM accepts EITHER explicit InstanceIds OR tag/attribute Targets; + // supplying neither is a ValidationException. + if len(cfg.InstanceIDs) == 0 && len(cfg.Targets) == 0 { + return errors.New(errors.InvalidArgument, "either instance IDs or targets must be specified") + } + + if cfg.DocumentName == "" { + return errors.New(errors.InvalidArgument, "DocumentName is required") + } + + switch { + case cfg.DocumentVersion != "" && !docVersionPattern.MatchString(cfg.DocumentVersion): + return patternErr(cfg.DocumentVersion, "documentVersion", "([$]LATEST|[$]DEFAULT|^[1-9][0-9]*$)") + case cfg.MaxConcurrency != "" && !maxConcurrencyPattern.MatchString(cfg.MaxConcurrency): + return patternErr(cfg.MaxConcurrency, "maxConcurrency", maxConcurrencyPattern.String()) + case cfg.MaxErrors != "" && !maxErrorsPattern.MatchString(cfg.MaxErrors): + return patternErr(cfg.MaxErrors, "maxErrors", maxErrorsPattern.String()) + case cfg.TimeoutSeconds != 0 && (cfg.TimeoutSeconds < minCommandTimeout || cfg.TimeoutSeconds > maxCommandTimeout): + return ssmErrf(excValidation, errors.InvalidArgument, + "1 validation error detected: Value '%d' at 'timeoutSeconds' failed to satisfy constraint: "+ + "Member must have value between %d and %d", cfg.TimeoutSeconds, minCommandTimeout, maxCommandTimeout) + case len(cfg.Comment) > maxCommandComment: + return lengthErr(cfg.Comment, "comment", maxCommandComment) + case len(cfg.InstanceIDs) > maxCommandInstances: + return ssmErrf(excValidation, errors.InvalidArgument, + "1 validation error detected: Value at 'instanceIds' failed to satisfy constraint: "+ + "Member must have length less than or equal to %d", maxCommandInstances) + case len(cfg.Targets) > maxCommandTargets: + return ssmErrf(excValidation, errors.InvalidArgument, + "1 validation error detected: Value at 'targets' failed to satisfy constraint: "+ + "Member must have length less than or equal to %d", maxCommandTargets) + } + + for _, id := range cfg.InstanceIDs { + if !instanceIDPattern.MatchString(id) { + return patternErr(id, "instanceIds", `(^i-(\w{8}|\w{17})$)|(^mi-\w{17}$)`) + } + } + + return nil +} + +func patternErr(value, field, pattern string) error { + return ssmErrf(excValidation, errors.InvalidArgument, + "1 validation error detected: Value '%s' at '%s' failed to satisfy constraint: "+ + "Member must satisfy regular expression pattern: %s", value, field, pattern) +} + +func lengthErr(value, field string, limit int) error { + return ssmErrf(excValidation, errors.InvalidArgument, + "1 validation error detected: Value '%s' at '%s' failed to satisfy constraint: "+ + "Member must have length less than or equal to %d", value, field, limit) +} + +// commandDocument resolves a SendCommand DocumentName (a name or an ARN) and +// DocumentVersion through the customer documents and the AWS-owned catalog. +// Only Command documents can be sent. +func (m *Mock) commandDocument(ref, version string) (commandDoc, error) { m.docMu.RLock() defer m.docMu.RUnlock() @@ -85,27 +266,45 @@ func (m *Mock) commandDocument(ref string) (string, error) { if err != nil { // The catalog holds the common AWS-owned documents, not all of them. // An unknown name in the AWS namespace is taken as an AWS-owned Command - // document so a real one the catalog lacks still runs. + // document so a real one the catalog lacks still runs. Its parameters + // are unknown, so they are not checked. if name := documentName(ref); awsOwnedName(name) { - return name, nil + if version != "" && version != versionDefault && version != versionLatest && version != "1" { + return commandDoc{}, ssmErrf(excInvalidDocumentVersion, errors.NotFound, + "The document version isn't valid or doesn't exist.") + } + + return commandDoc{name: name}, nil } - return "", err + return commandDoc{}, err } if d.docType != ssmdriver.DocumentTypeCommand { - return "", ssmErrf(excInvalidDocument, errors.InvalidArgument, + return commandDoc{}, ssmErrf(excInvalidDocument, errors.InvalidArgument, "Document %s of type %s can't be used with SendCommand.", d.name, d.docType) } - return d.name, nil + v, err := d.resolveVersion(version, "") + if err != nil { + return commandDoc{}, err + } + + return commandDoc{name: d.name, specs: v.meta.paramSpecs, steps: v.meta.steps, strict: true}, nil +} + +// resolvedTarget is an instance a tag/attribute Target selected. +type resolvedTarget struct { + ID string + State string } -// resolveTargets maps SSM Targets to the instance ids they select. Multiple -// targets are AND-combined, matching real SSM. Resolution needs the compute -// mock; without it (or on a lookup error) no ids are resolved, which still -// yields an accepted command. -func (m *Mock) resolveTargets(ctx context.Context, targets []ssmdriver.CommandTarget) []string { +// resolveTargets maps SSM Targets to the instances they select. Multiple +// targets are AND-combined, matching real SSM. Terminated instances are not +// managed nodes and are left out. Resolution needs the compute mock; without it +// (or on a lookup error) nothing is resolved, which still yields an accepted +// command. +func (m *Mock) resolveTargets(ctx context.Context, targets []ssmdriver.CommandTarget) []resolvedTarget { if len(targets) == 0 || m.instanceResolver == nil { return nil } @@ -135,12 +334,20 @@ func (m *Mock) resolveTargets(ctx context.Context, targets []ssmdriver.CommandTa return nil } - ids := make([]string, 0, len(found)) + out := make([]resolvedTarget, 0, len(found)) + for i := range found { - ids = append(ids, found[i].ID) + if managedState(found[i].State) { + out = append(out, resolvedTarget{ID: found[i].ID, State: found[i].State}) + } } - return ids + return out +} + +// managedState reports whether an instance in this state is a managed node. +func managedState(state string) bool { + return state != stateTerminated && state != stateShuttingDown } // targetFilterName maps a supported SSM Target Key to the equivalent EC2 @@ -162,44 +369,11 @@ func targetFilterName(key string) (string, bool) { } } -func dedupeStrings(in []string) []string { - seen := make(map[string]bool, len(in)) - out := make([]string, 0, len(in)) - - for _, s := range in { - if seen[s] { - continue - } - - seen[s] = true - - out = append(out, s) - } - - return out -} - -// GetCommandInvocation returns the recorded invocation for one instance. -// -// An unknown pair is InvocationDoesNotExist rather than a fabricated success: -// a caller polling a command it never sent has a real bug, and answering -// "Success" would bury it. -func (m *Mock) GetCommandInvocation( - _ context.Context, commandID, instanceID string, -) (*ssmdriver.CommandInvocation, error) { - inv, ok := m.commands.Get(commandKey(commandID, instanceID)) - if !ok { - return nil, errors.Newf(errors.NotFound, - "InvocationDoesNotExist: no invocation of command %q on instance %q", - commandID, instanceID) - } - - return &inv, nil -} - -// checkTargets rejects instance ids the compute mock does not know. +// checkTargets rejects instance ids the compute mock does not know, and +// instances that are not running: their agent is offline, which real SSM +// reports as not in a valid state. func (m *Mock) checkTargets(ctx context.Context, instanceIDs []string) error { - if m.instanceResolver == nil { + if m.instanceResolver == nil || len(instanceIDs) == 0 { return nil } @@ -213,21 +387,22 @@ func (m *Mock) checkTargets(ctx context.Context, instanceIDs []string) error { return ssmErrf(excInvalidInstanceID, errors.NotFound, "%v", err) } - known := make(map[string]bool, len(found)) + state := make(map[string]string, len(found)) for i := range found { - known[found[i].ID] = true + state[found[i].ID] = found[i].State } for _, id := range instanceIDs { - if !known[id] { - return ssmErrf(excInvalidInstanceID, errors.NotFound, - "Instance %s is not a managed instance.", id) + s, ok := state[id] + + switch { + case !ok || !managedState(s): + return ssmErrf(excInvalidInstanceID, errors.NotFound, "Instance %s is not a managed instance.", id) + case s != stateRunning: + return ssmErrf(excInvalidInstanceID, errors.FailedPrecondition, + "Instances [[%s]] not in a valid state for account %s", id, m.opts.AccountID) } } return nil } - -func commandKey(commandID, instanceID string) string { - return commandID + "|" + instanceID -} diff --git a/providers/aws/ssm/run_command_test.go b/providers/aws/ssm/run_command_test.go new file mode 100644 index 000000000..61a069e9d --- /dev/null +++ b/providers/aws/ssm/run_command_test.go @@ -0,0 +1,220 @@ +package ssm_test + +import ( + "context" + "slices" + "testing" + "time" + + "github.com/stackshy/cloudemu/v2/config" + "github.com/stackshy/cloudemu/v2/providers/aws/ssm" + ssmdriver "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" + computedriver "github.com/stackshy/cloudemu/v2/services/compute/driver" +) + +// fakeFleet answers DescribeInstances from a fixed instance list. +type fakeFleet []computedriver.Instance + +func (f fakeFleet) DescribeInstances(_ context.Context, ids []string, filters []computedriver.DescribeFilter, + _ ...computedriver.DescribeInstancesOptions, +) ([]computedriver.Instance, error) { + var out []computedriver.Instance + + for _, inst := range f { + if len(ids) > 0 && !slices.Contains(ids, inst.ID) { + continue + } + + if len(filters) > 0 && !slices.Contains(filters[0].Values, inst.Tags["Role"]) { + continue + } + + out = append(out, inst) + } + + return out, nil +} + +const ( + webA = "i-0aaaaaaaaaaaaaaaa" + webB = "i-0bbbbbbbbbbbbbbbb" + webC = "i-0cccccccccccccccc" +) + +func newSettledMock(t *testing.T) (*ssm.Mock, *config.FakeClock) { + t.Helper() + + fc := config.NewFakeClock(time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)) + m := ssm.New(config.NewOptions(config.WithClock(fc), config.WithAsyncSettle())) + m.SetInstanceResolver(fakeFleet{ + {ID: webA, State: "running", Tags: map[string]string{"Role": "web"}}, + {ID: webB, State: "running", Tags: map[string]string{"Role": "web"}}, + {ID: webC, State: "stopped", Tags: map[string]string{"Role": "web"}}, + }) + + return m, fc +} + +func invocationStatus(t *testing.T, m *ssm.Mock, commandID, instanceID string) string { + t.Helper() + + inv, err := m.GetCommandInvocation(context.Background(), commandID, instanceID, "") + if err != nil { + t.Fatalf("GetCommandInvocation %s: %v", instanceID, err) + } + + return inv.Status +} + +// MaxConcurrency 1 runs the invocations one after the other. +func TestRunCommandMaxConcurrencyStaggers(t *testing.T) { + m, fc := newSettledMock(t) + + cmd, err := m.SendCommand(context.Background(), ssmdriver.CommandConfig{ + InstanceIDs: []string{webA, webB}, DocumentName: "AWS-RunShellScript", MaxConcurrency: "1", + Parameters: map[string][]string{"commands": {"uptime"}}, + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + fc.Advance(2 * time.Second) + + if a, b := invocationStatus(t, m, cmd.CommandID, webA), invocationStatus(t, m, cmd.CommandID, webB); a != "InProgress" || b != "Pending" { + t.Fatalf("first batch: %s / %s, want InProgress / Pending", a, b) + } + + fc.Advance(5 * time.Second) + + if a, b := invocationStatus(t, m, cmd.CommandID, webA), invocationStatus(t, m, cmd.CommandID, webB); a != "Success" || b != "InProgress" { + t.Fatalf("second batch: %s / %s, want Success / InProgress", a, b) + } + + fc.Advance(time.Minute) + + cmds, err := m.ListCommands(context.Background(), ssmdriver.CommandQuery{CommandID: cmd.CommandID}) + if err != nil || len(cmds) != 1 || cmds[0].Status != "Success" || cmds[0].CompletedCount != 2 { + t.Fatalf("ListCommands = %+v, %v", cmds, err) + } +} + +// A tag target that selects a stopped instance waits for delivery and then +// times out, while the running ones succeed. +func TestRunCommandStoppedTargetTimesOut(t *testing.T) { + m, fc := newSettledMock(t) + + cmd, err := m.SendCommand(context.Background(), ssmdriver.CommandConfig{ + Targets: []ssmdriver.CommandTarget{{Key: "tag:Role", Values: []string{"web"}}}, DocumentName: "AWS-RunShellScript", + TimeoutSeconds: 60, Parameters: map[string][]string{"commands": {"uptime"}}, + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + if cmd.TargetCount != 3 { + t.Fatalf("TargetCount = %d, want 3", cmd.TargetCount) + } + + fc.Advance(30 * time.Second) + + if s := invocationStatus(t, m, cmd.CommandID, webC); s != "Delayed" { + t.Fatalf("stopped instance before timeout = %s, want Delayed", s) + } + + fc.Advance(time.Minute) + + inv, _ := m.GetCommandInvocation(context.Background(), cmd.CommandID, webC, "") + if inv.Status != "TimedOut" || inv.StatusDetails != "DeliveryTimedOut" { + t.Fatalf("stopped instance after timeout = %s / %s", inv.Status, inv.StatusDetails) + } + + cmds, _ := m.ListCommands(context.Background(), ssmdriver.CommandQuery{CommandID: cmd.CommandID}) + if c := cmds[0]; c.Status != "Failed" || c.StatusDetails != "Incomplete" || c.DeliveryTimedOutCount != 1 || c.ErrorCount != 0 { + t.Fatalf("command = %+v, want Failed/Incomplete with one delivery timeout", c) + } + + byStatus, err := m.ListCommandInvocations(context.Background(), ssmdriver.CommandQuery{ + Filters: []ssmdriver.CommandFilter{{Key: "Status", Value: "DeliveryTimedOut"}}, + }, false) + if err != nil || len(byStatus) != 1 || byStatus[0].InstanceID != webC { + t.Fatalf("Status filter = %+v, %v", byStatus, err) + } +} + +// Command history survives a snapshot: a restored mock reports the same +// commands, keeps settling them from their send time and honours cancels. +func TestRunCommandSnapshotRoundTrip(t *testing.T) { + ctx := context.Background() + src, fc := newSettledMock(t) + + done, err := src.SendCommand(ctx, ssmdriver.CommandConfig{ + InstanceIDs: []string{webA}, DocumentName: "AWS-RunShellScript", Comment: "done", + Parameters: map[string][]string{"commands": {"uptime"}}, + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + fc.Advance(time.Minute) + + running, err := src.SendCommand(ctx, ssmdriver.CommandConfig{ + InstanceIDs: []string{webA, webB}, DocumentName: "AWS-RunShellScript", + Parameters: map[string][]string{"commands": {"sleep 5"}}, + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + fc.Advance(2 * time.Second) + + if err := src.CancelCommand(ctx, running.CommandID, []string{webB}); err != nil { + t.Fatalf("CancelCommand: %v", err) + } + + raw, err := src.Snapshot(ctx, false) + if err != nil { + t.Fatalf("Snapshot: %v", err) + } + + dst := ssm.New(config.NewOptions(config.WithClock(fc), config.WithAsyncSettle())) + if err := dst.Restore(ctx, raw); err != nil { + t.Fatalf("Restore: %v", err) + } + + cmds, err := dst.ListCommands(ctx, ssmdriver.CommandQuery{}) + if err != nil || len(cmds) != 2 || cmds[0].CommandID != running.CommandID || cmds[1].Comment != "done" { + t.Fatalf("restored commands = %+v, %v", cmds, err) + } + + if s := invocationStatus(t, dst, running.CommandID, webA); s != "InProgress" { + t.Fatalf("restored running invocation = %s, want InProgress", s) + } + + fc.Advance(time.Minute) + + if a, b := invocationStatus(t, dst, running.CommandID, webA), invocationStatus(t, dst, running.CommandID, webB); a != "Success" || b != "Cancelled" { + t.Fatalf("restored invocations = %s / %s, want Success / Cancelled", a, b) + } + + if s := invocationStatus(t, dst, done.CommandID, webA); s != "Success" { + t.Fatalf("restored finished invocation = %s", s) + } +} + +// A snapshot written before the command history existed still restores its +// invocations as finished commands. +func TestRunCommandRestoresLegacySnapshot(t *testing.T) { + const legacy = `{"commands":{"11111111-2222-3333-4444-555555555555|i-0aaaaaaaaaaaaaaaa":` + + `{"CommandID":"11111111-2222-3333-4444-555555555555","InstanceID":"i-0aaaaaaaaaaaaaaaa",` + + `"DocumentName":"AWS-RunShellScript","Status":"Success","ResponseCode":0}}}` + + m := newMock() + if err := m.Restore(context.Background(), []byte(legacy)); err != nil { + t.Fatalf("Restore: %v", err) + } + + inv, err := m.GetCommandInvocation(context.Background(), "11111111-2222-3333-4444-555555555555", webA, "") + if err != nil || inv.Status != "Success" || inv.DocumentName != "AWS-RunShellScript" { + t.Fatalf("legacy invocation = %+v, %v", inv, err) + } +} diff --git a/providers/aws/ssm/snapshot.go b/providers/aws/ssm/snapshot.go index 4e243c854..d1414ddff 100644 --- a/providers/aws/ssm/snapshot.go +++ b/providers/aws/ssm/snapshot.go @@ -4,21 +4,26 @@ import ( "context" "encoding/json" "fmt" + "sort" "time" "github.com/stackshy/cloudemu/v2/internal/snapshot" + ssmdriver "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" ) var _ snapshot.Snapshottable = (*Mock)(nil) -// ssmSnapshot is the full serialized state of the SSM Parameter Store mock. -// params holds an unexported paramData (with a slice of unexported *version), so -// it is promoted to an exported snapshot form keyed by parameter name; commands -// holds a fully-exported ssmdriver.CommandInvocation and round-trips through the -// generic memstore helper. The wired instanceResolver and opts are not -// serialized. +// ssmSnapshot is the full serialized state of the SSM mock. params holds an +// unexported paramData (with a slice of unexported *version), so it is +// promoted to an exported snapshot form keyed by parameter name. The Run +// Command history is a map of exported command records keyed by command id. +// The wired instanceResolver, outputStore and opts are not serialized. type ssmSnapshot struct { - Params map[string]*paramSnapshot `json:"params,omitempty"` + Params map[string]*paramSnapshot `json:"params,omitempty"` + // CommandHistory holds every Run Command send. + CommandHistory json.RawMessage `json:"commandHistory,omitempty"` + // Commands is the per-invocation form older snapshots used. It is only + // read. Commands json.RawMessage `json:"commands,omitempty"` Settings map[string]settingSnapshot `json:"settings,omitempty"` // Documents holds the customer SSM documents keyed by name. @@ -69,12 +74,15 @@ func (m *Mock) Snapshot(_ context.Context, _ bool) (json.RawMessage, error) { Params: m.snapshotParams(), Settings: m.snapshotSettings(), Documents: m.snapshotDocuments(), } + m.cmdMu.RLock() cmds, err := m.commands.Snapshot() + m.cmdMu.RUnlock() + if err != nil { return nil, fmt.Errorf("ssm: snapshot commands: %w", err) } - snap.Commands = cmds + snap.CommandHistory = cmds return json.Marshal(snap) } @@ -144,12 +152,62 @@ func (m *Mock) Restore(_ context.Context, data json.RawMessage) error { return err } - if len(snap.Commands) > 0 { - if err := m.commands.LoadSnapshot(snap.Commands); err != nil { + return m.restoreCommands(snap.CommandHistory, snap.Commands) +} + +// legacyInvocation is one entry of the per-invocation command form older +// snapshots used. +type legacyInvocation struct { + CommandID string + InstanceID string + DocumentName string +} + +// restoreCommands loads the command history, converting the older +// per-invocation form into finished commands. +func (m *Mock) restoreCommands(history, legacy json.RawMessage) error { + m.cmdMu.Lock() + defer m.cmdMu.Unlock() + + if len(history) > 0 { + if err := m.commands.LoadSnapshot(history); err != nil { return fmt.Errorf("ssm: restore commands: %w", err) } } + if len(legacy) == 0 { + return nil + } + + var old map[string]legacyInvocation + if err := json.Unmarshal(legacy, &old); err != nil { + return fmt.Errorf("ssm: restore commands: %w", err) + } + + keys := make([]string, 0, len(old)) + for k := range old { + keys = append(keys, k) + } + + sort.Strings(keys) + + for _, k := range keys { + inv := old[k] + + rec, ok := m.commands.Get(inv.CommandID) + if !ok { + rec = &commandRecord{Command: ssmdriver.Command{ + CommandID: inv.CommandID, DocumentName: inv.DocumentName, DocumentVersion: versionDefault, + RequestedDateTime: m.opts.Clock.Now(), MaxConcurrency: defaultMaxConcurrency, + MaxErrors: defaultMaxErrors, TimeoutSeconds: defaultCommandTimeout, + }} + m.commands.Set(inv.CommandID, rec) + } + + rec.Command.InstanceIDs = append(rec.Command.InstanceIDs, inv.InstanceID) + rec.Invocations = append(rec.Invocations, &invocationRecord{InstanceID: inv.InstanceID, OutputWritten: true}) + } + return nil } diff --git a/providers/aws/ssm/ssm.go b/providers/aws/ssm/ssm.go index bf2029af7..e3e36a18a 100644 --- a/providers/aws/ssm/ssm.go +++ b/providers/aws/ssm/ssm.go @@ -31,7 +31,6 @@ import ( "github.com/stackshy/cloudemu/v2/internal/awsevents" "github.com/stackshy/cloudemu/v2/internal/idgen" "github.com/stackshy/cloudemu/v2/internal/memstore" - ssmdriver "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" "github.com/stackshy/cloudemu/v2/services/parameterstore/driver" ) @@ -77,9 +76,15 @@ type KMSCrypto interface { // Mock is an in-memory mock implementation of SSM Parameter Store. type Mock struct { - params *memstore.Store[*paramData] - commands *memstore.Store[ssmdriver.CommandInvocation] + params *memstore.Store[*paramData] + // commands holds every Run Command send, keyed by command id. cmdMu + // guards the records' mutable fields. + commands *memstore.Store[*commandRecord] + cmdMu sync.RWMutex instanceResolver InstanceResolver + // outputStore, when wired, receives Run Command output for commands that + // name an S3 bucket. + outputStore OutputStore // kmsCrypto, when wired via SetKMSCrypto, encrypts SecureString values through // real KMS. Nil stores them verbatim (library plaintext fallback). kmsCrypto KMSCrypto @@ -150,7 +155,7 @@ func (m *Mock) revealValue(ctx context.Context, v *version, withDecryption bool) func New(opts *config.Options) *Mock { return &Mock{ params: memstore.New[*paramData](), - commands: memstore.New[ssmdriver.CommandInvocation](), + commands: memstore.New[*commandRecord](), opts: opts, settings: newServiceSettings(opts.Clock.Now()), documents: memstore.New[*document](), diff --git a/server/aws/ssm/handler.go b/server/aws/ssm/handler.go index 90962a79d..94077b623 100644 --- a/server/aws/ssm/handler.go +++ b/server/aws/ssm/handler.go @@ -1,5 +1,6 @@ // Package ssm implements the AWS Systems Manager (SSM) JSON-RPC protocol as a -// server.Handler: Parameter Store, Run Command, Documents and tagging. Point +// server.Handler: Parameter Store, Run Command, managed nodes, Documents and +// tagging. Point // the real aws-sdk-go-v2 SSM client at a Server registered with this handler // and those operations work against the in-memory drivers. // @@ -64,13 +65,6 @@ func parameterOps() map[string]handlerFunc { } } -func runCommandOps() map[string]handlerFunc { - return map[string]handlerFunc{ - "SendCommand": (*Handler).sendCommand, - "GetCommandInvocation": (*Handler).getCommandInvocation, - } -} - // Matches returns true for SSM-shaped requests, identified by an X-Amz-Target // header of "AmazonSSM.". func (*Handler) Matches(r *http.Request) bool { diff --git a/server/aws/ssm/run_command.go b/server/aws/ssm/run_command.go index f0d5baf98..9fecdc9e9 100644 --- a/server/aws/ssm/run_command.go +++ b/server/aws/ssm/run_command.go @@ -2,67 +2,191 @@ package ssm import ( "net/http" + "strconv" + "time" - cerrors "github.com/stackshy/cloudemu/v2/errors" ssmnative "github.com/stackshy/cloudemu/v2/providers/aws/ssm/driver" "github.com/stackshy/cloudemu/v2/server/wire" ) +// Run Command page sizes. +const ( + maxResultsCommands = 50 // ListCommands / ListCommandInvocations: 1..50 + maxResultsNodes = 50 // DescribeInstanceInformation: 5..50 + defaultNodesPage = 10 + minNodesPage = 5 +) + +func runCommandOps() map[string]handlerFunc { + return map[string]handlerFunc{ + "SendCommand": (*Handler).sendCommand, + "GetCommandInvocation": (*Handler).getCommandInvocation, + "ListCommands": (*Handler).listCommands, + "ListCommandInvocations": (*Handler).listCommandInvocations, + "CancelCommand": (*Handler).cancelCommand, + "DescribeInstanceInformation": (*Handler).describeInstanceInformation, + } +} + type ssmTarget struct { Key string `json:"Key"` Values []string `json:"Values"` } +type notificationConfigJSON struct { + NotificationArn string `json:"NotificationArn,omitempty"` + NotificationEvents []string `json:"NotificationEvents,omitempty"` + NotificationType string `json:"NotificationType,omitempty"` +} + +type cloudWatchOutputJSON struct { + CloudWatchLogGroupName string `json:"CloudWatchLogGroupName,omitempty"` + CloudWatchOutputEnabled bool `json:"CloudWatchOutputEnabled"` +} + type sendCommandRequest struct { - InstanceIds []string `json:"InstanceIds"` - Targets []ssmTarget `json:"Targets"` - DocumentName string `json:"DocumentName"` - Comment string `json:"Comment"` - Parameters map[string][]string `json:"Parameters"` + InstanceIDs []string `json:"InstanceIds"` + Targets []ssmTarget `json:"Targets"` + DocumentName string `json:"DocumentName"` + DocumentVersion string `json:"DocumentVersion"` + Comment string `json:"Comment"` + Parameters map[string][]string `json:"Parameters"` + TimeoutSeconds int32 `json:"TimeoutSeconds"` + MaxConcurrency string `json:"MaxConcurrency"` + MaxErrors string `json:"MaxErrors"` + OutputS3Region string `json:"OutputS3Region"` + OutputS3BucketName string `json:"OutputS3BucketName"` + OutputS3KeyPrefix string `json:"OutputS3KeyPrefix"` + ServiceRoleArn string `json:"ServiceRoleArn"` + NotificationConfig *notificationConfigJSON `json:"NotificationConfig"` + CloudWatchOutputConfig *cloudWatchOutputJSON `json:"CloudWatchOutputConfig"` } type commandJSON struct { - CommandId string `json:"CommandId"` - DocumentName string `json:"DocumentName"` - Status string `json:"Status"` - InstanceIds []string `json:"InstanceIds"` - Targets []ssmTarget `json:"Targets,omitempty"` - Comment string `json:"Comment,omitempty"` + CommandID string `json:"CommandId"` + DocumentName string `json:"DocumentName"` + DocumentVersion string `json:"DocumentVersion"` + Comment string `json:"Comment"` + ExpiresAfter float64 `json:"ExpiresAfter"` + Parameters map[string][]string `json:"Parameters"` + InstanceIDs []string `json:"InstanceIds"` + Targets []ssmTarget `json:"Targets"` + RequestedDateTime float64 `json:"RequestedDateTime"` + Status string `json:"Status"` + StatusDetails string `json:"StatusDetails"` + OutputS3Region string `json:"OutputS3Region,omitempty"` + OutputS3BucketName string `json:"OutputS3BucketName"` + OutputS3KeyPrefix string `json:"OutputS3KeyPrefix"` + MaxConcurrency string `json:"MaxConcurrency"` + MaxErrors string `json:"MaxErrors"` + TargetCount int32 `json:"TargetCount"` + CompletedCount int32 `json:"CompletedCount"` + ErrorCount int32 `json:"ErrorCount"` + DeliveryTimedOutCount int32 `json:"DeliveryTimedOutCount"` + ServiceRole string `json:"ServiceRole"` + NotificationConfig notificationConfigJSON `json:"NotificationConfig"` + CloudWatchOutputConfig cloudWatchOutputJSON `json:"CloudWatchOutputConfig"` + TimeoutSeconds int32 `json:"TimeoutSeconds"` +} + +type commandPluginJSON struct { + Name string `json:"Name"` + Status string `json:"Status"` + StatusDetails string `json:"StatusDetails"` + ResponseCode int32 `json:"ResponseCode"` + ResponseStartDateTime *float64 `json:"ResponseStartDateTime,omitempty"` + ResponseFinishDateTime *float64 `json:"ResponseFinishDateTime,omitempty"` + Output string `json:"Output"` + StandardOutputURL string `json:"StandardOutputUrl"` + StandardErrorURL string `json:"StandardErrorUrl"` + OutputS3Region string `json:"OutputS3Region,omitempty"` + OutputS3BucketName string `json:"OutputS3BucketName"` + OutputS3KeyPrefix string `json:"OutputS3KeyPrefix"` } -type sendCommandResponse struct { - Command commandJSON `json:"Command"` +type commandInvocationJSON struct { + CommandID string `json:"CommandId"` + InstanceID string `json:"InstanceId"` + InstanceName string `json:"InstanceName"` + Comment string `json:"Comment"` + DocumentName string `json:"DocumentName"` + DocumentVersion string `json:"DocumentVersion"` + RequestedDateTime float64 `json:"RequestedDateTime"` + Status string `json:"Status"` + StatusDetails string `json:"StatusDetails"` + TraceOutput string `json:"TraceOutput"` + StandardOutputURL string `json:"StandardOutputUrl"` + StandardErrorURL string `json:"StandardErrorUrl"` + CommandPlugins []commandPluginJSON `json:"CommandPlugins"` + ServiceRole string `json:"ServiceRole"` + NotificationConfig notificationConfigJSON `json:"NotificationConfig"` + CloudWatchOutputConfig cloudWatchOutputJSON `json:"CloudWatchOutputConfig"` } type getCommandInvocationRequest struct { - CommandId string `json:"CommandId"` - InstanceId string `json:"InstanceId"` + CommandID string `json:"CommandId"` + InstanceID string `json:"InstanceId"` + PluginName string `json:"PluginName"` } type getCommandInvocationResponse struct { - CommandId string `json:"CommandId"` - InstanceId string `json:"InstanceId"` - DocumentName string `json:"DocumentName"` - Status string `json:"Status"` - StatusDetails string `json:"StatusDetails"` - ResponseCode int32 `json:"ResponseCode"` - StandardOutputContent string `json:"StandardOutputContent"` - StandardErrorContent string `json:"StandardErrorContent"` -} - -// runCommand reports whether the configured driver supports Run Command. -func (h *Handler) runCommand() (ssmnative.RunCommand, bool) { - rc, ok := h.store.(ssmnative.RunCommand) + CommandID string `json:"CommandId"` + InstanceID string `json:"InstanceId"` + Comment string `json:"Comment"` + DocumentName string `json:"DocumentName"` + DocumentVersion string `json:"DocumentVersion"` + PluginName string `json:"PluginName"` + ResponseCode int32 `json:"ResponseCode"` + ExecutionStartDateTime string `json:"ExecutionStartDateTime"` + ExecutionElapsedTime string `json:"ExecutionElapsedTime"` + ExecutionEndDateTime string `json:"ExecutionEndDateTime"` + Status string `json:"Status"` + StatusDetails string `json:"StatusDetails"` + StandardOutputContent string `json:"StandardOutputContent"` + StandardOutputURL string `json:"StandardOutputUrl"` + StandardErrorContent string `json:"StandardErrorContent"` + StandardErrorURL string `json:"StandardErrorUrl"` + CloudWatchOutputConfig cloudWatchOutputJSON `json:"CloudWatchOutputConfig"` +} - return rc, ok +type commandFilterJSON struct { + Key string `json:"key"` + Value string `json:"value"` } -func (h *Handler) sendCommand(w http.ResponseWriter, r *http.Request) { - store, ok := h.runCommand() +type listCommandsRequest struct { + CommandID string `json:"CommandId"` + InstanceID string `json:"InstanceId"` + MaxResults int32 `json:"MaxResults"` + NextToken string `json:"NextToken"` + Filters []commandFilterJSON `json:"Filters"` + Details bool `json:"Details"` +} + +func (q *listCommandsRequest) query() ssmnative.CommandQuery { + out := ssmnative.CommandQuery{CommandID: q.CommandID, InstanceID: q.InstanceID} + for _, f := range q.Filters { + out.Filters = append(out.Filters, ssmnative.CommandFilter{Key: f.Key, Value: f.Value}) + } + + return out +} + +// runCommand reports whether the configured driver supports Run Command, +// writing the error response when it does not. +func (h *Handler) runCommand(w http.ResponseWriter) (ssmnative.RunCommand, bool) { + rc, ok := h.store.(ssmnative.RunCommand) if !ok { wire.WriteJSONError(w, http.StatusBadRequest, "UnsupportedOperationException", "this driver does not support Run Command") + } + return rc, ok +} + +func (h *Handler) sendCommand(w http.ResponseWriter, r *http.Request) { + store, ok := h.runCommand(w) + if !ok { return } @@ -71,34 +195,74 @@ func (h *Handler) sendCommand(w http.ResponseWriter, r *http.Request) { return } - commandID, err := store.SendCommand(r.Context(), ssmnative.CommandConfig{ - InstanceIDs: req.InstanceIds, - Targets: toDriverTargets(req.Targets), - DocumentName: req.DocumentName, - Comment: req.Comment, - Parameters: req.Parameters, - }) + cfg := ssmnative.CommandConfig{ + InstanceIDs: req.InstanceIDs, Targets: toDriverTargets(req.Targets), DocumentName: req.DocumentName, + DocumentVersion: req.DocumentVersion, Comment: req.Comment, Parameters: req.Parameters, + TimeoutSeconds: req.TimeoutSeconds, MaxConcurrency: req.MaxConcurrency, MaxErrors: req.MaxErrors, + OutputS3Region: req.OutputS3Region, OutputS3BucketName: req.OutputS3BucketName, + OutputS3KeyPrefix: req.OutputS3KeyPrefix, ServiceRoleArn: req.ServiceRoleArn, + } + + if n := req.NotificationConfig; n != nil { + cfg.Notification = &ssmnative.NotificationConfig{ + NotificationArn: n.NotificationArn, NotificationEvents: n.NotificationEvents, NotificationType: n.NotificationType, + } + } + + if c := req.CloudWatchOutputConfig; c != nil { + cfg.CloudWatchOutput = &ssmnative.CloudWatchOutputConfig{ + LogGroupName: c.CloudWatchLogGroupName, OutputEnabled: c.CloudWatchOutputEnabled, + } + } + + cmd, err := store.SendCommand(r.Context(), cfg) if err != nil { - // The provider names the exception: InvalidInstanceId for a target - // that is not a managed instance, InvalidDocument for a document that - // does not resolve. + // The provider names the exception: InvalidInstanceId, InvalidDocument, + // InvalidDocumentVersion, InvalidParameters or ValidationException. writeErr(w, err) return } - // Real SSM reports the command as Pending here (it has been accepted, not - // finished), and the caller learns the outcome from GetCommandInvocation. - // Reporting Success would invite a caller to skip the poll it would need - // against the real service. - wire.WriteJSON(w, sendCommandResponse{Command: commandJSON{ - CommandId: commandID, - DocumentName: req.DocumentName, - Status: "Pending", - InstanceIds: req.InstanceIds, - Targets: req.Targets, - Comment: req.Comment, - }}) + wire.WriteJSON(w, map[string]any{"Command": toCommandJSON(cmd)}) +} + +func toCommandJSON(c *ssmnative.Command) commandJSON { + out := commandJSON{ + CommandID: c.CommandID, DocumentName: c.DocumentName, DocumentVersion: c.DocumentVersion, Comment: c.Comment, + ExpiresAfter: epoch(c.ExpiresAfter), Parameters: c.Parameters, InstanceIDs: nonNil(c.InstanceIDs), + Targets: toWireTargets(c.Targets), RequestedDateTime: epoch(c.RequestedDateTime), Status: c.Status, + StatusDetails: c.StatusDetails, OutputS3Region: c.OutputS3Region, OutputS3BucketName: c.OutputS3BucketName, + OutputS3KeyPrefix: c.OutputS3KeyPrefix, MaxConcurrency: c.MaxConcurrency, MaxErrors: c.MaxErrors, + TargetCount: c.TargetCount, CompletedCount: c.CompletedCount, ErrorCount: c.ErrorCount, + DeliveryTimedOutCount: c.DeliveryTimedOutCount, ServiceRole: c.ServiceRole, + NotificationConfig: toNotificationJSON(c.Notification), CloudWatchOutputConfig: toCloudWatchJSON(c.CloudWatchOutput), + TimeoutSeconds: c.TimeoutSeconds, + } + + if out.Parameters == nil { + out.Parameters = map[string][]string{} + } + + return out +} + +func toNotificationJSON(n *ssmnative.NotificationConfig) notificationConfigJSON { + if n == nil { + return notificationConfigJSON{NotificationArn: "", NotificationEvents: []string{}} + } + + return notificationConfigJSON{ + NotificationArn: n.NotificationArn, NotificationEvents: n.NotificationEvents, NotificationType: n.NotificationType, + } +} + +func toCloudWatchJSON(c *ssmnative.CloudWatchOutputConfig) cloudWatchOutputJSON { + if c == nil { + return cloudWatchOutputJSON{} + } + + return cloudWatchOutputJSON{CloudWatchLogGroupName: c.LogGroupName, CloudWatchOutputEnabled: c.OutputEnabled} } // toDriverTargets converts wire Targets to the driver's CommandTarget shape. @@ -115,12 +279,38 @@ func toDriverTargets(in []ssmTarget) []ssmnative.CommandTarget { return out } +func toWireTargets(in []ssmnative.CommandTarget) []ssmTarget { + out := make([]ssmTarget, 0, len(in)) + for _, t := range in { + out = append(out, ssmTarget{Key: t.Key, Values: t.Values}) + } + + return out +} + +// isoTime renders the ExecutionStart/EndDateTime strings; "" before the +// plugin starts. +func isoTime(t time.Time) string { + if t.IsZero() { + return "" + } + + return t.UTC().Format("2006-01-02T15:04:05.000Z") +} + +func optionalEpoch(t time.Time) *float64 { + if t.IsZero() { + return nil + } + + e := epoch(t) + + return &e +} + func (h *Handler) getCommandInvocation(w http.ResponseWriter, r *http.Request) { - store, ok := h.runCommand() + store, ok := h.runCommand(w) if !ok { - wire.WriteJSONError(w, http.StatusBadRequest, - "UnsupportedOperationException", "this driver does not support Run Command") - return } @@ -129,23 +319,266 @@ func (h *Handler) getCommandInvocation(w http.ResponseWriter, r *http.Request) { return } - inv, err := store.GetCommandInvocation(r.Context(), req.CommandId, req.InstanceId) + inv, err := store.GetCommandInvocation(r.Context(), req.CommandID, req.InstanceID, req.PluginName) if err != nil { - // AWS names this one specifically, and callers branch on it while - // polling a command that has not registered yet. - wire.WriteJSONError(w, http.StatusBadRequest, "InvocationDoesNotExist", cerrors.Message(err)) + // InvocationDoesNotExist for an unknown pair, which callers branch on + // while polling a command that has not registered yet. + writeErr(w, err) return } + elapsed := "" + if !inv.ExecutionStartTime.IsZero() && !inv.ExecutionEndTime.IsZero() { + elapsed = "PT" + strconv.FormatFloat(inv.ExecutionEndTime.Sub(inv.ExecutionStartTime).Seconds(), 'f', 3, 64) + "S" + } + wire.WriteJSON(w, getCommandInvocationResponse{ - CommandId: inv.CommandID, - InstanceId: inv.InstanceID, - DocumentName: inv.DocumentName, - Status: inv.Status, - StatusDetails: inv.Status, - ResponseCode: inv.ResponseCode, - StandardOutputContent: inv.Stdout, - StandardErrorContent: inv.Stderr, + CommandID: inv.CommandID, InstanceID: inv.InstanceID, Comment: inv.Comment, DocumentName: inv.DocumentName, + DocumentVersion: inv.DocumentVersion, PluginName: inv.PluginName, ResponseCode: inv.ResponseCode, + ExecutionStartDateTime: isoTime(inv.ExecutionStartTime), ExecutionElapsedTime: elapsed, + ExecutionEndDateTime: isoTime(inv.ExecutionEndTime), Status: inv.Status, StatusDetails: inv.StatusDetails, + StandardOutputContent: inv.Stdout, StandardOutputURL: inv.StandardOutputURL, + StandardErrorContent: inv.Stderr, StandardErrorURL: inv.StandardErrorURL, + CloudWatchOutputConfig: toCloudWatchJSON(inv.CloudWatchOutput), }) } + +// commandPage slices a result set by NextToken and MaxResults, answering +// InvalidNextToken for a token this server did not issue. +func commandPage(w http.ResponseWriter, token string, maxResults int32, limit, total int) (start, end int, next string, ok bool) { + if _, err := decodePageToken(token); err != nil { + wire.WriteJSONError(w, http.StatusBadRequest, "InvalidNextToken", "The specified token isn't valid.") + + return 0, 0, "", false + } + + start, end, next, err := pageWindow(token, maxResults, limit, total) + if err != nil { + writeErr(w, err) + + return 0, 0, "", false + } + + return start, end, next, true +} + +func (h *Handler) listCommands(w http.ResponseWriter, r *http.Request) { + store, ok := h.runCommand(w) + if !ok { + return + } + + var req listCommandsRequest + if !wire.DecodeJSON(w, r, &req) { + return + } + + cmds, err := store.ListCommands(r.Context(), req.query()) + if err != nil { + writeErr(w, err) + + return + } + + start, end, next, ok := commandPage(w, req.NextToken, req.MaxResults, maxResultsCommands, len(cmds)) + if !ok { + return + } + + out := make([]commandJSON, 0, end-start) + for i := start; i < end; i++ { + out = append(out, toCommandJSON(&cmds[i])) + } + + wire.WriteJSON(w, map[string]any{"Commands": out, "NextToken": omitEmpty(next)}) +} + +func (h *Handler) listCommandInvocations(w http.ResponseWriter, r *http.Request) { + store, ok := h.runCommand(w) + if !ok { + return + } + + var req listCommandsRequest + if !wire.DecodeJSON(w, r, &req) { + return + } + + invs, err := store.ListCommandInvocations(r.Context(), req.query(), req.Details) + if err != nil { + writeErr(w, err) + + return + } + + start, end, next, ok := commandPage(w, req.NextToken, req.MaxResults, maxResultsCommands, len(invs)) + if !ok { + return + } + + out := make([]commandInvocationJSON, 0, end-start) + for i := start; i < end; i++ { + out = append(out, toInvocationJSON(&invs[i])) + } + + wire.WriteJSON(w, map[string]any{"CommandInvocations": out, "NextToken": omitEmpty(next)}) +} + +func toInvocationJSON(inv *ssmnative.CommandInvocation) commandInvocationJSON { + out := commandInvocationJSON{ + CommandID: inv.CommandID, InstanceID: inv.InstanceID, InstanceName: inv.InstanceName, Comment: inv.Comment, + DocumentName: inv.DocumentName, DocumentVersion: inv.DocumentVersion, RequestedDateTime: epoch(inv.RequestedDateTime), + Status: inv.Status, StatusDetails: inv.StatusDetails, StandardOutputURL: inv.StandardOutputURL, + StandardErrorURL: inv.StandardErrorURL, CommandPlugins: make([]commandPluginJSON, 0, len(inv.Plugins)), + ServiceRole: inv.ServiceRole, NotificationConfig: toNotificationJSON(inv.Notification), + CloudWatchOutputConfig: toCloudWatchJSON(inv.CloudWatchOutput), + } + + for i := range inv.Plugins { + p := &inv.Plugins[i] + out.CommandPlugins = append(out.CommandPlugins, commandPluginJSON{ + Name: p.Name, Status: p.Status, StatusDetails: p.StatusDetails, ResponseCode: p.ResponseCode, + ResponseStartDateTime: optionalEpoch(p.ResponseStartDateTime), ResponseFinishDateTime: optionalEpoch(p.ResponseFinishDateTime), + Output: p.Output, StandardOutputURL: p.StandardOutputURL, StandardErrorURL: p.StandardErrorURL, + OutputS3Region: p.OutputS3Region, OutputS3BucketName: p.OutputS3BucketName, OutputS3KeyPrefix: p.OutputS3KeyPrefix, + }) + } + + return out +} + +func (h *Handler) cancelCommand(w http.ResponseWriter, r *http.Request) { + store, ok := h.runCommand(w) + if !ok { + return + } + + var req struct { + CommandID string `json:"CommandId"` + InstanceIDs []string `json:"InstanceIds"` + } + + if !wire.DecodeJSON(w, r, &req) { + return + } + + if err := store.CancelCommand(r.Context(), req.CommandID, req.InstanceIDs); err != nil { + writeErr(w, err) + + return + } + + wire.WriteJSON(w, map[string]any{}) +} + +type instanceInfoFilterJSON struct { + Key string `json:"key"` + ValueSet []string `json:"valueSet"` +} + +type instanceStringFilterJSON struct { + Key string `json:"Key"` + Values []string `json:"Values"` +} + +type instanceInformationJSON struct { + InstanceID string `json:"InstanceId"` + PingStatus string `json:"PingStatus"` + LastPingDateTime float64 `json:"LastPingDateTime"` + AgentVersion string `json:"AgentVersion"` + IsLatestVersion bool `json:"IsLatestVersion"` + PlatformType string `json:"PlatformType"` + PlatformName string `json:"PlatformName"` + PlatformVersion string `json:"PlatformVersion"` + ResourceType string `json:"ResourceType"` + IPAddress string `json:"IPAddress,omitempty"` + ComputerName string `json:"ComputerName,omitempty"` + SourceID string `json:"SourceId"` + SourceType string `json:"SourceType"` +} + +func (h *Handler) describeInstanceInformation(w http.ResponseWriter, r *http.Request) { + nodes, ok := h.store.(ssmnative.ManagedNodes) + if !ok { + wire.WriteJSONError(w, http.StatusBadRequest, + "UnsupportedOperationException", "this driver does not support managed nodes") + + return + } + + var req struct { + InstanceInformationFilterList []instanceInfoFilterJSON `json:"InstanceInformationFilterList"` + Filters []instanceStringFilterJSON `json:"Filters"` + MaxResults int32 `json:"MaxResults"` + NextToken string `json:"NextToken"` + } + + if !wire.DecodeJSON(w, r, &req) { + return + } + + filters, msg := nodeFilters(req.InstanceInformationFilterList, req.Filters, req.MaxResults) + if msg != "" { + wire.WriteJSONError(w, http.StatusBadRequest, "ValidationException", msg) + + return + } + + infos, err := nodes.DescribeInstanceInformation(r.Context(), filters) + if err != nil { + writeErr(w, err) + + return + } + + maxResults := req.MaxResults + if maxResults == 0 { + maxResults = defaultNodesPage + } + + start, end, next, ok := commandPage(w, req.NextToken, maxResults, maxResultsNodes, len(infos)) + if !ok { + return + } + + out := make([]instanceInformationJSON, 0, end-start) + + for i := start; i < end; i++ { + in := &infos[i] + out = append(out, instanceInformationJSON{ + InstanceID: in.InstanceID, PingStatus: in.PingStatus, LastPingDateTime: epoch(in.LastPingDateTime), + AgentVersion: in.AgentVersion, IsLatestVersion: in.IsLatestVersion, PlatformType: in.PlatformType, + PlatformName: in.PlatformName, PlatformVersion: in.PlatformVersion, ResourceType: in.ResourceType, + IPAddress: in.IPAddress, ComputerName: in.ComputerName, SourceID: in.SourceID, SourceType: in.SourceType, + }) + } + + wire.WriteJSON(w, map[string]any{"InstanceInformationList": out, "NextToken": omitEmpty(next)}) +} + +// nodeFilters merges the legacy and current DescribeInstanceInformation +// filter lists, returning a validation message for a request real SSM rejects. +func nodeFilters(legacy []instanceInfoFilterJSON, current []instanceStringFilterJSON, + maxResults int32, +) (filters []ssmnative.InstanceInformationFilter, msg string) { + if len(legacy) > 0 && len(current) > 0 { + return nil, "You can use either InstanceInformationFilterList or Filters, but not both." + } + + if maxResults != 0 && maxResults < minNodesPage { + return nil, "1 validation error detected: Value '" + strconv.Itoa(int(maxResults)) + + "' at 'maxResults' failed to satisfy constraint: Member must have value greater than or equal to 5" + } + + filters = make([]ssmnative.InstanceInformationFilter, 0, len(current)+len(legacy)) + for _, f := range current { + filters = append(filters, ssmnative.InstanceInformationFilter{Key: f.Key, Values: f.Values}) + } + + for _, f := range legacy { + filters = append(filters, ssmnative.InstanceInformationFilter{Key: f.Key, Values: f.ValueSet}) + } + + return filters, "" +} diff --git a/server/aws/ssm/run_command_lifecycle_test.go b/server/aws/ssm/run_command_lifecycle_test.go new file mode 100644 index 000000000..fb91d8f48 --- /dev/null +++ b/server/aws/ssm/run_command_lifecycle_test.go @@ -0,0 +1,624 @@ +package ssm_test + +import ( + "context" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/aws/aws-sdk-go-v2/aws" + awsconfig "github.com/aws/aws-sdk-go-v2/config" + "github.com/aws/aws-sdk-go-v2/credentials" + awsec2 "github.com/aws/aws-sdk-go-v2/service/ec2" + awss3 "github.com/aws/aws-sdk-go-v2/service/s3" + awsssm "github.com/aws/aws-sdk-go-v2/service/ssm" + ssmtypes "github.com/aws/aws-sdk-go-v2/service/ssm/types" + + "github.com/stackshy/cloudemu/v2" + "github.com/stackshy/cloudemu/v2/config" + awsserver "github.com/stackshy/cloudemu/v2/server/aws" +) + +type runCommandEnv struct { + ssm *awsssm.Client + ec2 *awsec2.Client + s3 *awss3.Client +} + +func newRunCommandEnv(t *testing.T, opts ...config.Option) runCommandEnv { + t.Helper() + + cloud := cloudemu.NewAWS(opts...) + ts := httptest.NewServer(awsserver.New(awsserver.DriversFrom(cloud))) + t.Cleanup(ts.Close) + + cfg, err := awsconfig.LoadDefaultConfig(context.Background(), + awsconfig.WithRegion("us-east-1"), + awsconfig.WithCredentialsProvider(credentials.NewStaticCredentialsProvider("test", "test", "")), + ) + if err != nil { + t.Fatalf("load config: %v", err) + } + + cfg.BaseEndpoint = aws.String(ts.URL) + + return runCommandEnv{ + ssm: awsssm.NewFromConfig(cfg), + ec2: awsec2.NewFromConfig(cfg), + s3: awss3.NewFromConfig(cfg, func(o *awss3.Options) { o.UsePathStyle = true }), + } +} + +func shellParams(cmd string) map[string][]string { + return map[string][]string{"commands": {cmd}} +} + +// typedDocJSON declares one parameter of each checked kind. +const typedDocJSON = `{ + "schemaVersion": "2.2", + "description": "typed parameters", + "parameters": { + "name": {"type": "String", "allowedPattern": "^[a-z]+$"}, + "color": {"type": "String", "default": "red", "allowedValues": ["red", "blue"]}, + "count": {"type": "Integer", "default": 1}, + "enabled": {"type": "Boolean", "default": true}, + "files": {"type": "StringList", "default": ["a"], "maxItems": 2}, + "labels": {"type": "StringMap", "default": {}} + }, + "mainSteps": [ + {"action": "aws:runShellScript", "name": "first", "inputs": {"runCommand": ["echo {{ name }}"]}}, + {"action": "aws:runShellScript", "name": "second", "inputs": {"runCommand": ["echo {{ color }}"]}} + ] +}` + +func TestSendCommandValidatesCatalogParameters(t *testing.T) { + ctx := context.Background() + env := newRunCommandEnv(t) + ids := runInstances(t, env.ec2, 1) + + cases := map[string]map[string][]string{ + "missing required commands": nil, + "undeclared parameter": {"commands": {"ls"}, "bogus": {"x"}}, + "pattern mismatch": {"commands": {"ls"}, "executionTimeout": {"999999"}}, + "two values for a String": {"commands": {"ls"}, "workingDirectory": {"/a", "/b"}}, + } + + for name, params := range cases { + _, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), Parameters: params, + }) + t.Run(name, func(t *testing.T) { wantAPIError(t, err, "InvalidParameters") }) + } + + if _, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), + Parameters: map[string][]string{"commands": {"ls", "pwd"}, "executionTimeout": {"600"}}, + }); err != nil { + t.Fatalf("valid send: %v", err) + } +} + +func TestSendCommandValidatesTypedParameters(t *testing.T) { + ctx := context.Background() + env := newRunCommandEnv(t) + ids := runInstances(t, env.ec2, 1) + + if _, err := env.ssm.CreateDocument(ctx, &awsssm.CreateDocumentInput{ + Name: aws.String("typed-doc"), Content: aws.String(typedDocJSON), + }); err != nil { + t.Fatalf("CreateDocument: %v", err) + } + + send := func(params map[string][]string) error { + _, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("typed-doc"), Parameters: params, + }) + + return err + } + + bad := map[string]map[string][]string{ + "required missing": {"color": {"red"}}, + "pattern": {"name": {"Bad1"}}, + "allowed values": {"name": {"ok"}, "color": {"green"}}, + "integer": {"name": {"ok"}, "count": {"many"}}, + "boolean": {"name": {"ok"}, "enabled": {"yes"}}, + "string list maxItems": {"name": {"ok"}, "files": {"a", "b", "c"}}, + "string map": {"name": {"ok"}, "labels": {"not-json"}}, + } + + for name, params := range bad { + err := send(params) + t.Run(name, func(t *testing.T) { wantAPIError(t, err, "InvalidParameters") }) + } + + good := map[string][]string{ + "name": {"ok"}, "color": {"blue"}, "count": {"7"}, "enabled": {"false"}, + "files": {"a", "b"}, "labels": {`{"k":"v"}`}, + } + if err := send(good); err != nil { + t.Fatalf("valid typed send: %v", err) + } + + if err := send(map[string][]string{"name": {"ok"}}); err != nil { + t.Fatalf("defaults should fill the optional parameters: %v", err) + } +} + +func TestSendCommandDocumentVersion(t *testing.T) { + ctx := context.Background() + env := newRunCommandEnv(t) + ids := runInstances(t, env.ec2, 1) + + if _, err := env.ssm.CreateDocument(ctx, &awsssm.CreateDocumentInput{ + Name: aws.String("versioned"), Content: aws.String(docJSON), + }); err != nil { + t.Fatalf("CreateDocument: %v", err) + } + + // Version 2 adds a required parameter; the default stays at 1. + if _, err := env.ssm.UpdateDocument(ctx, &awsssm.UpdateDocumentInput{ + Name: aws.String("versioned"), Content: aws.String(typedDocJSON), DocumentVersion: aws.String("$LATEST"), + }); err != nil { + t.Fatalf("UpdateDocument: %v", err) + } + + send := func(version string, params map[string][]string) (*awsssm.SendCommandOutput, error) { + in := &awsssm.SendCommandInput{InstanceIds: ids, DocumentName: aws.String("versioned"), Parameters: params} + if version != "" { + in.DocumentVersion = aws.String(version) + } + + return env.ssm.SendCommand(ctx, in) + } + + out, err := send("", nil) + if err != nil { + t.Fatalf("default version send: %v", err) + } + + if got := aws.ToString(out.Command.DocumentVersion); got != "$DEFAULT" { + t.Errorf("echoed DocumentVersion = %q, want $DEFAULT", got) + } + + if _, err := send("1", map[string][]string{"msg": {"yo"}}); err != nil { + t.Fatalf("version 1 send: %v", err) + } + + _, err = send("$LATEST", map[string][]string{"msg": {"yo"}}) + wantAPIError(t, err, "InvalidParameters") + + out, err = send("2", map[string][]string{"name": {"ok"}}) + if err != nil { + t.Fatalf("version 2 send: %v", err) + } + + if got := aws.ToString(out.Command.DocumentVersion); got != "2" { + t.Errorf("echoed DocumentVersion = %q, want 2", got) + } + + _, err = send("9", nil) + wantAPIError(t, err, "InvalidDocumentVersion") + + _, err = send("abc", nil) + wantAPIError(t, err, "ValidationException") +} + +func TestSendCommandEchoesRequest(t *testing.T) { + ctx := context.Background() + env := newRunCommandEnv(t) + ids := runInstances(t, env.ec2, 2) + + out, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), Parameters: shellParams("uptime"), + Comment: aws.String("nightly"), MaxConcurrency: aws.String("1"), MaxErrors: aws.String("10%"), + TimeoutSeconds: aws.Int32(120), OutputS3BucketName: aws.String("no-such-bucket"), + OutputS3KeyPrefix: aws.String("runs"), + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + c := out.Command + if aws.ToString(c.Comment) != "nightly" || aws.ToString(c.MaxConcurrency) != "1" || + aws.ToString(c.MaxErrors) != "10%" || c.TimeoutSeconds == nil || *c.TimeoutSeconds != 120 || + aws.ToString(c.OutputS3BucketName) != "no-such-bucket" || aws.ToString(c.OutputS3KeyPrefix) != "runs" || + c.TargetCount != 2 || c.RequestedDateTime == nil || c.ExpiresAfter == nil || + !c.ExpiresAfter.After(*c.RequestedDateTime) || len(c.Parameters["commands"]) != 1 { + t.Fatalf("SendCommand echo = %+v", c) + } + + if c.Status != ssmtypes.CommandStatusPending { + t.Errorf("SendCommand status = %s, want Pending", c.Status) + } + + for name, in := range map[string]*awsssm.SendCommandInput{ + "max concurrency": {MaxConcurrency: aws.String("0")}, + "max errors": {MaxErrors: aws.String("x")}, + "timeout": {TimeoutSeconds: aws.Int32(10)}, + "comment": {Comment: aws.String(strings.Repeat("c", 101))}, + "instance id": {InstanceIds: []string{"not-an-instance"}}, + } { + if in.InstanceIds == nil { + in.InstanceIds = ids + } + + in.DocumentName = aws.String("AWS-RunShellScript") + in.Parameters = shellParams("ls") + + _, err := env.ssm.SendCommand(ctx, in) + t.Run(name, func(t *testing.T) { wantAPIError(t, err, "ValidationException") }) + } +} + +func TestListCommandsAndInvocations(t *testing.T) { + ctx := context.Background() + env := newRunCommandEnv(t) + ids := runInstances(t, env.ec2, 2) + + first, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), Parameters: shellParams("one"), + }) + if err != nil { + t.Fatalf("send first: %v", err) + } + + if _, err := env.ssm.CreateDocument(ctx, &awsssm.CreateDocumentInput{ + Name: aws.String("typed-doc"), Content: aws.String(typedDocJSON), + }); err != nil { + t.Fatalf("CreateDocument: %v", err) + } + + second, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids[:1], DocumentName: aws.String("typed-doc"), Parameters: map[string][]string{"name": {"ok"}}, + }) + if err != nil { + t.Fatalf("send second: %v", err) + } + + all, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{}) + if err != nil { + t.Fatalf("ListCommands: %v", err) + } + + if len(all.Commands) != 2 { + t.Fatalf("ListCommands = %d commands, want 2", len(all.Commands)) + } + + for _, c := range all.Commands { + if c.Status != ssmtypes.CommandStatusSuccess || aws.ToString(c.StatusDetails) != "Success" || + c.CompletedCount != c.TargetCount { + t.Errorf("listed command = %+v, want settled Success", c) + } + } + + byID, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{CommandId: second.Command.CommandId}) + if err != nil || len(byID.Commands) != 1 || aws.ToString(byID.Commands[0].DocumentName) != "typed-doc" { + t.Fatalf("ListCommands by id = %+v, %v", byID, err) + } + + byInstance, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{InstanceId: aws.String(ids[1])}) + if err != nil || len(byInstance.Commands) != 1 || + aws.ToString(byInstance.Commands[0].CommandId) != aws.ToString(first.Command.CommandId) { + t.Fatalf("ListCommands by instance = %+v, %v", byInstance, err) + } + + byDoc, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{Filters: []ssmtypes.CommandFilter{ + {Key: ssmtypes.CommandFilterKeyDocumentName, Value: aws.String("AWS-RunShellScript")}, + {Key: ssmtypes.CommandFilterKeyExecutionStage, Value: aws.String("Complete")}, + }}) + if err != nil || len(byDoc.Commands) != 1 { + t.Fatalf("ListCommands by document = %+v, %v", byDoc, err) + } + + page, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{MaxResults: aws.Int32(1)}) + if err != nil || len(page.Commands) != 1 || page.NextToken == nil { + t.Fatalf("ListCommands page 1 = %+v, %v", page, err) + } + + page2, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{MaxResults: aws.Int32(1), NextToken: page.NextToken}) + if err != nil || len(page2.Commands) != 1 || page2.NextToken != nil || + aws.ToString(page2.Commands[0].CommandId) == aws.ToString(page.Commands[0].CommandId) { + t.Fatalf("ListCommands page 2 = %+v, %v", page2, err) + } + + _, err = env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{Filters: []ssmtypes.CommandFilter{ + {Key: "Nope", Value: aws.String("x")}, + }}) + wantAPIError(t, err, "InvalidFilterKey") + + _, err = env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{NextToken: aws.String("garbage!")}) + wantAPIError(t, err, "InvalidNextToken") + + invs, err := env.ssm.ListCommandInvocations(ctx, &awsssm.ListCommandInvocationsInput{ + CommandId: first.Command.CommandId, + }) + if err != nil || len(invs.CommandInvocations) != 2 { + t.Fatalf("ListCommandInvocations = %+v, %v", invs, err) + } + + for _, inv := range invs.CommandInvocations { + if inv.Status != ssmtypes.CommandInvocationStatusSuccess || len(inv.CommandPlugins) != 0 { + t.Errorf("invocation without details = %+v", inv) + } + } + + detailed, err := env.ssm.ListCommandInvocations(ctx, &awsssm.ListCommandInvocationsInput{ + CommandId: second.Command.CommandId, Details: true, + }) + if err != nil || len(detailed.CommandInvocations) != 1 { + t.Fatalf("ListCommandInvocations details = %+v, %v", detailed, err) + } + + plugins := detailed.CommandInvocations[0].CommandPlugins + if len(plugins) != 2 || aws.ToString(plugins[0].Name) != "first" || aws.ToString(plugins[1].Name) != "second" || + plugins[0].Status != ssmtypes.CommandPluginStatusSuccess || plugins[0].ResponseCode != 0 { + t.Fatalf("command plugins = %+v", plugins) + } + + _, err = env.ssm.ListCommandInvocations(ctx, &awsssm.ListCommandInvocationsInput{Filters: []ssmtypes.CommandFilter{ + {Key: ssmtypes.CommandFilterKeyExecutionStage, Value: aws.String("Complete")}, + }}) + wantAPIError(t, err, "InvalidFilterKey") + + step, err := env.ssm.GetCommandInvocation(ctx, &awsssm.GetCommandInvocationInput{ + CommandId: second.Command.CommandId, InstanceId: aws.String(ids[0]), PluginName: aws.String("second"), + }) + if err != nil || aws.ToString(step.PluginName) != "second" || step.Status != ssmtypes.CommandInvocationStatusSuccess { + t.Fatalf("GetCommandInvocation plugin = %+v, %v", step, err) + } + + _, err = env.ssm.GetCommandInvocation(ctx, &awsssm.GetCommandInvocationInput{ + CommandId: second.Command.CommandId, InstanceId: aws.String(ids[0]), PluginName: aws.String("third"), + }) + wantAPIError(t, err, "InvalidPluginName") +} + +// With async settle a command moves Pending -> InProgress -> Success, and +// every read reports the same state. +func TestRunCommandAsyncLifecycle(t *testing.T) { + ctx := context.Background() + fc := config.NewFakeClock(time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)) + env := newRunCommandEnv(t, config.WithClock(fc), config.WithAsyncSettle()) + ids := runInstances(t, env.ec2, 1) + fc.Advance(time.Minute) + + out, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), Parameters: shellParams("sleep 1"), + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + id := out.Command.CommandId + + observe := func(wantCmd ssmtypes.CommandStatus, wantInv ssmtypes.CommandInvocationStatus) *awsssm.GetCommandInvocationOutput { + t.Helper() + + inv, err := env.ssm.GetCommandInvocation(ctx, &awsssm.GetCommandInvocationInput{CommandId: id, InstanceId: aws.String(ids[0])}) + if err != nil || inv.Status != wantInv { + t.Fatalf("GetCommandInvocation = %+v, %v, want %s", inv, err, wantInv) + } + + listed, err := env.ssm.ListCommandInvocations(ctx, &awsssm.ListCommandInvocationsInput{CommandId: id}) + if err != nil || len(listed.CommandInvocations) != 1 || listed.CommandInvocations[0].Status != wantInv { + t.Fatalf("ListCommandInvocations = %+v, %v, want %s", listed, err, wantInv) + } + + cmds, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{CommandId: id}) + if err != nil || len(cmds.Commands) != 1 || cmds.Commands[0].Status != wantCmd { + t.Fatalf("ListCommands = %+v, %v, want %s", cmds, err, wantCmd) + } + + return inv + } + + inv := observe(ssmtypes.CommandStatusPending, ssmtypes.CommandInvocationStatusPending) + if inv.ResponseCode != -1 || aws.ToString(inv.ExecutionStartDateTime) != "" { + t.Errorf("pending invocation = %+v", inv) + } + + fc.Advance(1500 * time.Millisecond) + + inv = observe(ssmtypes.CommandStatusInProgress, ssmtypes.CommandInvocationStatusInProgress) + if inv.ResponseCode != -1 || aws.ToString(inv.ExecutionStartDateTime) == "" { + t.Errorf("in-progress invocation = %+v", inv) + } + + fc.Advance(time.Minute) + + inv = observe(ssmtypes.CommandStatusSuccess, ssmtypes.CommandInvocationStatusSuccess) + if inv.ResponseCode != 0 || aws.ToString(inv.ExecutionEndDateTime) == "" || aws.ToString(inv.ExecutionElapsedTime) == "" { + t.Errorf("finished invocation = %+v", inv) + } +} + +func TestCancelCommand(t *testing.T) { + ctx := context.Background() + fc := config.NewFakeClock(time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)) + env := newRunCommandEnv(t, config.WithClock(fc), config.WithAsyncSettle()) + ids := runInstances(t, env.ec2, 2) + fc.Advance(time.Minute) + + send := func() *string { + out, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), Parameters: shellParams("sleep 60"), + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + return out.Command.CommandId + } + + running := send() + fc.Advance(1500 * time.Millisecond) + + _, err := env.ssm.CancelCommand(ctx, &awsssm.CancelCommandInput{CommandId: aws.String("11111111-2222-3333-4444-555555555555")}) + wantAPIError(t, err, "InvalidCommandId") + + _, err = env.ssm.CancelCommand(ctx, &awsssm.CancelCommandInput{CommandId: running, InstanceIds: []string{"i-0123456789abcdef0"}}) + wantAPIError(t, err, "InvalidInstanceId") + + if _, err := env.ssm.CancelCommand(ctx, &awsssm.CancelCommandInput{CommandId: running}); err != nil { + t.Fatalf("CancelCommand: %v", err) + } + + fc.Advance(time.Minute) + + cmds, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{CommandId: running}) + if err != nil || cmds.Commands[0].Status != ssmtypes.CommandStatusCancelled || cmds.Commands[0].CompletedCount != 2 { + t.Fatalf("cancelled command = %+v, %v", cmds, err) + } + + for _, id := range ids { + inv, err := env.ssm.GetCommandInvocation(ctx, &awsssm.GetCommandInvocationInput{CommandId: running, InstanceId: aws.String(id)}) + if err != nil || inv.Status != ssmtypes.CommandInvocationStatusCancelled { + t.Fatalf("cancelled invocation = %+v, %v", inv, err) + } + } + + // A finished command is left as it is. + done := send() + fc.Advance(time.Minute) + + if _, err := env.ssm.CancelCommand(ctx, &awsssm.CancelCommandInput{CommandId: done}); err != nil { + t.Fatalf("CancelCommand on finished command: %v", err) + } + + cmds, err = env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{CommandId: done}) + if err != nil || cmds.Commands[0].Status != ssmtypes.CommandStatusSuccess { + t.Fatalf("finished command after cancel = %+v, %v", cmds, err) + } +} + +// executionTimeout shorter than the run makes the invocation time out. +func TestRunCommandExecutionTimeout(t *testing.T) { + ctx := context.Background() + fc := config.NewFakeClock(time.Date(2026, 1, 1, 0, 0, 0, 0, time.UTC)) + env := newRunCommandEnv(t, config.WithClock(fc), config.WithAsyncSettle()) + ids := runInstances(t, env.ec2, 1) + fc.Advance(time.Minute) + + out, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), + Parameters: map[string][]string{"commands": {"sleep 100"}, "executionTimeout": {"1"}}, + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + fc.Advance(time.Minute) + + inv, err := env.ssm.GetCommandInvocation(ctx, &awsssm.GetCommandInvocationInput{ + CommandId: out.Command.CommandId, InstanceId: aws.String(ids[0]), + }) + if err != nil || inv.Status != ssmtypes.CommandInvocationStatusTimedOut || + aws.ToString(inv.StatusDetails) != "ExecutionTimedOut" { + t.Fatalf("timed out invocation = %+v, %v", inv, err) + } + + cmds, err := env.ssm.ListCommands(ctx, &awsssm.ListCommandsInput{CommandId: out.Command.CommandId}) + if err != nil || cmds.Commands[0].Status != ssmtypes.CommandStatusTimedOut || cmds.Commands[0].ErrorCount != 1 { + t.Fatalf("timed out command = %+v, %v", cmds, err) + } +} + +func TestSendCommandWritesOutputToS3(t *testing.T) { + ctx := context.Background() + env := newRunCommandEnv(t) + ids := runInstances(t, env.ec2, 1) + + if _, err := env.s3.CreateBucket(ctx, &awss3.CreateBucketInput{Bucket: aws.String("run-output")}); err != nil { + t.Fatalf("CreateBucket: %v", err) + } + + out, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), Parameters: shellParams("hostname"), + OutputS3BucketName: aws.String("run-output"), OutputS3KeyPrefix: aws.String("logs"), + }) + if err != nil { + t.Fatalf("SendCommand: %v", err) + } + + id := aws.ToString(out.Command.CommandId) + wantKey := "logs/" + id + "/" + ids[0] + "/awsrunShellScript/0.awsrunShellScript/stdout" + + inv, err := env.ssm.GetCommandInvocation(ctx, &awsssm.GetCommandInvocationInput{ + CommandId: out.Command.CommandId, InstanceId: aws.String(ids[0]), + }) + if err != nil || !strings.HasSuffix(aws.ToString(inv.StandardOutputUrl), "/run-output/"+wantKey) { + t.Fatalf("StandardOutputUrl = %q, %v, want suffix %s", aws.ToString(inv.StandardOutputUrl), err, wantKey) + } + + objs, err := env.s3.ListObjectsV2(ctx, &awss3.ListObjectsV2Input{Bucket: aws.String("run-output"), Prefix: aws.String("logs/")}) + if err != nil { + t.Fatalf("ListObjectsV2: %v", err) + } + + found := false + for _, o := range objs.Contents { + found = found || aws.ToString(o.Key) == wantKey + } + + if !found { + t.Fatalf("output object %s not written; bucket has %+v", wantKey, objs.Contents) + } +} + +func TestSendCommandRejectsStoppedInstance(t *testing.T) { + ctx := context.Background() + env := newRunCommandEnv(t) + ids := runInstances(t, env.ec2, 1) + + if _, err := env.ec2.StopInstances(ctx, &awsec2.StopInstancesInput{InstanceIds: ids}); err != nil { + t.Fatalf("StopInstances: %v", err) + } + + _, err := env.ssm.SendCommand(ctx, &awsssm.SendCommandInput{ + InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), Parameters: shellParams("ls"), + }) + wantAPIError(t, err, "InvalidInstanceId") +} + +func TestDescribeInstanceInformation(t *testing.T) { + ctx := context.Background() + env := newRunCommandEnv(t) + ids := runInstances(t, env.ec2, 2) + + if _, err := env.ec2.StopInstances(ctx, &awsec2.StopInstancesInput{InstanceIds: ids[1:]}); err != nil { + t.Fatalf("StopInstances: %v", err) + } + + out, err := env.ssm.DescribeInstanceInformation(ctx, &awsssm.DescribeInstanceInformationInput{}) + if err != nil || len(out.InstanceInformationList) != 2 { + t.Fatalf("DescribeInstanceInformation = %+v, %v", out, err) + } + + status := map[string]ssmtypes.PingStatus{} + for _, info := range out.InstanceInformationList { + status[aws.ToString(info.InstanceId)] = info.PingStatus + + if info.ResourceType != ssmtypes.ResourceTypeEc2Instance || info.PlatformType != ssmtypes.PlatformTypeLinux || + aws.ToString(info.AgentVersion) == "" { + t.Errorf("instance information = %+v", info) + } + } + + if status[ids[0]] != ssmtypes.PingStatusOnline || status[ids[1]] != ssmtypes.PingStatusConnectionLost { + t.Fatalf("ping status = %v", status) + } + + online, err := env.ssm.DescribeInstanceInformation(ctx, &awsssm.DescribeInstanceInformationInput{ + Filters: []ssmtypes.InstanceInformationStringFilter{{Key: aws.String("PingStatus"), Values: []string{"Online"}}}, + }) + if err != nil || len(online.InstanceInformationList) != 1 || aws.ToString(online.InstanceInformationList[0].InstanceId) != ids[0] { + t.Fatalf("PingStatus filter = %+v, %v", online, err) + } + + _, err = env.ssm.DescribeInstanceInformation(ctx, &awsssm.DescribeInstanceInformationInput{ + Filters: []ssmtypes.InstanceInformationStringFilter{{Key: aws.String("Bogus"), Values: []string{"x"}}}, + }) + wantAPIError(t, err, "InvalidFilterKey") +} diff --git a/server/aws/ssm/runcommand_sdk_roundtrip_test.go b/server/aws/ssm/runcommand_sdk_roundtrip_test.go index df27dacd1..eb68a0b83 100644 --- a/server/aws/ssm/runcommand_sdk_roundtrip_test.go +++ b/server/aws/ssm/runcommand_sdk_roundtrip_test.go @@ -112,6 +112,7 @@ func TestRunCommandRegistersEveryTargetInstance(t *testing.T) { sent, err := c.SendCommand(ctx, &awsssm.SendCommandInput{ InstanceIds: ids, DocumentName: aws.String("AWS-RunShellScript"), + Parameters: map[string][]string{"commands": {"echo hi"}}, }) if err != nil { t.Fatalf("SendCommand: %v", err) @@ -183,6 +184,7 @@ func TestSendCommandTagTargets(t *testing.T) { sent, err := c.SendCommand(ctx, &awsssm.SendCommandInput{ DocumentName: aws.String("AWS-RunShellScript"), + Parameters: map[string][]string{"commands": {"echo hi"}}, Targets: []ssmtypes.Target{{ Key: aws.String("tag:Name"), Values: []string{"web"}, }}, @@ -224,6 +226,7 @@ func TestSendCommandTagTargetsNoMatch(t *testing.T) { sent, err := c.SendCommand(ctx, &awsssm.SendCommandInput{ DocumentName: aws.String("AWS-RunShellScript"), + Parameters: map[string][]string{"commands": {"echo hi"}}, Targets: []ssmtypes.Target{{ Key: aws.String("tag:Name"), Values: []string{"nonexistent"}, }}, @@ -251,6 +254,7 @@ func TestSendCommandUnsupportedTargetKeyMatchesNothing(t *testing.T) { sent, err := c.SendCommand(ctx, &awsssm.SendCommandInput{ DocumentName: aws.String("AWS-RunShellScript"), + Parameters: map[string][]string{"commands": {"echo hi"}}, Targets: []ssmtypes.Target{{ Key: aws.String("resource-groups:Name"), Values: []string{"my-group"}, }}, @@ -282,8 +286,9 @@ func TestSendCommandRejectsUnknownInstance(t *testing.T) { c, _ := newRunCommandClient(t) _, err := c.SendCommand(context.Background(), &awsssm.SendCommandInput{ - InstanceIds: []string{"i-doesnotexist"}, + InstanceIds: []string{"i-0000000000000dead"}, DocumentName: aws.String("AWS-RunShellScript"), + Parameters: map[string][]string{"commands": {"echo hi"}}, }) if err == nil { t.Fatal("SendCommand to an unknown instance should fail")