mirror of
https://github.com/multipleof4/ntfy.git
synced 2026-10-08 21:05:21 +00:00
WIP: Access cache
This commit is contained in:
@@ -7,6 +7,7 @@ import (
|
||||
"heckel.io/ntfy/v2/server"
|
||||
"heckel.io/ntfy/v2/test"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCLI_Access_Show(t *testing.T) {
|
||||
@@ -43,6 +44,12 @@ user * (role: anonymous, tier: none)
|
||||
`
|
||||
require.Equal(t, expected, stdout.String())
|
||||
|
||||
// The CLI commands above ran against a separate Manager instance (their own
|
||||
// process-equivalent), so the server's ACL cache hasn't seen the new grants
|
||||
// yet. Wait for the server's background reloader (interval set in
|
||||
// newTestServerWithAuth) to pick them up.
|
||||
time.Sleep(150 * time.Millisecond)
|
||||
|
||||
// See if access permissions match
|
||||
app, _, _, _ = newTestApp()
|
||||
require.Error(t, app.Run([]string{
|
||||
|
||||
@@ -378,6 +378,11 @@ func createUserManager(c *cli.Context) (*user.Manager, error) {
|
||||
ProvisionEnabled: false, // Hack: Do not re-provision users on manager initialization
|
||||
BcryptCost: user.DefaultUserPasswordBcryptCost,
|
||||
QueueWriterInterval: user.DefaultUserStatsQueueWriterInterval,
|
||||
// CLI Managers are short-lived; the background ACL cache poller would only
|
||||
// spam "database is closed" warnings after the subcommand returns. Mutations
|
||||
// still refresh the local cache synchronously; the running server (if any)
|
||||
// picks them up via its own poller.
|
||||
AccessCacheReloadInterval: -1,
|
||||
}
|
||||
if databaseURL != "" {
|
||||
host, dbErr := pg.Open(databaseURL)
|
||||
|
||||
@@ -9,6 +9,7 @@ import (
|
||||
"os"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestCLI_User_Add(t *testing.T) {
|
||||
@@ -128,6 +129,10 @@ func newTestServerWithAuth(t *testing.T) (s *server.Server, conf *server.Config,
|
||||
conf.File = configFile
|
||||
conf.AuthFile = filepath.Join(t.TempDir(), "user.db")
|
||||
conf.AuthDefault = user.PermissionDenyAll
|
||||
// Tight interval so cross-process writes from the `ntfy access`/`ntfy user`
|
||||
// CLI commands (which run via a separate Manager) propagate to the server's
|
||||
// ACL cache within tens of ms instead of the default 5s.
|
||||
conf.AuthAccessCacheReloadInterval = 25 * time.Millisecond
|
||||
s, port = test.StartServerWithConfig(t, conf)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -116,6 +116,7 @@ type Config struct {
|
||||
AuthTokens map[string][]*user.Token
|
||||
AuthBcryptCost int
|
||||
AuthStatsQueueWriterInterval time.Duration
|
||||
AuthAccessCacheReloadInterval time.Duration
|
||||
AttachmentCacheDir string
|
||||
AttachmentTotalSizeLimit int64
|
||||
AttachmentFileSizeLimit int64
|
||||
@@ -223,6 +224,7 @@ func NewConfig() *Config {
|
||||
AuthDefault: user.PermissionReadWrite,
|
||||
AuthBcryptCost: user.DefaultUserPasswordBcryptCost,
|
||||
AuthStatsQueueWriterInterval: user.DefaultUserStatsQueueWriterInterval,
|
||||
AuthAccessCacheReloadInterval: user.DefaultAccessCacheReloadInterval,
|
||||
AttachmentCacheDir: "",
|
||||
AttachmentTotalSizeLimit: DefaultAttachmentTotalSizeLimit,
|
||||
AttachmentFileSizeLimit: DefaultAttachmentFileSizeLimit,
|
||||
|
||||
@@ -257,6 +257,7 @@ func New(conf *Config) (*Server, error) {
|
||||
Tokens: conf.AuthTokens,
|
||||
BcryptCost: conf.AuthBcryptCost,
|
||||
QueueWriterInterval: conf.AuthStatsQueueWriterInterval,
|
||||
AccessCacheReloadInterval: conf.AuthAccessCacheReloadInterval,
|
||||
}
|
||||
if pool != nil {
|
||||
userManager, err = user.NewPostgresManager(pool, authConfig)
|
||||
|
||||
@@ -0,0 +1,180 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"regexp"
|
||||
"strings"
|
||||
"sync/atomic"
|
||||
|
||||
"heckel.io/ntfy/v2/db"
|
||||
)
|
||||
|
||||
// aclEntry mirrors one user_access row in the in-memory snapshot.
|
||||
//
|
||||
// topic is the raw stored value: it may contain \_ escapes (for literal underscores)
|
||||
// and % wildcards (translated from user-supplied *). For exact-match entries (no %)
|
||||
// matcher is nil and the entry is keyed by topic in aclSnapshot.exact. For wildcard
|
||||
// entries (with %) matcher is the pre-compiled regex equivalent of the LIKE pattern.
|
||||
type aclEntry struct {
|
||||
topic string
|
||||
read bool
|
||||
write bool
|
||||
matcher *regexp.Regexp
|
||||
}
|
||||
|
||||
// aclSnapshot is an immutable indexed form of the entire user_access table.
|
||||
//
|
||||
// exact[userName][escapedTopic] returns the matching entry in O(1) for the common
|
||||
// case where the requested topic appears verbatim in some rule. The key is the
|
||||
// stored form of the topic (i.e. with \_ escapes), so callers must pass topics
|
||||
// through escapeUnderscore before probing.
|
||||
//
|
||||
// wildcards[userName] is the linear scan list of %-bearing rules for that user.
|
||||
// Walked per request; trivially small in practice. Wildcards are NOT u_everyone-
|
||||
// only -- any user can create them.
|
||||
type aclSnapshot struct {
|
||||
exact map[string]map[string]aclEntry
|
||||
wildcards map[string][]aclEntry
|
||||
}
|
||||
|
||||
// aclCache holds the current snapshot behind an atomic pointer so that the hot
|
||||
// path (Lookup) is lock-free. reload builds a fresh snapshot off the request
|
||||
// path and atomically swaps the pointer; the old snapshot is GC'd once in-flight
|
||||
// Lookups release their references.
|
||||
//
|
||||
// A nil receiver behaves as if no snapshot were loaded -- Lookup returns
|
||||
// found=false, which the caller then resolves via DefaultAccess. This keeps
|
||||
// tests and edge cases (e.g. early-startup) safe.
|
||||
type aclCache struct {
|
||||
snap atomic.Pointer[aclSnapshot]
|
||||
}
|
||||
|
||||
func newAccessCache() *aclCache {
|
||||
return &aclCache{}
|
||||
}
|
||||
|
||||
// reload runs the bulk-load query against the primary and atomically swaps in
|
||||
// a fresh snapshot. The primary is used (not ReadOnly) so a reload immediately
|
||||
// after an ACL mutation sees the freshly-written rows without replica lag.
|
||||
func (c *aclCache) reload(d *db.DB, query string) error {
|
||||
rows, err := d.Query(query)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
defer rows.Close()
|
||||
snap := &aclSnapshot{
|
||||
exact: make(map[string]map[string]aclEntry),
|
||||
wildcards: make(map[string][]aclEntry),
|
||||
}
|
||||
for rows.Next() {
|
||||
var userName string
|
||||
var entry aclEntry
|
||||
if err := rows.Scan(&userName, &entry.topic, &entry.read, &entry.write); err != nil {
|
||||
return err
|
||||
}
|
||||
if strings.Contains(entry.topic, "%") {
|
||||
entry.matcher = compileLikeToRegex(entry.topic)
|
||||
snap.wildcards[userName] = append(snap.wildcards[userName], entry)
|
||||
} else {
|
||||
if snap.exact[userName] == nil {
|
||||
snap.exact[userName] = make(map[string]aclEntry)
|
||||
}
|
||||
snap.exact[userName][entry.topic] = entry
|
||||
}
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return err
|
||||
}
|
||||
c.snap.Store(snap)
|
||||
return nil
|
||||
}
|
||||
|
||||
// Lookup returns the effective (read, write, found) permission for the given
|
||||
// (username, topic), preserving the priority ordering of the original SQL query:
|
||||
// 1. specific user beats Everyone
|
||||
// 2. longer pattern beats shorter (more specific wins)
|
||||
// 3. write beats read at equal length (write is "stronger")
|
||||
func (c *aclCache) Lookup(usernameOrEveryone, topic string) (read, write, found bool) {
|
||||
if c == nil {
|
||||
return false, false, false
|
||||
}
|
||||
snap := c.snap.Load()
|
||||
if snap == nil {
|
||||
return false, false, false
|
||||
}
|
||||
// Pre-compute the escaped form once: exact-match keys in the snapshot are
|
||||
// stored as toSQLWildcard would emit them (literal _ -> \_), so the
|
||||
// incoming topic must be escaped the same way before map lookup.
|
||||
escaped := escapeUnderscore(topic)
|
||||
|
||||
// Specific user takes priority over Everyone. Skip the first lookup when
|
||||
// the request is already anonymous to avoid scanning the same map twice.
|
||||
if usernameOrEveryone != Everyone {
|
||||
if e, ok := pickBest(snap, usernameOrEveryone, topic, escaped); ok {
|
||||
return e.read, e.write, true
|
||||
}
|
||||
}
|
||||
if e, ok := pickBest(snap, Everyone, topic, escaped); ok {
|
||||
return e.read, e.write, true
|
||||
}
|
||||
return false, false, false
|
||||
}
|
||||
|
||||
// pickBest returns the highest-priority entry for a single user, combining the
|
||||
// exact-match O(1) probe with a linear scan over the (usually empty or tiny)
|
||||
// wildcard list. Priority within a user: longer pattern wins; write wins ties.
|
||||
func pickBest(snap *aclSnapshot, userName, topic, escaped string) (aclEntry, bool) {
|
||||
var best aclEntry
|
||||
var found bool
|
||||
if m, ok := snap.exact[userName]; ok {
|
||||
if e, ok := m[escaped]; ok {
|
||||
best, found = e, true
|
||||
}
|
||||
}
|
||||
for _, w := range snap.wildcards[userName] {
|
||||
if !w.matcher.MatchString(topic) {
|
||||
continue
|
||||
}
|
||||
if !found || better(w, best) {
|
||||
best, found = w, true
|
||||
}
|
||||
}
|
||||
return best, found
|
||||
}
|
||||
|
||||
// better implements the (length DESC, write DESC) tie-break used by the original
|
||||
// query's ORDER BY for entries owned by the same user.
|
||||
func better(a, b aclEntry) bool {
|
||||
if len(a.topic) != len(b.topic) {
|
||||
return len(a.topic) > len(b.topic)
|
||||
}
|
||||
if a.write != b.write {
|
||||
return a.write
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// compileLikeToRegex converts a stored ntfy LIKE pattern into an equivalent Go
|
||||
// regexp. In ntfy's stored form, % is the only wildcard (translated from *) and
|
||||
// \_ is a literal underscore; no other backslashes occur. Topics themselves are
|
||||
// restricted to [A-Za-z0-9_-] (see AllowedTopic), so neither % nor stray
|
||||
// backslashes appear in user-supplied input.
|
||||
func compileLikeToRegex(pattern string) *regexp.Regexp {
|
||||
var sb strings.Builder
|
||||
sb.WriteString("^")
|
||||
i := 0
|
||||
for i < len(pattern) {
|
||||
switch {
|
||||
case pattern[i] == '\\' && i+1 < len(pattern) && pattern[i+1] == '_':
|
||||
sb.WriteString(regexp.QuoteMeta("_"))
|
||||
i += 2
|
||||
case pattern[i] == '%':
|
||||
sb.WriteString(".*")
|
||||
i++
|
||||
default:
|
||||
sb.WriteString(regexp.QuoteMeta(string(pattern[i])))
|
||||
i++
|
||||
}
|
||||
}
|
||||
sb.WriteString("$")
|
||||
return regexp.MustCompile(sb.String())
|
||||
}
|
||||
@@ -0,0 +1,259 @@
|
||||
package user
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// Cache-only unit tests. Integration with the Manager (loading from the DB,
|
||||
// reload-after-mutation, end-to-end Authorize behavior) is covered by the
|
||||
// existing TestStoreAuthorizeTopicAccess* tests in manager_test.go via
|
||||
// forEachStoreBackend.
|
||||
|
||||
func TestCompileLikeToRegex_Exact(t *testing.T) {
|
||||
r := compileLikeToRegex("foo")
|
||||
require.True(t, r.MatchString("foo"))
|
||||
require.False(t, r.MatchString("foox"))
|
||||
require.False(t, r.MatchString("xfoo"))
|
||||
}
|
||||
|
||||
func TestCompileLikeToRegex_TrailingPercent(t *testing.T) {
|
||||
r := compileLikeToRegex("up%")
|
||||
require.True(t, r.MatchString("up"))
|
||||
require.True(t, r.MatchString("up123"))
|
||||
require.False(t, r.MatchString("xup"))
|
||||
}
|
||||
|
||||
func TestCompileLikeToRegex_LeadingAndEmbeddedPercent(t *testing.T) {
|
||||
r := compileLikeToRegex("%test%")
|
||||
require.True(t, r.MatchString("test"))
|
||||
require.True(t, r.MatchString("mytest"))
|
||||
require.True(t, r.MatchString("testxxx"))
|
||||
require.True(t, r.MatchString("xtestx"))
|
||||
require.False(t, r.MatchString("nope"))
|
||||
}
|
||||
|
||||
func TestCompileLikeToRegex_EscapedUnderscore(t *testing.T) {
|
||||
// "my\_topic" is the stored form of a literal "my_topic" -- the underscore
|
||||
// must match itself, NOT act as a SQL one-character wildcard.
|
||||
r := compileLikeToRegex(`my\_topic`)
|
||||
require.True(t, r.MatchString("my_topic"))
|
||||
require.False(t, r.MatchString("myXtopic"))
|
||||
require.False(t, r.MatchString("mytopic"))
|
||||
}
|
||||
|
||||
func TestCompileLikeToRegex_EscapedUnderscoreAdjacentToPercent(t *testing.T) {
|
||||
// "nz\_vip\_%" is the stored form of "nz_vip_*" -- literal "nz_vip_" prefix
|
||||
// followed by any suffix.
|
||||
r := compileLikeToRegex(`nz\_vip\_%`)
|
||||
require.True(t, r.MatchString("nz_vip_"))
|
||||
require.True(t, r.MatchString("nz_vip_alpha"))
|
||||
require.False(t, r.MatchString("nz_vipX"))
|
||||
require.False(t, r.MatchString("nzvip_alpha"))
|
||||
}
|
||||
|
||||
func TestCompileLikeToRegex_RegexMetaCharsInTopic(t *testing.T) {
|
||||
// Topics in ntfy can include '-', which is benign, but make sure
|
||||
// regex metacharacters in the pattern are escaped properly anyway.
|
||||
r := compileLikeToRegex("foo-bar")
|
||||
require.True(t, r.MatchString("foo-bar"))
|
||||
require.False(t, r.MatchString("foo.bar")) // would match if '-' leaked into a character class
|
||||
}
|
||||
|
||||
func TestACLCache_LookupOnNilReceiverSafe(t *testing.T) {
|
||||
var c *aclCache
|
||||
read, write, found := c.Lookup("phil", "mytopic")
|
||||
require.False(t, found)
|
||||
require.False(t, read)
|
||||
require.False(t, write)
|
||||
}
|
||||
|
||||
func TestACLCache_LookupBeforeReload(t *testing.T) {
|
||||
// Before reload the snapshot pointer is nil. The cache treats this as
|
||||
// "no rule found", which the caller resolves via DefaultAccess.
|
||||
c := newAccessCache()
|
||||
read, write, found := c.Lookup("phil", "mytopic")
|
||||
require.False(t, found)
|
||||
require.False(t, read)
|
||||
require.False(t, write)
|
||||
}
|
||||
|
||||
func TestACLCache_ExactMatchHit(t *testing.T) {
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: "phil", topic: "mytopic", read: true, write: true},
|
||||
}))
|
||||
read, write, found := c.Lookup("phil", "mytopic")
|
||||
require.True(t, found)
|
||||
require.True(t, read)
|
||||
require.True(t, write)
|
||||
}
|
||||
|
||||
func TestACLCache_ExactMatchMiss(t *testing.T) {
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: "phil", topic: "mytopic", read: true, write: true},
|
||||
}))
|
||||
_, _, found := c.Lookup("phil", "othertopic")
|
||||
require.False(t, found)
|
||||
}
|
||||
|
||||
func TestACLCache_LiteralUnderscoreExactMatch(t *testing.T) {
|
||||
// Stored as "my\_topic" (toSQLWildcard of "my_topic"). A literal underscore
|
||||
// in the requested topic must match, while any other single char must not.
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: "phil", topic: `my\_topic`, read: true, write: false},
|
||||
}))
|
||||
read, write, found := c.Lookup("phil", "my_topic")
|
||||
require.True(t, found)
|
||||
require.True(t, read)
|
||||
require.False(t, write)
|
||||
|
||||
_, _, found = c.Lookup("phil", "myXtopic")
|
||||
require.False(t, found)
|
||||
}
|
||||
|
||||
func TestACLCache_WildcardMatch(t *testing.T) {
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: Everyone, topic: "up%", read: false, write: true},
|
||||
}))
|
||||
read, write, found := c.Lookup("phil", "up42")
|
||||
require.True(t, found)
|
||||
require.False(t, read)
|
||||
require.True(t, write)
|
||||
}
|
||||
|
||||
func TestACLCache_SpecificUserBeatsEveryone(t *testing.T) {
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: Everyone, topic: "mytopic", read: true, write: false},
|
||||
{user: "phil", topic: "mytopic", read: false, write: false}, // deny-all for phil
|
||||
}))
|
||||
read, write, found := c.Lookup("phil", "mytopic")
|
||||
require.True(t, found)
|
||||
require.False(t, read)
|
||||
require.False(t, write)
|
||||
}
|
||||
|
||||
func TestACLCache_AnonymousReadsEveryone(t *testing.T) {
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: Everyone, topic: "announcements", read: true, write: false},
|
||||
}))
|
||||
read, write, found := c.Lookup(Everyone, "announcements")
|
||||
require.True(t, found)
|
||||
require.True(t, read)
|
||||
require.False(t, write)
|
||||
}
|
||||
|
||||
func TestACLCache_LongerPatternWinsForSameUser(t *testing.T) {
|
||||
// Both rules belong to the same user (Everyone). The more specific (longer)
|
||||
// "mytopic%" should beat the catch-all "%".
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: Everyone, topic: "%", read: true, write: false},
|
||||
{user: Everyone, topic: "mytopic%", read: true, write: true},
|
||||
}))
|
||||
read, write, found := c.Lookup(Everyone, "mytopicX")
|
||||
require.True(t, found)
|
||||
require.True(t, read)
|
||||
require.True(t, write)
|
||||
}
|
||||
|
||||
func TestACLCache_WriteBeatsReadAtEqualLength(t *testing.T) {
|
||||
// Two wildcard rules of identical length for the same user. The write rule
|
||||
// should win the tie-break.
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: Everyone, topic: "ab%", read: true, write: false},
|
||||
{user: Everyone, topic: "ab%", read: false, write: true}, // synthesized; impossible via real upsert but exercises the tie-break
|
||||
}))
|
||||
// One of the two will be the surviving exact-key entry (map collision keeps last);
|
||||
// but the wildcard slice is what we want to exercise. Inject two wildcard entries
|
||||
// directly to force the tie-break path.
|
||||
c.snap.Store(&aclSnapshot{
|
||||
exact: map[string]map[string]aclEntry{},
|
||||
wildcards: map[string][]aclEntry{
|
||||
Everyone: {
|
||||
{topic: "ab%", read: true, write: false, matcher: compileLikeToRegex("ab%")},
|
||||
{topic: "ab%", read: false, write: true, matcher: compileLikeToRegex("ab%")},
|
||||
},
|
||||
},
|
||||
})
|
||||
_, write, found := c.Lookup(Everyone, "abc")
|
||||
require.True(t, found)
|
||||
require.True(t, write)
|
||||
}
|
||||
|
||||
func TestACLCache_ConcurrentLookupAndReload(t *testing.T) {
|
||||
// Atomic-pointer swap must be safe under concurrent reads. The race detector
|
||||
// catches any unsafe shared mutation.
|
||||
c := newAccessCache()
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: Everyone, topic: "mytopic", read: true, write: true},
|
||||
}))
|
||||
|
||||
var stop atomic.Bool
|
||||
var wg sync.WaitGroup
|
||||
wg.Add(2)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for !stop.Load() {
|
||||
_, _, _ = c.Lookup(Everyone, "mytopic")
|
||||
}
|
||||
}()
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
for i := 0; i < 100; i++ {
|
||||
c.snap.Store(buildSnapshot(t, []rawACLRow{
|
||||
{user: Everyone, topic: "mytopic", read: i%2 == 0, write: i%2 == 1},
|
||||
}))
|
||||
}
|
||||
stop.Store(true)
|
||||
}()
|
||||
wg.Wait()
|
||||
}
|
||||
|
||||
// rawACLRow + buildSnapshot mirror the rows that reload would Scan from the DB
|
||||
// but avoid actually opening a DB for these unit tests.
|
||||
type rawACLRow struct {
|
||||
user string
|
||||
topic string
|
||||
read bool
|
||||
write bool
|
||||
}
|
||||
|
||||
func buildSnapshot(t *testing.T, rows []rawACLRow) *aclSnapshot {
|
||||
t.Helper()
|
||||
snap := &aclSnapshot{
|
||||
exact: make(map[string]map[string]aclEntry),
|
||||
wildcards: make(map[string][]aclEntry),
|
||||
}
|
||||
for _, r := range rows {
|
||||
e := aclEntry{topic: r.topic, read: r.read, write: r.write}
|
||||
if containsPercent(r.topic) {
|
||||
e.matcher = compileLikeToRegex(r.topic)
|
||||
snap.wildcards[r.user] = append(snap.wildcards[r.user], e)
|
||||
} else {
|
||||
if snap.exact[r.user] == nil {
|
||||
snap.exact[r.user] = make(map[string]aclEntry)
|
||||
}
|
||||
snap.exact[r.user][r.topic] = e
|
||||
}
|
||||
}
|
||||
return snap
|
||||
}
|
||||
|
||||
func containsPercent(s string) bool {
|
||||
for i := 0; i < len(s); i++ {
|
||||
if s[i] == '%' {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
+104
-32
@@ -38,6 +38,11 @@ const (
|
||||
const (
|
||||
DefaultUserStatsQueueWriterInterval = 33 * time.Second
|
||||
DefaultUserPasswordBcryptCost = 10
|
||||
// DefaultAccessCacheReloadInterval bounds how stale the in-memory ACL snapshot
|
||||
// can be relative to writes made by *other* processes (e.g. a separate `ntfy
|
||||
// access` CLI invocation modifying the same database). Mutations performed
|
||||
// by this Manager refresh the cache synchronously and do not depend on this.
|
||||
DefaultAccessCacheReloadInterval = 5 * time.Second
|
||||
)
|
||||
|
||||
var (
|
||||
@@ -53,6 +58,8 @@ type Manager struct {
|
||||
queries queries
|
||||
statsQueue map[string]*Stats // "Queue" to asynchronously write user stats to the database (UserID -> Stats)
|
||||
tokenQueue map[string]*TokenUpdate // "Queue" to asynchronously write token access stats to the database (Token ID -> TokenUpdate)
|
||||
accessCache *aclCache // In-memory snapshot of user_access; rebuilt after every ACL mutation
|
||||
quit chan struct{} // Closed by Close() to signal background goroutines to stop
|
||||
mu sync.Mutex
|
||||
}
|
||||
|
||||
@@ -65,20 +72,60 @@ func newManager(d *db.DB, queries queries, config *Config) (*Manager, error) {
|
||||
if config.QueueWriterInterval.Seconds() <= 0 {
|
||||
config.QueueWriterInterval = DefaultUserStatsQueueWriterInterval
|
||||
}
|
||||
if config.AccessCacheReloadInterval == 0 {
|
||||
config.AccessCacheReloadInterval = DefaultAccessCacheReloadInterval
|
||||
}
|
||||
manager := &Manager{
|
||||
config: config,
|
||||
db: d,
|
||||
statsQueue: make(map[string]*Stats),
|
||||
tokenQueue: make(map[string]*TokenUpdate),
|
||||
accessCache: newAccessCache(),
|
||||
quit: make(chan struct{}),
|
||||
queries: queries,
|
||||
}
|
||||
if err := manager.maybeProvisionUsersAccessAndTokens(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
// Populate the ACL cache after provisioning so the initial snapshot includes
|
||||
// any provisioned access rules. Subsequent mutations call reloadAccessCache.
|
||||
if err := manager.reloadAccessCache(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
go manager.asyncQueueWriter(manager.config.QueueWriterInterval)
|
||||
if manager.config.AccessCacheReloadInterval > 0 {
|
||||
go manager.asyncAccessCacheReloader(manager.config.AccessCacheReloadInterval)
|
||||
}
|
||||
return manager, nil
|
||||
}
|
||||
|
||||
// reloadAccessCache rebuilds the in-memory ACL snapshot from the primary
|
||||
// database. Called once at startup (after provisioning) and after every method
|
||||
// that mutates user_access (directly or by cascade from "user" deletion).
|
||||
func (a *Manager) reloadAccessCache() error {
|
||||
return a.accessCache.reload(a.db, a.queries.selectAllAccessForCache)
|
||||
}
|
||||
|
||||
// asyncAccessCacheReloader periodically refreshes the ACL snapshot so that
|
||||
// writes made by other processes against the same database (most notably the
|
||||
// `ntfy access` CLI subcommand running while a server holds the cache) become
|
||||
// visible within the configured interval. This Manager's own writes do not
|
||||
// depend on the poller -- they refresh the cache synchronously.
|
||||
func (a *Manager) asyncAccessCacheReloader(interval time.Duration) {
|
||||
ticker := time.NewTicker(interval)
|
||||
defer ticker.Stop()
|
||||
for {
|
||||
select {
|
||||
case <-a.quit:
|
||||
return
|
||||
case <-ticker.C:
|
||||
if err := a.reloadAccessCache(); err != nil {
|
||||
log.Tag(tag).Err(err).Warn("Reloading ACL cache failed")
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Authenticate checks username and password and returns a User if correct, and the user has not been
|
||||
// marked as deleted. The method returns in constant-ish time, regardless of whether the user exists or
|
||||
// the password is correct or incorrect.
|
||||
@@ -151,9 +198,13 @@ func (a *Manager) RemoveUser(username string) error {
|
||||
if err := a.CanChangeUser(username); err != nil {
|
||||
return err
|
||||
}
|
||||
return db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if err := db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
return a.removeUserTx(tx, username)
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
// user_access rows are cascade-deleted along with the user; refresh the snapshot.
|
||||
return a.reloadAccessCache()
|
||||
}
|
||||
|
||||
// removeUserTx deletes the user with the given username
|
||||
@@ -174,7 +225,7 @@ func (a *Manager) MarkUserRemoved(user *User) error {
|
||||
if !AllowedUsername(user.Name) {
|
||||
return ErrInvalidArgument
|
||||
}
|
||||
return db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if err := db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if err := a.resetUserAccessTx(tx, user.Name); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -186,7 +237,11 @@ func (a *Manager) MarkUserRemoved(user *User) error {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
// resetUserAccessTx wiped this user's user_access rows; refresh the snapshot.
|
||||
return a.reloadAccessCache()
|
||||
}
|
||||
|
||||
// RemoveDeletedUsers deletes all users that have been marked deleted
|
||||
@@ -194,7 +249,8 @@ func (a *Manager) RemoveDeletedUsers() error {
|
||||
if _, err := a.db.Exec(a.queries.deleteUsersMarked, time.Now().Unix()); err != nil {
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
// user_access rows are cascade-deleted with the users; refresh the snapshot.
|
||||
return a.reloadAccessCache()
|
||||
}
|
||||
|
||||
// ChangePassword changes a user's password
|
||||
@@ -225,9 +281,15 @@ func (a *Manager) ChangeRole(username string, role Role) error {
|
||||
if err := a.CanChangeUser(username); err != nil {
|
||||
return err
|
||||
}
|
||||
return db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if err := db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
return a.changeRoleTx(tx, username, role)
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
// Promotion to admin clears user_access rows for the user; refresh the snapshot.
|
||||
// Other role changes are no-ops for the cache but reloading is cheap and keeps
|
||||
// the code path uniform.
|
||||
return a.reloadAccessCache()
|
||||
}
|
||||
|
||||
// changeRoleTx changes a user's role
|
||||
@@ -588,9 +650,12 @@ func (a *Manager) resolvePerms(base, perm Permission) error {
|
||||
// read/write access to a topic. The parameter topicPattern may include wildcards (*). The ACL entry
|
||||
// owner may either be a user (username), or the system (empty).
|
||||
func (a *Manager) AllowAccess(username string, topicPattern string, permission Permission) error {
|
||||
return db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if err := db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
return a.allowAccessTx(tx, username, topicPattern, permission, false)
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return a.reloadAccessCache()
|
||||
}
|
||||
|
||||
func (a *Manager) allowAccessTx(tx *sql.Tx, username string, topicPattern string, permission Permission, provisioned bool) error {
|
||||
@@ -606,9 +671,12 @@ func (a *Manager) allowAccessTx(tx *sql.Tx, username string, topicPattern string
|
||||
// ResetAccess removes an access control list entry for a specific username/topic, or (if topic is
|
||||
// empty) for an entire user. The parameter topicPattern may include wildcards (*).
|
||||
func (a *Manager) ResetAccess(username string, topicPattern string) error {
|
||||
return db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if err := db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
return a.resetAccessTx(tx, username, topicPattern)
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return a.reloadAccessCache()
|
||||
}
|
||||
|
||||
func (a *Manager) resetAccessTx(tx *sql.Tx, username string, topicPattern string) error {
|
||||
@@ -650,24 +718,15 @@ func (a *Manager) AllowReservation(username string, topic string) error {
|
||||
// authorizeTopicAccess returns the read/write permissions for the given username and topic.
|
||||
// The found return value indicates whether an ACL entry was found at all.
|
||||
//
|
||||
// - The query may return two rows (one for everyone, and one for the user), but prioritizes the user.
|
||||
// - Furthermore, the query prioritizes more specific permissions (longer!) over more generic ones, e.g. "test*" > "*"
|
||||
// - The cache may contain two matching entries (one for everyone, and one for the user), but prioritizes the user.
|
||||
// - Furthermore, the lookup prioritizes more specific permissions (longer!) over more generic ones, e.g. "test*" > "*"
|
||||
// - It also prioritizes write permissions over read permissions
|
||||
//
|
||||
// The lookup is served entirely from the in-memory snapshot maintained by accessCache,
|
||||
// so this is on the hot path of every authenticatable HTTP request and must stay allocation-free.
|
||||
func (a *Manager) authorizeTopicAccess(usernameOrEveryone, topic string) (read, write, found bool, err error) {
|
||||
rows, err := a.db.ReadOnly().Query(a.queries.selectTopicPerms, Everyone, usernameOrEveryone, topic)
|
||||
if err != nil {
|
||||
return false, false, false, err
|
||||
}
|
||||
defer rows.Close()
|
||||
if !rows.Next() {
|
||||
return false, false, false, nil
|
||||
}
|
||||
if err := rows.Scan(&read, &write); err != nil {
|
||||
return false, false, false, err
|
||||
} else if err := rows.Err(); err != nil {
|
||||
return false, false, false, err
|
||||
}
|
||||
return read, write, true, nil
|
||||
read, write, found = a.accessCache.Lookup(usernameOrEveryone, topic)
|
||||
return read, write, found, nil
|
||||
}
|
||||
|
||||
// AllGrants returns all user-specific access control entries, mapped to their respective user IDs
|
||||
@@ -731,7 +790,7 @@ func (a *Manager) AddReservation(username string, topic string, everyone Permiss
|
||||
if !AllowedUsername(username) || username == Everyone || !AllowedTopic(topic) {
|
||||
return ErrInvalidArgument
|
||||
}
|
||||
return db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if err := db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if limit > 0 {
|
||||
hasReservation, err := a.hasReservationTx(tx, username, topic)
|
||||
if err != nil {
|
||||
@@ -754,7 +813,10 @@ func (a *Manager) AddReservation(username string, topic string, everyone Permiss
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return a.reloadAccessCache()
|
||||
}
|
||||
|
||||
// RemoveReservations deletes the access control entries associated with the given username/topic,
|
||||
@@ -769,14 +831,17 @@ func (a *Manager) RemoveReservations(username string, topics ...string) error {
|
||||
return ErrInvalidArgument
|
||||
}
|
||||
}
|
||||
return db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
if err := db.ExecTx(a.db, func(tx *sql.Tx) error {
|
||||
for _, topic := range topics {
|
||||
if err := a.removeReservationAccessTx(tx, username, topic); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}); err != nil {
|
||||
return err
|
||||
}
|
||||
return a.reloadAccessCache()
|
||||
}
|
||||
|
||||
// Reservations returns all user-owned topics, and the associated everyone-access
|
||||
@@ -1515,8 +1580,15 @@ func (a *Manager) maybeProvisionTokens(tx *sql.Tx, provisionUsernames []string,
|
||||
return nil
|
||||
}
|
||||
|
||||
// Close closes the underlying database
|
||||
// Close stops background goroutines and closes the underlying database.
|
||||
// Safe to call multiple times.
|
||||
func (a *Manager) Close() error {
|
||||
select {
|
||||
case <-a.quit:
|
||||
// already closed
|
||||
default:
|
||||
close(a.quit)
|
||||
}
|
||||
return a.db.Close()
|
||||
}
|
||||
|
||||
|
||||
@@ -70,12 +70,10 @@ const (
|
||||
postgresDeleteUsersProvisionedQuery = `DELETE FROM "user" WHERE provisioned = true`
|
||||
|
||||
// Access queries
|
||||
postgresSelectTopicPermsQuery = `
|
||||
SELECT read, write
|
||||
postgresSelectAllAccessForCacheQuery = `
|
||||
SELECT u.user_name, a.topic, a.read, a.write
|
||||
FROM user_access a
|
||||
JOIN "user" u ON u.id = a.user_id
|
||||
WHERE (u.user_name = $1 OR u.user_name = $2) AND $3 LIKE a.topic ESCAPE '\'
|
||||
ORDER BY u.user_name DESC, LENGTH(a.topic) DESC, CASE WHEN a.write THEN 1 ELSE 0 END DESC
|
||||
`
|
||||
postgresSelectUserAllAccessQuery = `
|
||||
SELECT user_id, topic, read, write, provisioned
|
||||
@@ -244,7 +242,7 @@ var postgresQueries = queries{
|
||||
deleteUserTier: postgresDeleteUserTierQuery,
|
||||
deleteUsersMarked: postgresDeleteUsersMarkedQuery,
|
||||
deleteUsersProvisioned: postgresDeleteUsersProvisionedQuery,
|
||||
selectTopicPerms: postgresSelectTopicPermsQuery,
|
||||
selectAllAccessForCache: postgresSelectAllAccessForCacheQuery,
|
||||
selectUserAllAccess: postgresSelectUserAllAccessQuery,
|
||||
selectUserAccess: postgresSelectUserAccessQuery,
|
||||
selectUserReservations: postgresSelectUserReservationsQuery,
|
||||
|
||||
@@ -76,12 +76,10 @@ const (
|
||||
sqliteDeleteUsersProvisionedQuery = `DELETE FROM user WHERE provisioned = 1`
|
||||
|
||||
// Access queries
|
||||
sqliteSelectTopicPermsQuery = `
|
||||
SELECT read, write
|
||||
sqliteSelectAllAccessForCacheQuery = `
|
||||
SELECT u.user, a.topic, a.read, a.write
|
||||
FROM user_access a
|
||||
JOIN user u ON u.id = a.user_id
|
||||
WHERE (u.user = ? OR u.user = ?) AND ? LIKE a.topic ESCAPE '\'
|
||||
ORDER BY u.user DESC, LENGTH(a.topic) DESC, a.write DESC
|
||||
`
|
||||
sqliteSelectUserAllAccessQuery = `
|
||||
SELECT user_id, topic, read, write, provisioned
|
||||
@@ -242,7 +240,7 @@ var sqliteQueries = queries{
|
||||
deleteUserTier: sqliteDeleteUserTierQuery,
|
||||
deleteUsersMarked: sqliteDeleteUsersMarkedQuery,
|
||||
deleteUsersProvisioned: sqliteDeleteUsersProvisionedQuery,
|
||||
selectTopicPerms: sqliteSelectTopicPermsQuery,
|
||||
selectAllAccessForCache: sqliteSelectAllAccessForCacheQuery,
|
||||
selectUserAllAccess: sqliteSelectUserAllAccessQuery,
|
||||
selectUserAccess: sqliteSelectUserAccessQuery,
|
||||
selectUserReservations: sqliteSelectUserReservationsQuery,
|
||||
|
||||
+11
-1
@@ -255,6 +255,16 @@ type Config struct {
|
||||
Tokens map[string][]*Token // Predefined users to create on startup (username -> []*Token)
|
||||
QueueWriterInterval time.Duration // Interval for the async queue writer to flush stats and token updates to the database
|
||||
BcryptCost int // Cost of generated passwords; lowering makes testing faster
|
||||
|
||||
// AccessCacheReloadInterval bounds the staleness of the in-memory ACL cache
|
||||
// relative to writes from other processes (e.g. `ntfy access` CLI against a
|
||||
// running server).
|
||||
// 0 -> use DefaultAccessCacheReloadInterval
|
||||
// negative -> disable the background poller; cache only refreshes on this
|
||||
// Manager's own ACL mutations. Use this for short-lived Managers
|
||||
// (e.g. the CLI subcommands) where polling is wasted work.
|
||||
// positive -> poll at the given interval
|
||||
AccessCacheReloadInterval time.Duration
|
||||
}
|
||||
|
||||
// Error constants used by the package
|
||||
@@ -303,7 +313,7 @@ type queries struct {
|
||||
deleteUsersProvisioned string
|
||||
|
||||
// Access queries
|
||||
selectTopicPerms string
|
||||
selectAllAccessForCache string // Bulk load: (user_name, topic, read, write) for the in-memory ACL cache
|
||||
selectUserAllAccess string
|
||||
selectUserAccess string
|
||||
selectUserReservations string
|
||||
|
||||
Reference in New Issue
Block a user