small-packages/mosdns/patches/206-add-auto-reload-rule-files-on-modification.patch
2026-05-10 09:48:30 +08:00

698 lines
18 KiB
Diff

From 9554c5b29aff4107504bdf45c35582e6d9743a93 Mon Sep 17 00:00:00 2001
From: sbwml <admin@cooluc.com>
Date: Mon, 20 Apr 2026 23:06:17 +0800
Subject: [PATCH] add auto-reload rule files on modification
Signed-off-by: sbwml <admin@cooluc.com>
---
pkg/utils/file_watcher.go | 73 +++++++++++++++
plugin/data_provider/domain_set/domain_set.go | 71 +++++++++++---
plugin/data_provider/ip_set/ip_set.go | 75 +++++++++++----
plugin/executable/arbitrary/arbitrary.go | 81 +++++++++++-----
plugin/executable/hosts/hosts.go | 87 ++++++++++++-----
plugin/executable/redirect/redirect.go | 93 +++++++++++++------
6 files changed, 374 insertions(+), 106 deletions(-)
create mode 100644 pkg/utils/file_watcher.go
--- /dev/null
+++ b/pkg/utils/file_watcher.go
@@ -0,0 +1,73 @@
+package utils
+
+import (
+ "os"
+ "time"
+)
+
+// FileWatcher periodically checks the modification time of the specified files.
+type FileWatcher struct {
+ done chan struct{}
+}
+
+// StartFileWatcher starts a new FileWatcher that checks files every 'interval'.
+// If any file modification time changes or a file is missing, onChange is triggered.
+// Note: onChange must be safe to run concurrently with other components.
+func StartFileWatcher(files []string, interval time.Duration, onChange func(changedFiles []string)) *FileWatcher {
+ if len(files) == 0 {
+ return nil
+ }
+ fw := &FileWatcher{
+ done: make(chan struct{}),
+ }
+ go func() {
+ modTimes := make(map[string]time.Time)
+ for _, f := range files {
+ stat, err := os.Stat(f)
+ if err == nil {
+ modTimes[f] = stat.ModTime()
+ }
+ }
+ ticker := time.NewTicker(interval)
+ defer ticker.Stop()
+
+ for {
+ select {
+ case <-fw.done:
+ return
+ case <-ticker.C:
+ var changedFiles []string
+ for _, f := range files {
+ stat, err := os.Stat(f)
+ if err != nil {
+ if os.IsNotExist(err) {
+ // File is missing or deleted.
+ if mt, ok := modTimes[f]; ok && !mt.IsZero() {
+ modTimes[f] = time.Time{} // mark as missing
+ changedFiles = append(changedFiles, f)
+ }
+ }
+ continue
+ }
+
+ if mt, ok := modTimes[f]; !ok || mt != stat.ModTime() {
+ modTimes[f] = stat.ModTime()
+ changedFiles = append(changedFiles, f)
+ }
+ }
+ if len(changedFiles) > 0 {
+ onChange(changedFiles)
+ }
+ }
+ }
+ }()
+ return fw
+}
+
+// Close stops the FileWatcher.
+func (fw *FileWatcher) Close() error {
+ if fw != nil {
+ close(fw.done)
+ }
+ return nil
+}
--- a/plugin/data_provider/domain_set/domain_set.go
+++ b/plugin/data_provider/domain_set/domain_set.go
@@ -22,10 +22,15 @@ package domain_set
import (
"bytes"
"fmt"
+ "os"
+ "sync/atomic"
+ "time"
+
"github.com/IrineSistiana/mosdns/v5/coremain"
"github.com/IrineSistiana/mosdns/v5/pkg/matcher/domain"
+ "github.com/IrineSistiana/mosdns/v5/pkg/utils"
"github.com/IrineSistiana/mosdns/v5/plugin/data_provider"
- "os"
+ "go.uber.org/zap"
)
const PluginType = "domain_set"
@@ -51,33 +56,69 @@ type Args struct {
var _ data_provider.DomainMatcherProvider = (*DomainSet)(nil)
type DomainSet struct {
- mg []domain.Matcher[struct{}]
+ mg atomic.Pointer[MatcherGroup]
+ fw *utils.FileWatcher
}
func (d *DomainSet) GetDomainMatcher() domain.Matcher[struct{}] {
- return MatcherGroup(d.mg)
+ return d
+}
+
+func (d *DomainSet) Match(s string) (struct{}, bool) {
+ if mg := d.mg.Load(); mg != nil {
+ return mg.Match(s)
+ }
+ return struct{}{}, false
+}
+
+func (d *DomainSet) Close() error {
+ if d.fw != nil {
+ return d.fw.Close()
+ }
+ return nil
}
// NewDomainSet inits a DomainSet from given args.
func NewDomainSet(bp *coremain.BP, args *Args) (*DomainSet, error) {
ds := &DomainSet{}
- m := domain.NewDomainMixMatcher()
- if err := LoadExpsAndFiles(args.Exps, args.Files, m); err != nil {
- return nil, err
+ loadInner := func() (*MatcherGroup, error) {
+ var mg MatcherGroup
+ m := domain.NewDomainMixMatcher()
+ if err := LoadExpsAndFiles(args.Exps, args.Files, m); err != nil {
+ return nil, err
+ }
+ if m.Len() > 0 {
+ mg = append(mg, m)
+ }
+
+ for _, tag := range args.Sets {
+ provider, _ := bp.M().GetPlugin(tag).(data_provider.DomainMatcherProvider)
+ if provider == nil {
+ return nil, fmt.Errorf("%s is not a DomainMatcherProvider", tag)
+ }
+ m := provider.GetDomainMatcher()
+ mg = append(mg, m)
+ }
+ return &mg, nil
}
- if m.Len() > 0 {
- ds.mg = append(ds.mg, m)
+
+ mg, err := loadInner()
+ if err != nil {
+ return nil, err
}
+ ds.mg.Store(mg)
- for _, tag := range args.Sets {
- provider, _ := bp.M().GetPlugin(tag).(data_provider.DomainMatcherProvider)
- if provider == nil {
- return nil, fmt.Errorf("%s is not a DomainMatcherProvider", tag)
- }
- m := provider.GetDomainMatcher()
- ds.mg = append(ds.mg, m)
+ if len(args.Files) > 0 {
+ ds.fw = utils.StartFileWatcher(args.Files, time.Second*3, func(changedFiles []string) {
+ newMg, err := loadInner()
+ if err == nil {
+ ds.mg.Store(newMg)
+ bp.L().Info("reloaded files", zap.Strings("files", changedFiles))
+ }
+ })
}
+
return ds, nil
}
--- a/plugin/data_provider/ip_set/ip_set.go
+++ b/plugin/data_provider/ip_set/ip_set.go
@@ -22,12 +22,17 @@ package ip_set
import (
"bytes"
"fmt"
- "github.com/IrineSistiana/mosdns/v5/coremain"
- "github.com/IrineSistiana/mosdns/v5/pkg/matcher/netlist"
- "github.com/IrineSistiana/mosdns/v5/plugin/data_provider"
"net/netip"
"os"
"strings"
+ "sync/atomic"
+ "time"
+
+ "github.com/IrineSistiana/mosdns/v5/coremain"
+ "github.com/IrineSistiana/mosdns/v5/pkg/matcher/netlist"
+ "github.com/IrineSistiana/mosdns/v5/pkg/utils"
+ "github.com/IrineSistiana/mosdns/v5/plugin/data_provider"
+ "go.uber.org/zap"
)
const PluginType = "ip_set"
@@ -49,31 +54,67 @@ type Args struct {
var _ data_provider.IPMatcherProvider = (*IPSet)(nil)
type IPSet struct {
- mg []netlist.Matcher
+ mg atomic.Pointer[MatcherGroup]
+ fw *utils.FileWatcher
}
func (d *IPSet) GetIPMatcher() netlist.Matcher {
- return MatcherGroup(d.mg)
+ return d
+}
+
+func (d *IPSet) Match(addr netip.Addr) bool {
+ if mg := d.mg.Load(); mg != nil {
+ return mg.Match(addr)
+ }
+ return false
+}
+
+func (d *IPSet) Close() error {
+ if d.fw != nil {
+ return d.fw.Close()
+ }
+ return nil
}
func NewIPSet(bp *coremain.BP, args *Args) (*IPSet, error) {
p := &IPSet{}
- l := netlist.NewList()
- if err := LoadFromIPsAndFiles(args.IPs, args.Files, l); err != nil {
- return nil, err
+ loadInner := func() (*MatcherGroup, error) {
+ var mg MatcherGroup
+ l := netlist.NewList()
+ if err := LoadFromIPsAndFiles(args.IPs, args.Files, l); err != nil {
+ return nil, err
+ }
+ l.Sort()
+ if l.Len() > 0 {
+ mg = append(mg, l)
+ }
+ for _, tag := range args.Sets {
+ provider, _ := bp.M().GetPlugin(tag).(data_provider.IPMatcherProvider)
+ if provider == nil {
+ return nil, fmt.Errorf("%s is not an IPMatcherProvider", tag)
+ }
+ mg = append(mg, provider.GetIPMatcher())
+ }
+ return &mg, nil
}
- l.Sort()
- if l.Len() > 0 {
- p.mg = append(p.mg, l)
+
+ mg, err := loadInner()
+ if err != nil {
+ return nil, err
}
- for _, tag := range args.Sets {
- provider, _ := bp.M().GetPlugin(tag).(data_provider.IPMatcherProvider)
- if provider == nil {
- return nil, fmt.Errorf("%s is not an IPMatcherProvider", tag)
- }
- p.mg = append(p.mg, provider.GetIPMatcher())
+ p.mg.Store(mg)
+
+ if len(args.Files) > 0 {
+ p.fw = utils.StartFileWatcher(args.Files, time.Second*3, func(changedFiles []string) {
+ newMg, err := loadInner()
+ if err == nil {
+ p.mg.Store(newMg)
+ bp.L().Info("reloaded files", zap.Strings("files", changedFiles))
+ }
+ })
}
+
return p, nil
}
--- a/plugin/executable/arbitrary/arbitrary.go
+++ b/plugin/executable/arbitrary/arbitrary.go
@@ -23,12 +23,17 @@ import (
"bytes"
"context"
"fmt"
+ "os"
+ "strings"
+ "sync/atomic"
+ "time"
+
"github.com/IrineSistiana/mosdns/v5/coremain"
"github.com/IrineSistiana/mosdns/v5/pkg/query_context"
+ "github.com/IrineSistiana/mosdns/v5/pkg/utils"
"github.com/IrineSistiana/mosdns/v5/pkg/zone_file"
"github.com/IrineSistiana/mosdns/v5/plugin/executable/sequence"
- "os"
- "strings"
+ "go.uber.org/zap"
)
const PluginType = "arbitrary"
@@ -45,38 +50,68 @@ type Args struct {
var _ sequence.Executable = (*Arbitrary)(nil)
type Arbitrary struct {
- m *zone_file.Matcher
+ m atomic.Pointer[zone_file.Matcher]
+ fw *utils.FileWatcher
}
-func NewArbitrary(args *Args) (*Arbitrary, error) {
- m := new(zone_file.Matcher)
- for i, s := range args.Rules {
- if err := m.Load(strings.NewReader(s)); err != nil {
- return nil, fmt.Errorf("failed to load rr #%d [%s], %w", i, s, err)
- }
- }
- for i, file := range args.Files {
- b, err := os.ReadFile(file)
- if err != nil {
- return nil, fmt.Errorf("failed to read file #%d [%s], %w", i, file, err)
+func NewArbitrary(bp *coremain.BP, args *Args) (*Arbitrary, error) {
+ a := &Arbitrary{}
+
+ loadInner := func() (*zone_file.Matcher, error) {
+ m := new(zone_file.Matcher)
+ for i, s := range args.Rules {
+ if err := m.Load(strings.NewReader(s)); err != nil {
+ return nil, fmt.Errorf("failed to load rr #%d [%s], %w", i, s, err)
+ }
}
- if err := m.Load(bytes.NewReader(b)); err != nil {
- return nil, fmt.Errorf("failed to load rr file #%d [%s], %w", i, file, err)
+ for i, file := range args.Files {
+ b, err := os.ReadFile(file)
+ if err != nil {
+ return nil, fmt.Errorf("failed to read file #%d [%s], %w", i, file, err)
+ }
+ if err := m.Load(bytes.NewReader(b)); err != nil {
+ return nil, fmt.Errorf("failed to load rr file #%d [%s], %w", i, file, err)
+ }
}
+ return m, nil
+ }
+
+ inner, err := loadInner()
+ if err != nil {
+ return nil, err
+ }
+ a.m.Store(inner)
+
+ if len(args.Files) > 0 {
+ a.fw = utils.StartFileWatcher(args.Files, time.Second*3, func(changedFiles []string) {
+ newInner, err := loadInner()
+ if err == nil {
+ a.m.Store(newInner)
+ bp.L().Info("reloaded files", zap.Strings("files", changedFiles))
+ }
+ })
}
- return &Arbitrary{
- m: m,
- }, nil
+
+ return a, nil
}
func (a *Arbitrary) Exec(_ context.Context, qCtx *query_context.Context) error {
- if r := a.m.Reply(qCtx.Q()); r != nil {
- qCtx.SetResponse(r)
+ if inner := a.m.Load(); inner != nil {
+ if r := inner.Reply(qCtx.Q()); r != nil {
+ qCtx.SetResponse(r)
+ }
+ }
+ return nil
+}
+
+func (a *Arbitrary) Close() error {
+ if a.fw != nil {
+ return a.fw.Close()
}
return nil
}
-func Init(_ *coremain.BP, v any) (any, error) {
+func Init(bp *coremain.BP, v any) (any, error) {
args := v.(*Args)
- return NewArbitrary(args)
+ return NewArbitrary(bp, args)
}
--- a/plugin/executable/hosts/hosts.go
+++ b/plugin/executable/hosts/hosts.go
@@ -23,13 +23,18 @@ import (
"bytes"
"context"
"fmt"
+ "os"
+ "sync/atomic"
+ "time"
+
"github.com/IrineSistiana/mosdns/v5/coremain"
"github.com/IrineSistiana/mosdns/v5/pkg/hosts"
"github.com/IrineSistiana/mosdns/v5/pkg/matcher/domain"
"github.com/IrineSistiana/mosdns/v5/pkg/query_context"
+ "github.com/IrineSistiana/mosdns/v5/pkg/utils"
"github.com/IrineSistiana/mosdns/v5/plugin/executable/sequence"
"github.com/miekg/dns"
- "os"
+ "go.uber.org/zap"
)
const PluginType = "hosts"
@@ -46,44 +51,76 @@ type Args struct {
}
type Hosts struct {
- h *hosts.Hosts
+ h atomic.Pointer[hosts.Hosts]
+ fw *utils.FileWatcher
}
-func Init(_ *coremain.BP, args any) (any, error) {
- return NewHosts(args.(*Args))
+func Init(bp *coremain.BP, args any) (any, error) {
+ return NewHosts(bp, args.(*Args))
}
-func NewHosts(args *Args) (*Hosts, error) {
- m := domain.NewMixMatcher[*hosts.IPs]()
- m.SetDefaultMatcher(domain.MatcherFull)
- for i, entry := range args.Entries {
- if err := domain.Load[*hosts.IPs](m, entry, hosts.ParseIPs); err != nil {
- return nil, fmt.Errorf("failed to load entry #%d %s, %w", i, entry, err)
- }
- }
- for i, file := range args.Files {
- b, err := os.ReadFile(file)
- if err != nil {
- return nil, fmt.Errorf("failed to read file #%d %s, %w", i, file, err)
+func NewHosts(bp *coremain.BP, args *Args) (*Hosts, error) {
+ h := &Hosts{}
+
+ loadInner := func() (*hosts.Hosts, error) {
+ m := domain.NewMixMatcher[*hosts.IPs]()
+ m.SetDefaultMatcher(domain.MatcherFull)
+ for i, entry := range args.Entries {
+ if err := domain.Load[*hosts.IPs](m, entry, hosts.ParseIPs); err != nil {
+ return nil, fmt.Errorf("failed to load entry #%d %s, %w", i, entry, err)
+ }
}
- if err := domain.LoadFromTextReader[*hosts.IPs](m, bytes.NewReader(b), hosts.ParseIPs); err != nil {
- return nil, fmt.Errorf("failed to load file #%d %s, %w", i, file, err)
+ for i, file := range args.Files {
+ b, err := os.ReadFile(file)
+ if err != nil {
+ return nil, fmt.Errorf("failed to read file #%d %s, %w", i, file, err)
+ }
+ if err := domain.LoadFromTextReader[*hosts.IPs](m, bytes.NewReader(b), hosts.ParseIPs); err != nil {
+ return nil, fmt.Errorf("failed to load file #%d %s, %w", i, file, err)
+ }
}
+ return hosts.NewHosts(m), nil
+ }
+
+ inner, err := loadInner()
+ if err != nil {
+ return nil, err
}
+ h.h.Store(inner)
- return &Hosts{
- h: hosts.NewHosts(m),
- }, nil
+ if len(args.Files) > 0 {
+ h.fw = utils.StartFileWatcher(args.Files, time.Second*3, func(changedFiles []string) {
+ newInner, err := loadInner()
+ if err == nil {
+ h.h.Store(newInner)
+ bp.L().Info("reloaded files", zap.Strings("files", changedFiles))
+ }
+ })
+ }
+
+ return h, nil
}
func (h *Hosts) Response(q *dns.Msg) *dns.Msg {
- return h.h.LookupMsg(q)
+ if inner := h.h.Load(); inner != nil {
+ return inner.LookupMsg(q)
+ }
+ return nil
}
func (h *Hosts) Exec(_ context.Context, qCtx *query_context.Context) error {
- r := h.h.LookupMsg(qCtx.Q())
- if r != nil {
- qCtx.SetResponse(r)
+ if inner := h.h.Load(); inner != nil {
+ r := inner.LookupMsg(qCtx.Q())
+ if r != nil {
+ qCtx.SetResponse(r)
+ }
+ }
+ return nil
+}
+
+func (h *Hosts) Close() error {
+ if h.fw != nil {
+ return h.fw.Close()
}
return nil
}
--- a/plugin/executable/redirect/redirect.go
+++ b/plugin/executable/redirect/redirect.go
@@ -25,10 +25,13 @@ import (
"fmt"
"os"
"strings"
+ "sync/atomic"
+ "time"
"github.com/IrineSistiana/mosdns/v5/coremain"
"github.com/IrineSistiana/mosdns/v5/pkg/matcher/domain"
"github.com/IrineSistiana/mosdns/v5/pkg/query_context"
+ "github.com/IrineSistiana/mosdns/v5/pkg/utils"
"github.com/IrineSistiana/mosdns/v5/plugin/executable/sequence"
"github.com/miekg/dns"
"go.uber.org/zap"
@@ -48,11 +51,12 @@ type Args struct {
}
type Redirect struct {
- m *domain.MixMatcher[string]
+ m atomic.Pointer[domain.MixMatcher[string]]
+ fw *utils.FileWatcher
}
func Init(bp *coremain.BP, args any) (any, error) {
- r, err := NewRedirect(args.(*Args))
+ r, err := NewRedirect(bp, args.(*Args))
if err != nil {
return nil, err
}
@@ -60,7 +64,8 @@ func Init(bp *coremain.BP, args any) (an
return r, nil
}
-func NewRedirect(args *Args) (*Redirect, error) {
+func NewRedirect(bp *coremain.BP, args *Args) (*Redirect, error) {
+ r := &Redirect{}
parseFunc := func(s string) (p, v string, err error) {
f := strings.Fields(s)
if len(f) != 2 {
@@ -68,23 +73,44 @@ func NewRedirect(args *Args) (*Redirect,
}
return f[0], dns.Fqdn(f[1]), nil
}
- m := domain.NewMixMatcher[string]()
- m.SetDefaultMatcher(domain.MatcherFull)
- for i, rule := range args.Rules {
- if err := domain.Load[string](m, rule, parseFunc); err != nil {
- return nil, fmt.Errorf("failed to load rule #%d %s, %w", i, rule, err)
- }
- }
- for i, file := range args.Files {
- b, err := os.ReadFile(file)
- if err != nil {
- return nil, fmt.Errorf("failed to read file #%d %s, %w", i, file, err)
+
+ loadInner := func() (*domain.MixMatcher[string], error) {
+ m := domain.NewMixMatcher[string]()
+ m.SetDefaultMatcher(domain.MatcherFull)
+ for i, rule := range args.Rules {
+ if err := domain.Load[string](m, rule, parseFunc); err != nil {
+ return nil, fmt.Errorf("failed to load rule #%d %s, %w", i, rule, err)
+ }
}
- if err := domain.LoadFromTextReader[string](m, bytes.NewReader(b), parseFunc); err != nil {
- return nil, fmt.Errorf("failed to load file #%d %s, %w", i, file, err)
+ for i, file := range args.Files {
+ b, err := os.ReadFile(file)
+ if err != nil {
+ return nil, fmt.Errorf("failed to read file #%d %s, %w", i, file, err)
+ }
+ if err := domain.LoadFromTextReader[string](m, bytes.NewReader(b), parseFunc); err != nil {
+ return nil, fmt.Errorf("failed to load file #%d %s, %w", i, file, err)
+ }
}
+ return m, nil
+ }
+
+ inner, err := loadInner()
+ if err != nil {
+ return nil, err
+ }
+ r.m.Store(inner)
+
+ if len(args.Files) > 0 {
+ r.fw = utils.StartFileWatcher(args.Files, time.Second*3, func(changedFiles []string) {
+ newInner, err := loadInner()
+ if err == nil {
+ r.m.Store(newInner)
+ bp.L().Info("reloaded files", zap.Strings("files", changedFiles), zap.Int("length", newInner.Len()))
+ }
+ })
}
- return &Redirect{m: m}, nil
+
+ return r, nil
}
func (r *Redirect) Exec(ctx context.Context, qCtx *query_context.Context, next sequence.ChainWalker) error {
@@ -94,7 +120,12 @@ func (r *Redirect) Exec(ctx context.Cont
}
orgQName := q.Question[0].Name
- redirectTarget, ok := r.m.Match(orgQName)
+ inner := r.m.Load()
+ if inner == nil {
+ return next.ExecNext(ctx, qCtx)
+ }
+
+ redirectTarget, ok := inner.Match(orgQName)
if !ok {
return next.ExecNext(ctx, qCtx)
}
@@ -104,16 +135,16 @@ func (r *Redirect) Exec(ctx context.Cont
q.Question[0].Name = orgQName
}()
err := next.ExecNext(ctx, qCtx)
- if r := qCtx.R(); r != nil {
+ if rResp := qCtx.R(); rResp != nil {
// Restore original query name.
- for i := range r.Question {
- if r.Question[i].Name == redirectTarget {
- r.Question[i].Name = orgQName
+ for i := range rResp.Question {
+ if rResp.Question[i].Name == redirectTarget {
+ rResp.Question[i].Name = orgQName
}
}
// Insert a CNAME record.
- newAns := make([]dns.RR, 1, len(r.Answer)+1)
+ newAns := make([]dns.RR, 1, len(rResp.Answer)+1)
newAns[0] = &dns.CNAME{
Hdr: dns.RR_Header{
Name: orgQName,
@@ -123,12 +154,22 @@ func (r *Redirect) Exec(ctx context.Cont
},
Target: redirectTarget,
}
- newAns = append(newAns, r.Answer...)
- r.Answer = newAns
+ newAns = append(newAns, rResp.Answer...)
+ rResp.Answer = newAns
}
return err
}
func (r *Redirect) Len() int {
- return r.m.Len()
+ if inner := r.m.Load(); inner != nil {
+ return inner.Len()
+ }
+ return 0
+}
+
+func (r *Redirect) Close() error {
+ if r.fw != nil {
+ return r.fw.Close()
+ }
+ return nil
}