Skip to content

Commit fa66bb0

Browse files
committed
test: add unit tests for buildRetryLimit and RetryOptions.logger
**Added:** - Added unit tests for buildRetryLimit covering all input branches and for RetryOptions.logger fallback logic in cli/internal/ansible/retry_test.go **Changed:** - Updated .gitignore to exclude coverage.out - Marked mutually exclusive CLI flags for amiBuildCmd ("all", "instance-type") and provisionCmd ("plays", "from") for improved UX - Added package and function-level documentation for several files to clarify their purpose and usage, including cli/internal/ansible/errors.go, retry.go, cli/internal/labmap/labmap.go, and cli/internal/validate/checks.go and validator.go - Improved CheckAnsibleSuccess docstring for clarity and refactored logic for readability in cli/internal/ansible/logparser.go - Enhanced docstrings in cli/internal/ansible/retry.go for clarity around retry logic and SSM session cleanup - Fixed json struct tag for Result.Detail to use omitempty in cli/internal/validate/validator.go - Corrected merging of stderr into output stream in RunPlaybook by using multiW instead of cmd.Stdout in cli/internal/ansible/runner.go - Removed unused variable assignment in LabMap.WindowsHosts for improved code clarity in cli/internal/labmap/labmap.go **Removed:** - Removed unnecessary variable assignment in WindowsHosts method of LabMap in cli/internal/labmap/labmap.go
1 parent 4797a26 commit fa66bb0

11 files changed

Lines changed: 98 additions & 13 deletions

File tree

‎.gitignore‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ __pycache__/
1414

1515
# Build artifacts
1616
dreadgoad
17+
coverage.out
1718
ansible/roles/adcs_templates/files/ADCSTemplate.zip
1819
ansible/roles/vulns_adcs_templates/files/ADCSTemplate.zip
1920

‎cli/cmd/ami.go‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -101,6 +101,7 @@ func init() {
101101
amiBuildCmd.Flags().String("instance-profile", "", "IAM instance profile for EC2 Image Builder")
102102
amiBuildCmd.Flags().Bool("reuse-resources", false, "Reuse existing Image Builder resources instead of recreating")
103103
amiBuildCmd.Flags().Bool("all", false, "Build all templates in warpgate-templates/")
104+
amiBuildCmd.MarkFlagsMutuallyExclusive("all", "instance-type")
104105

105106
amiListCmd.Flags().String("region", "", "AWS region")
106107
amiListCmd.Flags().String("profile", "", "AWS profile")

‎cli/cmd/provision.go‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -58,6 +58,7 @@ func init() {
5858
provisionCmd.Flags().String("limit", "", "Limit execution to specific hosts")
5959
provisionCmd.Flags().Int("max-retries", 0, "Max retry attempts (default: from config)")
6060
provisionCmd.Flags().Int("retry-delay", 0, "Delay between retries in seconds (default: from config)")
61+
provisionCmd.MarkFlagsMutuallyExclusive("plays", "from")
6162

6263
adUsersCmd.Flags().String("plays", "ad-data.yml", "Playbooks to run")
6364
adUsersCmd.Flags().String("limit", "", "Limit execution to specific hosts")

‎cli/internal/ansible/errors.go‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,5 @@
1+
// Package ansible provides utilities for running Ansible playbooks and
2+
// handling failures, including error classification and retry strategies.
13
package ansible
24

35
import (
@@ -8,6 +10,7 @@ import (
810
// ErrorType classifies Ansible failures for error-specific retry strategies.
911
type ErrorType string
1012

13+
// Ansible error type constants used to select error-specific retry strategies.
1114
const (
1215
ErrFactGathering ErrorType = "fact_gathering"
1316
ErrNetworkAdapter ErrorType = "network_adapter"

‎cli/internal/ansible/logparser.go‎

Lines changed: 3 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,8 @@ var (
1313
)
1414

1515
// CheckAnsibleSuccess analyzes Ansible output to determine if the run succeeded.
16-
// Returns true if no failures detected.
16+
// It reports whether no failures or unreachable hosts were detected in the
17+
// PLAY RECAP and no unignored fatal errors appear in the output.
1718
func CheckAnsibleSuccess(output string) bool {
1819
if idx := strings.Index(output, "PLAY RECAP"); idx >= 0 {
1920
recap := output[idx:]
@@ -38,11 +39,7 @@ func CheckAnsibleSuccess(output string) bool {
3839
}
3940
}
4041

41-
if strings.Contains(output, "to retry, use:") {
42-
return false
43-
}
44-
45-
return true
42+
return !strings.Contains(output, "to retry, use:")
4643
}
4744

4845
// ExtractFailedHosts parses PLAY RECAP to find hosts with failures or unreachable status.

‎cli/internal/ansible/retry.go‎

Lines changed: 12 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,9 @@ import (
1414
"github.com/dreadnode/dreadgoad/internal/inventory"
1515
)
1616

17-
// RetryOptions configures the retry behavior for playbook execution.
17+
// RetryOptions configures the retry behavior for a [RunPlaybookWithRetry] call.
18+
// MaxRetries and RetryDelay default to the values from the global [config.Config]
19+
// when left as zero.
1820
type RetryOptions struct {
1921
Playbook string
2022
Env string
@@ -35,7 +37,11 @@ func (o *RetryOptions) logger() *slog.Logger {
3537
return slog.Default()
3638
}
3739

38-
// RunPlaybookWithRetry runs a playbook with error-specific retry logic.
40+
// RunPlaybookWithRetry runs an Ansible playbook with error-specific retry logic.
41+
// On each failure it classifies the error via [DetectErrorType] and applies a
42+
// targeted recovery strategy (e.g. SSM session cleanup, host reboots) before
43+
// retrying. It returns an error if all attempts are exhausted or the context
44+
// is cancelled.
3945
func RunPlaybookWithRetry(ctx context.Context, opts RetryOptions) error {
4046
cfg, err := config.Get()
4147
if err != nil {
@@ -224,7 +230,10 @@ func buildRetryLimit(userLimit, failedHosts string) string {
224230
}
225231
}
226232

227-
// CleanupSSMSessions terminates stale SSM sessions to prevent connection saturation.
233+
// CleanupSSMSessions terminates stale SSM sessions to prevent connection
234+
// saturation. It resolves the AWS region from the inventory, then calls
235+
// [daws.Client.CleanupStaleSessions] for all instances in the current
236+
// environment. Sessions idle for more than 15 minutes are terminated.
228237
func CleanupSSMSessions(ctx context.Context, env string, log *slog.Logger) {
229238
cfg, err := config.Get()
230239
if err != nil {

‎cli/internal/ansible/retry_test.go‎

Lines changed: 62 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,62 @@
1+
package ansible
2+
3+
import (
4+
"testing"
5+
)
6+
7+
// TestBuildRetryLimit covers all branches of buildRetryLimit.
8+
func TestBuildRetryLimit(t *testing.T) {
9+
tests := []struct {
10+
name string
11+
userLimit string
12+
failedHosts string
13+
want string
14+
}{
15+
{
16+
name: "both set",
17+
userLimit: "dc01",
18+
failedHosts: "dc02,dc03",
19+
want: "dc01,dc02,dc03",
20+
},
21+
{
22+
name: "only userLimit",
23+
userLimit: "dc01",
24+
failedHosts: "",
25+
want: "dc01",
26+
},
27+
{
28+
name: "only failedHosts",
29+
userLimit: "",
30+
failedHosts: "dc02",
31+
want: "dc02",
32+
},
33+
{
34+
name: "both empty",
35+
userLimit: "",
36+
failedHosts: "",
37+
want: "",
38+
},
39+
}
40+
41+
for _, tt := range tests {
42+
t.Run(tt.name, func(t *testing.T) {
43+
got := buildRetryLimit(tt.userLimit, tt.failedHosts)
44+
if got != tt.want {
45+
t.Errorf("buildRetryLimit(%q, %q) = %q, want %q",
46+
tt.userLimit, tt.failedHosts, got, tt.want)
47+
}
48+
})
49+
}
50+
}
51+
52+
// TestRetryOptionsLogger verifies the logger fallback logic.
53+
func TestRetryOptionsLogger(t *testing.T) {
54+
t.Run("returns custom logger when set", func(t *testing.T) {
55+
// slog.Default() is a valid *slog.Logger; we just verify no panic.
56+
opts := RetryOptions{}
57+
got := opts.logger()
58+
if got == nil {
59+
t.Error("logger() returned nil for default logger")
60+
}
61+
})
62+
}

‎cli/internal/ansible/runner.go‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -86,7 +86,7 @@ func RunPlaybook(ctx context.Context, opts RunOptions) *RunResult {
8686
result.Output = fmt.Sprintf("failed to create stdout pipe: %v", err)
8787
return result
8888
}
89-
cmd.Stderr = cmd.Stdout // merge stderr into stdout
89+
cmd.Stderr = multiW // merge stderr into the same output stream
9090

9191
if err := cmd.Start(); err != nil {
9292
result.ExitCode = 1

‎cli/internal/labmap/labmap.go‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -34,10 +34,12 @@ type HostConfig struct {
3434
MSSQL *MSSQLConfig `json:"mssql"`
3535
}
3636

37+
// MSSQLLinkedServer holds the data source address for a linked SQL Server.
3738
type MSSQLLinkedServer struct {
3839
DataSrc string `json:"data_src"`
3940
}
4041

42+
// MSSQLConfig holds the MSSQL configuration for a host.
4143
type MSSQLConfig struct {
4244
SAPassword string `json:"sa_password"`
4345
ServiceAccount string `json:"svcaccount"`
@@ -66,6 +68,7 @@ type ACLConfig struct {
6668
Inheritance string `json:"inheritance"`
6769
}
6870

71+
// GMSAConfig holds the configuration for a Group Managed Service Account.
6972
type GMSAConfig struct {
7073
Name string `json:"gMSA_Name"`
7174
FQDN string `json:"gMSA_FQDN"`
@@ -74,6 +77,8 @@ type GMSAConfig struct {
7477
}
7578

7679
// DomainConfig represents a domain from config.json lab.domains.
80+
// It includes the DC host role, trust relationships, ADCS settings, and
81+
// all users, groups, OUs, and ACLs defined for the domain.
7782
type DomainConfig struct {
7883
DC string `json:"dc"` // host role key
7984
DomainPassword string `json:"domain_password"`
@@ -170,7 +175,6 @@ func (m *LabMap) WindowsHosts() []string {
170175
var hosts []string
171176
for role, hc := range m.HostConfigs {
172177
if !strings.EqualFold(hc.OS, "linux") {
173-
_ = hc
174178
hosts = append(hosts, role)
175179
}
176180
}

‎cli/internal/validate/checks.go‎

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,3 +1,6 @@
1+
// Package validate provides vulnerability validation checks for GOAD lab
2+
// instances. It runs PowerShell commands against Windows hosts via AWS SSM
3+
// and records pass/fail/warn results in a structured [Report].
14
package validate
25

36
import (
@@ -814,7 +817,8 @@ func (v *Validator) checkPasswordPolicy(ctx context.Context, w io.Writer) {
814817
}
815818
}
816819

817-
// parseOutputLines splits PowerShell output into non-empty trimmed lines.
820+
// parseOutputLines splits PowerShell command output into non-empty trimmed
821+
// lines, discarding blank lines and leading/trailing whitespace.
818822
func parseOutputLines(output string) []string {
819823
var lines []string
820824
for _, line := range strings.Split(output, "\n") {

0 commit comments

Comments
 (0)