Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 15 additions & 17 deletions internal/database/config.go
Original file line number Diff line number Diff line change
Expand Up @@ -147,14 +147,10 @@ func (c *Config) validatePoolSettings() error {
return nil
}

// DSN generates a PostgreSQL connection string
// If passwordOverride is provided, it's used instead of config.Password
func (c *Config) DSN(passwordOverride string) string {
password := c.Password
if passwordOverride != "" {
password = passwordOverride
}

// dsn formats a PostgreSQL connection string with the supplied password. It is
// the single source of the DSN field layout so DSN and RedactedDSN can never
// drift (a field added here applies to both the real and the redacted form).
func (c *Config) dsn(password string) string {
return fmt.Sprintf(
"host=%s port=%d user=%s password=%s dbname=%s sslmode=%s connect_timeout=%d",
c.Host,
Expand All @@ -167,17 +163,19 @@ func (c *Config) DSN(passwordOverride string) string {
)
}

// DSN generates a PostgreSQL connection string
// If passwordOverride is provided, it's used instead of config.Password
func (c *Config) DSN(passwordOverride string) string {
password := c.Password
if passwordOverride != "" {
password = passwordOverride
}
return c.dsn(password)
}

// RedactedDSN returns a DSN string with the password masked, safe for logging
func (c *Config) RedactedDSN() string {
return fmt.Sprintf(
"host=%s port=%d user=%s password=***** dbname=%s sslmode=%s connect_timeout=%d",
c.Host,
c.Port,
c.User,
c.Database,
c.SSLMode,
int(c.ConnectTimeout.Seconds()),
)
return c.dsn("*****")
}

// Helper functions for environment variable parsing
Expand Down
23 changes: 23 additions & 0 deletions internal/database/coverage_extra_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package database

import (
"context"
"strings"
"testing"
"time"

Expand Down Expand Up @@ -32,6 +33,28 @@ func TestRedactedDSN(t *testing.T) {
assert.NotContains(t, dsn, "supersecret")
}

// TestRedactedDSN_SharesLayoutWithDSN guards the 06-N2 dedup: RedactedDSN and
// DSN both derive from the single dsn() formatter, so the redacted string must
// equal the real DSN with only the password swapped for "*****". If a future
// field is added to dsn(), this stays green; if someone re-forks the format
// string for only one of the two, it fails.
func TestRedactedDSN_SharesLayoutWithDSN(t *testing.T) {
cfg := &Config{
Host: "db.example.com",
Port: 5432,
User: "admin",
Password: "supersecret",
Database: "cudly",
SSLMode: "require",
ConnectTimeout: 10 * time.Second,
}

real := cfg.DSN("")
redacted := cfg.RedactedDSN()
expected := strings.Replace(real, "password=supersecret", "password=*****", 1)
assert.Equal(t, expected, redacted)
}

// Tests for extractPasswordFromSecret
func TestExtractPasswordFromSecret_JSONWithPassword(t *testing.T) {
secret := `{"username":"admin","password":"db-pass-123","host":"db.example.com"}`
Expand Down
8 changes: 5 additions & 3 deletions pkg/common/tokens.go
Original file line number Diff line number Diff line change
Expand Up @@ -48,14 +48,16 @@ func DeriveIdempotencyToken(executionID string, recIndex int) string {
// keeps just enough of the prefix to correlate log lines for a single purchase
// while avoiding emitting the whole caller-supplied token into persistent logs
// (a stable per-execution identifier that should not leak verbatim). An empty
// token yields "(none)"; a token of 8 chars or fewer is returned unchanged
// since there is nothing left to redact.
// token yields "(none)". A token of 8 chars or fewer is fully redacted to
// "(redacted)" rather than echoed: an 8-char prefix of an 8-char input is the
// whole value, so for short inputs (e.g. a short secret a future caller might
// pass) nothing of the token is emitted.
func MaskToken(token string) string {
if token == "" {
return "(none)"
}
if len(token) <= 8 {
return token
return "(redacted)"
}
return token[:8] + "..."
}
Expand Down
7 changes: 5 additions & 2 deletions pkg/common/tokens_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -81,8 +81,11 @@ func TestMaskToken_NeverEmitsFullToken(t *testing.T) {

func TestMaskToken_EmptyAndShort(t *testing.T) {
assert.Equal(t, "(none)", MaskToken(""), "empty token must be reported as (none)")
assert.Equal(t, "abc", MaskToken("abc"), "tokens of <=8 chars have nothing to redact")
assert.Equal(t, "12345678", MaskToken("12345678"), "exactly 8 chars is returned unchanged")
// Short inputs (<=8 chars) are fully redacted, never echoed: an 8-char prefix
// of an 8-char value would leak the whole secret (CodeRabbit 10-L6).
assert.Equal(t, "(redacted)", MaskToken("abc"), "short tokens must be fully redacted, not echoed")
assert.Equal(t, "(redacted)", MaskToken("12345678"), "exactly 8 chars is fully redacted")
assert.NotContains(t, MaskToken("secret77"), "secret", "no part of a short secret may appear in the masked form")
assert.Equal(t, "12345678...", MaskToken("123456789"), "9 chars is truncated to 8 + ellipsis")
}

Expand Down
31 changes: 24 additions & 7 deletions pkg/errors/errors.go
Original file line number Diff line number Diff line change
@@ -1,4 +1,14 @@
// Package errors provides custom error types for CUDly.
//
// Type-level Is matching: every error type in this package implements Is by
// matching purely on the dynamic type of the target, ignoring the target's
// struct fields. That makes the zero-value pointer of each type a usable
// sentinel, e.g. errors.Is(err, &NotFoundError{}) reports whether err is (or
// wraps) any *NotFoundError. The corollary is that the target's fields are
// decorative for comparison: errors.Is(someNotFound, &NotFoundError{ID: "x"})
// is true regardless of whether someNotFound.ID == "x". To assert on specific
// fields, use errors.As to extract the concrete value and inspect it directly,
// or one of the Is*Error helpers below (which also use errors.As).
package errors

import (
Expand All @@ -24,7 +34,8 @@ func (e *NotFoundError) Error() string {
return fmt.Sprintf("%s not found", e.Resource)
}

// Is implements error comparison
// Is reports whether target is the same error type (field-insensitive; see
// the package doc on type-level Is matching).
func (e *NotFoundError) Is(target error) bool {
_, ok := target.(*NotFoundError)
return ok
Expand Down Expand Up @@ -65,7 +76,8 @@ func (e *ValidationError) Error() string {
return fmt.Sprintf("validation error: %s", e.Message)
}

// Is implements error comparison
// Is reports whether target is the same error type (field-insensitive; see
// the package doc on type-level Is matching).
func (e *ValidationError) Is(target error) bool {
_, ok := target.(*ValidationError)
return ok
Expand Down Expand Up @@ -98,7 +110,8 @@ func (e *AuthenticationError) Error() string {
return "authentication failed"
}

// Is implements error comparison
// Is reports whether target is the same error type (field-insensitive; see
// the package doc on type-level Is matching).
func (e *AuthenticationError) Is(target error) bool {
_, ok := target.(*AuthenticationError)
return ok
Expand Down Expand Up @@ -131,7 +144,8 @@ func (e *AuthorizationError) Error() string {
return "not authorized"
}

// Is implements error comparison
// Is reports whether target is the same error type (field-insensitive; see
// the package doc on type-level Is matching).
func (e *AuthorizationError) Is(target error) bool {
_, ok := target.(*AuthorizationError)
return ok
Expand Down Expand Up @@ -169,7 +183,8 @@ func (e *ConflictError) Error() string {
return fmt.Sprintf("%s already exists", e.Resource)
}

// Is implements error comparison
// Is reports whether target is the same error type (field-insensitive; see
// the package doc on type-level Is matching).
func (e *ConflictError) Is(target error) bool {
_, ok := target.(*ConflictError)
return ok
Expand Down Expand Up @@ -202,7 +217,8 @@ func (e *RateLimitError) Error() string {
return "rate limit exceeded"
}

// Is implements error comparison
// Is reports whether target is the same error type (field-insensitive; see
// the package doc on type-level Is matching).
func (e *RateLimitError) Is(target error) bool {
_, ok := target.(*RateLimitError)
return ok
Expand Down Expand Up @@ -240,7 +256,8 @@ func (e *ServiceError) Unwrap() error {
return e.Err
}

// Is implements error comparison
// Is reports whether target is the same error type (field-insensitive; see
// the package doc on type-level Is matching).
func (e *ServiceError) Is(target error) bool {
_, ok := target.(*ServiceError)
return ok
Expand Down
12 changes: 12 additions & 0 deletions pkg/errors/errors_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,18 @@ func TestNotFoundError(t *testing.T) {
assert.True(t, errors.Is(err, target))
})

t.Run("Is is field-insensitive (type-level matching, 10-N4)", func(t *testing.T) {
t.Parallel()
// Documented contract: Is matches on type only, ignoring the target's
// fields. A target with a mismatching ID still reports true.
err := NewNotFoundError("User", "123")
assert.True(t, errors.Is(err, &NotFoundError{ID: "different"}),
"Is must match on type alone, regardless of target fields")
assert.True(t, errors.Is(err, &NotFoundError{Resource: "Other", ID: "x", Message: "y"}))
// A different type must not match.
assert.False(t, errors.Is(err, &ValidationError{}))
})

t.Run("IsNotFoundError helper", func(t *testing.T) {
t.Parallel()
err := NewNotFoundError("User", "123")
Expand Down
49 changes: 33 additions & 16 deletions pkg/logging/logger.go
Original file line number Diff line number Diff line change
Expand Up @@ -10,6 +10,7 @@ import (
"os"
"sort"
"strings"
"sync/atomic"
)

// Level represents a logging level
Expand All @@ -26,14 +27,29 @@ const (
LevelError
)

// Logger provides structured logging capabilities
// Logger provides structured logging capabilities.
//
// level is stored as an atomic.Int32 so that SetLevel/SetLevelValue (called at
// runtime, sometimes after worker goroutines have launched) and the level reads
// performed by every Debug/Info/Warn/Error call are race-free under the
// concurrent fan-out, which logs from many goroutines.
type Logger struct {
level Level
level atomic.Int32
logger *log.Logger
prefix string
metadata map[string]interface{}
}

// getLevel returns the logger's current level, read atomically.
func (l *Logger) getLevel() Level {
return Level(l.level.Load())
}

// setLevel stores the logger's level atomically.
func (l *Logger) setLevel(level Level) {
l.level.Store(int32(level))
}

// Config holds logger configuration
type Config struct {
Level string
Expand Down Expand Up @@ -79,37 +95,38 @@ func New(cfg Config) *Logger {
flags = 0 // Use custom time format
}

return &Logger{
level: ParseLevel(cfg.Level),
l := &Logger{
logger: log.New(output, cfg.Prefix, flags),
prefix: cfg.Prefix,
metadata: make(map[string]interface{}),
}
l.setLevel(ParseLevel(cfg.Level))
return l
}

// SetLevel sets the logging level
func SetLevel(level string) {
defaultLogger.level = ParseLevel(level)
defaultLogger.setLevel(ParseLevel(level))
}

// SetLevelValue sets the logging level using a Level value
func SetLevelValue(level Level) {
defaultLogger.level = level
defaultLogger.setLevel(level)
}

// GetLevel returns the current log level
func GetLevel() Level {
return defaultLogger.level
return defaultLogger.getLevel()
}

// With creates a new logger with additional metadata
func (l *Logger) With(key string, value interface{}) *Logger {
newLogger := &Logger{
level: l.level,
logger: l.logger,
prefix: l.prefix,
metadata: make(map[string]interface{}),
}
newLogger.setLevel(l.getLevel())
for k, v := range l.metadata {
newLogger.metadata[k] = v
}
Expand Down Expand Up @@ -138,56 +155,56 @@ func (l *Logger) formatMessage(msg string) string {

// Debug logs a debug message
func (l *Logger) Debug(msg string) {
if l.level <= LevelDebug {
if l.getLevel() <= LevelDebug {
l.logger.Printf("[DEBUG] %s", l.formatMessage(msg))
}
}

// Debugf logs a formatted debug message
func (l *Logger) Debugf(format string, args ...interface{}) {
if l.level <= LevelDebug {
if l.getLevel() <= LevelDebug {
l.logger.Printf("[DEBUG] %s", l.formatMessage(fmt.Sprintf(format, args...)))
}
}

// Info logs an info message
func (l *Logger) Info(msg string) {
if l.level <= LevelInfo {
if l.getLevel() <= LevelInfo {
l.logger.Printf("[INFO] %s", l.formatMessage(msg))
}
}

// Infof logs a formatted info message
func (l *Logger) Infof(format string, args ...interface{}) {
if l.level <= LevelInfo {
if l.getLevel() <= LevelInfo {
l.logger.Printf("[INFO] %s", l.formatMessage(fmt.Sprintf(format, args...)))
}
}

// Warn logs a warning message
func (l *Logger) Warn(msg string) {
if l.level <= LevelWarn {
if l.getLevel() <= LevelWarn {
l.logger.Printf("[WARN] %s", l.formatMessage(msg))
}
}

// Warnf logs a formatted warning message
func (l *Logger) Warnf(format string, args ...interface{}) {
if l.level <= LevelWarn {
if l.getLevel() <= LevelWarn {
l.logger.Printf("[WARN] %s", l.formatMessage(fmt.Sprintf(format, args...)))
}
}

// Error logs an error message
func (l *Logger) Error(msg string) {
if l.level <= LevelError {
if l.getLevel() <= LevelError {
l.logger.Printf("[ERROR] %s", l.formatMessage(msg))
}
}

// Errorf logs a formatted error message
func (l *Logger) Errorf(format string, args ...interface{}) {
if l.level <= LevelError {
if l.getLevel() <= LevelError {
l.logger.Printf("[ERROR] %s", l.formatMessage(fmt.Sprintf(format, args...)))
}
}
Expand Down
Loading
Loading