mirror of
https://github.com/caiwx86/small-packages.git
synced 2026-08-01 04:17:53 +08:00
698 lines
18 KiB
Diff
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
|
|
}
|