diff --git a/go.mod b/go.mod index 55aed2e..5d5755d 100644 --- a/go.mod +++ b/go.mod @@ -27,6 +27,7 @@ require ( go.opentelemetry.io/otel/trace v1.34.0 golang.org/x/sync v0.12.0 golang.org/x/term v0.30.0 + gopkg.in/yaml.v3 v3.0.1 ) require ( @@ -76,5 +77,4 @@ require ( google.golang.org/genproto/googleapis/rpc v0.0.0-20250303144028-a0af3efb3deb // indirect google.golang.org/grpc v1.71.0 // indirect google.golang.org/protobuf v1.36.5 // indirect - gopkg.in/yaml.v3 v3.0.1 // indirect ) diff --git a/logger/console_test.go b/logger/console_test.go index 7cdb8b4..2e624f1 100644 --- a/logger/console_test.go +++ b/logger/console_test.go @@ -4,17 +4,19 @@ import ( "bytes" "log" "os" + "strings" "testing" "github.com/stretchr/testify/assert" ) func captureOutput(f func()) string { + prev := log.Writer() var buf bytes.Buffer log.SetOutput(&buf) + defer log.SetOutput(prev) f() - log.SetOutput(nil) - return buf.String() + return strings.TrimSpace(buf.String()) } func TestConsoleLogger(t *testing.T) { diff --git a/string/mask.go b/string/mask.go index 641399a..f6e0626 100644 --- a/string/mask.go +++ b/string/mask.go @@ -1,6 +1,7 @@ package string import ( + "encoding/json" "fmt" "net/url" "regexp" @@ -70,24 +71,75 @@ var isURL = regexp.MustCompile(`^(\w+)://`) var isEmail = regexp.MustCompile(`^[a-zA-Z0-9._%+-]+@[a-zA-Z0-9.-]+\.[a-zA-Z]{2,}$`) var isJWT = regexp.MustCompile(`^[a-zA-Z0-9-_]+\.[a-zA-Z0-9-_]+\.[a-zA-Z0-9-_]+$`) +// MaskValue masks sensitive information in the given argument. +func MaskValue(arg string) string { + if isURL.MatchString(arg) { + u, err := MaskURL(arg) + if err == nil { + return u + } else { + return Mask(arg) + } + } else if isEmail.MatchString(arg) { + return MaskEmail(arg) + } else if isJWT.MatchString(arg) { + return Mask(arg) + } else { + return arg + } +} + // MaskArguments masks sensitive information in the given arguments. func MaskArguments(args []string) []string { masked := make([]string, len(args)) for i, arg := range args { - if isURL.MatchString(arg) { - u, err := MaskURL(arg) - if err == nil { - masked[i] = u - } else { - masked[i] = Mask(arg) - } - } else if isEmail.MatchString(arg) { - masked[i] = MaskEmail(arg) - } else if isJWT.MatchString(arg) { - masked[i] = Mask(arg) - } else { - masked[i] = arg - } + masked[i] = MaskValue(arg) } return masked } + +// MaskedString is a custom string type that masks its value when formatted or text-marshaled. +type MaskedString string + +// Text returns the unmasked text value. +func (ms MaskedString) Text() string { + return string(ms) +} + +// Bytes returns the unmasked byte slice value. +func (ms MaskedString) Bytes() []byte { + return []byte(ms.Text()) +} + +// String implements fmt.Stringer to return a masked representation. +func (ms MaskedString) String() string { + if len(ms) == 0 { + return "" + } + return Mask(string(ms)) +} + +// MarshalText implements encoding.TextMarshaler for masked text output. +func (ms MaskedString) MarshalText() ([]byte, error) { + return []byte(ms.String()), nil +} + +// MarshalJSON implements json.Marshaler for real (unmasked) JSON output. +func (ms MaskedString) MarshalJSON() ([]byte, error) { + return json.Marshal(string(ms)) +} + +// MarshalYAML implements yaml.Marshaler for real (unmasked) YAML output. +func (ms MaskedString) MarshalYAML() (any, error) { + return string(ms), nil +} + +// GoString implements fmt.GoStringer so %#v also prints masked. +func (ms MaskedString) GoString() string { + return ms.String() +} + +// NewMaskedString returns a string using the special type MaskedString. +func NewMaskedString(s string) MaskedString { + return MaskedString(s) +} diff --git a/string/mask_test.go b/string/mask_test.go index 33a68ef..b4db6ed 100644 --- a/string/mask_test.go +++ b/string/mask_test.go @@ -1,10 +1,16 @@ package string import ( + "bytes" + "encoding/json" "fmt" + "log" + "strings" "testing" + "github.com/agentuity/go-common/logger" "github.com/stretchr/testify/assert" + "gopkg.in/yaml.v3" ) func TestMasking(t *testing.T) { @@ -95,3 +101,406 @@ func TestMaskedEmail(t *testing.T) { assert.Equal(t, test.expected, result) } } + +func TestMaskedString(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + {"empty string", "", ""}, + {"single char", "a", "*"}, + {"short string", "abc", "a**"}, + {"medium string", "password", "pass****"}, + {"long string", "verylongpassword", "verylong********"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ms := NewMaskedString(tt.input) + assert.Equal(t, tt.want, ms.String()) + }) + } +} + +func TestMaskedString_MarshalText(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + {"empty string", "", ""}, + {"single char", "a", "*"}, + {"password", "secret123", "secr*****"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ms := NewMaskedString(tt.input) + text, err := ms.MarshalText() + assert.NoError(t, err) + assert.Equal(t, tt.want, string(text)) + }) + } +} + +func TestMaskedString_MarshalJSON(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + {"empty string", "", `""`}, + {"password", "secret123", `"secret123"`}, + {"with quotes", `test"quote`, `"test\"quote"`}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ms := NewMaskedString(tt.input) + json, err := ms.MarshalJSON() + assert.NoError(t, err) + assert.Equal(t, tt.want, string(json)) + }) + } +} + +func TestMaskedString_MarshalYAML(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + {"empty string", "", ""}, + {"password", "secret123", "secret123"}, + {"with special chars", "test@value", "test@value"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ms := NewMaskedString(tt.input) + yaml, err := ms.MarshalYAML() + assert.NoError(t, err) + assert.Equal(t, tt.want, yaml) + }) + } +} + +func TestMaskedString_Text(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + {"empty string", "", ""}, + {"single char", "a", "a"}, + {"password", "secret123", "secret123"}, + {"long text", "this is a very long secret password", "this is a very long secret password"}, + {"with special chars", "test@#$%^&*()", "test@#$%^&*()"}, + {"with quotes", `test"quote'mixed`, `test"quote'mixed`}, + {"with newlines", "line1\nline2\nline3", "line1\nline2\nline3"}, + {"unicode", "café🔒", "café🔒"}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + ms := NewMaskedString(tt.input) + text := ms.Text() + assert.Equal(t, tt.want, text) + }) + } +} + +func TestMaskedString_Behavior_Comparison(t *testing.T) { + testCases := []string{ + "", + "a", + "password123", + "very_long_secret_value_that_should_be_masked", + } + + for _, input := range testCases { + t.Run(fmt.Sprintf("input_%s", input), func(t *testing.T) { + ms := NewMaskedString(input) + + // Text() should return unmasked original value + assert.Equal(t, input, ms.Text()) + + // String() should return masked value (unless empty) + if input == "" { + assert.Equal(t, "", ms.String()) + } else { + assert.Equal(t, Mask(input), ms.String()) + } + + // JSON should serialize unmasked value + jsonBytes, err := ms.MarshalJSON() + assert.NoError(t, err) + expectedJSON, _ := json.Marshal(input) + assert.Equal(t, expectedJSON, jsonBytes) + + // YAML should return unmasked value + yamlVal, err := ms.MarshalYAML() + assert.NoError(t, err) + assert.Equal(t, input, yamlVal) + }) + } +} + +func TestMaskedGoString_Behavior_Comparison(t *testing.T) { + testCases := []string{ + "", + "a", + "password123", + "very_long_secret_value_that_should_be_masked", + } + + for _, input := range testCases { + t.Run(fmt.Sprintf("input_%#v", input), func(t *testing.T) { + ms := NewMaskedString(input) + val := fmt.Sprintf("%#v", ms) + assert.Equal(t, ms.String(), val) + }) + } +} + +func TestMaskedSprintf_Behavior_Comparison(t *testing.T) { + testCases := []string{ + "", + "a", + "password123", + "very_long_secret_value_that_should_be_masked", + } + + for _, input := range testCases { + t.Run(fmt.Sprintf("input_%v", input), func(t *testing.T) { + ms := NewMaskedString(input) + val := fmt.Sprintf("%v", ms) + assert.Equal(t, ms.String(), val) + }) + } +} + +func TestMaskedSprintfs_Behavior_Comparison(t *testing.T) { + testCases := []string{ + "", + "a", + "password123", + "very_long_secret_value_that_should_be_masked", + } + + for _, input := range testCases { + t.Run(fmt.Sprintf("input_%s", input), func(t *testing.T) { + ms := NewMaskedString(input) + val := fmt.Sprintf("%s", ms) + assert.Equal(t, ms.String(), val) + }) + } +} + +func captureOutput(f func()) string { + prev := log.Writer() + var buf bytes.Buffer + log.SetOutput(&buf) + defer log.SetOutput(prev) + f() + return strings.TrimSpace(buf.String()) +} + +func TestMaskedLogger(t *testing.T) { + testCases := []string{ + "", + "a", + "password123", + "very_long_secret_value_that_should_be_masked", + } + + for _, input := range testCases { + t.Run(fmt.Sprintf("logger %s", input), func(t *testing.T) { + ms := NewMaskedString(input) + log := logger.NewConsoleLogger() + output := captureOutput(func() { + log.Info("msg: %s", ms) + }) + assert.Contains(t, output, ms.String()) + }) + } +} + +type TestStruct struct { + Value MaskedString `json:"value" yaml:"value"` + Name string `json:"name" yaml:"name"` + Password MaskedString `json:"password" yaml:"password"` +} + +func TestMaskedString_JSON_Unmarshaling(t *testing.T) { + tests := []struct { + name string + jsonStr string + expected TestStruct + }{ + { + name: "basic values", + jsonStr: `{"value":"secret123","name":"test","password":"mypassword"}`, + expected: TestStruct{ + Value: NewMaskedString("secret123"), + Name: "test", + Password: NewMaskedString("mypassword"), + }, + }, + { + name: "empty values", + jsonStr: `{"value":"","name":"","password":""}`, + expected: TestStruct{ + Value: NewMaskedString(""), + Name: "", + Password: NewMaskedString(""), + }, + }, + { + name: "with special characters", + jsonStr: `{"value":"test@#$%","name":"user","password":"p@ssw0rd!"}`, + expected: TestStruct{ + Value: NewMaskedString("test@#$%"), + Name: "user", + Password: NewMaskedString("p@ssw0rd!"), + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var result TestStruct + err := json.Unmarshal([]byte(tt.jsonStr), &result) + assert.NoError(t, err) + + // Check that values were unmarshaled correctly + assert.Equal(t, tt.expected.Value.Text(), result.Value.Text()) + assert.Equal(t, tt.expected.Name, result.Name) + assert.Equal(t, tt.expected.Password.Text(), result.Password.Text()) + + // Verify that String() still returns masked values + if tt.expected.Value.Text() != "" { + assert.Equal(t, Mask(tt.expected.Value.Text()), result.Value.String()) + } + if tt.expected.Password.Text() != "" { + assert.Equal(t, Mask(tt.expected.Password.Text()), result.Password.String()) + } + }) + } +} + +func TestMaskedString_YAML_Unmarshaling(t *testing.T) { + tests := []struct { + name string + yamlStr string + expected TestStruct + }{ + { + name: "basic values", + yamlStr: `value: secret123 +name: test +password: mypassword`, + expected: TestStruct{ + Value: NewMaskedString("secret123"), + Name: "test", + Password: NewMaskedString("mypassword"), + }, + }, + { + name: "empty values", + yamlStr: `value: "" +name: "" +password: ""`, + expected: TestStruct{ + Value: NewMaskedString(""), + Name: "", + Password: NewMaskedString(""), + }, + }, + { + name: "with special characters", + yamlStr: `value: "test@#$%" +name: user +password: "p@ssw0rd!"`, + expected: TestStruct{ + Value: NewMaskedString("test@#$%"), + Name: "user", + Password: NewMaskedString("p@ssw0rd!"), + }, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + var result TestStruct + err := yaml.Unmarshal([]byte(tt.yamlStr), &result) + assert.NoError(t, err) + + // Check that values were unmarshaled correctly + assert.Equal(t, tt.expected.Value.Text(), result.Value.Text()) + assert.Equal(t, tt.expected.Name, result.Name) + assert.Equal(t, tt.expected.Password.Text(), result.Password.Text()) + + // Verify that String() still returns masked values + if tt.expected.Value.Text() != "" { + assert.Equal(t, Mask(tt.expected.Value.Text()), result.Value.String()) + } + if tt.expected.Password.Text() != "" { + assert.Equal(t, Mask(tt.expected.Password.Text()), result.Password.String()) + } + }) + } +} + +func TestMaskedString_RoundTrip_JSON(t *testing.T) { + original := TestStruct{ + Value: NewMaskedString("secret_value"), + Name: "testuser", + Password: NewMaskedString("super_secret_password"), + } + + // Marshal to JSON + jsonData, err := json.Marshal(original) + assert.NoError(t, err) + + // Unmarshal back + var unmarshaled TestStruct + err = json.Unmarshal(jsonData, &unmarshaled) + assert.NoError(t, err) + + // Verify round-trip worked correctly + assert.Equal(t, original.Value.Text(), unmarshaled.Value.Text()) + assert.Equal(t, original.Name, unmarshaled.Name) + assert.Equal(t, original.Password.Text(), unmarshaled.Password.Text()) + + // Verify masking still works + assert.Equal(t, original.Value.String(), unmarshaled.Value.String()) + assert.Equal(t, original.Password.String(), unmarshaled.Password.String()) +} + +func TestMaskedString_RoundTrip_YAML(t *testing.T) { + original := TestStruct{ + Value: NewMaskedString("secret_value"), + Name: "testuser", + Password: NewMaskedString("super_secret_password"), + } + + // Marshal to YAML + yamlData, err := yaml.Marshal(original) + assert.NoError(t, err) + + // Unmarshal back + var unmarshaled TestStruct + err = yaml.Unmarshal(yamlData, &unmarshaled) + assert.NoError(t, err) + + // Verify round-trip worked correctly + assert.Equal(t, original.Value.Text(), unmarshaled.Value.Text()) + assert.Equal(t, original.Name, unmarshaled.Name) + assert.Equal(t, original.Password.Text(), unmarshaled.Password.Text()) + + // Verify masking still works + assert.Equal(t, original.Value.String(), unmarshaled.Value.String()) + assert.Equal(t, original.Password.String(), unmarshaled.Password.String()) +} diff --git a/string/string.go b/string/string.go index 581ac4f..546767d 100644 --- a/string/string.go +++ b/string/string.go @@ -1,16 +1,16 @@ package string -import "strings" +import ( + "strings" + + "github.com/agentuity/go-common/sys" +) // StringPointer will set the pointer to nil if the string is not nil but an empty string func StringPointer(v string) *string { if v != "" { nv := strings.TrimSpace(v) - if nv == "" { - return nil - } else { - return &nv - } + return sys.Ptr(nv) } return nil } @@ -22,7 +22,7 @@ func ClearEmptyStringPointer(v *string) *string { if nv == "" { return nil } else { - return &nv + return sys.Ptr(nv) } } return nil diff --git a/sys/net.go b/sys/net.go index 3df151c..f2ffa90 100644 --- a/sys/net.go +++ b/sys/net.go @@ -2,6 +2,7 @@ package sys import ( "net" + neturl "net/url" "strings" ) @@ -18,8 +19,22 @@ func GetFreePort() (port int, err error) { return } -// IsLocalhost returns true if the URL is localhost or 127.0.0.1 or 0.0.0.0. +// IsLocalhost returns true if the input points to localhost/loopback or an unspecified address. func IsLocalhost(url string) bool { - // technically 127.0.0.0 – 127.255.255.255 is the loopback range but most people use 127.0.0.1 - return strings.Contains(url, "localhost") || strings.Contains(url, "127.0.0.1") || strings.Contains(url, "0.0.0.0") + // Accept either a full URL or a bare host[:port]. + host := url + if u, err := neturl.Parse(url); err == nil && u.Host != "" { + host = u.Hostname() // strips [] for IPv6 + } else if h, _, err := net.SplitHostPort(url); err == nil { + host = h + } else { + host = strings.Trim(host, "[]") + } + if strings.EqualFold(host, "localhost") { + return true + } + if ip := net.ParseIP(host); ip != nil { + return ip.IsLoopback() || ip.IsUnspecified() // 127/8, ::1, 0.0.0.0, :: + } + return false } diff --git a/sys/net_test.go b/sys/net_test.go index f019910..870b674 100644 --- a/sys/net_test.go +++ b/sys/net_test.go @@ -31,6 +31,8 @@ func TestIsLocalhost(t *testing.T) { {"https://127.0.0.1", true}, {"http://0.0.0.0:8000", true}, {"https://0.0.0.0", true}, + {"https://[::1]", true}, + {"https://[::]", true}, {"http://example.com", false}, {"https://192.168.1.1", false}, {"", false}, diff --git a/sys/pointer.go b/sys/pointer.go new file mode 100644 index 0000000..1913451 --- /dev/null +++ b/sys/pointer.go @@ -0,0 +1,54 @@ +package sys + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +// Ptr returns a pointer to the given value. +func Ptr[T any](v T) *T { + return &v +} + +func TestPtr(t *testing.T) { + // Test with string + str := "hello" + strPtr := Ptr(str) + assert.NotNil(t, strPtr) + assert.Equal(t, str, *strPtr) + + // Test with int + num := 42 + numPtr := Ptr(num) + assert.NotNil(t, numPtr) + assert.Equal(t, num, *numPtr) + + // Test with bool + flag := true + flagPtr := Ptr(flag) + assert.NotNil(t, flagPtr) + assert.Equal(t, flag, *flagPtr) + + // Test with struct + type testStruct struct { + Name string + Age int + } + s := testStruct{Name: "test", Age: 25} + sPtr := Ptr(s) + assert.NotNil(t, sPtr) + assert.Equal(t, s, *sPtr) + + // Test with slice + slice := []int{1, 2, 3} + slicePtr := Ptr(slice) + assert.NotNil(t, slicePtr) + assert.Equal(t, slice, *slicePtr) + + // Test with nil interface + var nilInterface interface{} + nilPtr := Ptr(nilInterface) + assert.NotNil(t, nilPtr) + assert.Equal(t, nilInterface, *nilPtr) +} diff --git a/sys/result.go b/sys/result.go new file mode 100644 index 0000000..728309e --- /dev/null +++ b/sys/result.go @@ -0,0 +1,68 @@ +package sys + +import ( + "errors" + "strings" +) + +// Result represents a value that can be either successful (Ok) or an error (Err), +// similar to Rust's Result type. +type Result[T any] struct { + Ok T + Err error +} + +// IsOk returns true if the Result contains a successful value (no error). +func (r Result[T]) IsOk() bool { + return r.Err == nil +} + +// IsErr returns true if the Result contains an error. +func (r Result[T]) IsErr(checks ...error) bool { + // Fast-path: if no error, return false immediately + if r.Err == nil { + return false + } + + if len(checks) == 0 { + return r.Err != nil + } + for _, err := range checks { + // Skip nil checks to avoid calling errors.Is with nil target + if err == nil { + continue + } + if errors.Is(r.Err, err) { + return true + } + } + return false +} + +// IsErrMatches returns true if the Result contains an error containing any of the given string values. +func (r Result[T]) IsErrMatches(checks ...string) bool { + if len(checks) == 0 { + return r.Err != nil + } + if r.Err == nil { + return false + } + val := r.Err.Error() + for _, err := range checks { + if strings.Contains(val, err) { + return true + } + } + return false +} + +// Ok creates a new Result with a successful value. +func Ok[T any](value T) Result[T] { + return Result[T]{Ok: value, Err: nil} +} + +// Err creates a new Result with an error. +func Err[T any](err error) Result[T] { + var zero T + return Result[T]{Ok: zero, Err: err} +} diff --git a/sys/result_test.go b/sys/result_test.go new file mode 100644 index 0000000..cc9f142 --- /dev/null +++ b/sys/result_test.go @@ -0,0 +1,134 @@ +package sys + +import ( + "errors" + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestResult_IsOk(t *testing.T) { + tests := []struct { + name string + result Result[string] + expected bool + }{ + { + name: "Ok result", + result: Result[string]{Ok: "success", Err: nil}, + expected: true, + }, + { + name: "Error result", + result: Result[string]{Ok: "", Err: errors.New("error")}, + expected: false, + }, + { + name: "Empty result with nil error", + result: Result[string]{Ok: "", Err: nil}, + expected: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.expected, tt.result.IsOk()) + }) + } +} + +func TestResult_IsErr(t *testing.T) { + tests := []struct { + name string + result Result[int] + expected bool + }{ + { + name: "Ok result", + result: Result[int]{Ok: 42, Err: nil}, + expected: false, + }, + { + name: "Error result", + result: Result[int]{Ok: 0, Err: errors.New("error")}, + expected: true, + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + assert.Equal(t, tt.expected, tt.result.IsErr()) + }) + } +} + +// Helper functions for creating Results +func TestOk(t *testing.T) { + result := Ok(42) + assert.True(t, result.IsOk()) + assert.Equal(t, 42, result.Ok) + assert.Nil(t, result.Err) +} + +func TestErr(t *testing.T) { + err := errors.New("test error") + result := Err[string](err) + assert.True(t, result.IsErr()) + assert.Equal(t, "", result.Ok) + assert.Equal(t, err, result.Err) +} + +func TestErrMultiple(t *testing.T) { + err := errors.New("test error") + err2 := errors.New("test error2") + result := Err[string](err) + assert.True(t, result.IsErr(err2, err)) + assert.Equal(t, "", result.Ok) + assert.Equal(t, err, result.Err) + result2 := Err[string](err2) + assert.True(t, result2.IsErr(err, err2)) + assert.Equal(t, "", result2.Ok) + assert.Equal(t, err2, result2.Err) +} + +func TestErrMatch(t *testing.T) { + err := errors.New("test error") + result := Err[string](err) + assert.True(t, result.IsErrMatches("test")) + assert.Equal(t, "", result.Ok) + assert.Equal(t, err, result.Err) +} + +func TestErrMatchMultiple(t *testing.T) { + err := errors.New("test error") + result := Err[string](err) + assert.True(t, result.IsErrMatches("foo", "error")) + assert.Equal(t, "", result.Ok) + assert.Equal(t, err, result.Err) +} + +// Test with different types +func TestResult_WithDifferentTypes(t *testing.T) { + // Test with struct + type Person struct { + Name string + Age int + } + + person := Person{Name: "John", Age: 30} + result := Ok(person) + assert.True(t, result.IsOk()) + assert.Equal(t, person, result.Ok) + + // Test with slice + numbers := []int{1, 2, 3} + sliceResult := Ok(numbers) + assert.True(t, sliceResult.IsOk()) + assert.Equal(t, numbers, sliceResult.Ok) + + // Test with pointer + str := "hello" + ptrResult := Ok(&str) + assert.True(t, ptrResult.IsOk()) + assert.Equal(t, &str, ptrResult.Ok) +}