Merge branch 'main' of github.com:binwiederhier/ntfy into 1511-ampersand

This commit is contained in:
binwiederhier
2026-06-24 05:42:12 -04:00
260 changed files with 38921 additions and 14501 deletions
+55
View File
@@ -152,6 +152,61 @@ func (l *RateLimiter) Reset() {
l.value = 0
}
// CountingReader wraps an io.Reader and counts the number of bytes read through it.
type CountingReader struct {
r io.Reader
total int64
}
// NewCountingReader creates a new CountingReader
func NewCountingReader(r io.Reader) *CountingReader {
return &CountingReader{r: r}
}
// Read passes through to the underlying reader and counts the bytes read
func (r *CountingReader) Read(p []byte) (n int, err error) {
n, err = r.r.Read(p)
r.total += int64(n)
return
}
// Total returns the total number of bytes read so far
func (r *CountingReader) Total() int64 {
return r.total
}
// LimitReader implements an io.Reader that will pass through all Read calls to the underlying
// reader r until any of the limiter's limit is reached, at which point a Read will return ErrLimitReached.
// Each limiter's value is increased after every read based on the number of bytes actually read.
type LimitReader struct {
r io.Reader
limiters []Limiter
}
// NewLimitReader creates a new LimitReader
func NewLimitReader(r io.Reader, limiters ...Limiter) *LimitReader {
return &LimitReader{
r: r,
limiters: limiters,
}
}
// Read passes through all reads to the underlying reader until any of the given limiter's limit is reached
func (r *LimitReader) Read(p []byte) (n int, err error) {
n, err = r.r.Read(p)
if n > 0 {
for i := 0; i < len(r.limiters); i++ {
if !r.limiters[i].AllowN(int64(n)) {
for j := i - 1; j >= 0; j-- {
r.limiters[j].AllowN(-int64(n)) // Revert limiters if not allowed
}
return 0, ErrLimitReached
}
}
}
return
}
// LimitWriter implements an io.Writer that will pass through all Write calls to the underlying
// writer w until any of the limiter's limit is reached, at which point a Write will return ErrLimitReached.
// Each limiter's value is increased with every write.
+99 -1
View File
@@ -2,9 +2,12 @@ package util
import (
"bytes"
"github.com/stretchr/testify/require"
"io"
"strings"
"testing"
"time"
"github.com/stretchr/testify/require"
)
func TestFixedLimiter_AllowValueReset(t *testing.T) {
@@ -147,3 +150,98 @@ func TestLimitWriter_WriteTwoDifferentLimiters_Wait_FixedLimiterFail(t *testing.
_, err = lw.Write(make([]byte, 8)) // <<< FixedLimiter fails
require.Equal(t, ErrLimitReached, err)
}
func TestCountingReader_Total(t *testing.T) {
cr := NewCountingReader(strings.NewReader("hello world"))
buf := make([]byte, 5)
n, err := cr.Read(buf)
require.Nil(t, err)
require.Equal(t, 5, n)
require.Equal(t, int64(5), cr.Total())
n, err = cr.Read(buf)
require.Nil(t, err)
require.Equal(t, 5, n)
require.Equal(t, int64(10), cr.Total())
n, err = cr.Read(buf)
require.Nil(t, err)
require.Equal(t, 1, n)
require.Equal(t, int64(11), cr.Total())
_, err = cr.Read(buf)
require.Equal(t, io.EOF, err)
require.Equal(t, int64(11), cr.Total())
}
func TestCountingReader_Empty(t *testing.T) {
cr := NewCountingReader(strings.NewReader(""))
require.Equal(t, int64(0), cr.Total())
_, err := cr.Read(make([]byte, 10))
require.Equal(t, io.EOF, err)
require.Equal(t, int64(0), cr.Total())
}
func TestLimitReader_ReadNoLimiter(t *testing.T) {
lr := NewLimitReader(strings.NewReader("hello"))
data, err := io.ReadAll(lr)
require.Nil(t, err)
require.Equal(t, "hello", string(data))
}
func TestLimitReader_ReadOneLimiter(t *testing.T) {
l := NewFixedLimiter(10)
lr := NewLimitReader(strings.NewReader("hello world!"), l)
buf := make([]byte, 5)
n, err := lr.Read(buf)
require.Nil(t, err)
require.Equal(t, 5, n)
require.Equal(t, int64(5), l.Value())
n, err = lr.Read(buf)
require.Nil(t, err)
require.Equal(t, 5, n)
require.Equal(t, int64(10), l.Value())
_, err = lr.Read(buf)
require.Equal(t, ErrLimitReached, err)
}
func TestLimitReader_ReadTwoLimiters(t *testing.T) {
l1 := NewFixedLimiter(11)
l2 := NewFixedLimiter(8)
lr := NewLimitReader(strings.NewReader("hello world!"), l1, l2)
buf := make([]byte, 5)
n, err := lr.Read(buf)
require.Nil(t, err)
require.Equal(t, 5, n)
// Second read: l2 (limit 8) should reject 5 more bytes
_, err = lr.Read(buf)
require.Equal(t, ErrLimitReached, err)
// l1 should have been reverted
require.Equal(t, int64(5), l1.Value())
require.Equal(t, int64(5), l2.Value())
}
func TestLimitReader_ReadAll(t *testing.T) {
l := NewFixedLimiter(100)
lr := NewLimitReader(strings.NewReader("hello"), l)
data, err := io.ReadAll(lr)
require.Nil(t, err)
require.Equal(t, "hello", string(data))
require.Equal(t, int64(5), l.Value())
}
func TestLimitReader_ReadExactLimit(t *testing.T) {
l := NewFixedLimiter(5)
lr := NewLimitReader(bytes.NewReader(make([]byte, 5)), l)
data, err := io.ReadAll(lr)
require.Nil(t, err)
require.Equal(t, 5, len(data))
require.Equal(t, int64(5), l.Value())
}
+2 -6
View File
@@ -4,6 +4,7 @@ import (
"bytes"
"encoding/json"
"reflect"
"slices"
"strings"
)
@@ -95,12 +96,7 @@ func coalesce(v ...any) any {
// Returns:
// - bool: True if all values are non-empty, false otherwise
func all(v ...any) bool {
for _, val := range v {
if empty(val) {
return false
}
}
return true
return !slices.ContainsFunc(v, empty)
}
// anyNonEmpty checks if at least one value in a list is non-empty.
+58 -16
View File
@@ -2,20 +2,21 @@ package util
import (
"bytes"
crand "crypto/rand"
"encoding/base64"
"encoding/json"
"errors"
"fmt"
"io"
"math"
"math/rand"
"net/netip"
"os"
"regexp"
"slices"
"strconv"
"strings"
"sync"
"time"
"unicode/utf8"
"github.com/gabriel-vasile/mimetype"
"golang.org/x/term"
@@ -28,8 +29,6 @@ const (
)
var (
random = rand.New(rand.NewSource(time.Now().UnixNano()))
randomMutex = sync.Mutex{}
sizeStrRegex = regexp.MustCompile(`(?i)^(\d+)([gmkb])?$`)
errInvalidPriority = errors.New("invalid priority")
noQuotesRegex = regexp.MustCompile(`^[-_./:@a-zA-Z0-9]+$`)
@@ -49,12 +48,7 @@ func FileExists(filename string) bool {
// Contains returns true if needle is contained in haystack
func Contains[T comparable](haystack []T, needle T) bool {
for _, s := range haystack {
if s == needle {
return true
}
}
return false
return slices.Contains(haystack, needle)
}
// ContainsIP returns true if any one of the of prefixes contains the ip.
@@ -147,14 +141,33 @@ func RandomLowerStringPrefix(prefix string, length int) string {
return randomStringPrefixWithCharset(prefix, length, randomStringLowerCaseCharset)
}
// randomStringPrefixWithCharset builds a random string from charset using crypto/rand.
// We use rejection sampling (dropping the few highest byte values that would skew the
// distribution) so every character is uniformly distributed -- important because these
// strings back security tokens (access tokens, magic-link tokens, IDs), not just labels.
func randomStringPrefixWithCharset(prefix string, length int, charset string) string {
randomMutex.Lock() // Who would have thought that random.Intn() is not thread-safe?!
defer randomMutex.Unlock()
b := make([]byte, length-len(prefix))
for i := range b {
b[i] = charset[random.Intn(len(charset))]
n := length - len(prefix)
if n <= 0 {
return prefix[:length]
}
return prefix + string(b)
result := make([]byte, n)
limit := 256 - (256 % len(charset)) // reject byte values >= limit to avoid modulo bias
buf := make([]byte, n)
for i := 0; i < n; {
if _, err := crand.Read(buf); err != nil {
panic("crypto/rand failed: " + err.Error()) // Should never happen on a sane system
}
for _, c := range buf {
if i >= n {
break
}
if int(c) < limit {
result[i] = charset[int(c)%len(charset)]
i++
}
}
}
return prefix + string(result)
}
// ValidRandomString returns true if the given string matches the format created by RandomString
@@ -346,6 +359,16 @@ func MaybeMarshalJSON(v any) string {
return string(jsonBytes)
}
// EncodeJSON writes the JSON encoding of v to w, without escaping HTML-significant
// characters (<, >, &). Unlike the standard library's default, ntfy does not embed its
// JSON responses in HTML, so escaping these characters only makes the raw output harder
// to read (see #1511).
func EncodeJSON(w io.Writer, v any) error {
encoder := json.NewEncoder(w)
encoder.SetEscapeHTML(false)
return encoder.Encode(v)
}
// QuoteCommand combines a command array to a string, quoting arguments that need quoting.
// This function is naive, and sometimes wrong. It is only meant for lo pretty-printing a command.
//
@@ -438,3 +461,22 @@ func Int(v int) *int {
func Time(v time.Time) *time.Time {
return &v
}
// SanitizeUTF8 ensures a string is safe to store in PostgreSQL by handling two cases:
//
// 1. Invalid UTF-8 sequences: Some clients send Latin-1/ISO-8859-1 encoded text (e.g. accented
// characters like é, ñ, ß) in HTTP headers or SMTP messages. Go treats these as raw bytes in
// strings, but PostgreSQL rejects them. Any invalid UTF-8 byte is replaced with the Unicode
// replacement character (U+FFFD, "�") so the message is still delivered rather than lost.
//
// 2. NUL bytes (0x00): These are valid in UTF-8 but PostgreSQL TEXT columns reject them.
// They are stripped entirely.
func SanitizeUTF8(s string) string {
if !utf8.ValidString(s) {
s = strings.ToValidUTF8(s, "\xef\xbf\xbd") // U+FFFD
}
if strings.ContainsRune(s, 0) {
s = strings.ReplaceAll(s, "\x00", "")
}
return s
}
+32
View File
@@ -1,6 +1,7 @@
package util
import (
"bytes"
"errors"
"io"
"net/netip"
@@ -25,6 +26,30 @@ func TestRandomString(t *testing.T) {
require.NotEqual(t, s1, s2)
}
// TestRandomString_CSPRNG guards the crypto/rand-backed generator: every character must come
// from the expected charset (rejection sampling correctness) and a large batch must be unique
// (no clock-seeded PRNG collapsing to a predictable stream).
func TestRandomString_CSPRNG(t *testing.T) {
const charset = "abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ0123456789"
seen := make(map[string]bool)
charCounts := make(map[rune]int)
for i := 0; i < 5000; i++ {
s := RandomString(48)
require.Equal(t, 48, len(s))
require.False(t, seen[s], "duplicate random string generated")
seen[s] = true
for _, c := range s {
require.Contains(t, charset, string(c))
charCounts[c]++
}
}
// Every charset character should appear at least once across 5000*48 draws; a heavily
// biased or broken generator would leave gaps.
for _, c := range charset {
require.Greater(t, charCounts[c], 0, "character %q never appeared", string(c))
}
}
func TestFileExists(t *testing.T) {
filename := filepath.Join(t.TempDir(), "somefile.txt")
require.Nil(t, os.WriteFile(filename, []byte{0x25, 0x86}, 0600))
@@ -275,3 +300,10 @@ func TestMaybeMarshalJSON(t *testing.T) {
require.Equal(t, `"`+strings.Repeat("x", 4999), MaybeMarshalJSON(strings.Repeat("x", 6000)))
}
func TestEncodeJSON(t *testing.T) {
// HTML-significant characters (<, >, &) must NOT be escaped, see #1511
var buf bytes.Buffer
require.Nil(t, EncodeJSON(&buf, map[string]string{"message": "<b>a&b</b>"}))
require.Equal(t, `{"message":"<b>a&b</b>"}`+"\n", buf.String())
}