diff --git a/helper_test.go b/helper_test.go index 283306a..db6bdb9 100644 --- a/helper_test.go +++ b/helper_test.go @@ -38,5 +38,6 @@ func resetMocks() { issuesMu.Lock() issues = nil notices = nil + errorCount = 0 issuesMu.Unlock() } diff --git a/issues.go b/issues.go index f1a0076..356054d 100644 --- a/issues.go +++ b/issues.go @@ -13,9 +13,10 @@ import ( // when there's something to report. var ( - issuesMu sync.Mutex - issues []string - notices []string + issuesMu sync.Mutex + issues []string + notices []string + errorCount int ) func logIssue(level, msg string) { @@ -23,11 +24,20 @@ func logIssue(level, msg string) { defer issuesMu.Unlock() fmt.Printf(" [%s] %s\n", level, msg) issues = append(issues, fmt.Sprintf("[%s] %s", level, msg)) + if level == "ERROR" { + errorCount++ + } } func warn(msg string) { logIssue("WARN", msg) } func errLog(msg string) { logIssue("ERROR", msg) } +func hasErrors() bool { + issuesMu.Lock() + defer issuesMu.Unlock() + return errorCount > 0 +} + func notice(msg string) { issuesMu.Lock() defer issuesMu.Unlock() diff --git a/issues_test.go b/issues_test.go index 0ed99db..f8905d6 100644 --- a/issues_test.go +++ b/issues_test.go @@ -38,6 +38,7 @@ func TestIssuesLogging(t *testing.T) { issuesMu.Lock() issueLen := len(issues) noticeLen := len(notices) + errCount := errorCount issuesMu.Unlock() if issueLen != 2 { @@ -46,6 +47,12 @@ func TestIssuesLogging(t *testing.T) { if noticeLen != 1 { t.Errorf("expected 1 notice, got %d", noticeLen) } + if errCount != 1 { + t.Errorf("expected 1 error count, got %d", errCount) + } + if !hasErrors() { + t.Error("expected hasErrors() to return true after errLog call") + } } func TestWriteRunLog(t *testing.T) { @@ -83,6 +90,24 @@ func TestWriteRunLog(t *testing.T) { } } +func TestHasErrors(t *testing.T) { + defer resetMocks() + + if hasErrors() { + t.Error("expected hasErrors() false with no errors logged") + } + + warn("just a warning") + if hasErrors() { + t.Error("expected hasErrors() false after only a warning") + } + + errLog("a real error") + if !hasErrors() { + t.Error("expected hasErrors() true after errLog call") + } +} + func TestPrintNotices(t *testing.T) { defer resetMocks() diff --git a/main.go b/main.go index c5e19af..4b27ba0 100644 --- a/main.go +++ b/main.go @@ -169,6 +169,11 @@ func runMain(args []string) { printNotices() fmt.Println("\nDone.") + if hasErrors() { + osExit(1) + return + } + home, _ := os.UserHomeDir() zshrc := filepath.Join(home, ".zshrc") if hasCmd("zsh") {