From 4ab27f0758bd24fb4fd970970ff8fc462e5b6fc8 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Andri=20=C3=93skarsson?= Date: Wed, 25 Feb 2026 20:14:18 +0100 Subject: [PATCH] Add URL pattern matching with custom segment-based matcher Replace Blocklist with RuleSet that combines hostname fast-path (map lookup) with compiled URL pattern rules supporting wildcards (*), separator characters (^), address anchors (|), and domain anchors (||). URL patterns are evaluated after MITM for HTTPS requests, enabling path-specific blocking inside TLS tunnels. Blocked requests return 204. --- connect.go | 18 ++- http.go | 2 +- main.go | 10 +- pkg/blocklist/blocklist.go | 152 +----------------- pkg/blocklist/blocklist_test.go | 167 -------------------- pkg/blocklist/pattern.go | 270 ++++++++++++++++++++++++++++++++ pkg/blocklist/pattern_test.go | 95 +++++++++++ pkg/blocklist/ruleset.go | 212 +++++++++++++++++++++++++ pkg/blocklist/ruleset_test.go | 180 +++++++++++++++++++++ proxy.go | 6 +- proxy_test.go | 111 +++++++++++-- 11 files changed, 888 insertions(+), 335 deletions(-) delete mode 100644 pkg/blocklist/blocklist_test.go create mode 100644 pkg/blocklist/pattern.go create mode 100644 pkg/blocklist/pattern_test.go create mode 100644 pkg/blocklist/ruleset.go create mode 100644 pkg/blocklist/ruleset_test.go diff --git a/connect.go b/connect.go index 1445037..4288e40 100644 --- a/connect.go +++ b/connect.go @@ -16,7 +16,7 @@ func (p *proxyHandler) handleConnect(w http.ResponseWriter, r *http.Request) { host = r.Host port = "443" } - if p.blocklist.IsBlocked(host) { + if p.rules.IsHostBlocked(host) { w.WriteHeader(http.StatusNoContent) return } @@ -71,8 +71,22 @@ func (p *proxyHandler) proxyTLSRequests(clientTLS *tls.Conn, host, port string) return } - start := time.Now() targetURL := "https://" + host + req.URL.String() + + // URL-level blocking for pattern rules (hostname was already + // checked at CONNECT time; this catches path-specific rules) + if p.rules.ShouldBlock(targetURL) { + blocked := &http.Response{ + StatusCode: http.StatusNoContent, + ProtoMajor: 1, + ProtoMinor: 1, + Header: make(http.Header), + } + blocked.Write(clientTLS) + continue + } + + start := time.Now() req.URL.Scheme = "https" req.URL.Host = net.JoinHostPort(host, port) req.RequestURI = "" diff --git a/http.go b/http.go index e73b9a4..1fa2933 100644 --- a/http.go +++ b/http.go @@ -20,7 +20,7 @@ var hopByHopHeaders = []string{ } func (p *proxyHandler) handleHTTP(w http.ResponseWriter, r *http.Request) { - if p.blocklist.IsBlocked(r.URL.Hostname()) { + if p.rules.ShouldBlock(r.URL.String()) { w.WriteHeader(http.StatusNoContent) return } diff --git a/main.go b/main.go index 7525851..bc957fc 100644 --- a/main.go +++ b/main.go @@ -39,21 +39,21 @@ func main() { os.Exit(1) } - var bl *blocklist.Blocklist + var rules *blocklist.RuleSet if len(blocklistPaths) > 0 { - bl = blocklist.New() + rules = blocklist.NewRuleSet() for _, path := range blocklistPaths { - if err := bl.LoadFile(path); err != nil { + if err := rules.LoadFile(path); err != nil { fmt.Fprintf(os.Stderr, "failed to load blocklist: %v\n", err) os.Exit(1) } } - fmt.Fprintf(os.Stderr, "Loaded %d blocked hostnames\n", bl.Len()) + fmt.Fprintf(os.Stderr, "Loaded %d blocked hostnames, %d URL rules\n", rules.HostCount(), rules.RuleCount()) } certs := newCertCache(caCert, caKey) caCertPEM := encodeCertPEM(caCert) - handler := newProxyHandler(certs, caCertPEM, bl) + handler := newProxyHandler(certs, caCertPEM, rules) listenAddr := fmt.Sprintf("%s:%d", *addr, *port) fmt.Fprintf(os.Stderr, "ublproxy listening on http://%s\n", listenAddr) diff --git a/pkg/blocklist/blocklist.go b/pkg/blocklist/blocklist.go index c83d890..89804bf 100644 --- a/pkg/blocklist/blocklist.go +++ b/pkg/blocklist/blocklist.go @@ -1,157 +1,13 @@ package blocklist -import ( - "bufio" - "fmt" - "os" - "strings" -) +import "strings" -type Blocklist struct { - hosts map[string]struct{} -} - -func New() *Blocklist { - return &Blocklist{hosts: make(map[string]struct{})} -} - -func (b *Blocklist) Add(host string) { - b.hosts[host] = struct{}{} -} - -// LoadFile reads a blocklist file and adds all parsed hostnames. -// Supports adblock-style (||host^) and hosts-file formats. -func (b *Blocklist) LoadFile(path string) error { - f, err := os.Open(path) - if err != nil { - return fmt.Errorf("open blocklist %s: %w", path, err) - } - defer f.Close() - - scanner := bufio.NewScanner(f) - for scanner.Scan() { - if host, ok := ParseLine(scanner.Text()); ok { - b.Add(host) - } - } - - return scanner.Err() -} - -// IsBlocked returns true if the host or any of its parent domains are in the -// blocklist. Safe to call on a nil receiver (returns false). -func (b *Blocklist) IsBlocked(host string) bool { - if b == nil { - return false - } - - // Walk up domain segments: "a.b.c.com" -> "b.c.com" -> "c.com" -> "com" - for { - if _, ok := b.hosts[host]; ok { - return true - } - dot := strings.IndexByte(host, '.') - if dot < 0 { - return false - } - host = host[dot+1:] - } -} - -// Len returns the number of hostnames in the blocklist. -func (b *Blocklist) Len() int { - if b == nil { - return 0 - } - return len(b.hosts) -} - -// ParseLine parses a single line from a blocklist file. Returns the hostname -// and true if the line is a valid hostname block rule, or ("", false) if the -// line should be skipped. -func ParseLine(line string) (string, bool) { - line = strings.TrimSpace(line) - - if line == "" { - return "", false - } - - // Comments - if line[0] == '!' || line[0] == '#' { - return "", false - } - - // Adblock header - if line[0] == '[' { - return "", false - } - - // Exception rules (not supported yet) - if strings.HasPrefix(line, "@@") { - return "", false - } - - // Element hiding / snippet filters - if strings.Contains(line, "##") || strings.Contains(line, "#$#") || strings.Contains(line, "#?#") || strings.Contains(line, "#@#") { - return "", false - } - - // Adblock domain anchor: ||hostname^ - if strings.HasPrefix(line, "||") { - return parseAdblockDomainAnchor(line) - } - - // Hosts-file format: "0.0.0.0 hostname" or "127.0.0.1 hostname" - if strings.HasPrefix(line, "0.0.0.0 ") || strings.HasPrefix(line, "127.0.0.1 ") { - return parseHostsFileLine(line) - } - - // IPv6 loopback in hosts files - if strings.HasPrefix(line, "::1 ") { - return "", false - } - - // URL pattern rules (contains / but no ||) - if strings.ContainsAny(line, "/*") { - return "", false - } - - // Plain hostname (one per line) - if isValidHostname(line) { - return line, true - } - - return "", false -} - -func parseAdblockDomainAnchor(line string) (string, bool) { - // Strip "||" prefix - host := line[2:] - - // Strip options after $ - if idx := strings.IndexByte(host, '$'); idx >= 0 { - host = host[:idx] - } - - // Must end with ^ (separator) for a hostname-only rule - if !strings.HasSuffix(host, "^") { - return "", false - } - host = host[:len(host)-1] - - // If it contains a path separator, it's a URL pattern, not hostname-only - if strings.ContainsAny(host, "/:") { - return "", false - } - - if !isValidHostname(host) { +func parseHostsFileLine(line string) (string, bool) { + // Must start with a known loopback prefix + if !strings.HasPrefix(line, "0.0.0.0 ") && !strings.HasPrefix(line, "127.0.0.1 ") { return "", false } - return host, true -} - -func parseHostsFileLine(line string) (string, bool) { // Strip inline comments if idx := strings.IndexByte(line, '#'); idx >= 0 { line = strings.TrimSpace(line[:idx]) diff --git a/pkg/blocklist/blocklist_test.go b/pkg/blocklist/blocklist_test.go deleted file mode 100644 index 3c31573..0000000 --- a/pkg/blocklist/blocklist_test.go +++ /dev/null @@ -1,167 +0,0 @@ -package blocklist_test - -import ( - "os" - "path/filepath" - "testing" - - "ublproxy/pkg/blocklist" -) - -func TestParseLine(t *testing.T) { - tests := []struct { - input string - wantHost string - wantOK bool - }{ - // Adblock domain anchor format - {"||ads.example.com^", "ads.example.com", true}, - {"||tracker.net^", "tracker.net", true}, - {"||ads.example.com^$third-party", "ads.example.com", true}, - {"||ads.example.com^$script,image,domain=example.com", "ads.example.com", true}, - - // Hosts-file format - {"0.0.0.0 ads.example.com", "ads.example.com", true}, - {"127.0.0.1 ads.example.com", "ads.example.com", true}, - {"0.0.0.0 ads.example.com # inline comment", "ads.example.com", true}, - {"127.0.0.1 tracker.net", "tracker.net", true}, - - // Plain hostname - {"ads.example.com", "ads.example.com", true}, - {"tracker.net", "tracker.net", true}, - - // Comments — should be skipped - {"! this is an adblock comment", "", false}, - {"# this is a hosts-file comment", "", false}, - - // Blank / whitespace - {"", "", false}, - {" ", "", false}, - {"\t", "", false}, - - // Adblock header - {"[Adblock Plus 2.0]", "", false}, - {"[Adblock]", "", false}, - - // Exception rules — not supported yet - {"@@||ads.example.com^", "", false}, - {"@@||ads.example.com^$document", "", false}, - - // URL pattern rules — not a hostname block, skip - {"/banner/*/img^", "", false}, - {"||ads.example.com/path^", "", false}, - - // Element hiding / snippet rules — skip - {"example.com##.advert", "", false}, - {"example.com#$#log Hello", "", false}, - - // Hosts-file loopback entries — skip - {"0.0.0.0 0.0.0.0", "", false}, - {"127.0.0.1 localhost", "", false}, - {"::1 localhost", "", false}, - {"0.0.0.0 local", "", false}, - } - - for _, tt := range tests { - host, ok := blocklist.ParseLine(tt.input) - if ok != tt.wantOK || host != tt.wantHost { - t.Errorf("ParseLine(%q) = (%q, %v), want (%q, %v)", - tt.input, host, ok, tt.wantHost, tt.wantOK) - } - } -} - -func TestIsBlocked(t *testing.T) { - bl := blocklist.New() - bl.Add("ads.example.com") - bl.Add("tracker.net") - - tests := []struct { - host string - want bool - }{ - // Exact matches - {"ads.example.com", true}, - {"tracker.net", true}, - - // Subdomain matches - {"foo.ads.example.com", true}, - {"bar.foo.ads.example.com", true}, - {"sub.tracker.net", true}, - - // Non-matching - {"example.com", false}, - {"other.com", false}, - {"notads.example.com", false}, - {"adsexample.com", false}, - - // Empty - {"", false}, - } - - for _, tt := range tests { - if got := bl.IsBlocked(tt.host); got != tt.want { - t.Errorf("IsBlocked(%q) = %v, want %v", tt.host, got, tt.want) - } - } -} - -func TestIsBlockedNilSafe(t *testing.T) { - var bl *blocklist.Blocklist - - if bl.IsBlocked("anything.com") { - t.Error("nil Blocklist.IsBlocked should return false") - } -} - -func TestLoadFile(t *testing.T) { - content := `! EasyList header -[Adblock Plus 2.0] -! Homepage: https://easylist.to/ - -||ads.example.com^ -||tracker.net^$third-party -0.0.0.0 malware.example.org -127.0.0.1 spyware.test -# a comment -@@||allowed.example.com^ - -example.com##.ad-banner -/banner/*/img^ -` - - dir := t.TempDir() - path := filepath.Join(dir, "blocklist.txt") - if err := os.WriteFile(path, []byte(content), 0644); err != nil { - t.Fatalf("write temp file: %v", err) - } - - bl := blocklist.New() - if err := bl.LoadFile(path); err != nil { - t.Fatalf("LoadFile: %v", err) - } - - blocked := []string{ - "ads.example.com", - "sub.ads.example.com", - "tracker.net", - "malware.example.org", - "spyware.test", - } - for _, host := range blocked { - if !bl.IsBlocked(host) { - t.Errorf("expected %q to be blocked", host) - } - } - - notBlocked := []string{ - "example.com", - "allowed.example.com", - "other.com", - } - for _, host := range notBlocked { - if bl.IsBlocked(host) { - t.Errorf("expected %q to NOT be blocked", host) - } - } -} diff --git a/pkg/blocklist/pattern.go b/pkg/blocklist/pattern.go new file mode 100644 index 0000000..a76b216 --- /dev/null +++ b/pkg/blocklist/pattern.go @@ -0,0 +1,270 @@ +package blocklist + +import ( + "errors" + "strings" +) + +// segment represents one piece of a compiled adblock pattern. +type segment struct { + kind segmentKind + literal string // for segLiteral +} + +type segmentKind int + +const ( + segLiteral segmentKind = iota // match exact string + segWildcard // match any sequence of characters (including empty) + segSeparator // match a single separator char or end of string +) + +// Rule represents a compiled adblock filter pattern. +type Rule struct { + segments []segment + anchorStart bool // | at start: must match beginning of URL + anchorEnd bool // | at end: must match end of URL + domainAnchor bool // || at start: must match at domain boundary + domainSuffix string // the domain part after || (e.g. "example.com") + pathPattern []segment // the pattern after the domain part (e.g. "/ads/*.gif") +} + +// Compile parses an adblock filter pattern string into a Rule. +func Compile(pattern string) (*Rule, error) { + if pattern == "" { + return nil, errors.New("empty pattern") + } + + r := &Rule{} + + // Handle domain anchor (||) + if strings.HasPrefix(pattern, "||") { + r.domainAnchor = true + pattern = pattern[2:] + return compileDomainAnchor(r, pattern) + } + + // Handle start anchor (|) + if strings.HasPrefix(pattern, "|") { + r.anchorStart = true + pattern = pattern[1:] + } + + // Handle end anchor (|) + if strings.HasSuffix(pattern, "|") { + r.anchorEnd = true + pattern = pattern[:len(pattern)-1] + } + + r.segments = compileSegments(pattern) + return r, nil +} + +// compileDomainAnchor handles ||domain.com/path patterns. +// Splits the pattern into a domain part and an optional path pattern. +func compileDomainAnchor(r *Rule, pattern string) (*Rule, error) { + // Find where the domain ends: first ^, /, :, or * after || + domainEnd := len(pattern) + for i, c := range pattern { + if c == '/' || c == ':' || c == '^' || c == '*' { + domainEnd = i + break + } + } + + r.domainSuffix = strings.ToLower(pattern[:domainEnd]) + if domainEnd < len(pattern) { + // Handle end anchor on the remaining part + rest := pattern[domainEnd:] + if strings.HasSuffix(rest, "|") { + r.anchorEnd = true + rest = rest[:len(rest)-1] + } + r.pathPattern = compileSegments(rest) + } + + return r, nil +} + +func compileSegments(pattern string) []segment { + // Case-insensitive: lowercase the pattern (we'll lowercase the URL during match) + pattern = strings.ToLower(pattern) + + var segments []segment + i := 0 + for i < len(pattern) { + switch pattern[i] { + case '*': + // Collapse consecutive wildcards + if len(segments) == 0 || segments[len(segments)-1].kind != segWildcard { + segments = append(segments, segment{kind: segWildcard}) + } + i++ + case '^': + segments = append(segments, segment{kind: segSeparator}) + i++ + default: + // Collect literal characters until we hit * or ^ + start := i + for i < len(pattern) && pattern[i] != '*' && pattern[i] != '^' { + i++ + } + segments = append(segments, segment{kind: segLiteral, literal: pattern[start:i]}) + } + } + return segments +} + +// Match reports whether the rule matches the given URL. +func (r *Rule) Match(rawURL string) bool { + url := strings.ToLower(rawURL) + + if r.domainAnchor { + return r.matchDomainAnchor(url) + } + + if r.anchorStart { + return matchSegments(r.segments, url, r.anchorEnd) + } + + // No start anchor: try matching at every position + for i := range len(url) { + if matchSegments(r.segments, url[i:], r.anchorEnd) { + return true + } + } + return false +} + +// matchDomainAnchor checks if the URL contains the domain at a domain boundary +// (after :// or after a dot), then matches the remaining path pattern. +func (r *Rule) matchDomainAnchor(url string) bool { + // Find the host portion of the URL + hostStart, hostEnd := findHost(url) + if hostStart < 0 { + return false + } + + host := url[hostStart:hostEnd] + + // Check if host matches domain suffix at a domain boundary + if !matchesDomainSuffix(host, r.domainSuffix) { + return false + } + + // If no path pattern, the domain match alone is sufficient + if len(r.pathPattern) == 0 { + return true + } + + // Match the path pattern against the rest of the URL after the host + rest := url[hostEnd:] + return matchSegments(r.pathPattern, rest, r.anchorEnd) +} + +// findHost extracts the host portion from a URL string. +// Returns start and end indices, or (-1, -1) if no host found. +func findHost(url string) (int, int) { + // Find :// + schemeEnd := strings.Index(url, "://") + if schemeEnd < 0 { + return -1, -1 + } + hostStart := schemeEnd + 3 + + // Host ends at /, :, or end of string + hostEnd := len(url) + for i := hostStart; i < len(url); i++ { + if url[i] == '/' || url[i] == ':' { + hostEnd = i + break + } + } + + return hostStart, hostEnd +} + +// matchesDomainSuffix checks if host equals domain or ends with .domain. +func matchesDomainSuffix(host, domain string) bool { + if host == domain { + return true + } + return strings.HasSuffix(host, "."+domain) +} + +// matchSegments matches a sequence of segments against text. +// If anchorEnd is true, the match must consume the entire text. +func matchSegments(segments []segment, text string, anchorEnd bool) bool { + return matchSegmentsAt(segments, 0, text, 0, anchorEnd) +} + +func matchSegmentsAt(segments []segment, si int, text string, ti int, anchorEnd bool) bool { + for si < len(segments) { + seg := segments[si] + switch seg.kind { + case segLiteral: + if ti+len(seg.literal) > len(text) { + return false + } + if text[ti:ti+len(seg.literal)] != seg.literal { + return false + } + ti += len(seg.literal) + si++ + + case segSeparator: + // Match a single separator character, or end of string + if ti == len(text) { + // End of string counts as separator + si++ + continue + } + if !isSeparator(text[ti]) { + return false + } + ti++ + si++ + + case segWildcard: + si++ + // If wildcard is the last segment, it matches everything + if si == len(segments) { + if anchorEnd { + return true + } + return true + } + // Try matching the rest of the pattern at every position + for pos := ti; pos <= len(text); pos++ { + if matchSegmentsAt(segments, si, text, pos, anchorEnd) { + return true + } + } + return false + } + } + + // All segments consumed + if anchorEnd { + return ti == len(text) + } + return true +} + +// isSeparator returns true if the byte is a separator character per adblock spec. +// A separator is anything except a letter, digit, _, -, ., or %. +func isSeparator(b byte) bool { + if b >= 'a' && b <= 'z' { + return false + } + if b >= 'A' && b <= 'Z' { + return false + } + if b >= '0' && b <= '9' { + return false + } + if b == '_' || b == '-' || b == '.' || b == '%' { + return false + } + return true +} diff --git a/pkg/blocklist/pattern_test.go b/pkg/blocklist/pattern_test.go new file mode 100644 index 0000000..54e783c --- /dev/null +++ b/pkg/blocklist/pattern_test.go @@ -0,0 +1,95 @@ +package blocklist_test + +import ( + "testing" + + "ublproxy/pkg/blocklist" +) + +func TestPatternMatch(t *testing.T) { + tests := []struct { + pattern string + url string + want bool + }{ + // Plain substring match + {"ad.gif", "http://example.com/ad.gif", true}, + {"ad.gif", "http://example.com/ad.gif?q=1", true}, + {"ad.gif", "http://example.com/page.html", false}, + {"banner", "http://example.com/ads/banner123.gif", true}, + {"banner", "http://example.com/page.html", false}, + + // Wildcard + {"/ads/banner*.gif", "http://example.com/ads/banner123.gif", true}, + {"/ads/banner*.gif", "http://example.com/ads/banner.gif", true}, + {"/ads/banner*.gif", "http://example.com/ads/bannerXYZ.gif", true}, + {"/ads/banner*.gif", "http://example.com/ads/tracking.js", false}, + {"ad*banner", "http://example.com/ad-and-banner", true}, + {"ad*banner", "http://example.com/adbanner", true}, + {"ad*banner", "http://example.com/banner-ad", false}, + + // Separator character (^) + // ^ matches any non-alphanumeric except _ - . % + {"example.com^", "http://example.com/", true}, + {"example.com^", "http://example.com:8000/", true}, + {"example.com^", "http://example.com.ar/", false}, + {"^foo.bar^", "http://example.com/foo.bar?a=1", true}, + {"^foo.bar^", "http://example.com/foo.bar", true}, // end of string counts as separator + + // Address start anchor (|) + {"|http://example.com", "http://example.com/banner.gif", true}, + {"|http://example.com", "https://example.com/banner.gif", false}, + {"|http://bad.example/", "http://bad.example/banner.gif", true}, + {"|http://bad.example/", "http://good.example/analyze?http://bad.example/", false}, + + // Address end anchor (|) + {"swf|", "http://example.com/annoyingflash.swf", true}, + {"swf|", "http://example.com/swf/index.html", false}, + {".gif|", "http://example.com/banner.gif", true}, + {".gif|", "http://example.com/banner.gif?q=1", false}, + + // Domain anchor (||) + {"||example.com", "http://example.com/banner.gif", true}, + {"||example.com", "https://example.com/banner.gif", true}, + {"||example.com", "http://www.example.com/banner.gif", true}, + {"||example.com", "http://badexample.com/banner.gif", false}, + {"||example.com/banner.gif", "http://example.com/banner.gif", true}, + {"||example.com/banner.gif", "http://example.com/other.gif", false}, + + // Domain anchor with separator + {"||ads.example.com^", "http://ads.example.com/tracking.js", true}, + {"||ads.example.com^", "http://ads.example.com:8080/tracking.js", true}, + {"||ads.example.com^", "http://ads.example.com.ar/tracking.js", false}, + + // Combined features + {"||example.com/ads/*.gif", "http://example.com/ads/banner123.gif", true}, + {"||example.com/ads/*.gif", "http://example.com/ads/tracker.js", false}, + {"||example.com^*/tracking", "http://example.com/path/tracking", true}, + + // Case insensitivity (default) + {"AdVeRt", "http://example.com/advert.js", true}, + {"BANNER", "http://example.com/banner.gif", true}, + + // Empty / edge cases + {"", "http://example.com/", false}, + } + + for _, tt := range tests { + t.Run(tt.pattern+"_"+tt.url, func(t *testing.T) { + rule, err := blocklist.Compile(tt.pattern) + if tt.pattern == "" { + if err == nil { + t.Fatal("expected error for empty pattern, got nil") + } + return + } + if err != nil { + t.Fatalf("Compile(%q): %v", tt.pattern, err) + } + got := rule.Match(tt.url) + if got != tt.want { + t.Errorf("Compile(%q).Match(%q) = %v, want %v", tt.pattern, tt.url, got, tt.want) + } + }) + } +} diff --git a/pkg/blocklist/ruleset.go b/pkg/blocklist/ruleset.go new file mode 100644 index 0000000..29e2023 --- /dev/null +++ b/pkg/blocklist/ruleset.go @@ -0,0 +1,212 @@ +package blocklist + +import ( + "bufio" + "fmt" + "os" + "strings" +) + +// RuleSet holds blocking rules for URL filtering. It combines a hostname map +// (fast path for ||hostname^ rules) with compiled URL pattern rules. +type RuleSet struct { + hosts map[string]struct{} + rules []*Rule +} + +func NewRuleSet() *RuleSet { + return &RuleSet{hosts: make(map[string]struct{})} +} + +// AddHostname adds a hostname to the fast-path blocklist. +// Matches the hostname and all its subdomains. +func (rs *RuleSet) AddHostname(host string) { + rs.hosts[strings.ToLower(host)] = struct{}{} +} + +// AddRule compiles an adblock URL pattern and adds it to the rule list. +// Returns an error if the pattern is invalid. +func (rs *RuleSet) AddRule(pattern string) error { + rule, err := Compile(pattern) + if err != nil { + return err + } + rs.rules = append(rs.rules, rule) + return nil +} + +// LoadFile reads a blocklist file and adds all parsed rules. +// Hostname-only rules go to the fast path; URL patterns become compiled rules. +func (rs *RuleSet) LoadFile(path string) error { + f, err := os.Open(path) + if err != nil { + return fmt.Errorf("open blocklist %s: %w", path, err) + } + defer f.Close() + + scanner := bufio.NewScanner(f) + for scanner.Scan() { + rs.addLine(scanner.Text()) + } + return scanner.Err() +} + +// addLine parses a single line from a blocklist file and adds it to the +// appropriate data structure (hostname map or compiled rule list). +func (rs *RuleSet) addLine(line string) { + line = strings.TrimSpace(line) + + if line == "" || line[0] == '!' || line[0] == '#' || line[0] == '[' { + return + } + + // Exception rules (Phase 2) + if strings.HasPrefix(line, "@@") { + return + } + + // Element hiding / snippet filters + if strings.Contains(line, "##") || strings.Contains(line, "#$#") || strings.Contains(line, "#?#") || strings.Contains(line, "#@#") { + return + } + + // Strip options after $ for now (Phase 3 will handle them) + rawPattern := line + if idx := strings.IndexByte(rawPattern, '$'); idx >= 0 { + rawPattern = rawPattern[:idx] + } + + // Try to extract a hostname-only rule for the fast path + if host, ok := extractHostnameRule(rawPattern); ok { + rs.AddHostname(host) + return + } + + // Hosts-file format: "0.0.0.0 hostname" or "127.0.0.1 hostname" + if host, ok := parseHostsFileLine(rawPattern); ok { + rs.AddHostname(host) + return + } + + // IPv6 loopback in hosts files + if strings.HasPrefix(rawPattern, "::1 ") { + return + } + + // Compile as a URL pattern rule + rule, err := Compile(rawPattern) + if err != nil { + return // skip invalid patterns silently + } + rs.rules = append(rs.rules, rule) +} + +// extractHostnameRule checks if the pattern is a hostname-only rule +// (||hostname^ with no path or wildcards). Returns the hostname and true +// if it qualifies for the fast-path map. +func extractHostnameRule(pattern string) (string, bool) { + if !strings.HasPrefix(pattern, "||") { + return "", false + } + + host := pattern[2:] + + // Must end with ^ for a hostname-only rule + if !strings.HasSuffix(host, "^") { + return "", false + } + host = host[:len(host)-1] + + // If it contains path or wildcard characters, it's a URL pattern + if strings.ContainsAny(host, "/*:") { + return "", false + } + + if !isValidHostname(host) { + return "", false + } + + return host, true +} + +// ShouldBlock returns true if the URL matches any blocking rule. +// Safe to call on a nil receiver (returns false). +func (rs *RuleSet) ShouldBlock(rawURL string) bool { + if rs == nil { + return false + } + + // Fast path: check hostname against the hostname map + url := strings.ToLower(rawURL) + host := extractHostFromURL(url) + if rs.isHostBlocked(host) { + return true + } + + // Slow path: check URL against compiled pattern rules + for _, rule := range rs.rules { + if rule.Match(rawURL) { + return true + } + } + + return false +} + +// IsHostBlocked returns true if the hostname (or any parent domain) is in +// the fast-path hostname map. Safe to call on a nil receiver. +// Use this for CONNECT-level blocking where only the hostname is available. +func (rs *RuleSet) IsHostBlocked(host string) bool { + if rs == nil { + return false + } + return rs.isHostBlocked(strings.ToLower(host)) +} + +func (rs *RuleSet) isHostBlocked(host string) bool { + for { + if _, ok := rs.hosts[host]; ok { + return true + } + dot := strings.IndexByte(host, '.') + if dot < 0 { + return false + } + host = host[dot+1:] + } +} + +// extractHostFromURL pulls the hostname from a URL string. +func extractHostFromURL(url string) string { + schemeEnd := strings.Index(url, "://") + if schemeEnd < 0 { + return "" + } + hostStart := schemeEnd + 3 + + hostEnd := len(url) + for i := hostStart; i < len(url); i++ { + if url[i] == '/' || url[i] == ':' { + hostEnd = i + break + } + } + + return url[hostStart:hostEnd] +} + +// HostCount returns the number of hostnames in the fast-path map. +func (rs *RuleSet) HostCount() int { + if rs == nil { + return 0 + } + return len(rs.hosts) +} + +// RuleCount returns the number of compiled URL pattern rules. +func (rs *RuleSet) RuleCount() int { + if rs == nil { + return 0 + } + return len(rs.rules) +} diff --git a/pkg/blocklist/ruleset_test.go b/pkg/blocklist/ruleset_test.go new file mode 100644 index 0000000..0064a6a --- /dev/null +++ b/pkg/blocklist/ruleset_test.go @@ -0,0 +1,180 @@ +package blocklist_test + +import ( + "os" + "testing" + + "ublproxy/pkg/blocklist" +) + +func TestRuleSetHostnameBlocking(t *testing.T) { + rs := blocklist.NewRuleSet() + rs.AddHostname("ads.example.com") + rs.AddHostname("tracker.net") + + tests := []struct { + url string + want bool + }{ + // Exact hostname match + {"http://ads.example.com/tracking.js", true}, + {"https://ads.example.com/pixel.gif", true}, + + // Subdomain match + {"http://cdn.ads.example.com/banner.gif", true}, + + // Non-matching + {"http://example.com/page.html", false}, + {"http://good.example.com/page.html", false}, + + // Second hostname + {"http://tracker.net/collect", true}, + {"http://sub.tracker.net/event", true}, + {"http://nottracker.net/page", false}, + } + + for _, tt := range tests { + t.Run(tt.url, func(t *testing.T) { + got := rs.ShouldBlock(tt.url) + if got != tt.want { + t.Errorf("ShouldBlock(%q) = %v, want %v", tt.url, got, tt.want) + } + }) + } +} + +func TestRuleSetPatternBlocking(t *testing.T) { + rs := blocklist.NewRuleSet() + rs.AddRule("||example.com/ads/*.gif") + rs.AddRule("/tracking.js") + + tests := []struct { + url string + want bool + }{ + // URL pattern match + {"http://example.com/ads/banner123.gif", true}, + {"https://example.com/ads/small.gif", true}, + {"http://example.com/ads/tracker.js", false}, + {"http://example.com/images/photo.gif", false}, + + // Substring pattern match + {"http://other.com/tracking.js", true}, + {"http://example.com/tracking.js?v=1", true}, + {"http://example.com/page.html", false}, + } + + for _, tt := range tests { + t.Run(tt.url, func(t *testing.T) { + got := rs.ShouldBlock(tt.url) + if got != tt.want { + t.Errorf("ShouldBlock(%q) = %v, want %v", tt.url, got, tt.want) + } + }) + } +} + +func TestRuleSetMixed(t *testing.T) { + rs := blocklist.NewRuleSet() + rs.AddHostname("ads.example.com") + rs.AddRule("/tracking.js") + + // Hostname rule blocks whole domain + if !rs.ShouldBlock("http://ads.example.com/page.html") { + t.Error("hostname rule should block ads.example.com") + } + + // Pattern rule blocks matching URLs on any domain + if !rs.ShouldBlock("http://good.example.com/tracking.js") { + t.Error("pattern rule should block tracking.js on any domain") + } + + // Neither rule matches + if rs.ShouldBlock("http://good.example.com/page.html") { + t.Error("should not block unrelated URL") + } +} + +func TestRuleSetNilSafe(t *testing.T) { + var rs *blocklist.RuleSet + if rs.ShouldBlock("http://example.com/") { + t.Error("nil RuleSet should not block anything") + } +} + +func TestRuleSetLoadFile(t *testing.T) { + content := `! Comment line +[Adblock Plus 2.0] +||ads.example.com^ +0.0.0.0 tracker.net +/banner*.gif +||cdn.example.com/ads/* +example.org##.ad-class +@@||allowed.com^ +` + f, err := os.CreateTemp("", "ruleset-test-*.txt") + if err != nil { + t.Fatal(err) + } + defer os.Remove(f.Name()) + + if _, err := f.WriteString(content); err != nil { + t.Fatal(err) + } + f.Close() + + rs := blocklist.NewRuleSet() + if err := rs.LoadFile(f.Name()); err != nil { + t.Fatalf("LoadFile: %v", err) + } + + tests := []struct { + url string + want bool + }{ + // Hostname from adblock format + {"http://ads.example.com/tracking.js", true}, + {"http://sub.ads.example.com/pixel.gif", true}, + + // Hostname from hosts-file format + {"http://tracker.net/collect", true}, + + // URL pattern rule + {"http://example.com/banner123.gif", true}, + {"http://example.com/bannerXYZ.gif", true}, + {"http://example.com/page.html", false}, + + // Domain-anchored URL pattern + {"http://cdn.example.com/ads/popup.js", true}, + {"http://cdn.example.com/images/photo.gif", false}, + + // Element hiding rules are skipped (not blocking rules) + {"http://example.org/page.html", false}, + } + + for _, tt := range tests { + t.Run(tt.url, func(t *testing.T) { + got := rs.ShouldBlock(tt.url) + if got != tt.want { + t.Errorf("ShouldBlock(%q) = %v, want %v", tt.url, got, tt.want) + } + }) + } +} + +func TestRuleSetConnectBlocking(t *testing.T) { + rs := blocklist.NewRuleSet() + rs.AddHostname("ads.example.com") + + // For CONNECT requests, the proxy synthesizes a URL from host:port + // The proxy should check IsHostBlocked for CONNECT efficiency + if !rs.IsHostBlocked("ads.example.com") { + t.Error("IsHostBlocked should return true for blocked hostname") + } + if !rs.IsHostBlocked("sub.ads.example.com") { + t.Error("IsHostBlocked should return true for subdomain of blocked hostname") + } + if rs.IsHostBlocked("example.com") { + t.Error("IsHostBlocked should return false for non-blocked hostname") + } +} diff --git a/proxy.go b/proxy.go index 32287f3..ccfac07 100644 --- a/proxy.go +++ b/proxy.go @@ -10,15 +10,15 @@ import ( type proxyHandler struct { certs *certCache caCertPEM []byte - blocklist *blocklist.Blocklist + rules *blocklist.RuleSet transport *http.Transport } -func newProxyHandler(certs *certCache, caCertPEM []byte, bl *blocklist.Blocklist) *proxyHandler { +func newProxyHandler(certs *certCache, caCertPEM []byte, rules *blocklist.RuleSet) *proxyHandler { return &proxyHandler{ certs: certs, caCertPEM: caCertPEM, - blocklist: bl, + rules: rules, transport: &http.Transport{ // Skip verification when connecting to upstream servers since // we are acting as a proxy, not validating end-server identity diff --git a/proxy_test.go b/proxy_test.go index f7f7c0c..1fc7243 100644 --- a/proxy_test.go +++ b/proxy_test.go @@ -25,8 +25,8 @@ type testEnv struct { // startTestEnv spins up an upstream HTTP server, an upstream HTTPS server, and // the proxy itself. The CA is generated in-memory — no disk I/O. -// Pass nil for bl to create a proxy without blocking. -func startTestEnv(t *testing.T, upstreamHandler http.Handler, bl *blocklist.Blocklist) *testEnv { +// Pass nil for rules to create a proxy without blocking. +func startTestEnv(t *testing.T, upstreamHandler http.Handler, rules *blocklist.RuleSet) *testEnv { t.Helper() // Upstream servers @@ -41,7 +41,7 @@ func startTestEnv(t *testing.T, upstreamHandler http.Handler, bl *blocklist.Bloc certs := newCertCache(caCert, caKey) caCertPEM := encodeCertPEM(caCert) - handler := newProxyHandler(certs, caCertPEM, bl) + handler := newProxyHandler(certs, caCertPEM, rules) // The proxy needs a real http.Server (not httptest) so that Hijack works // on the ResponseWriter. httptest.NewServer wraps net/http.Server, which @@ -349,14 +349,14 @@ func TestHTTPSProxyPOST(t *testing.T) { func TestBlocksHTTPByHostname(t *testing.T) { var upstreamHit atomic.Bool - bl := blocklist.New() - bl.Add("127.0.0.1") + rs := blocklist.NewRuleSet() + rs.AddHostname("127.0.0.1") env := startTestEnv(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { upstreamHit.Store(true) w.WriteHeader(http.StatusOK) w.Write([]byte("should not see this")) - }), bl) + }), rs) client := env.httpClient(t) resp, err := client.Get(env.httpURL + "/tracking.js") @@ -382,14 +382,14 @@ func TestBlocksHTTPByHostname(t *testing.T) { func TestBlocksHTTPSByHostname(t *testing.T) { var upstreamHit atomic.Bool - bl := blocklist.New() - bl.Add("127.0.0.1") + rs := blocklist.NewRuleSet() + rs.AddHostname("127.0.0.1") env := startTestEnv(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { upstreamHit.Store(true) w.WriteHeader(http.StatusOK) w.Write([]byte("should not see this")) - }), bl) + }), rs) client := env.httpClient(t) @@ -404,3 +404,96 @@ func TestBlocksHTTPSByHostname(t *testing.T) { t.Error("upstream was hit, but request should have been blocked") } } + +func TestBlocksHTTPByURLPattern(t *testing.T) { + var lastPath string + + rs := blocklist.NewRuleSet() + rs.AddRule("/ads/banner*.gif") + + env := startTestEnv(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + lastPath = r.URL.Path + w.WriteHeader(http.StatusOK) + w.Write([]byte("upstream response")) + }), rs) + + client := env.httpClient(t) + + // Matching URL pattern should be blocked + resp, err := client.Get(env.httpURL + "/ads/banner123.gif") + if err != nil { + t.Fatalf("GET blocked URL: %v", err) + } + resp.Body.Close() + + if resp.StatusCode != http.StatusNoContent { + t.Errorf("blocked status = %d, want %d", resp.StatusCode, http.StatusNoContent) + } + + // Non-matching URL should pass through + resp, err = client.Get(env.httpURL + "/page.html") + if err != nil { + t.Fatalf("GET allowed URL: %v", err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Errorf("allowed status = %d, want %d", resp.StatusCode, http.StatusOK) + } + + body, _ := io.ReadAll(resp.Body) + if string(body) != "upstream response" { + t.Errorf("body = %q, want %q", body, "upstream response") + } + + if lastPath != "/page.html" { + t.Errorf("upstream saw path = %q, want %q", lastPath, "/page.html") + } +} + +func TestBlocksHTTPSByURLPattern(t *testing.T) { + var requestPaths []string + + rs := blocklist.NewRuleSet() + rs.AddRule("/ads/tracking.js") + + env := startTestEnv(t, http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + requestPaths = append(requestPaths, r.URL.Path) + w.WriteHeader(http.StatusOK) + w.Write([]byte("upstream response")) + }), rs) + + client := env.httpClient(t) + + // Non-matching path should pass through (CONNECT succeeds, MITM proxies request) + resp, err := client.Get(env.httpsURL + "/page.html") + if err != nil { + t.Fatalf("GET allowed HTTPS URL: %v", err) + } + body, _ := io.ReadAll(resp.Body) + resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Errorf("allowed status = %d, want %d", resp.StatusCode, http.StatusOK) + } + if string(body) != "upstream response" { + t.Errorf("body = %q, want %q", body, "upstream response") + } + + // Matching URL pattern should be blocked after MITM (CONNECT succeeds, + // but the individual request inside the tunnel is blocked) + resp, err = client.Get(env.httpsURL + "/ads/tracking.js") + if err != nil { + t.Fatalf("GET blocked HTTPS URL: %v", err) + } + resp.Body.Close() + + if resp.StatusCode != http.StatusNoContent { + t.Errorf("blocked status = %d, want %d", resp.StatusCode, http.StatusNoContent) + } + + // Only the allowed request should have reached upstream + if len(requestPaths) != 1 || requestPaths[0] != "/page.html" { + t.Errorf("upstream saw paths = %v, want [/page.html]", requestPaths) + } +} -- 2.51.2