From f6fbd330be6e308d528e1f0bd7972c358b41370f Mon Sep 17 00:00:00 2001 From: Marc Brooks Date: Mon, 12 Jan 2026 09:03:52 -0600 Subject: [PATCH] Test coverage improvements (#748) * Test coverage improvements Implements #397 and #591, might help with #473 Test coverage for core is 97.6% with only real fringe cases remaining. Improve `Fields` `isNilValue()` portability. Added a new global handler `FatalExitFunc` to allow intercepting `Fatal()` messages. Fixed CBOR float constants (removed the CBOR prefix) Added tests for: - global `FatalExitFunc` to allow intercepting Fatal messages (both for testing and public use). - `Logger` - `DisableSampling` - `.With()` copying existing context if present. - `.With().Fields()` all forms of `ErrorStackMarshaler` returns. - `.WithLevel()` for `FatalLevel`, `PanicLevel`, and `DisabledLevel`. - `.Err()` with `nil` and non-`nil` `error` and test the resulting log level. - `.should()` covering `nil` writer. - `.Output()` gets context values. - `.UpdateContext()` on a disabled logger doesn't panic and is a nop. - `.With()` all forms of `ErrorStackMarshaler` returns. - with a `nil`writer. - `.Hook()` passing no hooks. - `Array` - `.MarshalZerologArray` is a nop that won't panic. - `Context` - ` .Err()` and `.AnErr()` for `nil` errors and all forms of `ErrorStackMarshaler` returns. - `Event` - `.Caller()` to ensure we don't panic or add invalid information if `runtime.Caller()` fails -` .Err()` and `.AnErr()` for `nil` errors and all forms of `ErrorStackMarshaler` returns. - `Fields` - ` .appendFields()` all forms of `ErrorStackMarshaler` returns. - `HookLevel` - `.Run()` methods. - `LevelSampler` - `.Sample() methods. - `Syslog` - `.Write()`, `.WriteLevel()`, and `.Close()` methods. - `.WriteLevel()` with an `InvalidLevel`. - `Writer` - `.Write()` short write and error cases. - `MultiLevelWriter` `.WriteLevel()` and for `.Write()` error and `.Close()` cases. - test of unmarshalling a level byte returns correct error. - CBOR decodeStream - `.decodeFloat()`, `.binaryFmt()`, `.DecodeIfBinaryToString()`, `.DecodeObjectToStr()`, `.DecodeIfBinaryToBytes()`, `.decodeTagData()`, and `.decodeSimpleFloat()` - handling of invalid UTF-8 sequences - handling of UTC times. - handling of timestamps - handling of various map lengths Restructure `Event` `.caller()` so we test for ok and eliminate untestable coverage hole. Restructure `Context` `.Err()` when the `ErrorStackMarshaler` returns a `nil` so there's code to cover. Inverted logic for`Event` `.Caller()`'s call to `runtime.Caller()` for simpler testing. Inverted logic for `Array` `.putArray()` and `Event` `.putEvent()` so there isn't uncoverable code. Restructure `Field` `.appendFieldList()` to early return when `ErrorStackMarshaler` returns a `nil` Added comments for things we can't get coverage on. Did a go fmt ./... Coverage of core is now 100% on `Array`, `Context`, `Ctx`, `Event`, `Field`, `Hook`, and `Syslog`. Coverage of `Globals`, `Log`, `Sampler`, and `Writer` is almost all except some real edge-cases. JSON encoder coverage is 100% CBOR encoder coverage is 96.3% with base, cbor, string, time and types at 100% and decode_stream (which is lacks coverage on some panic states, and two incorrect coverage-tool lapses) * Fix CBOR tests for StackMarshaler Forgot to use the `decodeIfBinaryToString()` --- array.go | 6 +- array_test.go | 5 + console.go | 6 +- context.go | 2 +- context_test.go | 143 ++++++++++++++++ ctx.go | 11 +- event.go | 13 +- event_test.go | 255 +++++++++++++++++++++++++++++ fields.go | 13 +- globals.go | 4 + hook_test.go | 97 +++++++++++ internal/cbor/cbor.go | 12 +- internal/cbor/decoder_test.go | 281 ++++++++++++++++++++++++++++++++ internal/cbor/string_test.go | 69 +++++++- internal/testcases.go | 2 + log.go | 128 ++++++++------- log_test.go | 258 ++++++++++++++++++++++++++++- sampler_test.go | 45 +++++ syslog_test.go | 130 ++++++++++++++- writer_test.go | 298 ++++++++++++++++++++++++++++++++++ 20 files changed, 1681 insertions(+), 97 deletions(-) create mode 100644 context_test.go diff --git a/array.go b/array.go index 2220c06..a75a572 100644 --- a/array.go +++ b/array.go @@ -38,10 +38,9 @@ func putArray(a *Array) { // // See https://golang.org/issue/23199 const maxSize = 1 << 16 // 64KiB - if cap(a.buf) > maxSize { - return + if cap(a.buf) <= maxSize { + arrayPool.Put(a) } - arrayPool.Put(a) } // Arr creates an array to be added to an Event or Context. @@ -60,6 +59,7 @@ func Arr() *Array { // MarshalZerologArray method here is no-op - since data is // already in the needed format. func (*Array) MarshalZerologArray(*Array) { + // untestable: there's no code to be covered } func (a *Array) write(dst []byte) []byte { diff --git a/array_test.go b/array_test.go index 1d1517e..2a8c5b0 100644 --- a/array_test.go +++ b/array_test.go @@ -59,3 +59,8 @@ func TestArray(t *testing.T) { t.Errorf("Array.write()\ngot: %s\nwant: %s", got, want) } } + +func TestArray_MarshalZerologArray(t *testing.T) { + a := Arr() + a.MarshalZerologArray(nil) // no-op method, should not panic +} diff --git a/console.go b/console.go index 6c881ef..39bdad1 100644 --- a/console.go +++ b/console.go @@ -101,9 +101,9 @@ type ConsoleWriter struct { // NewConsoleWriter creates and initializes a new ConsoleWriter. func NewConsoleWriter(options ...func(w *ConsoleWriter)) ConsoleWriter { w := ConsoleWriter{ - Out: os.Stdout, - TimeFormat: consoleDefaultTimeFormat, - PartsOrder: consoleDefaultPartsOrder(), + Out: os.Stdout, + TimeFormat: consoleDefaultTimeFormat, + PartsOrder: consoleDefaultPartsOrder(), } for _, opt := range options { diff --git a/context.go b/context.go index 0b6029f..b56563c 100644 --- a/context.go +++ b/context.go @@ -184,7 +184,7 @@ func (c Context) Err(err error) Context { if c.l.stack && ErrorStackMarshaler != nil { switch m := ErrorStackMarshaler(err).(type) { case nil: - // do nothing + return c // do nothing with nil errors case LogObjectMarshaler: c = c.Object(ErrorStackFieldName, m) case error: diff --git a/context_test.go b/context_test.go new file mode 100644 index 0000000..9d1e54a --- /dev/null +++ b/context_test.go @@ -0,0 +1,143 @@ +package zerolog + +import ( + "bytes" + "errors" + "testing" +) + +type myError struct{} + +func (e *myError) Error() string { return "test" } + +func TestContext_ErrWithStackMarshaler(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler + ErrorStackMarshaler = func(err error) interface{} { + return "stack-trace" + } + + var buf bytes.Buffer + log := New(&buf).With().Stack().Err(errors.New("test error")).Logger() + + log.Info().Msg("test message") + + got := decodeIfBinaryToString(buf.Bytes()) + want := `{"level":"info","stack":"stack-trace","error":"test error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Context.Err() with stack marshaler = %q, want %q", got, want) + } +} + +func TestContext_AnErrWithNilErrorMarshal(t *testing.T) { + // Save original + original := ErrorMarshalFunc + defer func() { ErrorMarshalFunc = original }() + + // Set marshaler to return a nil error pointer + ErrorMarshalFunc = func(err error) interface{} { + return (*myError)(nil) // nil pointer of error type + } + + var buf bytes.Buffer + log := New(&buf).With().AnErr("test", errors.New("some error")).Logger() + + log.Info().Msg("test message") + + got := decodeIfBinaryToString(buf.Bytes()) + want := `{"level":"info","message":"test message"}` + "\n" // No "test" field because isNilValue returned true + if got != want { + t.Errorf("Context.AnErr() with nil error marshal = %q, want %q", got, want) + } +} + +func TestContext_ErrWithNilStackMarshaler(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set marshaler to return nil + ErrorStackMarshaler = func(err error) interface{} { + return nil + } + + var buf bytes.Buffer + log := New(&buf).With().Stack().Err(errors.New("test error")).Logger() + + log.Info().Msg("test message") + + got := decodeIfBinaryToString(buf.Bytes()) + want := `{"level":"info","message":"test message"}` + "\n" // No stack or error field because stack marshaler returned nil + if got != want { + t.Errorf("Context.Err() with nil stack marshaler = %q, want %q", got, want) + } +} + +func TestContext_ErrWithStackMarshalerObject(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns LogObjectMarshaler + ErrorStackMarshaler = func(err error) interface{} { + return logObjectMarshalerImpl{name: "user", age: 30} + } + + var buf bytes.Buffer + log := New(&buf).With().Stack().Err(errors.New("test error")).Logger() + + log.Info().Msg("test message") + + got := decodeIfBinaryToString(buf.Bytes()) + want := `{"level":"info","stack":{"name":"user","age":-30},"error":"test error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Context.Err() with stack marshaler object = %q, want %q", got, want) + } +} + +func TestContext_ErrWithStackMarshalerError(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns an error + ErrorStackMarshaler = func(err error) interface{} { + return errors.New("stack error") + } + + var buf bytes.Buffer + log := New(&buf).With().Stack().Err(errors.New("test error")).Logger() + + log.Info().Msg("test message") + + got := decodeIfBinaryToString(buf.Bytes()) + want := `{"level":"info","stack":"stack error","error":"test error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Context.Err() with stack marshaler error = %q, want %q", got, want) + } +} + +func TestContext_ErrWithStackMarshalerInterface(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns an int + ErrorStackMarshaler = func(err error) interface{} { + return 42 + } + + var buf bytes.Buffer + log := New(&buf).With().Stack().Err(errors.New("test error")).Logger() + + log.Info().Msg("test message") + + got := decodeIfBinaryToString(buf.Bytes()) + want := `{"level":"info","stack":42,"error":"test error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Context.Err() with stack marshaler interface = %q, want %q", got, want) + } +} diff --git a/ctx.go b/ctx.go index 60432d1..649191e 100644 --- a/ctx.go +++ b/ctx.go @@ -25,12 +25,11 @@ type ctxKey struct{} // replacing it in a new Context), use UpdateContext with the following // notation: // -// ctx := r.Context() -// l := zerolog.Ctx(ctx) -// l.UpdateContext(func(c Context) Context { -// return c.Str("bar", "baz") -// }) -// +// ctx := r.Context() +// l := zerolog.Ctx(ctx) +// l.UpdateContext(func(c Context) Context { +// return c.Str("bar", "baz") +// }) func (l Logger) WithContext(ctx context.Context) context.Context { if _, ok := ctx.Value(ctxKey{}).(*Logger); !ok && l.level == Disabled { // Do not store disabled logger. diff --git a/event.go b/event.go index 9e6a46c..0e05347 100644 --- a/event.go +++ b/event.go @@ -48,10 +48,9 @@ func putEvent(e *Event) { // // See https://golang.org/issue/23199 const maxSize = 1 << 16 // 64KiB - if cap(e.buf) > maxSize { - return + if cap(e.buf) <= maxSize { + eventPool.Put(e) } - eventPool.Put(e) } // LogObjectMarshaler provides a strongly-typed and encoding-agnostic interface @@ -435,7 +434,7 @@ func (e *Event) Err(err error) *Event { if e.stack && ErrorStackMarshaler != nil { switch m := ErrorStackMarshaler(err).(type) { case nil: - // do nothing + return e case LogObjectMarshaler: e = e.Object(ErrorStackFieldName, m) case error: @@ -835,11 +834,9 @@ func (e *Event) caller(skip int) *Event { if e == nil { return e } - pc, file, line, ok := runtime.Caller(skip + e.skipFrame) - if !ok { - return e + if pc, file, line, ok := runtime.Caller(skip + e.skipFrame); ok { + e.buf = enc.AppendString(enc.AppendKey(e.buf, CallerFieldName), CallerMarshalFunc(pc, file, line)) } - e.buf = enc.AppendString(enc.AppendKey(e.buf, CallerFieldName), CallerMarshalFunc(pc, file, line)) return e } diff --git a/event_test.go b/event_test.go index f3f8261..6c26fe1 100644 --- a/event_test.go +++ b/event_test.go @@ -387,6 +387,23 @@ func TestEvent_MsgFunc(t *testing.T) { } } +func TestEvent_CallerRuntimeFail(t *testing.T) { + var buf bytes.Buffer + e := newEvent(LevelWriterAdapter{&buf}, DebugLevel, false, nil, nil) + + // Set a very large skipFrame to make runtime.Caller fail + e.CallerSkipFrame(1000) + e.Caller() + + e.Msg("test") + + got := strings.TrimSpace(buf.String()) + want := `{"message":"test"}` // No caller field because runtime.Caller failed + if got != want { + t.Errorf("Event.Caller() with failed runtime.Caller = %q, want %q", got, want) + } +} + func TestEvent_DoneHandler(t *testing.T) { e := newEvent(nil, InfoLevel, false, nil, nil) @@ -461,3 +478,241 @@ func TestEvent_Msg_ErrorHandlerNil(t *testing.T) { t.Errorf("Expected stderr output %q, got %q", expected, string(captured)) } } + +type mockLogObjectMarshaler struct { + data string +} + +func (m mockLogObjectMarshaler) MarshalZerologObject(e *Event) { + e.Str("stack_func", m.data) +} + +func TestEvent_ErrWithStackMarshaler(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler + ErrorStackMarshaler = func(err error) interface{} { + return "stack-trace" + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Err(err).Msg("test message") + + got := buf.String() + want := `{"stack":"stack-trace","error":"test error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Event.Err() with stack marshaler = %q, want %q", got, want) + } +} + +func TestEvent_FieldsWithErrorAndStackMarshaler(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler + ErrorStackMarshaler = func(err error) interface{} { + return "stack-trace" + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Fields([]interface{}{"error", err}).Msg("test message") + + got := buf.String() + want := `{"error":"test error","stack":"stack-trace","message":"test message"}` + "\n" + if got != want { + t.Errorf("Event.Fields() with error and stack marshaler = %q, want %q", got, want) + } +} + +func TestEvent_FieldsWithErrorAndStackMarshalerObject(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns LogObjectMarshaler + ErrorStackMarshaler = func(err error) interface{} { + return mockLogObjectMarshaler{data: "stack-data"} + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Fields([]interface{}{"error", err}).Msg("test message") + + got := buf.String() + want := `{"error":"test error","stack":{"stack_func":"stack-data"},"message":"test message"}` + "\n" + if got != want { + t.Errorf("Event.Fields() with error and stack marshaler object = %q, want %q", got, want) + } +} + +func TestEvent_FieldsWithErrorAndStackMarshalerError(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns an error + ErrorStackMarshaler = func(err error) interface{} { + return errors.New("stack error") + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Fields([]interface{}{"error", err}).Msg("test message") + + got := buf.String() + want := `{"error":"test error","stack":"stack error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Event.Fields() with error and stack marshaler error = %q, want %q", got, want) + } +} + +func TestEvent_FieldsWithErrorAndStackMarshalerInterface(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns an int + ErrorStackMarshaler = func(err error) interface{} { + return 42 + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Fields([]interface{}{"error", err}).Msg("test message") + + got := buf.String() + want := `{"error":"test error","stack":42,"message":"test message"}` + "\n" + if got != want { + t.Errorf("Event.Fields() with error and stack marshaler interface = %q, want %q", got, want) + } +} + +func TestEvent_FieldsWithErrorAndStackMarshalerNil(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set marshaler to return nil + ErrorStackMarshaler = func(err error) interface{} { + return nil + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Fields([]interface{}{"error", err}).Msg("test message") + + got := buf.String() + want := `{"error":"test error","message":"test message"}` + "\n" // No stack field because marshaler returned nil + if got != want { + t.Errorf("Event.Fields() with error and nil stack marshaler = %q, want %q", got, want) + } +} + +func TestEvent_ErrWithStackMarshalerObject(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns LogObjectMarshaler + ErrorStackMarshaler = func(err error) interface{} { + return mockLogObjectMarshaler{data: "stack-data"} + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Err(err).Msg("test message") + + got := buf.String() + want := `{"stack":{"stack_func":"stack-data"},"error":"test error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Event.Err() with stack marshaler object = %q, want %q", got, want) + } +} + +func TestEvent_ErrWithStackMarshalerError(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns an error + ErrorStackMarshaler = func(err error) interface{} { + return errors.New("stack error") + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Err(err).Msg("test message") + + got := buf.String() + want := `{"stack":"stack error","error":"test error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Event.Err() with stack marshaler error = %q, want %q", got, want) + } +} + +func TestEvent_ErrWithStackMarshalerInterface(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set a mock marshaler that returns an int + ErrorStackMarshaler = func(err error) interface{} { + return 42 + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Err(err).Msg("test message") + + got := buf.String() + want := `{"stack":42,"error":"test error","message":"test message"}` + "\n" + if got != want { + t.Errorf("Event.Err() with stack marshaler interface = %q, want %q", got, want) + } +} + +func TestEvent_ErrWithStackMarshalerNil(t *testing.T) { + // Save original + original := ErrorStackMarshaler + defer func() { ErrorStackMarshaler = original }() + + // Set marshaler to return nil + ErrorStackMarshaler = func(err error) interface{} { + return nil + } + + var buf bytes.Buffer + log := New(&buf) + + err := errors.New("test error") + log.Log().Stack().Err(err).Msg("test message") + + got := buf.String() + want := `{"message":"test message"}` + "\n" // No fields because stack marshaler returned nil + if got != want { + t.Errorf("Event.Err() with nil stack marshaler = %q, want %q", got, want) + } +} diff --git a/fields.go b/fields.go index b77a57a..08759ef 100644 --- a/fields.go +++ b/fields.go @@ -5,13 +5,18 @@ import ( "encoding/json" "io" "net" + "reflect" "sort" "time" - "unsafe" ) -func isNilValue(i interface{}) bool { - return (*[2]uintptr)(unsafe.Pointer(&i))[1] == 0 +func isNilValue(e error) bool { + switch reflect.TypeOf(e).Kind() { + case reflect.Ptr: + return reflect.ValueOf(e).IsNil() + default: + return false + } } func appendFields(dst []byte, fields interface{}, stack bool, ctx context.Context, hooks []Hook) []byte { @@ -77,7 +82,7 @@ func appendFieldList(dst []byte, kvList []interface{}, stack bool, ctx context.C if stack && ErrorStackMarshaler != nil { switch m := ErrorStackMarshaler(val).(type) { case nil: - // do nothing + return dst // do nothing with nil errors case LogObjectMarshaler: dst = enc.AppendKey(dst, ErrorStackFieldName) dst = appendObject(dst, m, stack, ctx, hooks) diff --git a/globals.go b/globals.go index e34d0fc..d23c127 100644 --- a/globals.go +++ b/globals.go @@ -134,6 +134,10 @@ var ( // be thread safe and non-blocking. ErrorHandler func(err error) + // FatalExitFunc is called by log.Fatal() instead of os.Exit(1). If not set, + // os.Exit(1) is called. + FatalExitFunc func() + // DefaultContextLogger is returned from Ctx() if there is no logger associated // with the context. DefaultContextLogger *Logger diff --git a/hook_test.go b/hook_test.go index 100b71e..3a36eb2 100644 --- a/hook_test.go +++ b/hook_test.go @@ -50,6 +50,10 @@ func TestHook(t *testing.T) { want string test func(log Logger) }{ + {"Message", `{"message":"test message"}` + "\n", func(log Logger) { + log = log.Hook() + log.Log().Msg("test message") + }}, {"Message", `{"level_name":"nolevel","message":"test message"}` + "\n", func(log Logger) { log = log.Hook(levelNameHook) log.Log().Msg("test message") @@ -171,6 +175,99 @@ func TestHook(t *testing.T) { } } +func TestLevelHook(t *testing.T) { + var called []string + + traceHook := HookFunc(func(e *Event, level Level, msg string) { + called = append(called, "trace") + }) + debugHook := HookFunc(func(e *Event, level Level, msg string) { + called = append(called, "debug") + }) + infoHook := HookFunc(func(e *Event, level Level, msg string) { + called = append(called, "info") + }) + warnHook := HookFunc(func(e *Event, level Level, msg string) { + called = append(called, "warn") + }) + errorHook := HookFunc(func(e *Event, level Level, msg string) { + called = append(called, "error") + }) + fatalHook := HookFunc(func(e *Event, level Level, msg string) { + called = append(called, "fatal") + }) + panicHook := HookFunc(func(e *Event, level Level, msg string) { + called = append(called, "panic") + }) + noLevelHook := HookFunc(func(e *Event, level Level, msg string) { + called = append(called, "nolevel") + }) + + hook := LevelHook{ + TraceHook: traceHook, + DebugHook: debugHook, + InfoHook: infoHook, + WarnHook: warnHook, + ErrorHook: errorHook, + FatalHook: fatalHook, + PanicHook: panicHook, + NoLevelHook: noLevelHook, + } + + e := &Event{} + + // Test each level + hook.Run(e, TraceLevel, "") + if len(called) != 1 || called[0] != "trace" { + t.Errorf("TraceLevel hook not called correctly: %v", called) + } + + called = nil + hook.Run(e, DebugLevel, "") + if len(called) != 1 || called[0] != "debug" { + t.Errorf("DebugLevel hook not called correctly: %v", called) + } + + called = nil + hook.Run(e, InfoLevel, "") + if len(called) != 1 || called[0] != "info" { + t.Errorf("InfoLevel hook not called correctly: %v", called) + } + + called = nil + hook.Run(e, WarnLevel, "") + if len(called) != 1 || called[0] != "warn" { + t.Errorf("WarnLevel hook not called correctly: %v", called) + } + + called = nil + hook.Run(e, ErrorLevel, "") + if len(called) != 1 || called[0] != "error" { + t.Errorf("ErrorLevel hook not called correctly: %v", called) + } + + called = nil + hook.Run(e, FatalLevel, "") + if len(called) != 1 || called[0] != "fatal" { + t.Errorf("FatalLevel hook not called correctly: %v", called) + } + + called = nil + hook.Run(e, PanicLevel, "") + if len(called) != 1 || called[0] != "panic" { + t.Errorf("PanicLevel hook not called correctly: %v", called) + } + + called = nil + hook.Run(e, NoLevel, "") + if len(called) != 1 || called[0] != "nolevel" { + t.Errorf("NoLevel hook not called correctly: %v", called) + } + + // Test NewLevelHook + _ = NewLevelHook() +} + func BenchmarkHooks(b *testing.B) { logger := New(io.Discard) b.ResetTimer() diff --git a/internal/cbor/cbor.go b/internal/cbor/cbor.go index 15088d2..e7509bd 100644 --- a/internal/cbor/cbor.go +++ b/internal/cbor/cbor.go @@ -55,12 +55,12 @@ const ( ) const ( - float32Nan = "\xfa\x7f\xc0\x00\x00" - float32PosInfinity = "\xfa\x7f\x80\x00\x00" - float32NegInfinity = "\xfa\xff\x80\x00\x00" - float64Nan = "\xfb\x7f\xf8\x00\x00\x00\x00\x00\x00" - float64PosInfinity = "\xfb\x7f\xf0\x00\x00\x00\x00\x00\x00" - float64NegInfinity = "\xfb\xff\xf0\x00\x00\x00\x00\x00\x00" + float32Nan = "\x7f\xc0\x00\x00" + float32PosInfinity = "\x7f\x80\x00\x00" + float32NegInfinity = "\xff\x80\x00\x00" + float64Nan = "\x7f\xf8\x00\x00\x00\x00\x00\x00" + float64PosInfinity = "\x7f\xf0\x00\x00\x00\x00\x00\x00" + float64NegInfinity = "\xff\xf0\x00\x00\x00\x00\x00\x00" ) // IntegerTimeFieldFormat indicates the format of timestamp decoded diff --git a/internal/cbor/decoder_test.go b/internal/cbor/decoder_test.go index 20e7fc0..6d91431 100644 --- a/internal/cbor/decoder_test.go +++ b/internal/cbor/decoder_test.go @@ -80,6 +80,8 @@ var mapDecodeTestCases = []struct { }{ {[]byte("\xa2\x64IETF\x20"), "{\"IETF\":-1}"}, {[]byte("\xa2\x65Array\x84\x20\x00\x18\xc8\x14"), "{\"Array\":[-1,0,200,20]}"}, + {[]byte("\xa6\x61\x61\x01\x61\x62\x02\x61\x63\x03"), "{\"a\":1,\"b\":2,\"c\":3}"}, + {[]byte("\xbf\x61a\x01\x61b\x02\xff"), "{\"a\":1,\"b\":2}"}, } func TestDecodeMap(t *testing.T) { @@ -115,6 +117,45 @@ func TestDecodeFloat(t *testing.T) { t.Errorf("decodeFloat(0x%s)=%f, want:%f", hex.EncodeToString([]byte(tc.Binary)), got, tc.Val) } } + for _, tc := range internal.Float64TestCases { + got, _ := decodeFloat(getReader(tc.Binary)) + if got != tc.Val && math.IsNaN(got) != math.IsNaN(tc.Val) { + t.Errorf("decodeFloat(0x%s)=%f, want:%f", hex.EncodeToString([]byte(tc.Binary)), got, tc.Val) + } + } + + // Test float64 special values with correct CBOR encoding + float64Tests := []struct { + name string + input string + want float64 + }{ + {"float64 NaN", "\xfb\x7f\xf8\x00\x00\x00\x00\x00\x00", math.NaN()}, + {"float64 +Inf", "\xfb\x7f\xf0\x00\x00\x00\x00\x00\x00", math.Inf(0)}, + {"float64 -Inf", "\xfb\xff\xf0\x00\x00\x00\x00\x00\x00", math.Inf(-1)}, + {"float64 1.0", "\xfb\x3f\xf0\x00\x00\x00\x00\x00\x00", 1.0}, + } + + for _, tt := range float64Tests { + t.Run(tt.name, func(t *testing.T) { + got, _ := decodeFloat(getReader(tt.input)) + if math.IsNaN(tt.want) { + if !math.IsNaN(got) { + t.Errorf("decodeFloat(%q) = %f, want NaN", tt.input, got) + } + } else if math.IsInf(tt.want, 0) { + if !math.IsInf(got, 0) { + t.Errorf("decodeFloat(%q) = %f, want +Inf", tt.input, got) + } + } else if math.IsInf(tt.want, -1) { + if !math.IsInf(got, -1) { + t.Errorf("decodeFloat(%q) = %f, want -Inf", tt.input, got) + } + } else if got != tt.want { + t.Errorf("decodeFloat(%q) = %f, want %f", tt.input, got, tt.want) + } + }) + } } func TestDecodeTimestamp(t *testing.T) { @@ -137,6 +178,26 @@ func TestDecodeTimestamp(t *testing.T) { t.Errorf("decodeFloat(0x%s)=%s, want:%s", hex.EncodeToString([]byte(tc.Out)), tm, tc.RfcStr) } } + + // Test with decodeTimeZone = nil to cover the else branches + oldTimeZone := decodeTimeZone + decodeTimeZone = nil + defer func() { decodeTimeZone = oldTimeZone }() + + for _, tc := range internal.TimeIntegerTestcases { + tm := decodeTagData(getReader(tc.Binary)) + if string(tm) != "\""+tc.RfcStr+"\"" { + t.Errorf("decodeFloat(0x%s)=%s, want:%s", hex.EncodeToString([]byte(tc.Binary)), tm, tc.RfcStr) + } + } + for _, tc := range internal.TimeFloatTestcases { + tm := decodeTagData(getReader(tc.Out)) + got, _ := time.Parse(string(tm), string(tm)) + want, _ := time.Parse(tc.RfcStr, tc.RfcStr) + if got.Sub(want) > time.Microsecond { + t.Errorf("decodeFloat(0x%s)=%s, want:%s", hex.EncodeToString([]byte(tc.Out)), tm, tc.RfcStr) + } + } } func TestDecodeNetworkAddr(t *testing.T) { @@ -172,6 +233,8 @@ var compositeCborTestCases = []struct { }{ {[]byte("\xbf\x64IETF\x20\x65Array\x9f\x20\x00\x18\xc8\x14\xff\xff"), "{\"IETF\":-1,\"Array\":[-1,0,200,20]}\n"}, {[]byte("\xbf\x64IETF\x64YES!\x65Array\x9f\x20\x00\x18\xc8\x14\xff\xff"), "{\"IETF\":\"YES!\",\"Array\":[-1,0,200,20]}\n"}, + {[]byte("\xbf\x61a\x01\x61b\x02\x61c\x03\xff"), "{\"a\":1,\"b\":2,\"c\":3}\n"}, + {[]byte("\xc1\x1a\x51\x0f\x30\xd8"), "\"2013-02-04T03:54:00Z\"\n"}, } func TestDecodeCbor2Json(t *testing.T) { @@ -206,3 +269,221 @@ func TestDecodeNegativeCbor2Json(t *testing.T) { } } } + +func TestBinaryFmt(t *testing.T) { + tests := []struct { + input []byte + want bool + }{ + {[]byte{}, false}, + {[]byte{0x00}, false}, + {[]byte{0x7F}, false}, + {[]byte{0x80}, true}, + {[]byte{0xFF}, true}, + {[]byte{0x00, 0x80}, false}, // Only checks first byte + } + + for _, tt := range tests { + got := binaryFmt(tt.input) + if got != tt.want { + t.Errorf("binaryFmt(%v) = %v, want %v", tt.input, got, tt.want) + } + } +} + +func TestDecodeIfBinaryToString(t *testing.T) { + tests := []struct { + name string + input []byte + want string + }{ + { + name: "non-binary input", + input: []byte(`{"key":"value"}`), + want: `{"key":"value"}`, + }, + { + name: "binary input - simple object", + input: []byte("\xbf\x64IETF\x20\xff"), // {"IETF": -1} in indefinite length CBOR + want: "{\"IETF\":-1}\n", + }, + { + name: "binary input - multiple objects", + input: []byte("\xbf\x64IETF\x20\xff\xbf\x65Array\x84\x20\x00\x18\xc8\x14\xff"), // Two objects + want: "{\"IETF\":-1}\n{\"Array\":[-1,0,200,20]}\n", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := DecodeIfBinaryToString(tt.input) + if got != tt.want { + t.Errorf("DecodeIfBinaryToString() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestDecodeObjectToStr(t *testing.T) { + tests := []struct { + name string + input []byte + want string + }{ + { + name: "non-binary input", + input: []byte(`{"key":"value"}`), + want: `{"key":"value"}`, + }, + { + name: "binary input - simple object", + input: []byte("\xbf\x64IETF\x20\xff"), // {"IETF": -1} in indefinite length CBOR + want: "{\"IETF\":-1}", + }, + { + name: "binary input - array", + input: []byte("\x84\x20\x00\x18\xc8\x14"), // [-1, 0, 200, 20] + want: "[-1,0,200,20]", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := DecodeObjectToStr(tt.input) + if got != tt.want { + t.Errorf("DecodeObjectToStr() = %q, want %q", got, tt.want) + } + }) + } +} + +func TestDecodeIfBinaryToBytes(t *testing.T) { + tests := []struct { + name string + input []byte + want []byte + }{ + { + name: "non-binary input", + input: []byte(`{"key":"value"}`), + want: []byte(`{"key":"value"}`), + }, + { + name: "binary input - simple object", + input: []byte("\xbf\x64IETF\x20\xff"), // {"IETF": -1} in indefinite length CBOR + want: []byte("{\"IETF\":-1}\n"), + }, + { + name: "binary input - multiple objects", + input: []byte("\xbf\x64IETF\x20\xff\xbf\x65Array\x84\x20\x00\x18\xc8\x14\xff"), // Two objects + want: []byte("{\"IETF\":-1}\n{\"Array\":[-1,0,200,20]}\n"), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := DecodeIfBinaryToBytes(tt.input) + if !bytes.Equal(got, tt.want) { + t.Errorf("DecodeIfBinaryToBytes() = %q, want %q", string(got), string(tt.want)) + } + }) + } +} + +func TestDecodeEmbeddedCBOR(t *testing.T) { + // Test embedded CBOR tag: 0xD8 0x3F (tag 63) followed by byte string + // 0xD8 = major type 6 (tags) + additional type 24 (uint8 follows) + // 0x3F = 63 (additionalTypeEmbeddedCBOR) + // 0x43 = major type 2 (byte string) + length 3 + // 0x01 0x02 0x03 = the embedded CBOR data + + embeddedCBOR := []byte("\xd8\x3f\x43\x01\x02\x03") + expected := "\"data:application/cbor;base64,AQID\"" + + got := decodeTagData(getReader(string(embeddedCBOR))) + if string(got) != expected { + t.Errorf("decodeTagData(embedded CBOR) = %q, want %q", string(got), expected) + } +} + +func TestDecodeEmbeddedJSON(t *testing.T) { + t.Run("valid embedded JSON", func(t *testing.T) { + // Test embedded JSON tag: 0xD9 0x01 0x06 (tag 262) followed by byte string. + // 0xD9 = major type 6 (tags) + additional type 25 (uint16 follows) + // 0x01 0x06 = 262 (additionalTypeEmbeddedJSON) + // 0x47 = major type 2 (byte string) + length 7 + // {"a":1} = embedded JSON payload (no surrounding quotes expected) + embeddedJSON := []byte("\xd9\x01\x06\x47{\"a\":1}") + expected := "{\"a\":1}" + + got := decodeTagData(getReader(string(embeddedJSON))) + if string(got) != expected { + t.Errorf("decodeTagData(embedded JSON) = %q, want %q", string(got), expected) + } + }) + + t.Run("unsupported embedded type panics", func(t *testing.T) { + // Same embedded JSON tag, but followed by a UTF-8 string instead of a byte string. + // This should hit the "Unsupported embedded Type" panic branch. + bad := []byte("\xd9\x01\x06\x61x") + + defer func() { + if r := recover(); r == nil { + t.Fatalf("expected panic, got none") + } + }() + + _ = decodeTagData(getReader(string(bad))) + }) +} + +func TestDecodeHexString(t *testing.T) { + // Test hex string tag: 0xD9 0x01 0x07 (tag 263) followed by byte string + // 0xD9 = major type 6 (tags) + additional type 25 (uint16 follows) + // 0x01 0x07 = 263 (additionalTypeTagHexString) + // 0x43 = major type 2 (byte string) + length 3 + // 0x01 0x02 0x03 = the byte data to hex encode + + hexString := []byte("\xd9\x01\x07\x43\x01\x02\x03") + expected := "\"010203\"" + + got := decodeTagData(getReader(string(hexString))) + if string(got) != expected { + t.Errorf("decodeTagData(hex string) = %q, want %q", string(got), expected) + } +} + +func TestDecodeSimpleFloat(t *testing.T) { + tests := []struct { + name string + input string + want string + }{ + // Boolean and null cases (already covered) + {"true", "\xf5", "true"}, + {"false", "\xf4", "false"}, + {"null", "\xf6", "null"}, + + // Float32 cases + {"float32 1.0", "\xfa\x3f\x80\x00\x00", "1"}, + {"float32 1.5", "\xfa\x3f\xc0\x00\x00", "1.5"}, + {"float32 +Inf", "\xfa\x7f\x80\x00\x00", "\"+Inf\""}, + {"float32 -Inf", "\xfa\xff\x80\x00\x00", "\"-Inf\""}, + {"float32 NaN", "\xfa\x7f\xc0\x00\x00", "\"NaN\""}, + + // Float64 cases + {"float64 1.0", "\xfb\x3f\xf0\x00\x00\x00\x00\x00\x00", "1"}, + {"float64 +Inf", "\xfb\x7f\xf0\x00\x00\x00\x00\x00\x00", "\"+Inf\""}, + {"float64 -Inf", "\xfb\xff\xf0\x00\x00\x00\x00\x00\x00", "\"-Inf\""}, + {"float64 NaN", "\xfb\x7f\xf8\x00\x00\x00\x00\x00\x00", "\"NaN\""}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := decodeSimpleFloat(getReader(tt.input)) + if string(got) != tt.want { + t.Errorf("decodeSimpleFloat(%q) = %q, want %q", tt.input, string(got), tt.want) + } + }) + } +} diff --git a/internal/cbor/string_test.go b/internal/cbor/string_test.go index aa88d34..b29907b 100644 --- a/internal/cbor/string_test.go +++ b/internal/cbor/string_test.go @@ -41,6 +41,7 @@ var encodeStringTests = []struct { "<------------------------------------ This is a 100 character string ----------------------------->" + "<------------------------------------ This is a 100 character string ----------------------------->"}, {"emoji \u2764\ufe0f!", "\x6demoji ❤️!", "emoji \u2764\ufe0f!"}, + {"invalid utf8 \xff", "\x6einvalid utf8 \xff", "invalid utf8 \\ufffd"}, } var encodeByteTests = []struct { @@ -96,7 +97,7 @@ func TestAppendStrings(t *testing.T) { array = append(array, tt.plain) } want := make([]byte, 0) - want = append(want, 0x94) // start array length 24 + want = append(want, 0x95) // start array for _, tt := range encodeStringTests { want = append(want, []byte(tt.binary)...) } @@ -214,3 +215,69 @@ func BenchmarkAppendString(b *testing.B) { }) } } + +func TestAppendEmbeddedJSON(t *testing.T) { + tests := []struct { + name string + input []byte + want string + }{ + { + name: "empty JSON", + input: []byte{}, + want: "\xd9\x01\x06@", // tag 0xd9 + empty byte string + }, + { + name: "small JSON", + input: []byte(`{"key":"value"}`), + want: "\xd9\x01\x06O{\"key\":\"value\"}", // tag 0xd9 + byte string with content + }, + { + name: "large JSON (>23 bytes)", + input: []byte(`{"key":"this is a very long value that exceeds the 23 byte limit for direct encoding"}`), + want: "\xd9\x01\x06XV{\"key\":\"this is a very long value that exceeds the 23 byte limit for direct encoding\"}", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := AppendEmbeddedJSON([]byte{}, tt.input) + if string(got) != tt.want { + t.Errorf("AppendEmbeddedJSON() = %q, want %q", string(got), tt.want) + } + }) + } +} + +func TestAppendEmbeddedCBOR(t *testing.T) { + tests := []struct { + name string + input []byte + want string + }{ + { + name: "empty CBOR", + input: []byte{}, + want: "\xd8?@", // tag 0xd8 + empty byte string + }, + { + name: "small CBOR", + input: []byte{0x01, 0x02, 0x03}, + want: "\xd8?C\x01\x02\x03", // tag 0xd8 + byte string with 3 bytes + }, + { + name: "large CBOR (>23 bytes)", + input: make([]byte, 30), // 30 bytes of zeros + want: "\xd8?X\x1e" + string(make([]byte, 30)), // tag 0xd8 + byte string with length prefix + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got := AppendEmbeddedCBOR([]byte{}, tt.input) + if string(got) != tt.want { + t.Errorf("AppendEmbeddedCBOR() = %q, want %q", string(got), tt.want) + } + }) + } +} diff --git a/internal/testcases.go b/internal/testcases.go index 63ca307..914bc60 100644 --- a/internal/testcases.go +++ b/internal/testcases.go @@ -217,6 +217,8 @@ var IPPrefixTestCases = []struct { {net.IPNet{IP: net.IP{0, 0, 0, 0}, Mask: net.CIDRMask(0, 32)}, "\"0.0.0.0/0\"", "\xd9\x01\x05\xa1\x44\x00\x00\x00\x00\x00"}, {net.IPNet{IP: net.IP{192, 168, 0, 100}, Mask: net.CIDRMask(24, 32)}, "\"192.168.0.100/24\"", "\xd9\x01\x05\xa1\x44\xc0\xa8\x00\x64\x18\x18"}, + {net.IPNet{IP: net.IP{0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1}, Mask: net.CIDRMask(128, 128)}, "\"::1/128\"", + "\xd9\x01\x05\xa1\x50\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x00\x01\x18\x80"}, } var IPPrefixArrayTestCases = []struct { diff --git a/log.go b/log.go index f615ee1..9ff669e 100644 --- a/log.go +++ b/log.go @@ -2,85 +2,85 @@ // // A global Logger can be use for simple logging: // -// import "github.com/rs/zerolog/log" +// import "github.com/rs/zerolog/log" // -// log.Info().Msg("hello world") -// // Output: {"time":1494567715,"level":"info","message":"hello world"} +// log.Info().Msg("hello world") +// // Output: {"time":1494567715,"level":"info","message":"hello world"} // // NOTE: To import the global logger, import the "log" subpackage "github.com/rs/zerolog/log". // // Fields can be added to log messages: // -// log.Info().Str("foo", "bar").Msg("hello world") -// // Output: {"time":1494567715,"level":"info","message":"hello world","foo":"bar"} +// log.Info().Str("foo", "bar").Msg("hello world") +// // Output: {"time":1494567715,"level":"info","message":"hello world","foo":"bar"} // // Create logger instance to manage different outputs: // -// logger := zerolog.New(os.Stderr).With().Timestamp().Logger() -// logger.Info(). -// Str("foo", "bar"). -// Msg("hello world") -// // Output: {"time":1494567715,"level":"info","message":"hello world","foo":"bar"} +// logger := zerolog.New(os.Stderr).With().Timestamp().Logger() +// logger.Info(). +// Str("foo", "bar"). +// Msg("hello world") +// // Output: {"time":1494567715,"level":"info","message":"hello world","foo":"bar"} // // Sub-loggers let you chain loggers with additional context: // -// sublogger := log.With().Str("component", "foo").Logger() -// sublogger.Info().Msg("hello world") -// // Output: {"time":1494567715,"level":"info","message":"hello world","component":"foo"} +// sublogger := log.With().Str("component", "foo").Logger() +// sublogger.Info().Msg("hello world") +// // Output: {"time":1494567715,"level":"info","message":"hello world","component":"foo"} // // Level logging // -// zerolog.SetGlobalLevel(zerolog.InfoLevel) +// zerolog.SetGlobalLevel(zerolog.InfoLevel) // -// log.Debug().Msg("filtered out message") -// log.Info().Msg("routed message") +// log.Debug().Msg("filtered out message") +// log.Info().Msg("routed message") // -// if e := log.Debug(); e.Enabled() { -// // Compute log output only if enabled. -// value := compute() -// e.Str("foo": value).Msg("some debug message") -// } -// // Output: {"level":"info","time":1494567715,"routed message"} +// if e := log.Debug(); e.Enabled() { +// // Compute log output only if enabled. +// value := compute() +// e.Str("foo": value).Msg("some debug message") +// } +// // Output: {"level":"info","time":1494567715,"routed message"} // // Customize automatic field names: // -// log.TimestampFieldName = "t" -// log.LevelFieldName = "p" -// log.MessageFieldName = "m" +// log.TimestampFieldName = "t" +// log.LevelFieldName = "p" +// log.MessageFieldName = "m" // -// log.Info().Msg("hello world") -// // Output: {"t":1494567715,"p":"info","m":"hello world"} +// log.Info().Msg("hello world") +// // Output: {"t":1494567715,"p":"info","m":"hello world"} // // Log with no level and message: // -// log.Log().Str("foo","bar").Msg("") -// // Output: {"time":1494567715,"foo":"bar"} +// log.Log().Str("foo","bar").Msg("") +// // Output: {"time":1494567715,"foo":"bar"} // // Add contextual fields to global Logger: // -// log.Logger = log.With().Str("foo", "bar").Logger() +// log.Logger = log.With().Str("foo", "bar").Logger() // // Sample logs: // -// sampled := log.Sample(&zerolog.BasicSampler{N: 10}) -// sampled.Info().Msg("will be logged every 10 messages") +// sampled := log.Sample(&zerolog.BasicSampler{N: 10}) +// sampled.Info().Msg("will be logged every 10 messages") // // Log with contextual hooks: // -// // Create the hook: -// type SeverityHook struct{} +// // Create the hook: +// type SeverityHook struct{} // -// func (h SeverityHook) Run(e *zerolog.Event, level zerolog.Level, msg string) { -// if level != zerolog.NoLevel { -// e.Str("severity", level.String()) -// } -// } +// func (h SeverityHook) Run(e *zerolog.Event, level zerolog.Level, msg string) { +// if level != zerolog.NoLevel { +// e.Str("severity", level.String()) +// } +// } // -// // And use it: -// var h SeverityHook -// log := zerolog.New(os.Stdout).Hook(h) -// log.Warn().Msg("") -// // Output: {"level":"warn","severity":"warn"} +// // And use it: +// var h SeverityHook +// log := zerolog.New(os.Stdout).Hook(h) +// log.Warn().Msg("") +// // Output: {"level":"warn","severity":"warn"} // // # Caveats // @@ -89,11 +89,11 @@ // There is no fields deduplication out-of-the-box. // Using the same key multiple times creates new key in final JSON each time. // -// logger := zerolog.New(os.Stderr).With().Timestamp().Logger() -// logger.Info(). -// Timestamp(). -// Msg("dup") -// // Output: {"level":"info","time":1494567715,"time":1494567715,"message":"dup"} +// logger := zerolog.New(os.Stderr).With().Timestamp().Logger() +// logger.Info(). +// Timestamp(). +// Msg("dup") +// // Output: {"level":"info","time":1494567715,"time":1494567715,"message":"dup"} // // In this case, many consumers will take the last value, // but this is not guaranteed; check yours if in doubt. @@ -102,15 +102,15 @@ // // Be careful when calling UpdateContext. It is not concurrency safe. Use the With method to create a child logger: // -// func handler(w http.ResponseWriter, r *http.Request) { -// // Create a child logger for concurrency safety -// logger := log.Logger.With().Logger() +// func handler(w http.ResponseWriter, r *http.Request) { +// // Create a child logger for concurrency safety +// logger := log.Logger.With().Logger() // -// // Add context fields, for example User-Agent from HTTP headers -// logger.UpdateContext(func(c zerolog.Context) zerolog.Context { -// ... -// }) -// } +// // Add context fields, for example User-Agent from HTTP headers +// logger.UpdateContext(func(c zerolog.Context) zerolog.Context { +// ... +// }) +// } package zerolog import ( @@ -382,18 +382,24 @@ func (l *Logger) Err(err error) *Event { return l.Info() } -// Fatal starts a new message with fatal level. The os.Exit(1) function -// is called by the Msg method, which terminates the program immediately. +// Fatal starts a new message with fatal level. The FatalExitFunc interceptor function +// is called by the Msg method, which by default terminates the program immediately +// using os.Exit(1), any desired behavior can be implemented by setting FatalExitFunc. // // You must call Msg on the returned event in order to send the event. func (l *Logger) Fatal() *Event { return l.newEvent(FatalLevel, func(msg string) { if closer, ok := l.w.(io.Closer); ok { // Close the writer to flush any buffered message. Otherwise the message - // will be lost as os.Exit() terminates the program immediately. + // could be lost if FatalExitFunc() terminates the program immediately or + // os.Exit(1) is called if not FatalExitFunc isn't set (default). closer.Close() } - os.Exit(1) + if FatalExitFunc != nil { + FatalExitFunc() + } else { + os.Exit(1) // untestable: terminates the program, cannot be covered + } }) } diff --git a/log_test.go b/log_test.go index 2456634..9ca35d6 100644 --- a/log_test.go +++ b/log_test.go @@ -156,6 +156,20 @@ func TestWith(t *testing.T) { } } +func TestStackedWiths(t *testing.T) { + out := &bytes.Buffer{} + ctx := New(out).With(). + Bool("bool", true) + ctx = ctx.Logger().With(). + Int("int", 1) + log := ctx.Logger() + log.Log().Msg("") + if got, want := decodeIfBinaryToString(out.Bytes()), + `{"bool":true,"int":1}`+"\n"; got != want { + t.Errorf("invalid log output:\ngot: %v\nwant: %v", got, want) + } +} + func TestWithPlurals(t *testing.T) { out := &bytes.Buffer{} ctx := New(out).With(). @@ -273,6 +287,25 @@ func TestFieldsMap_Arrays(t *testing.T) { t.Errorf("invalid log output:\ngot: %v\nwant: %v", got, want) } } + +func TestWithErr(t *testing.T) { + var err error = nil + out := &bytes.Buffer{} + ctx := New(out).With(). + Fields(map[string]interface{}{ + "nil": nil, + "nilerror": err, + "error": errors.New("some error"), + "loggable": loggableError{errors.New("loggable")}, + "non-loggable": nonLoggableError{fmt.Errorf("oops"), 401}, + }) + log := ctx.Logger() + log.Log().Msg("") + if got, want := decodeIfBinaryToString(out.Bytes()), `{"error":"some error","loggable":{"l":"LOGGABLE"},"nil":null,"nilerror":null,"non-loggable":"oops"}`+"\n"; got != want { + t.Errorf("invalid log output:\ngot: %v\nwant: %v", got, want) + } +} + func TestFieldsErr(t *testing.T) { var err error = nil out := &bytes.Buffer{} @@ -717,6 +750,28 @@ func TestSampling(t *testing.T) { } } +func TestDisableSampling(t *testing.T) { + // Save original state + original := samplingDisabled() + defer DisableSampling(original) + + out := &bytes.Buffer{} + log := New(out).Sample(&BasicSampler{N: 2}) + + // Enable sampling disable + DisableSampling(true) + + log.Log().Int("i", 1).Msg("") + log.Log().Int("i", 2).Msg("") + log.Log().Int("i", 3).Msg("") + log.Log().Int("i", 4).Msg("") + + // All messages should be logged since sampling is disabled + if got, want := decodeIfBinaryToString(out.Bytes()), "{\"i\":1}\n{\"i\":2}\n{\"i\":3}\n{\"i\":4}\n"; got != want { + t.Errorf("invalid log output:\ngot: %v\nwant: %v", got, want) + } +} + func TestDiscard(t *testing.T) { out := &bytes.Buffer{} log := New(out) @@ -752,6 +807,10 @@ func (lw *levelWriter) WriteLevel(lvl Level, p []byte) (int, error) { return len(p), nil } +func (lw *levelWriter) Close() error { + return nil +} + func TestLevelWriter(t *testing.T) { lw := &levelWriter{ ops: []struct { @@ -775,10 +834,15 @@ func TestLevelWriter(t *testing.T) { log.WithLevel(InfoLevel).Msg("7") log.WithLevel(WarnLevel).Msg("8") log.WithLevel(ErrorLevel).Msg("9") + log.WithLevel(FatalLevel).Msg("10") + log.WithLevel(PanicLevel).Msg("11") log.WithLevel(NoLevel).Msg("nolevel-2") - log.WithLevel(-1).Msg("-1") // Same as TraceLevel - log.WithLevel(-2).Msg("-2") // Will log - log.WithLevel(-3).Msg("-3") // Will not log + log.WithLevel(-1).Msg("-1") // Same as TraceLevel + log.WithLevel(-2).Msg("-2") // Will log + log.WithLevel(-3).Msg("-3") // Will not log + log.WithLevel(Disabled).Msg("Disabled") // Will not log + log.Err(nil).Msg("e-1") // Will log at InfoLevel + log.Err(errors.New("some error")).Msg("e-2") // Will log at ErrorLevel want := []struct { l Level @@ -795,15 +859,168 @@ func TestLevelWriter(t *testing.T) { {InfoLevel, `{"level":"info","message":"7"}` + "\n"}, {WarnLevel, `{"level":"warn","message":"8"}` + "\n"}, {ErrorLevel, `{"level":"error","message":"9"}` + "\n"}, + {FatalLevel, `{"level":"fatal","message":"10"}` + "\n"}, + {PanicLevel, `{"level":"panic","message":"11"}` + "\n"}, {NoLevel, `{"message":"nolevel-2"}` + "\n"}, {Level(-1), `{"level":"trace","message":"-1"}` + "\n"}, {Level(-2), `{"level":"-2","message":"-2"}` + "\n"}, + {InfoLevel, `{"level":"info","message":"e-1"}` + "\n"}, + {ErrorLevel, `{"level":"error","error":"some error","message":"e-2"}` + "\n"}, } if got := lw.ops; !reflect.DeepEqual(got, want) { t.Errorf("invalid ops:\ngot:\n%v\nwant:\n%v", got, want) } } +func TestDisabledLevel(t *testing.T) { + lw := &levelWriter{ + ops: []struct { + l Level + p string + }{}, + } + + // Allow extra-verbose logs. + SetGlobalLevel(TraceLevel - 1) + log := New(lw).Level(Disabled) + log.Error().Msg("0") // will not log + log.Log().Msg("nolevel-1") // will not log + log.WithLevel(ErrorLevel).Msg("3") // will not log + log.WithLevel(NoLevel).Msg("nolevel-2") // will not log + log.WithLevel(Disabled).Msg("Disabled") // will not log + + want := []struct { + l Level + p string + }{} + if got := lw.ops; !reflect.DeepEqual(got, want) { + t.Errorf("invalid ops:\ngot:\n%v\nwant:\n%v", got, want) + } +} + +func TestPanicLevel(t *testing.T) { + lw := &levelWriter{ + ops: []struct { + l Level + p string + }{}, + } + + // Allow extra-verbose logs. + SetGlobalLevel(TraceLevel - 1) + log := New(lw).Level(TraceLevel - 1) + + // Catch the panic from log.Panic().Msg("1") + defer func() { + if r := recover(); r == nil { + t.Error("expected panic from log.Panic()") + } + }() + log.Panic().Msg("1") + log.WithLevel(PanicLevel).Msg("2") + + want := []struct { + l Level + p string + }{ + {PanicLevel, `{"level":"panic","message":"1"}` + "\n"}, + {PanicLevel, `{"level":"panic","message":"2"}` + "\n"}, + } + if got := lw.ops; !reflect.DeepEqual(got, want) { + t.Errorf("invalid ops:\ngot:\n%v\nwant:\n%v", got, want) + } +} + +func TestFatalLevel(t *testing.T) { + lw := &levelWriter{ + ops: []struct { + l Level + p string + }{}, + } + + // Allow extra-verbose logs. + SetGlobalLevel(TraceLevel - 1) + log := New(lw).Level(TraceLevel - 1) + + // Set FatalExitFunc to panic so we can catch it + oldFatalExitFunc := FatalExitFunc + FatalExitFunc = func() { panic("fatal exit") } + defer func() { FatalExitFunc = oldFatalExitFunc }() + + // Catch the panic from log.Fatal().Msg("1") + defer func() { + if r := recover(); r == nil || r != "fatal exit" { + t.Errorf("expected panic 'fatal exit' from log.Fatal(), got %v", r) + } + }() + log.Fatal().Msg("1") + log.WithLevel(FatalLevel).Msg("2") + + want := []struct { + l Level + p string + }{ + {FatalLevel, `{"level":"fatal","message":"1"}` + "\n"}, + {FatalLevel, `{"level":"fatal","message":"2"}` + "\n"}, + } + if got := lw.ops; !reflect.DeepEqual(got, want) { + t.Errorf("invalid ops:\ngot:\n%v\nwant:\n%v", got, want) + } +} + +func TestFatalDisabled(t *testing.T) { + out := &bytes.Buffer{} + log := New(out).Level(PanicLevel) // Disable FatalLevel + + // Set FatalExitFunc to set a flag + var fatalCalled bool + oldFatalExitFunc := FatalExitFunc + FatalExitFunc = func() { fatalCalled = true } + defer func() { FatalExitFunc = oldFatalExitFunc }() + + // Call Fatal, which should be disabled, call done, and return nil + e := log.Fatal() + if e != nil { + t.Error("Expected nil event when Fatal is disabled") + } + if !fatalCalled { + t.Error("Expected FatalExitFunc to be called when Fatal is disabled") + } + if out.Len() > 0 { + t.Errorf("Expected no output when Fatal is disabled, got: %s", out.String()) + } +} + +func TestPanicDisabled(t *testing.T) { + out := &bytes.Buffer{} + log := New(out).Level(Disabled) // Disable all levels + + // Call Panic, which should be disabled, call done, and panic with "" + defer func() { + if r := recover(); r == nil || r != "" { + t.Errorf("Expected panic with empty string when Panic is disabled, got %v", r) + } + }() + e := log.Panic() + if e != nil { + t.Error("Expected nil event when Panic is disabled") + } + if out.Len() > 0 { + t.Errorf("Expected no output when Panic is disabled, got: %s", out.String()) + } +} + +func TestLoggerShouldWithNilWriter(t *testing.T) { + // Create a logger with nil writer to test the should method's nil check + log := Logger{w: nil, level: TraceLevel} + + e := log.Info() + if e != nil { + t.Error("Expected nil event when writer is nil") + } +} + func TestContextTimestamp(t *testing.T) { TimestampFunc = func() time.Time { return time.Date(2001, time.February, 3, 4, 5, 6, 7, time.UTC) @@ -864,6 +1081,17 @@ func TestOutputWithTimestamp(t *testing.T) { } } +func TestOutputWithContext(t *testing.T) { + ignoredOut := &bytes.Buffer{} + out := &bytes.Buffer{} + log := New(ignoredOut).With().Str("foo", "bar").Logger().Output(out) + log.Log().Msg("hello world") + + if got, want := decodeIfBinaryToString(out.Bytes()), `{"foo":"bar","message":"hello world"}`+"\n"; got != want { + t.Errorf("invalid log output:\ngot: %v\nwant: %v", got, want) + } +} + func TestCallerMarshalFunc(t *testing.T) { out := &bytes.Buffer{} log := New(out) @@ -970,6 +1198,22 @@ func TestUpdateEmptyContext(t *testing.T) { } } +func TestUpdateContextOnDisabledLogger(t *testing.T) { + var buf bytes.Buffer + log := disabledLogger + + log.UpdateContext(func(c Context) Context { + return c.Str("foo", "bar") + }) + log.Info().Msg("no panic") + + want := "" + + if got := decodeIfBinaryToString(buf.Bytes()); got != want { + t.Errorf("invalid log output:\ngot: %q\nwant: %q", got, want) + } +} + func TestLevel_String(t *testing.T) { tests := []struct { name string @@ -1097,6 +1341,14 @@ func TestUnmarshalTextLevel(t *testing.T) { } } +func TestUnmarshalTextLevelNil(t *testing.T) { + var l *Level + err := l.UnmarshalText([]byte("info")) + if err == nil || err.Error() != "can't unmarshal a nil *Level" { + t.Errorf("UnmarshalText() on nil *Level error = %v, want 'can't unmarshal a nil *Level'", err) + } +} + func TestHTMLNoEscaping(t *testing.T) { out := &bytes.Buffer{} log := New(out) diff --git a/sampler_test.go b/sampler_test.go index b986375..a26082f 100644 --- a/sampler_test.go +++ b/sampler_test.go @@ -135,3 +135,48 @@ func TestBurst(t *testing.T) { } } } + +func TestLevelSampler(t *testing.T) { + // Create mock samplers that return true for specific levels + traceSampler := &BasicSampler{N: 1} // Always sample + debugSampler := &BasicSampler{N: 0} // Never sample + infoSampler := &BasicSampler{N: 1} // Always sample + warnSampler := &BasicSampler{N: 0} // Never sample + errorSampler := &BasicSampler{N: 1} // Always sample + + sampler := LevelSampler{ + TraceSampler: traceSampler, + DebugSampler: debugSampler, + InfoSampler: infoSampler, + WarnSampler: warnSampler, + ErrorSampler: errorSampler, + } + + // Test each level + if !sampler.Sample(TraceLevel) { + t.Error("TraceLevel should be sampled") + } + if sampler.Sample(DebugLevel) { + t.Error("DebugLevel should not be sampled") + } + if !sampler.Sample(InfoLevel) { + t.Error("InfoLevel should be sampled") + } + if sampler.Sample(WarnLevel) { + t.Error("WarnLevel should not be sampled") + } + if !sampler.Sample(ErrorLevel) { + t.Error("ErrorLevel should be sampled") + } + + // Test levels not covered by the LevelSampler sampler (FatalLevel, PanicLevel, NoLevel) - should return true + if !sampler.Sample(FatalLevel) { + t.Error("FatalLevel should return true when no sampler is set") + } + if !sampler.Sample(PanicLevel) { + t.Error("PanicLevel should return true when no sampler is set") + } + if !sampler.Sample(NoLevel) { + t.Error("NoLevel should return true when no sampler is set") + } +} diff --git a/syslog_test.go b/syslog_test.go index e889b01..df1aa85 100644 --- a/syslog_test.go +++ b/syslog_test.go @@ -5,6 +5,7 @@ package zerolog import ( "bytes" + "io" "reflect" "strings" "testing" @@ -19,7 +20,7 @@ type syslogTestWriter struct { } func (w *syslogTestWriter) Write(p []byte) (int, error) { - return 0, nil + return len(p), nil } func (w *syslogTestWriter) Trace(m string) error { w.events = append(w.events, syslogEvent{"Trace", m}) @@ -106,3 +107,130 @@ func TestSyslogWriter_WithCEE(t *testing.T) { t.Errorf("Bad CEE message start: want %v, got %v", want, got) } } + +type errorSyslogWriter struct { + *syslogTestWriter + writeError error +} + +func (w *errorSyslogWriter) Write(p []byte) (int, error) { + if w.writeError != nil { + return 0, w.writeError + } + return len(p), nil +} + +func TestSyslogWriter_Write(t *testing.T) { + // Test Write method without prefix + sw := &syslogTestWriter{} + writer := SyslogLevelWriter(sw) + + data := []byte("test message") + n, err := writer.Write(data) + if err != nil { + t.Errorf("Write failed: %v", err) + } + if n != len(data) { + t.Errorf("Write returned wrong length: got %d, want %d", n, len(data)) + } + + // Test Write method with CEE prefix + sw2 := &syslogTestWriter{} + writer2 := SyslogCEEWriter(sw2) + + data2 := []byte("test message") + n2, err2 := writer2.Write(data2) + if err2 != nil { + t.Errorf("Write with CEE failed: %v", err2) + } + expectedLen := len(ceePrefix) + len(data2) + if n2 != expectedLen { + t.Errorf("Write with CEE returned wrong length: got %d, want %d", n2, expectedLen) + } + + // Test Write method with CEE prefix and error on prefix write + sw3 := &errorSyslogWriter{syslogTestWriter: &syslogTestWriter{}, writeError: io.EOF} + writer3 := SyslogCEEWriter(sw3) + + _, err3 := writer3.Write(data2) + if err3 != io.EOF { + t.Errorf("Write with CEE error failed: got %v, want %v", err3, io.EOF) + } +} + +func TestSyslogWriter_WriteLevel_AllLevels(t *testing.T) { + sw := &syslogTestWriter{} + writer := SyslogLevelWriter(sw) + + // Test all levels to ensure full coverage + writer.WriteLevel(TraceLevel, []byte(`{"level":"trace","message":"trace"}`+"\n")) + writer.WriteLevel(DebugLevel, []byte(`{"level":"debug","message":"debug"}`+"\n")) + writer.WriteLevel(InfoLevel, []byte(`{"level":"info","message":"info"}`+"\n")) + writer.WriteLevel(WarnLevel, []byte(`{"level":"warn","message":"warn"}`+"\n")) + writer.WriteLevel(ErrorLevel, []byte(`{"level":"error","message":"error"}`+"\n")) + writer.WriteLevel(FatalLevel, []byte(`{"level":"fatal","message":"fatal"}`+"\n")) + writer.WriteLevel(PanicLevel, []byte(`{"level":"panic","message":"panic"}`+"\n")) + writer.WriteLevel(NoLevel, []byte(`{"message":"nolevel"}`+"\n")) + + want := []syslogEvent{ + {"Debug", `{"level":"debug","message":"debug"}` + "\n"}, + {"Info", `{"level":"info","message":"info"}` + "\n"}, + {"Warning", `{"level":"warn","message":"warn"}` + "\n"}, + {"Err", `{"level":"error","message":"error"}` + "\n"}, + {"Emerg", `{"level":"fatal","message":"fatal"}` + "\n"}, + {"Crit", `{"level":"panic","message":"panic"}` + "\n"}, + {"Info", `{"message":"nolevel"}` + "\n"}, + } + if got := sw.events; !reflect.DeepEqual(got, want) { + t.Errorf("Invalid syslog message routing: want %v, got %v", want, got) + } +} + +type closableSyslogWriter struct { + *syslogTestWriter + closed bool +} + +func (w *closableSyslogWriter) Close() error { + w.closed = true + return nil +} + +func TestSyslogWriter_Close(t *testing.T) { + // Test with closable writer + sw := &closableSyslogWriter{syslogTestWriter: &syslogTestWriter{}} + writer := SyslogLevelWriter(sw).(syslogWriter) // Cast to concrete type to access Close + + err := writer.Close() + if err != nil { + t.Errorf("Close failed: %v", err) + } + if !sw.closed { + t.Error("Close was not called on underlying writer") + } + + // Test with non-closable writer + sw2 := &syslogTestWriter{} + writer2 := SyslogLevelWriter(sw2).(syslogWriter) // Cast to concrete type to access Close + + err = writer2.Close() + if err != nil { + t.Errorf("Close failed for non-closable writer: %v", err) + } +} + +func TestSyslogWriter_WriteLevel_InvalidLevel(t *testing.T) { + sw := &syslogTestWriter{} + writer := SyslogLevelWriter(sw) + + // Test invalid level - should panic + defer func() { + if r := recover(); r == nil { + t.Error("Expected panic for invalid level") + } else if r != "invalid level" { + t.Errorf("Expected panic 'invalid level', got %v", r) + } + }() + + writer.WriteLevel(Level(100), []byte("test")) +} diff --git a/writer_test.go b/writer_test.go index f2a61df..82780b4 100644 --- a/writer_test.go +++ b/writer_test.go @@ -12,6 +12,36 @@ import ( "testing" ) +type closableBuffer struct { + *bytes.Buffer + closed bool + closeError error +} + +func (cb *closableBuffer) Close() error { + cb.closed = true + return cb.closeError +} + +type errorWriter struct { + writeError error + shortWrite bool +} + +func (ew *errorWriter) Write(p []byte) (int, error) { + if ew.writeError != nil { + return 0, ew.writeError + } + if ew.shortWrite { + return len(p) - 1, nil // Return short write + } + return len(p), nil +} + +func (ew *errorWriter) WriteLevel(level Level, p []byte) (int, error) { + return ew.Write(p) +} + func TestMultiSyslogWriter(t *testing.T) { sw := &syslogTestWriter{} log := New(MultiLevelWriter(SyslogLevelWriter(sw))) @@ -250,3 +280,271 @@ func TestTriggerLevelWriter(t *testing.T) { }) } } + +func TestLevelWriterAdapter_Close(t *testing.T) { + // Test with closable writer + buf := &bytes.Buffer{} + adapter := LevelWriterAdapter{Writer: buf} + + // bytes.Buffer doesn't implement io.Closer, so Close should return nil + err := adapter.Close() + if err != nil { + t.Errorf("Close should not return error for non-closable writer: %v", err) + } + + // Test with closable writer + closableBuf := &closableBuffer{Buffer: &bytes.Buffer{}} + adapter2 := LevelWriterAdapter{Writer: closableBuf} + + err = adapter2.Close() + if err != nil { + t.Errorf("Close should not return error: %v", err) + } + if !closableBuf.closed { + t.Error("Close should have been called on closable writer") + } +} + +func TestSyncWriter(t *testing.T) { + buf := &bytes.Buffer{} + + // Test SyncWriter with regular io.Writer + syncWriter := SyncWriter(buf) + + // Test Write + data := []byte("test data") + n, err := syncWriter.Write(data) + if err != nil { + t.Errorf("Write failed: %v", err) + } + if n != len(data) { + t.Errorf("Write returned wrong length: got %d, want %d", n, len(data)) + } + if got := buf.String(); got != string(data) { + t.Errorf("Write wrote wrong data: got %q, want %q", got, string(data)) + } + + // Test SyncWriter with LevelWriter - use it with a logger + levelBuf := &bytes.Buffer{} + levelWriter := LevelWriterAdapter{levelBuf} + syncLevelWriter := SyncWriter(levelWriter) + + logger := New(syncLevelWriter) + logger.Info().Msg("test message") + + expected := `{"level":"info","message":"test message"}` + "\n" + if got := levelBuf.String(); got != expected { + t.Errorf("SyncWriter with LevelWriter failed: got %q, want %q", got, expected) + } + + // Test SyncWriter Close with closable writer + closableBuf := &closableBuffer{Buffer: &bytes.Buffer{}, closed: false} + closableSyncWriter := SyncWriter(closableBuf) + + if closeable, ok := closableSyncWriter.(io.Closer); !ok { + t.Error("SyncWriter should implement Close method") + } else { + err := closeable.Close() + if err != nil { + t.Errorf("Close failed: %v", err) + } + } + if !closableBuf.closed { + t.Error("Close should have been called on closable writer") + } + + // Test SyncWriter Close with closable writer that returns error + errorBuf := &closableBuffer{Buffer: &bytes.Buffer{}, closed: false, closeError: io.EOF} + errorSyncWriter := SyncWriter(errorBuf) + + if closeable, ok := errorSyncWriter.(io.Closer); !ok { + t.Error("SyncWriter should implement Close method") + } else { + err := closeable.Close() + if err != io.EOF { + t.Errorf("Close should have returned EOF error, got: %v", err) + } + } + if !errorBuf.closed { + t.Error("Close should have been called on closable writer") + } +} + +func TestMultiLevelWriter_Write(t *testing.T) { + // Test successful writes + buf1 := &bytes.Buffer{} + buf2 := &bytes.Buffer{} + + multiWriter := MultiLevelWriter(buf1, buf2) + + data := []byte("test data") + n, err := multiWriter.Write(data) + if err != nil { + t.Errorf("Write failed: %v", err) + } + if n != len(data) { + t.Errorf("Write returned wrong length: got %d, want %d", n, len(data)) + } + + if got1 := buf1.String(); got1 != string(data) { + t.Errorf("First writer got wrong data: got %q, want %q", got1, string(data)) + } + if got2 := buf2.String(); got2 != string(data) { + t.Errorf("Second writer got wrong data: got %q, want %q", got2, string(data)) + } + + // Test with error writer + errorWriter1 := &errorWriter{writeError: io.EOF} + buf3 := &bytes.Buffer{} + + errorMultiWriter := MultiLevelWriter(errorWriter1, buf3) + + _, err = errorMultiWriter.Write(data) + if err != io.EOF { + t.Errorf("Write should have returned EOF error, got: %v", err) + } + + // Test with short write + shortWriter := &errorWriter{shortWrite: true} + buf4 := &bytes.Buffer{} + + shortMultiWriter := MultiLevelWriter(shortWriter, buf4) + + _, err = shortMultiWriter.Write(data) + if err != io.ErrShortWrite { + t.Errorf("Write should have returned ErrShortWrite, got: %v", err) + } +} + +func TestMultiLevelWriter_WriteLevel(t *testing.T) { + // Test successful writes + buf1 := &bytes.Buffer{} + buf2 := &bytes.Buffer{} + + multiWriter := MultiLevelWriter(buf1, buf2) + + data := []byte("test level data") + n, err := multiWriter.WriteLevel(InfoLevel, data) + if err != nil { + t.Errorf("WriteLevel failed: %v", err) + } + if n != len(data) { + t.Errorf("WriteLevel returned wrong length: got %d, want %d", n, len(data)) + } + + if got1 := buf1.String(); got1 != string(data) { + t.Errorf("First writer got wrong data: got %q, want %q", got1, string(data)) + } + if got2 := buf2.String(); got2 != string(data) { + t.Errorf("Second writer got wrong data: got %q, want %q", got2, string(data)) + } + + // Test with error writer + errorWriter1 := &errorWriter{writeError: io.EOF} + buf3 := &bytes.Buffer{} + + errorMultiWriter := MultiLevelWriter(errorWriter1, buf3) + + _, err = errorMultiWriter.WriteLevel(InfoLevel, data) + if err != io.EOF { + t.Errorf("WriteLevel should have returned EOF error, got: %v", err) + } +} + +func TestMultiLevelWriter_Close(t *testing.T) { + buf1 := &closableBuffer{Buffer: &bytes.Buffer{}, closed: false} + buf2 := &bytes.Buffer{} // non-closable + + multiWriter := MultiLevelWriter(buf1, buf2) + + // Cast to concrete type to access Close + mw := multiWriter.(multiLevelWriter) + + err := mw.Close() + if err != nil { + t.Errorf("Close failed: %v", err) + } + + if !buf1.closed { + t.Error("First closable writer should have been closed") + } + + // Test multiLevelWriter Close with error + errorBuf1 := &closableBuffer{Buffer: &bytes.Buffer{}, closed: false, closeError: io.EOF} + errorBuf2 := &bytes.Buffer{} // non-closable + + errorMultiWriter := MultiLevelWriter(errorBuf1, errorBuf2) + emw := errorMultiWriter.(multiLevelWriter) + + err = emw.Close() + if err != io.EOF { + t.Errorf("Close should have returned EOF error, got: %v", err) + } + if !errorBuf1.closed { + t.Error("First closable writer should have been closed") + } +} + +func TestNewTestWriter(t *testing.T) { + writer := NewTestWriter(t) + + if writer.T != t { + t.Error("NewTestWriter should set the testing interface") + } + if writer.Frame != 0 { + t.Errorf("NewTestWriter should set Frame to 0, got %d", writer.Frame) + } +} + +func TestFilteredLevelWriter_Write(t *testing.T) { + buf := &bytes.Buffer{} + filteredWriter := FilteredLevelWriter{ + Writer: LevelWriterAdapter{buf}, + Level: InfoLevel, + } + + data := []byte("test data") + n, err := filteredWriter.Write(data) + if err != nil { + t.Errorf("Write failed: %v", err) + } + if n != len(data) { + t.Errorf("Write returned wrong length: got %d, want %d", n, len(data)) + } + + if got := buf.String(); got != string(data) { + t.Errorf("Write should always write: got %q, want %q", got, string(data)) + } +} + +func TestFilteredLevelWriter_Close(t *testing.T) { + buf := &closableBuffer{Buffer: &bytes.Buffer{}, closed: false} + filteredWriter := FilteredLevelWriter{ + Writer: LevelWriterAdapter{buf}, + Level: InfoLevel, + } + + err := filteredWriter.Close() + if err != nil { + t.Errorf("Close failed: %v", err) + } + + if !buf.closed { + t.Error("Underlying closable writer should have been closed") + } + + // Test FilteredLevelWriter Close with error + errorBuf := &closableBuffer{Buffer: &bytes.Buffer{}, closed: false, closeError: io.EOF} + errorFilteredWriter := FilteredLevelWriter{ + Writer: LevelWriterAdapter{errorBuf}, + Level: InfoLevel, + } + + err = errorFilteredWriter.Close() + if err != io.EOF { + t.Errorf("Close should have returned EOF error, got: %v", err) + } + if !errorBuf.closed { + t.Error("Underlying closable writer should have been closed") + } +}