diff --git a/DECISIONS.md b/DECISIONS.md index 1d859a2..5aa70db 100644 --- a/DECISIONS.md +++ b/DECISIONS.md @@ -47,3 +47,7 @@ - 2026-02-26 m+git@andri.dk — Per-user blocklist subscriptions stored in `blocklist_subscriptions` table. Users can add/remove/toggle blocklist URLs via the portal UI. Subscription URLs are merged with CLI `--blocklist` flags during rule reload. CLI flags remain as system-wide defaults; user subscriptions are additive. - 2026-02-26 m+git@andri.dk — Blocklist download cache in `blocklist_cache` SQLite table. Remote URLs are downloaded once and cached for 24 hours. `reloadRules()` reads from cache instead of re-fetching on every rule change. Stale cache is used as fallback if re-download fails. "Refresh All" button clears cache and triggers re-download. Cache stored in SQLite (not filesystem) to keep everything in one place. - 2026-02-26 m+git@andri.dk — Portal page now shows authenticated users a rule management panel (list/add/toggle/delete rules) and a blocklist subscription panel (list/add/toggle/delete subscriptions, refresh all, suggested lists). Both panels are hidden when not authenticated. No external JS dependencies. +- 2026-02-26 m+git@andri.dk — Per-user layered rule architecture. Split `rules atomic.Pointer[RuleSet]` into `baselineRules` (always-active, from `--blocklist` + `--default-subscription`) and `userRules sync.Map` (per-credential, lazily loaded from DB). Layered evaluation: user `@@` exceptions override baseline blocks, then user blocks, then baseline blocks. User `#@#` element hiding exceptions suppress baseline `##` selectors. API mutations invalidate only the affected user's cached RuleSet. +- 2026-02-26 m+git@andri.dk — Session map extended from `IP→token` to `IP→{Token, CredentialID}`. The credential ID identifies which user is on which IP, enabling per-user rule lookup at all proxy layers (CONNECT, HTTP, element hiding injection). +- 2026-02-26 m+git@andri.dk — `--default-subscription` CLI flag (repeatable). Defaults to EasyList + EasyPrivacy if none specified. These are loaded into the baseline at startup and apply to all traffic. Users cannot disable the baseline, only add personal rules/exceptions on top. +- 2026-02-26 m+git@andri.dk — Baseline rules loaded synchronously at startup (`reloadBaseline()` called after DB setup). Ensures rules are ready before any traffic arrives. Remote subscriptions use DB cache with 24h TTL, so subsequent starts are fast. diff --git a/api.go b/api.go index d2a42f7..2e688f2 100644 --- a/api.go +++ b/api.go @@ -33,8 +33,9 @@ type apiHandler struct { sessions *sessionMap // onRulesChanged is called after any rule mutation (create/delete/patch) - // to trigger a hot-reload of the in-memory RuleSet. - onRulesChanged func() + // to invalidate the cached per-user RuleSet. The argument is the + // credential ID whose rules changed. + onRulesChanged func(credentialID string) // challenges stores pending WebAuthn challenges keyed by base64url // challenge value. Challenges are single-use and expire after challengeTTL. diff --git a/api_auth.go b/api_auth.go index 0ff0ffc..20a4eee 100644 --- a/api_auth.go +++ b/api_auth.go @@ -161,9 +161,12 @@ func (a *apiHandler) handleRegisterFinish(w http.ResponseWriter, r *http.Request return } - // Associate session with client IP for script injection + // Associate session with client IP for script injection and per-user rules if a.sessions != nil { - a.sessions.Set(clientIPFromRequest(r), sess.Token) + a.sessions.Set(clientIPFromRequest(r), sessionEntry{ + Token: sess.Token, + CredentialID: credID, + }) } writeJSON(w, http.StatusOK, map[string]string{"token": sess.Token}) @@ -291,9 +294,12 @@ func (a *apiHandler) handleLoginFinish(w http.ResponseWriter, r *http.Request) { return } - // Associate session with client IP for script injection + // Associate session with client IP for script injection and per-user rules if a.sessions != nil { - a.sessions.Set(clientIPFromRequest(r), sess.Token) + a.sessions.Set(clientIPFromRequest(r), sessionEntry{ + Token: sess.Token, + CredentialID: req.CredentialID, + }) } writeJSON(w, http.StatusOK, map[string]string{"token": sess.Token}) diff --git a/api_rules.go b/api_rules.go index 4c08af7..565da4e 100644 --- a/api_rules.go +++ b/api_rules.go @@ -81,7 +81,7 @@ func (a *apiHandler) handleCreateRule(w http.ResponseWriter, r *http.Request, se return } - a.triggerReload() + a.triggerReload(sess.CredentialID) writeJSON(w, http.StatusCreated, toRuleResponse(*rule)) } @@ -97,7 +97,7 @@ func (a *apiHandler) handleDeleteRule(w http.ResponseWriter, r *http.Request, pa return } - a.triggerReload() + a.triggerReload(sess.CredentialID) writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) } @@ -128,12 +128,13 @@ func (a *apiHandler) handlePatchRule(w http.ResponseWriter, r *http.Request, pat return } - a.triggerReload() + a.triggerReload(sess.CredentialID) writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) } -// triggerReload calls the onRulesChanged callback if set. -func (a *apiHandler) triggerReload() { +// triggerReload calls the onRulesChanged callback if set, passing the +// credential ID whose rules changed so only that user's cache is invalidated. +func (a *apiHandler) triggerReload(credentialID string) { if a.onRulesChanged != nil { go func() { defer func() { @@ -141,7 +142,7 @@ func (a *apiHandler) triggerReload() { fmt.Fprintf(os.Stderr, "panic in onRulesChanged: %v\n", r) } }() - a.onRulesChanged() + a.onRulesChanged(credentialID) }() } } diff --git a/api_subscriptions.go b/api_subscriptions.go index 0d14b9d..34349f0 100644 --- a/api_subscriptions.go +++ b/api_subscriptions.go @@ -90,7 +90,7 @@ func (a *apiHandler) handleCreateSubscription(w http.ResponseWriter, r *http.Req return } - a.triggerReload() + a.triggerReload(sess.CredentialID) writeJSON(w, http.StatusCreated, toSubscriptionResponse(*sub)) } @@ -106,7 +106,7 @@ func (a *apiHandler) handleDeleteSubscription(w http.ResponseWriter, r *http.Req return } - a.triggerReload() + a.triggerReload(sess.CredentialID) writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) } @@ -137,7 +137,7 @@ func (a *apiHandler) handlePatchSubscription(w http.ResponseWriter, r *http.Requ return } - a.triggerReload() + a.triggerReload(sess.CredentialID) writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) } @@ -148,7 +148,7 @@ func (a *apiHandler) handleRefreshSubscriptions(w http.ResponseWriter, r *http.R return } - a.triggerReload() + a.triggerReload(sess.CredentialID) writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) } diff --git a/connect.go b/connect.go index 41d1c62..e801810 100644 --- a/connect.go +++ b/connect.go @@ -17,8 +17,8 @@ func (p *proxyHandler) handleConnect(w http.ResponseWriter, r *http.Request) { host = r.Host port = "443" } - rules := p.getRules() - if rules != nil && rules.IsHostBlocked(host) { + clientIP, _, _ := net.SplitHostPort(r.RemoteAddr) + if p.shouldBlockHost(clientIP, host) { http.Error(w, "blocked", http.StatusForbidden) return } @@ -85,8 +85,7 @@ func (p *proxyHandler) proxyTLSRequests(clientTLS *tls.Conn, host, port, clientI // URL-level blocking for pattern rules (hostname was already // checked at CONNECT time; this catches path-specific rules) ctx := matchContextFromRequest(req) - rules := p.getRules() - if rules != nil && rules.ShouldBlockRequest(targetURL, ctx) { + if p.shouldBlock(clientIP, targetURL, ctx) { req.Body.Close() blocked := &http.Response{ StatusCode: http.StatusNoContent, diff --git a/elemhide_inject.go b/elemhide_inject.go index 348343a..65c183f 100644 --- a/elemhide_inject.go +++ b/elemhide_inject.go @@ -37,11 +37,13 @@ var srcBlockableTags = map[string]string{ } // srcBlockContext carries the page context needed to resolve relative src -// attributes and check them against URL blocking rules. +// attributes and check them against URL blocking rules. Uses the proxy's +// layered shouldBlock for per-user rule evaluation. type srcBlockContext struct { - scheme string - host string - rules *blocklist.RuleSet + scheme string + host string + proxy *proxyHandler + clientIP string } // resolveSrc resolves an element's src attribute to an absolute URL. @@ -77,19 +79,27 @@ func (p *proxyHandler) applyElementHiding(resp *http.Response, host, clientIP st return nil, false } - var eh *blocklist.ElementHiding - var hasURLRules bool - rules := p.getRules() - if rules != nil { - eh = rules.ElementHidingForDomain(host) - hasURLRules = rules.HostCount() > 0 || rules.RuleCount() > 0 + baseline := p.getBaselineRules() + credID := p.credentialForIP(clientIP) + userRS := p.getUserRules(credID) + + // Merge element hiding from baseline and user rules + var baselineEH, userEH *blocklist.ElementHiding + if baseline != nil { + baselineEH = baseline.ElementHidingForDomain(host) + } + if userRS != nil { + userEH = userRS.ElementHidingForDomain(host) } + hasURLRules := (baseline != nil && (baseline.HostCount() > 0 || baseline.RuleCount() > 0)) || + (userRS != nil && (userRS.HostCount() > 0 || userRS.RuleCount() > 0)) + // Generate bootstrap script tag (empty string if no session) scriptTag := p.bootstrapScriptTag(clientIP, host) // Nothing to do if there are no rules AND no script to inject - if eh == nil && !hasURLRules && scriptTag == "" { + if baselineEH == nil && userEH == nil && !hasURLRules && scriptTag == "" { return nil, false } @@ -116,12 +126,15 @@ func (p *proxyHandler) applyElementHiding(resp *http.Response, host, clientIP st return nil, false } - sc := srcBlockContext{scheme: "https", host: host, rules: rules} + // For src-based resource stripping, use the layered shouldBlock + // approach via a proxy-aware srcBlockContext. + sc := srcBlockContext{scheme: "https", host: host, proxy: p, clientIP: clientIP} modified := stripBlockedResources(body, sc) - if eh != nil && eh.CSS != "" { - // Sanitize CSS to prevent XSS via injection - safeCSS := styleCloseRe.ReplaceAllString(eh.CSS, `<\/style`) + // Merge baseline + user element hiding CSS, applying user #@# exceptions + css := mergeElementHidingCSS(baselineEH, userEH, userRS, host) + if css != "" { + safeCSS := styleCloseRe.ReplaceAllString(css, `<\/style`) styleTag := []byte("") modified = injectStyleTag(modified, styleTag) } @@ -138,6 +151,34 @@ func (p *proxyHandler) applyElementHiding(resp *http.Response, host, clientIP st return modified, true } +// mergeElementHidingCSS combines element hiding selectors from baseline and +// user RuleSets for a specific domain. User #@# exception rules suppress +// matching baseline ## selectors. Returns empty string if no CSS to inject. +func mergeElementHidingCSS(baseline, user *blocklist.ElementHiding, userRS *blocklist.RuleSet, domain string) string { + var selectors []string + + // Add baseline selectors, filtering out any excepted by user #@# rules + if baseline != nil { + for _, sel := range baseline.Selectors { + if userRS != nil && userRS.IsElementHideExcepted(sel, domain) { + continue + } + selectors = append(selectors, sel) + } + } + + // Add user selectors (their own internal exceptions already applied) + if user != nil { + selectors = append(selectors, user.Selectors...) + } + + if len(selectors) == 0 { + return "" + } + + return strings.Join(selectors, ",\n") + " {\n display: none !important;\n}\n" +} + // injectBeforeClose inserts content before the first found closing tag, // or appends if none is found. Tags are tried in order. func injectBeforeClose(htmlDoc, content []byte, tags ...[]byte) []byte { @@ -154,7 +195,7 @@ func injectBeforeClose(htmlDoc, content []byte, tags ...[]byte) []byte { // resolves to a blocked address. Other elements are passed through unchanged — // element hiding for those is handled by CSS injection only. func stripBlockedResources(src []byte, sc srcBlockContext) []byte { - if sc.rules == nil { + if sc.proxy == nil { return src } @@ -198,7 +239,7 @@ func stripBlockedResources(src []byte, sc srcBlockContext) []byte { resolved := sc.resolveSrc(urlVal) ctx := blocklist.MatchContext{PageDomain: sc.host} - if !sc.rules.ShouldBlockRequest(resolved, ctx) { + if !sc.proxy.shouldBlock(sc.clientIP, resolved, ctx) { buf.Write(rawBytes) continue } diff --git a/http.go b/http.go index 3e61515..ed07d4d 100644 --- a/http.go +++ b/http.go @@ -34,8 +34,8 @@ var hopByHopHeaders = []string{ func (p *proxyHandler) handleHTTP(w http.ResponseWriter, r *http.Request) { ctx := matchContextFromRequest(r) - rules := p.getRules() - if rules != nil && rules.ShouldBlockRequest(r.URL.String(), ctx) { + clientIP := clientIPFromRequest(r) + if p.shouldBlock(clientIP, r.URL.String(), ctx) { w.WriteHeader(http.StatusNoContent) return } diff --git a/inject.go b/inject.go index 202fb5b..65abd35 100644 --- a/inject.go +++ b/inject.go @@ -19,14 +19,14 @@ func (p *proxyHandler) bootstrapScriptTag(clientIP, host string) string { return "" } - token := p.sessions.Get(clientIP) - if token == "" { + entry := p.sessions.Get(clientIP) + if entry == nil { return "" } script := bootstrapJS script = strings.ReplaceAll(script, "__UBLPROXY_PORTAL__", p.portalOrigin) - script = strings.ReplaceAll(script, "__UBLPROXY_TOKEN__", token) + script = strings.ReplaceAll(script, "__UBLPROXY_TOKEN__", entry.Token) script = strings.ReplaceAll(script, "__UBLPROXY_HOST__", host) return "" diff --git a/inject_test.go b/inject_test.go index b4f61ec..bf91633 100644 --- a/inject_test.go +++ b/inject_test.go @@ -13,7 +13,7 @@ import ( func TestBootstrapScriptTagWithSession(t *testing.T) { sm := newSessionMap() - sm.Set("192.168.1.10", "test-token-abc") + sm.Set("192.168.1.10", sessionEntry{Token: "test-token-abc", CredentialID: "cred-1"}) p := &proxyHandler{ sessions: sm, @@ -69,7 +69,7 @@ func TestBootstrapScriptTagNoSessionMap(t *testing.T) { func TestScriptInjectionInHTML(t *testing.T) { sm := newSessionMap() - sm.Set("127.0.0.1", "my-token") + sm.Set("127.0.0.1", sessionEntry{Token: "my-token", CredentialID: "cred-1"}) p := &proxyHandler{ sessions: sm, @@ -130,7 +130,7 @@ func TestNoScriptInjectionWithoutSession(t *testing.T) { func TestNoScriptInjectionForNonHTML(t *testing.T) { sm := newSessionMap() - sm.Set("127.0.0.1", "my-token") + sm.Set("127.0.0.1", sessionEntry{Token: "my-token", CredentialID: "cred-1"}) p := &proxyHandler{ sessions: sm, @@ -151,7 +151,7 @@ func TestNoScriptInjectionForNonHTML(t *testing.T) { func TestScriptInjectionWithGzip(t *testing.T) { sm := newSessionMap() - sm.Set("127.0.0.1", "gzip-token") + sm.Set("127.0.0.1", sessionEntry{Token: "gzip-token", CredentialID: "cred-1"}) p := &proxyHandler{ sessions: sm, @@ -186,7 +186,7 @@ func TestScriptInjectionWithGzip(t *testing.T) { func TestScriptInjectionWithRules(t *testing.T) { sm := newSessionMap() - sm.Set("127.0.0.1", "rules-token") + sm.Set("127.0.0.1", sessionEntry{Token: "rules-token", CredentialID: "cred-1"}) rs := blocklist.NewRuleSet() rs.AddLine("##.ad-banner") @@ -195,7 +195,7 @@ func TestScriptInjectionWithRules(t *testing.T) { sessions: sm, portalOrigin: "https://127.0.0.1:8443", } - p.rules.Store(rs) + p.baselineRules.Store(rs) htmlBody := `
Content
` resp := &http.Response{ @@ -234,30 +234,30 @@ func TestSessionMap(t *testing.T) { sm := newSessionMap() // Get from empty map - if got := sm.Get("1.2.3.4"); got != "" { - t.Errorf("Get empty = %q, want empty", got) + if got := sm.Get("1.2.3.4"); got != nil { + t.Errorf("Get empty = %v, want nil", got) } // Set and get - sm.Set("1.2.3.4", "token-a") - if got := sm.Get("1.2.3.4"); got != "token-a" { - t.Errorf("Get = %q, want %q", got, "token-a") + sm.Set("1.2.3.4", sessionEntry{Token: "token-a", CredentialID: "cred-1"}) + if got := sm.Get("1.2.3.4"); got == nil || got.Token != "token-a" || got.CredentialID != "cred-1" { + t.Errorf("Get = %v, want token-a/cred-1", got) } // Overwrite - sm.Set("1.2.3.4", "token-b") - if got := sm.Get("1.2.3.4"); got != "token-b" { - t.Errorf("Get after overwrite = %q, want %q", got, "token-b") + sm.Set("1.2.3.4", sessionEntry{Token: "token-b", CredentialID: "cred-2"}) + if got := sm.Get("1.2.3.4"); got == nil || got.Token != "token-b" || got.CredentialID != "cred-2" { + t.Errorf("Get after overwrite = %v, want token-b/cred-2", got) } // Different IP - if got := sm.Get("5.6.7.8"); got != "" { - t.Errorf("Get different IP = %q, want empty", got) + if got := sm.Get("5.6.7.8"); got != nil { + t.Errorf("Get different IP = %v, want nil", got) } // Delete sm.Delete("1.2.3.4") - if got := sm.Get("1.2.3.4"); got != "" { - t.Errorf("Get after delete = %q, want empty", got) + if got := sm.Get("1.2.3.4"); got != nil { + t.Errorf("Get after delete = %v, want nil", got) } } diff --git a/main.go b/main.go index ba9b540..035a327 100644 --- a/main.go +++ b/main.go @@ -9,7 +9,6 @@ import ( "path/filepath" "strings" - "ublproxy/pkg/blocklist" "ublproxy/pkg/store" "ublproxy/pkg/webauthn" ) @@ -39,36 +38,30 @@ func main() { var blocklistSources stringSlice flag.Var(&blocklistSources, "blocklist", "path or URL to a blocklist file (can be specified multiple times)") + var defaultSubs stringSlice + flag.Var(&defaultSubs, "default-subscription", "default blocklist subscription URL, always active for all users (can be specified multiple times; defaults to EasyList + EasyPrivacy if none specified)") + flag.Parse() + // Built-in defaults if no --default-subscription flags were given + if len(defaultSubs) == 0 { + defaultSubs = stringSlice{ + "https://easylist.to/easylist/easylist.txt", + "https://easylist.to/easylist/easyprivacy.txt", + } + } + caCert, caKey, err := loadOrGenerateCA(*caDir) if err != nil { fmt.Fprintf(os.Stderr, "CA setup failed: %v\n", err) os.Exit(1) } - var rules *blocklist.RuleSet - if len(blocklistSources) > 0 { - rules = blocklist.NewRuleSet() - for _, src := range blocklistSources { - var loadErr error - if strings.HasPrefix(src, "http://") || strings.HasPrefix(src, "https://") { - fmt.Fprintf(os.Stderr, "Fetching %s\n", src) - loadErr = rules.LoadURL(src) - } else { - loadErr = rules.LoadFile(src) - } - if loadErr != nil { - fmt.Fprintf(os.Stderr, "failed to load blocklist: %v\n", loadErr) - os.Exit(1) - } - } - 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, rules) + handler := newProxyHandler(certs, caCertPEM) + handler.blocklistSources = blocklistSources + handler.defaultSubscriptions = defaultSubs // Open SQLite database for credential/session/rule storage db, err := store.Open(*dbPath) @@ -78,7 +71,6 @@ func main() { } defer db.Close() handler.store = db - handler.blocklistSources = blocklistSources // Configure WebAuthn with the portal hostname. WebAuthn requires a // domain name as RP ID — IP addresses are not allowed by the spec. @@ -91,11 +83,15 @@ func main() { sm := newSessionMap() api := newAPIHandler(db, webauthnCfg, sm) - api.onRulesChanged = handler.reloadRules + api.onRulesChanged = handler.invalidateUserRules handler.api = api handler.sessions = sm handler.portalOrigin = portalOrigin + // Load baseline rules (--blocklist + --default-subscription) at startup. + // This runs synchronously so rules are ready before traffic arrives. + handler.reloadBaseline() + // Auto-detect LAN IP for the portal TLS cert so it covers both the // hostname and the LAN IP (useful for CA cert download, etc.) var extraIPs []net.IP diff --git a/pkg/blocklist/elemhide.go b/pkg/blocklist/elemhide.go index a351457..e905f3b 100644 --- a/pkg/blocklist/elemhide.go +++ b/pkg/blocklist/elemhide.go @@ -8,7 +8,8 @@ import ( // ElementHiding holds the CSS for hiding ad elements on a specific domain. // All selectors are combined into a single display:none stylesheet. type ElementHiding struct { - CSS string // display:none CSS for all selectors (may be empty) + CSS string // display:none CSS for all selectors (may be empty) + Selectors []string // individual CSS selectors before joining } // ElementHideRule represents a CSS element hiding rule from an adblock filter. @@ -164,5 +165,5 @@ func (rs *RuleSet) computeElementHiding(domain string) *ElementHiding { } css := strings.Join(selectors, ",\n") + " {\n display: none !important;\n}\n" - return &ElementHiding{CSS: css} + return &ElementHiding{CSS: css, Selectors: selectors} } diff --git a/pkg/blocklist/ruleset.go b/pkg/blocklist/ruleset.go index c034668..60358b4 100644 --- a/pkg/blocklist/ruleset.go +++ b/pkg/blocklist/ruleset.go @@ -309,6 +309,56 @@ func (rs *RuleSet) isHostBlocked(host string) bool { } } +// MatchesException returns true if the URL matches any exception rule (@@) +// in this RuleSet. This is used for per-user exception checking where user +// exceptions need to override baseline blocking rules. +// Safe to call on a nil receiver (returns false). +func (rs *RuleSet) MatchesException(rawURL string, ctx MatchContext) bool { + if rs == nil { + return false + } + + lowerURL := strings.ToLower(rawURL) + lowerCtx := MatchContext{ + PageDomain: strings.ToLower(ctx.PageDomain), + ResourceType: ctx.ResourceType, + } + host := extractHostFromURL(lowerURL) + + if rs.matchDomainIndexed(rs.domainExc, host, rawURL, lowerURL, lowerCtx) { + return true + } + for _, exc := range rs.exceptions { + if exc.matchWithContextLower(rawURL, lowerURL, lowerCtx) { + return true + } + } + return false +} + +// MatchesExceptionHost returns true if the hostname matches any hostname-level +// exception rule (@@||hostname^). Used at CONNECT time where only the host is +// known. Safe to call on a nil receiver (returns false). +func (rs *RuleSet) MatchesExceptionHost(host string) bool { + if rs == nil { + return false + } + // Build a synthetic URL so the exception rule matching works + syntheticURL := "https://" + host + "/" + return rs.MatchesException(syntheticURL, MatchContext{}) +} + +// IsElementHideExcepted returns true if this RuleSet contains an element +// hiding exception (#@#) for the given CSS selector on the given domain. +// Used to let user #@# exceptions suppress baseline ## rules. +// Safe to call on a nil receiver (returns false). +func (rs *RuleSet) IsElementHideExcepted(selector, domain string) bool { + if rs == nil || rs.elemHideIdx == nil { + return false + } + return rs.elemHideIdx.isExcepted(selector, strings.ToLower(domain)) +} + // extractHostFromURL pulls the hostname from a URL string. func extractHostFromURL(url string) string { schemeEnd := strings.Index(url, "://") diff --git a/pkg/blocklist/ruleset_test.go b/pkg/blocklist/ruleset_test.go index 73985f8..d188b0a 100644 --- a/pkg/blocklist/ruleset_test.go +++ b/pkg/blocklist/ruleset_test.go @@ -229,6 +229,42 @@ func TestRuleSetExceptionHostname(t *testing.T) { } } +func TestMatchesException(t *testing.T) { + rs := blocklist.NewRuleSet() + rs.AddException("@@||safe.example.com^") + rs.AddException("@@/ads/acceptable*") + + // Hostname-level exception + if !rs.MatchesException("https://safe.example.com/page", blocklist.MatchContext{}) { + t.Error("should match hostname exception") + } + // Path-level exception + if !rs.MatchesException("https://example.com/ads/acceptable-banner.png", blocklist.MatchContext{}) { + t.Error("should match path exception") + } + // No exception for unrelated URL + if rs.MatchesException("https://ads.example.com/tracker.js", blocklist.MatchContext{}) { + t.Error("should not match unrelated URL") + } + // Nil receiver + var nilRS *blocklist.RuleSet + if nilRS.MatchesException("https://safe.example.com/", blocklist.MatchContext{}) { + t.Error("nil receiver should return false") + } +} + +func TestMatchesExceptionHost(t *testing.T) { + rs := blocklist.NewRuleSet() + rs.AddException("@@||safe.example.com^") + + if !rs.MatchesExceptionHost("safe.example.com") { + t.Error("should match excepted host") + } + if rs.MatchesExceptionHost("ads.example.com") { + t.Error("should not match non-excepted host") + } +} + func TestRuleSetLoadFileWithExceptions(t *testing.T) { content := `||ads.example.com^ /tracking.js diff --git a/pkg/store/rules.go b/pkg/store/rules.go index 108b16e..fb11219 100644 --- a/pkg/store/rules.go +++ b/pkg/store/rules.go @@ -67,6 +67,37 @@ func (s *Store) ListRules(credentialID string) ([]Rule, error) { return rules, rows.Err() } +// ListEnabledRules returns all enabled rules for a single credential. +// Used when building a per-user RuleSet. +func (s *Store) ListEnabledRules(credentialID string) ([]Rule, error) { + rows, err := s.db.Query( + "SELECT id, credential_id, rule, domain, enabled, created_at FROM rules WHERE credential_id = ? AND enabled = 1", + credentialID, + ) + if err != nil { + return nil, fmt.Errorf("list enabled rules: %w", err) + } + defer rows.Close() + + var rules []Rule + for rows.Next() { + var r Rule + var createdAt string + var enabled int + if err := rows.Scan(&r.ID, &r.CredentialID, &r.Rule, &r.Domain, &enabled, &createdAt); err != nil { + return nil, fmt.Errorf("scan rule: %w", err) + } + r.Enabled = enabled != 0 + var parseErr error + r.CreatedAt, parseErr = time.Parse("2006-01-02 15:04:05", createdAt) + if parseErr != nil { + return nil, fmt.Errorf("parse created_at: %w", parseErr) + } + rules = append(rules, r) + } + return rules, rows.Err() +} + // ListAllEnabledRules returns all enabled rules across all credentials. // Used when rebuilding the in-memory RuleSet. func (s *Store) ListAllEnabledRules() ([]Rule, error) { diff --git a/pkg/store/subscriptions.go b/pkg/store/subscriptions.go index 9d6b230..2ce27e9 100644 --- a/pkg/store/subscriptions.go +++ b/pkg/store/subscriptions.go @@ -50,6 +50,29 @@ func (s *Store) ListSubscriptions(credentialID string) ([]Subscription, error) { return scanSubscriptions(rows) } +// ListEnabledSubscriptionURLs returns all enabled subscription URLs for a +// single credential. Used when building a per-user RuleSet. +func (s *Store) ListEnabledSubscriptionURLs(credentialID string) ([]string, error) { + rows, err := s.db.Query( + "SELECT url FROM blocklist_subscriptions WHERE credential_id = ? AND enabled = 1", + credentialID, + ) + if err != nil { + return nil, fmt.Errorf("list enabled subscription urls: %w", err) + } + defer rows.Close() + + var urls []string + for rows.Next() { + var url string + if err := rows.Scan(&url); err != nil { + return nil, fmt.Errorf("scan url: %w", err) + } + urls = append(urls, url) + } + return urls, rows.Err() +} + // ListAllEnabledSubscriptionURLs returns all unique enabled subscription URLs // across all users. Used when rebuilding the in-memory RuleSet. func (s *Store) ListAllEnabledSubscriptionURLs() ([]string, error) { diff --git a/proxy.go b/proxy.go index 51b1a55..369da61 100644 --- a/proxy.go +++ b/proxy.go @@ -16,25 +16,39 @@ import ( ) type proxyHandler struct { - certs *certCache - caCertPEM []byte - rules atomic.Pointer[blocklist.RuleSet] - transport *http.Transport - store *store.Store - api *apiHandler - sessions *sessionMap + certs *certCache + caCertPEM []byte + transport *http.Transport + store *store.Store + api *apiHandler + sessions *sessionMap + portalOrigin string + // baselineRules are the always-active rules loaded from --blocklist + // sources and --default-subscription lists. These apply to all traffic + // regardless of user. + baselineRules atomic.Pointer[blocklist.RuleSet] + + // userRules caches per-user RuleSets keyed by credential ID. Each + // user's RuleSet contains their custom rules and subscription lists. + // Loaded lazily on first proxied request for that user. + userRules sync.Map // credential ID -> *blocklist.RuleSet + // blocklistSources are the static blocklist file paths/URLs loaded at - // startup. Needed to rebuild the RuleSet when user rules change. + // startup (from CLI --blocklist flags). blocklistSources []string - // reloadMu serializes rule reloads to prevent concurrent rebuilds. + // defaultSubscriptions are the default blocklist subscription URLs + // (from CLI --default-subscription flags or the built-in defaults). + defaultSubscriptions []string + + // reloadMu serializes baseline rule reloads to prevent concurrent rebuilds. reloadMu sync.Mutex } -func newProxyHandler(certs *certCache, caCertPEM []byte, rules *blocklist.RuleSet) *proxyHandler { - p := &proxyHandler{ +func newProxyHandler(certs *certCache, caCertPEM []byte) *proxyHandler { + return &proxyHandler{ certs: certs, caCertPEM: caCertPEM, // Force HTTP/1.1 upstream. Go's default HTTP/2 support causes hangs @@ -46,63 +60,157 @@ func newProxyHandler(certs *certCache, caCertPEM []byte, rules *blocklist.RuleSe IdleConnTimeout: 90 * time.Second, }, } - if rules != nil { - p.rules.Store(rules) +} + +// getBaselineRules returns the baseline RuleSet, or nil if none is loaded. +func (p *proxyHandler) getBaselineRules() *blocklist.RuleSet { + return p.baselineRules.Load() +} + +// getUserRules returns the cached per-user RuleSet for the given credential, +// loading it lazily from the database on first access. Returns nil if the +// user has no custom rules or subscriptions, or if the store is unavailable. +func (p *proxyHandler) getUserRules(credentialID string) *blocklist.RuleSet { + if credentialID == "" || p.store == nil { + return nil + } + + if cached, ok := p.userRules.Load(credentialID); ok { + return cached.(*blocklist.RuleSet) + } + + rs := p.loadUserRules(credentialID) + p.userRules.Store(credentialID, rs) + return rs +} + +// loadUserRules builds a RuleSet from a user's DB rules and subscriptions. +func (p *proxyHandler) loadUserRules(credentialID string) *blocklist.RuleSet { + rs := blocklist.NewRuleSet() + + subURLs, err := p.store.ListEnabledSubscriptionURLs(credentialID) + if err != nil { + fmt.Fprintf(os.Stderr, "user-rules: failed to load subscriptions for %s: %v\n", credentialID, err) + } else { + for _, url := range subURLs { + if err := p.loadBlocklistSource(rs, url); err != nil { + fmt.Fprintf(os.Stderr, "user-rules: failed to load subscription %s: %v\n", url, err) + } + } + } + + dbRules, err := p.store.ListEnabledRules(credentialID) + if err != nil { + fmt.Fprintf(os.Stderr, "user-rules: failed to load rules for %s: %v\n", credentialID, err) + } else { + for _, r := range dbRules { + rs.AddLine(r.Rule) + } } - return p + + return rs } -// getRules returns the current RuleSet, or nil if none is loaded. -func (p *proxyHandler) getRules() *blocklist.RuleSet { - return p.rules.Load() +// invalidateUserRules evicts the cached RuleSet for a user so it will be +// reloaded from the database on the next proxied request. +func (p *proxyHandler) invalidateUserRules(credentialID string) { + p.userRules.Delete(credentialID) +} + +// credentialForIP returns the credential ID for the authenticated session +// on the given client IP, or empty string if not authenticated. +func (p *proxyHandler) credentialForIP(clientIP string) string { + if p.sessions == nil { + return "" + } + entry := p.sessions.Get(clientIP) + if entry == nil { + return "" + } + return entry.CredentialID +} + +// shouldBlock checks whether a URL should be blocked, applying layered +// evaluation: user exceptions override baseline blocks, then user blocks +// are checked, then baseline blocks. +func (p *proxyHandler) shouldBlock(clientIP, url string, ctx blocklist.MatchContext) bool { + baseline := p.getBaselineRules() + credID := p.credentialForIP(clientIP) + userRS := p.getUserRules(credID) + + // User exception overrides baseline block + if userRS != nil && userRS.MatchesException(url, ctx) { + return false + } + + // User-specific block + if userRS != nil && userRS.ShouldBlockRequest(url, ctx) { + return true + } + + // Baseline block + if baseline != nil && baseline.ShouldBlockRequest(url, ctx) { + return true + } + + return false +} + +// shouldBlockHost checks whether a host should be blocked at the CONNECT +// level, applying layered evaluation like shouldBlock. +func (p *proxyHandler) shouldBlockHost(clientIP, host string) bool { + baseline := p.getBaselineRules() + credID := p.credentialForIP(clientIP) + userRS := p.getUserRules(credID) + + // User exception overrides baseline block + if userRS != nil && userRS.MatchesExceptionHost(host) { + return false + } + + // User-specific block + if userRS != nil && userRS.IsHostBlocked(host) { + return true + } + + // Baseline block + if baseline != nil && baseline.IsHostBlocked(host) { + return true + } + + return false } // blocklistCacheTTL is how long a cached blocklist download is considered // fresh. Stale entries are re-downloaded on next reload. const blocklistCacheTTL = 24 * time.Hour -// reloadRules rebuilds the in-memory RuleSet from static blocklists, -// user-subscribed blocklists, and user-created rules. Remote URLs are -// cached in the database to avoid re-downloading on every rule change. -func (p *proxyHandler) reloadRules() { +// reloadBaseline rebuilds the baseline RuleSet from --blocklist sources +// and --default-subscription URLs. Remote URLs are cached in the database +// to avoid re-downloading on every reload. This does not include per-user +// rules — those are loaded lazily via getUserRules. +func (p *proxyHandler) reloadBaseline() { p.reloadMu.Lock() defer p.reloadMu.Unlock() rs := blocklist.NewRuleSet() - // Reload static blocklists (CLI --blocklist flags) + // Load static blocklists (CLI --blocklist flags) for _, src := range p.blocklistSources { if err := p.loadBlocklistSource(rs, src); err != nil { - fmt.Fprintf(os.Stderr, "reload: failed to load blocklist %s: %v\n", src, err) + fmt.Fprintf(os.Stderr, "baseline: failed to load blocklist %s: %v\n", src, err) } } - if p.store != nil { - // Load user-subscribed blocklist URLs - subURLs, err := p.store.ListAllEnabledSubscriptionURLs() - if err != nil { - fmt.Fprintf(os.Stderr, "reload: failed to load subscription urls: %v\n", err) - } else { - for _, url := range subURLs { - if err := p.loadBlocklistSource(rs, url); err != nil { - fmt.Fprintf(os.Stderr, "reload: failed to load subscription %s: %v\n", url, err) - } - } - } - - // Add user-created rules from the database - dbRules, err := p.store.ListAllEnabledRules() - if err != nil { - fmt.Fprintf(os.Stderr, "reload: failed to load user rules: %v\n", err) - } else { - for _, r := range dbRules { - rs.AddLine(r.Rule) - } + // Load default subscription URLs (EasyList, EasyPrivacy, etc.) + for _, url := range p.defaultSubscriptions { + if err := p.loadBlocklistSource(rs, url); err != nil { + fmt.Fprintf(os.Stderr, "baseline: failed to load subscription %s: %v\n", url, err) } } - // Atomic swap — all concurrent readers immediately see the new rules - p.rules.Store(rs) + p.baselineRules.Store(rs) + fmt.Fprintf(os.Stderr, "baseline: loaded %d hostnames, %d URL rules\n", rs.HostCount(), rs.RuleCount()) } // loadBlocklistSource loads a blocklist from a file path or URL into the diff --git a/proxy_test.go b/proxy_test.go index b779c5f..c370389 100644 --- a/proxy_test.go +++ b/proxy_test.go @@ -56,7 +56,10 @@ func startTestEnv(t *testing.T, upstreamHandler http.Handler, rules *blocklist.R certs := newCertCache(caCert, caKey) caCertPEM := encodeCertPEM(caCert) - handler := newProxyHandler(certs, caCertPEM, rules) + handler := newProxyHandler(certs, caCertPEM) + if rules != nil { + handler.baselineRules.Store(rules) + } // Trust the test HTTPS server's certificate so the proxy can validate // upstream TLS connections (proxy validates upstream certs in production) @@ -1781,7 +1784,7 @@ func TestAllBlockableElementsTogether(t *testing.T) { // --- Hot-reload tests --- -func TestReloadRulesPicksUpUserRules(t *testing.T) { +func TestGetUserRulesLoadsFromDB(t *testing.T) { dbPath := filepath.Join(t.TempDir(), "reload.db") db, err := store.Open(dbPath) if err != nil { @@ -1801,33 +1804,36 @@ func TestReloadRulesPicksUpUserRules(t *testing.T) { if err != nil { t.Fatalf("generateCA: %v", err) } - handler := newProxyHandler(newCertCache(caCert, caKey), encodeCertPEM(caCert), nil) + handler := newProxyHandler(newCertCache(caCert, caKey), encodeCertPEM(caCert)) handler.store = db - // Before reload — no rules loaded - if rules := handler.getRules(); rules != nil { - t.Errorf("before reload: getRules() should be nil, got non-nil") - } - - // Reload rules — should pick up the user rule from the database - handler.reloadRules() - - rules := handler.getRules() + // No user rules cached yet — getUserRules lazily loads them + rules := handler.getUserRules("test-cred") if rules == nil { - t.Fatal("after reload: getRules() returned nil") + t.Fatal("getUserRules returned nil") } // The rule "/ads/*" should now be loaded ctx := blocklist.MatchContext{} if !rules.ShouldBlockRequest("http://example.com/ads/banner.gif", ctx) { - t.Error("after reload: /ads/banner.gif should be blocked") + t.Error("/ads/banner.gif should be blocked by user rules") } if rules.ShouldBlockRequest("http://example.com/page.html", ctx) { - t.Error("after reload: /page.html should not be blocked") + t.Error("/page.html should not be blocked") + } + + // Invalidate and verify re-load works + handler.invalidateUserRules("test-cred") + rules2 := handler.getUserRules("test-cred") + if rules2 == nil { + t.Fatal("getUserRules returned nil after invalidation") + } + if !rules2.ShouldBlockRequest("http://example.com/ads/banner.gif", ctx) { + t.Error("re-loaded rules should still block /ads/*") } } -func TestReloadRulesIncludesStaticBlocklists(t *testing.T) { +func TestReloadBaselineAndUserRulesSeparate(t *testing.T) { // Create a temporary blocklist file tmpDir := t.TempDir() blocklistPath := filepath.Join(tmpDir, "blocklist.txt") @@ -1840,7 +1846,7 @@ func TestReloadRulesIncludesStaticBlocklists(t *testing.T) { } t.Cleanup(func() { db.Close() }) - // Add a user rule too + // Add a user rule if err := db.SaveCredential("cred1", []byte("pubkey1")); err != nil { t.Fatalf("SaveCredential: %v", err) } @@ -1853,26 +1859,37 @@ func TestReloadRulesIncludesStaticBlocklists(t *testing.T) { if err != nil { t.Fatalf("generateCA: %v", err) } - handler := newProxyHandler(newCertCache(caCert, caKey), encodeCertPEM(caCert), nil) + handler := newProxyHandler(newCertCache(caCert, caKey), encodeCertPEM(caCert)) handler.store = db handler.blocklistSources = []string{blocklistPath} - handler.reloadRules() + // Reload baseline — should load static blocklist but NOT user rules + handler.reloadBaseline() - rules := handler.getRules() - if rules == nil { - t.Fatal("getRules() returned nil after reload") + baseline := handler.getBaselineRules() + if baseline == nil { + t.Fatal("getBaselineRules() returned nil after reloadBaseline") } - // Static blocklist rule should be loaded - if !rules.IsHostBlocked("blocked-host.example.com") { - t.Error("static blocklist hostname should be blocked") + // Static blocklist rule should be in baseline + if !baseline.IsHostBlocked("blocked-host.example.com") { + t.Error("static blocklist hostname should be blocked in baseline") } - // User rule should also be loaded (element hiding rule) - eh := rules.ElementHidingForDomain("example.com") + // User element hiding rule should NOT be in baseline + eh := baseline.ElementHidingForDomain("example.com") + if eh != nil && eh.CSS != "" { + t.Error("user element hiding rule should NOT be in baseline") + } + + // User rules should be loadable separately + userRules := handler.getUserRules("cred1") + if userRules == nil { + t.Fatal("getUserRules returned nil") + } + eh = userRules.ElementHidingForDomain("example.com") if eh == nil || eh.CSS == "" { - t.Error("user element hiding rule should be loaded") + t.Error("user element hiding rule should be loaded in user rules") } } @@ -1888,7 +1905,7 @@ func TestAPIRuleMutationTriggersReload(t *testing.T) { api := newAPIHandler(db, testAPIConfig, sm) var reloadCount atomic.Int32 - api.onRulesChanged = func() { + api.onRulesChanged = func(credentialID string) { reloadCount.Add(1) } diff --git a/session_map.go b/session_map.go index c2b39bd..d37ee0d 100644 --- a/session_map.go +++ b/session_map.go @@ -2,37 +2,48 @@ package main import "sync" -// sessionMap maps client IP addresses to their active session tokens. +// sessionEntry holds the session token and credential ID for an +// authenticated client. The credential ID identifies the user (passkey) +// and is used to look up per-user rules. +type sessionEntry struct { + Token string + CredentialID string +} + +// sessionMap maps client IP addresses to their active sessions. // When a user authenticates on the portal (via passkey), their session -// token is associated with their IP. The proxy uses this to embed the -// correct token in the bootstrap script injected into HTML responses. +// is associated with their IP. The proxy uses this to embed the +// correct token in the bootstrap script and to resolve per-user rules. type sessionMap struct { - mu sync.RWMutex - tokens map[string]string // client IP -> session token + mu sync.RWMutex + entries map[string]sessionEntry // client IP -> session } func newSessionMap() *sessionMap { - return &sessionMap{tokens: make(map[string]string)} + return &sessionMap{entries: make(map[string]sessionEntry)} } -// Set associates a session token with a client IP. -func (m *sessionMap) Set(clientIP, token string) { +// Set associates a session with a client IP. +func (m *sessionMap) Set(clientIP string, entry sessionEntry) { m.mu.Lock() - m.tokens[clientIP] = token + m.entries[clientIP] = entry m.mu.Unlock() } -// Get returns the session token for a client IP, or empty string if none. -func (m *sessionMap) Get(clientIP string) string { +// Get returns the session entry for a client IP, or nil if none. +func (m *sessionMap) Get(clientIP string) *sessionEntry { m.mu.RLock() - token := m.tokens[clientIP] + entry, ok := m.entries[clientIP] m.mu.RUnlock() - return token + if !ok { + return nil + } + return &entry } // Delete removes the session for a client IP. func (m *sessionMap) Delete(clientIP string) { m.mu.Lock() - delete(m.tokens, clientIP) + delete(m.entries, clientIP) m.mu.Unlock() }