mirror of
https://github.com/multipleof4/ntfy.git
synced 2026-10-08 21:05:21 +00:00
Merge branch 'main' of github.com:binwiederhier/ntfy into 1511-ampersand
This commit is contained in:
@@ -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
@@ -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())
|
||||
}
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
@@ -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())
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user