[READ-ONLY] Mirror of https://github.com/andrioid/ublproxy.
andrioid.github.io/ublproxy/
adblock
adblock-plus-list
adblocker
privacy-tools
proxy-server
self-hosted
1package main
2
3import (
4 "net"
5 "testing"
6 "time"
7
8 "github.com/miekg/dns"
9
10 "ublproxy/internal/blocklist"
11)
12
13// startTestDNSServer starts a DNS server on a random UDP port and returns
14// its address. The caller must call shutdown to stop the server.
15func startTestDNSServer(t *testing.T, ds *dnsServer) (addr string, shutdown func()) {
16 t.Helper()
17
18 pc, err := net.ListenPacket("udp", "127.0.0.1:0")
19 if err != nil {
20 t.Fatal(err)
21 }
22 addr = pc.LocalAddr().String()
23
24 srv := &dns.Server{PacketConn: pc, Handler: ds}
25 go func() { _ = srv.ActivateAndServe() }()
26
27 // Wait for the server to be ready.
28 deadline := time.Now().Add(2 * time.Second)
29 for time.Now().Before(deadline) {
30 m := new(dns.Msg)
31 m.SetQuestion("test.invalid.", dns.TypeA)
32 if _, err := dns.Exchange(m, addr); err == nil {
33 break
34 }
35 time.Sleep(10 * time.Millisecond)
36 }
37
38 return addr, func() { _ = srv.Shutdown() }
39}
40
41// startMockUpstream starts a mock DNS server that answers A queries with
42// the given IP address.
43func startMockUpstream(t *testing.T, answerIP string) (addr string, shutdown func()) {
44 t.Helper()
45
46 pc, err := net.ListenPacket("udp", "127.0.0.1:0")
47 if err != nil {
48 t.Fatal(err)
49 }
50 addr = pc.LocalAddr().String()
51
52 handler := dns.HandlerFunc(func(w dns.ResponseWriter, r *dns.Msg) {
53 m := new(dns.Msg)
54 m.SetReply(r)
55 if len(r.Question) > 0 && r.Question[0].Qtype == dns.TypeA {
56 m.Answer = append(m.Answer, &dns.A{
57 Hdr: dns.RR_Header{
58 Name: r.Question[0].Name,
59 Rrtype: dns.TypeA,
60 Class: dns.ClassINET,
61 Ttl: 300,
62 },
63 A: net.ParseIP(answerIP),
64 })
65 }
66 _ = w.WriteMsg(m)
67 })
68
69 srv := &dns.Server{PacketConn: pc, Handler: handler}
70 go func() { _ = srv.ActivateAndServe() }()
71
72 deadline := time.Now().Add(2 * time.Second)
73 for time.Now().Before(deadline) {
74 m := new(dns.Msg)
75 m.SetQuestion("test.invalid.", dns.TypeA)
76 if _, err := dns.Exchange(m, addr); err == nil {
77 break
78 }
79 time.Sleep(10 * time.Millisecond)
80 }
81
82 return addr, func() { _ = srv.Shutdown() }
83}
84
85func newTestProxyHandler() *proxyHandler {
86 return &proxyHandler{
87 sessions: newSessionMap(),
88 }
89}
90
91func TestDNSBlockedHostReturnsNullIP(t *testing.T) {
92 proxy := newTestProxyHandler()
93 rs := blocklist.NewRuleSet()
94 rs.AddLine("||ads.example.com^")
95 proxy.baselineRules.Store(rs)
96
97 upstreamAddr, upstreamShutdown := startMockUpstream(t, "93.184.216.34")
98 defer upstreamShutdown()
99
100 ds := &dnsServer{
101 proxy: proxy,
102 upstream: upstreamAddr,
103 }
104 addr, shutdown := startTestDNSServer(t, ds)
105 defer shutdown()
106
107 // A record for blocked host should return 0.0.0.0
108 m := new(dns.Msg)
109 m.SetQuestion("ads.example.com.", dns.TypeA)
110 resp, err := dns.Exchange(m, addr)
111 if err != nil {
112 t.Fatal(err)
113 }
114 if len(resp.Answer) != 1 {
115 t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
116 }
117 a, ok := resp.Answer[0].(*dns.A)
118 if !ok {
119 t.Fatalf("expected A record, got %T", resp.Answer[0])
120 }
121 if !a.A.Equal(net.IPv4zero) {
122 t.Errorf("expected 0.0.0.0, got %s", a.A)
123 }
124}
125
126func TestDNSBlockedHostAAAAReturnsNullIPv6(t *testing.T) {
127 proxy := newTestProxyHandler()
128 rs := blocklist.NewRuleSet()
129 rs.AddLine("||ads.example.com^")
130 proxy.baselineRules.Store(rs)
131
132 upstreamAddr, upstreamShutdown := startMockUpstream(t, "93.184.216.34")
133 defer upstreamShutdown()
134
135 ds := &dnsServer{
136 proxy: proxy,
137 upstream: upstreamAddr,
138 }
139 addr, shutdown := startTestDNSServer(t, ds)
140 defer shutdown()
141
142 // AAAA record for blocked host should return ::
143 m := new(dns.Msg)
144 m.SetQuestion("ads.example.com.", dns.TypeAAAA)
145 resp, err := dns.Exchange(m, addr)
146 if err != nil {
147 t.Fatal(err)
148 }
149 if len(resp.Answer) != 1 {
150 t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
151 }
152 aaaa, ok := resp.Answer[0].(*dns.AAAA)
153 if !ok {
154 t.Fatalf("expected AAAA record, got %T", resp.Answer[0])
155 }
156 if !aaaa.AAAA.Equal(net.IPv6zero) {
157 t.Errorf("expected ::, got %s", aaaa.AAAA)
158 }
159}
160
161func TestDNSNonBlockedHostForwardsUpstream(t *testing.T) {
162 proxy := newTestProxyHandler()
163 rs := blocklist.NewRuleSet()
164 rs.AddLine("||ads.example.com^")
165 proxy.baselineRules.Store(rs)
166
167 upstreamAddr, upstreamShutdown := startMockUpstream(t, "93.184.216.34")
168 defer upstreamShutdown()
169
170 ds := &dnsServer{
171 proxy: proxy,
172 upstream: upstreamAddr,
173 }
174 addr, shutdown := startTestDNSServer(t, ds)
175 defer shutdown()
176
177 // Non-blocked host should be forwarded to upstream
178 m := new(dns.Msg)
179 m.SetQuestion("example.com.", dns.TypeA)
180 resp, err := dns.Exchange(m, addr)
181 if err != nil {
182 t.Fatal(err)
183 }
184 if len(resp.Answer) != 1 {
185 t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
186 }
187 a, ok := resp.Answer[0].(*dns.A)
188 if !ok {
189 t.Fatalf("expected A record, got %T", resp.Answer[0])
190 }
191 if !a.A.Equal(net.ParseIP("93.184.216.34")) {
192 t.Errorf("expected 93.184.216.34, got %s", a.A)
193 }
194}
195
196func TestDNSDomainHierarchyBlocking(t *testing.T) {
197 proxy := newTestProxyHandler()
198 rs := blocklist.NewRuleSet()
199 rs.AddLine("||example.com^")
200 proxy.baselineRules.Store(rs)
201
202 upstreamAddr, upstreamShutdown := startMockUpstream(t, "1.2.3.4")
203 defer upstreamShutdown()
204
205 ds := &dnsServer{
206 proxy: proxy,
207 upstream: upstreamAddr,
208 }
209 addr, shutdown := startTestDNSServer(t, ds)
210 defer shutdown()
211
212 // Subdomain should also be blocked (domain hierarchy walk)
213 m := new(dns.Msg)
214 m.SetQuestion("sub.example.com.", dns.TypeA)
215 resp, err := dns.Exchange(m, addr)
216 if err != nil {
217 t.Fatal(err)
218 }
219 if len(resp.Answer) != 1 {
220 t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
221 }
222 a := resp.Answer[0].(*dns.A)
223 if !a.A.Equal(net.IPv4zero) {
224 t.Errorf("expected 0.0.0.0 for subdomain of blocked host, got %s", a.A)
225 }
226}
227
228func TestDNSPerUserExceptionOverridesBaseline(t *testing.T) {
229 proxy := newTestProxyHandler()
230
231 // Baseline blocks ads.example.com
232 baseline := blocklist.NewRuleSet()
233 baseline.AddLine("||ads.example.com^")
234 proxy.baselineRules.Store(baseline)
235
236 // User has an exception for ads.example.com
237 userRS := blocklist.NewRuleSet()
238 userRS.AddLine("@@||ads.example.com^")
239 proxy.userRules.Store("user-cred-1", userRS)
240
241 // Register the user's session from a specific IP
242 proxy.sessions.Set("127.0.0.1", sessionEntry{
243 CredentialID: "user-cred-1",
244 })
245
246 upstreamAddr, upstreamShutdown := startMockUpstream(t, "93.184.216.34")
247 defer upstreamShutdown()
248
249 ds := &dnsServer{
250 proxy: proxy,
251 upstream: upstreamAddr,
252 }
253 addr, shutdown := startTestDNSServer(t, ds)
254 defer shutdown()
255
256 // This user's exception should override the baseline block
257 m := new(dns.Msg)
258 m.SetQuestion("ads.example.com.", dns.TypeA)
259 resp, err := dns.Exchange(m, addr)
260 if err != nil {
261 t.Fatal(err)
262 }
263 if len(resp.Answer) != 1 {
264 t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
265 }
266 a, ok := resp.Answer[0].(*dns.A)
267 if !ok {
268 t.Fatalf("expected A record, got %T", resp.Answer[0])
269 }
270 // Should NOT be 0.0.0.0 — the user excepted this host
271 if a.A.Equal(net.IPv4zero) {
272 t.Error("expected upstream response (user exception), got 0.0.0.0")
273 }
274}
275
276func TestDNSActivityLogging(t *testing.T) {
277 proxy := newTestProxyHandler()
278 rs := blocklist.NewRuleSet()
279 rs.AddLine("||blocked.example.com^")
280 proxy.baselineRules.Store(rs)
281
282 activity := NewActivityLog(100)
283
284 upstreamAddr, upstreamShutdown := startMockUpstream(t, "1.2.3.4")
285 defer upstreamShutdown()
286
287 ds := &dnsServer{
288 proxy: proxy,
289 upstream: upstreamAddr,
290 activity: activity,
291 }
292 addr, shutdown := startTestDNSServer(t, ds)
293 defer shutdown()
294
295 // Query a blocked host
296 m := new(dns.Msg)
297 m.SetQuestion("blocked.example.com.", dns.TypeA)
298 if _, err := dns.Exchange(m, addr); err != nil {
299 t.Fatal(err)
300 }
301
302 // Query a non-blocked host
303 m2 := new(dns.Msg)
304 m2.SetQuestion("allowed.example.com.", dns.TypeA)
305 if _, err := dns.Exchange(m2, addr); err != nil {
306 t.Fatal(err)
307 }
308
309 entries := activity.Recent(10)
310 if len(entries) < 1 {
311 t.Fatal("expected at least 1 activity entry")
312 }
313
314 // Most recent should be the blocked query
315 var foundBlock bool
316 for _, e := range entries {
317 if e.Type == ActivityDNSBlocked && e.Host == "blocked.example.com" {
318 foundBlock = true
319 break
320 }
321 }
322 if !foundBlock {
323 t.Errorf("expected dns-blocked activity entry for blocked.example.com, got: %+v", entries)
324 }
325}
326
327func TestDNSNonAddressQtypeForwardedRegardless(t *testing.T) {
328 proxy := newTestProxyHandler()
329 rs := blocklist.NewRuleSet()
330 rs.AddLine("||blocked.example.com^")
331 proxy.baselineRules.Store(rs)
332
333 // Mock upstream that answers MX queries
334 pc, err := net.ListenPacket("udp", "127.0.0.1:0")
335 if err != nil {
336 t.Fatal(err)
337 }
338 upstreamAddr := pc.LocalAddr().String()
339 handler := dns.HandlerFunc(func(w dns.ResponseWriter, r *dns.Msg) {
340 m := new(dns.Msg)
341 m.SetReply(r)
342 if len(r.Question) > 0 && r.Question[0].Qtype == dns.TypeMX {
343 m.Answer = append(m.Answer, &dns.MX{
344 Hdr: dns.RR_Header{
345 Name: r.Question[0].Name,
346 Rrtype: dns.TypeMX,
347 Class: dns.ClassINET,
348 Ttl: 300,
349 },
350 Preference: 10,
351 Mx: "mail.example.com.",
352 })
353 }
354 _ = w.WriteMsg(m)
355 })
356 srv := &dns.Server{PacketConn: pc, Handler: handler}
357 go func() { _ = srv.ActivateAndServe() }()
358 defer func() { _ = srv.Shutdown() }()
359
360 ds := &dnsServer{
361 proxy: proxy,
362 upstream: upstreamAddr,
363 }
364 addr, shutdown := startTestDNSServer(t, ds)
365 defer shutdown()
366
367 // MX query for a blocked host should still be blocked (A/AAAA only get null-routed)
368 // For non-address types on a blocked host, we return an empty answer
369 // to prevent the host from being reachable via other record types.
370 m := new(dns.Msg)
371 m.SetQuestion("blocked.example.com.", dns.TypeMX)
372 resp, err := dns.Exchange(m, addr)
373 if err != nil {
374 t.Fatal(err)
375 }
376 // Blocked host: should get empty response, not forwarded to upstream
377 if len(resp.Answer) != 0 {
378 t.Errorf("expected empty answer for non-A/AAAA query on blocked host, got %d answers", len(resp.Answer))
379 }
380}
381
382func TestDNSTCPSupport(t *testing.T) {
383 proxy := newTestProxyHandler()
384 rs := blocklist.NewRuleSet()
385 rs.AddLine("||ads.example.com^")
386 proxy.baselineRules.Store(rs)
387
388 upstreamAddr, upstreamShutdown := startMockUpstream(t, "93.184.216.34")
389 defer upstreamShutdown()
390
391 ds := &dnsServer{
392 proxy: proxy,
393 upstream: upstreamAddr,
394 }
395
396 // Start a TCP DNS server
397 ln, err := net.Listen("tcp", "127.0.0.1:0")
398 if err != nil {
399 t.Fatal(err)
400 }
401 tcpAddr := ln.Addr().String()
402
403 srv := &dns.Server{Listener: ln, Handler: ds, Net: "tcp"}
404 go func() { _ = srv.ActivateAndServe() }()
405 defer func() { _ = srv.Shutdown() }()
406
407 // Wait for TCP server to be ready
408 deadline := time.Now().Add(2 * time.Second)
409 for time.Now().Before(deadline) {
410 c := new(dns.Client)
411 c.Net = "tcp"
412 m := new(dns.Msg)
413 m.SetQuestion("test.invalid.", dns.TypeA)
414 if _, _, err := c.Exchange(m, tcpAddr); err == nil {
415 break
416 }
417 time.Sleep(10 * time.Millisecond)
418 }
419
420 // Query blocked host over TCP
421 c := new(dns.Client)
422 c.Net = "tcp"
423 m := new(dns.Msg)
424 m.SetQuestion("ads.example.com.", dns.TypeA)
425 resp, _, err := c.Exchange(m, tcpAddr)
426 if err != nil {
427 t.Fatal(err)
428 }
429 if len(resp.Answer) != 1 {
430 t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
431 }
432 a, ok := resp.Answer[0].(*dns.A)
433 if !ok {
434 t.Fatalf("expected A record, got %T", resp.Answer[0])
435 }
436 if !a.A.Equal(net.IPv4zero) {
437 t.Errorf("expected 0.0.0.0 over TCP, got %s", a.A)
438 }
439}
440
441func TestDNSEmptyQuestion(t *testing.T) {
442 proxy := newTestProxyHandler()
443 ds := &dnsServer{
444 proxy: proxy,
445 upstream: "127.0.0.1:0",
446 }
447 addr, shutdown := startTestDNSServer(t, ds)
448 defer shutdown()
449
450 // Send a message with no question section
451 m := new(dns.Msg)
452 m.Id = dns.Id()
453 resp, err := dns.Exchange(m, addr)
454 if err != nil {
455 t.Fatal(err)
456 }
457 if resp.Rcode != dns.RcodeFormatError {
458 t.Errorf("expected FORMERR for empty question, got rcode %d", resp.Rcode)
459 }
460}
461
462func TestDNSHostsFileFormatBlocking(t *testing.T) {
463 proxy := newTestProxyHandler()
464 rs := blocklist.NewRuleSet()
465 rs.AddLine("0.0.0.0 malware.example.com")
466 proxy.baselineRules.Store(rs)
467
468 upstreamAddr, upstreamShutdown := startMockUpstream(t, "1.2.3.4")
469 defer upstreamShutdown()
470
471 ds := &dnsServer{
472 proxy: proxy,
473 upstream: upstreamAddr,
474 }
475 addr, shutdown := startTestDNSServer(t, ds)
476 defer shutdown()
477
478 m := new(dns.Msg)
479 m.SetQuestion("malware.example.com.", dns.TypeA)
480 resp, err := dns.Exchange(m, addr)
481 if err != nil {
482 t.Fatal(err)
483 }
484 if len(resp.Answer) != 1 {
485 t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
486 }
487 a := resp.Answer[0].(*dns.A)
488 if !a.A.Equal(net.IPv4zero) {
489 t.Errorf("expected 0.0.0.0 for hosts-file blocked entry, got %s", a.A)
490 }
491}
492
493func TestDNSNoBaselineRulesForwardsAll(t *testing.T) {
494 proxy := newTestProxyHandler()
495 // No baseline rules loaded
496
497 upstreamAddr, upstreamShutdown := startMockUpstream(t, "1.2.3.4")
498 defer upstreamShutdown()
499
500 ds := &dnsServer{
501 proxy: proxy,
502 upstream: upstreamAddr,
503 }
504 addr, shutdown := startTestDNSServer(t, ds)
505 defer shutdown()
506
507 m := new(dns.Msg)
508 m.SetQuestion("anything.example.com.", dns.TypeA)
509 resp, err := dns.Exchange(m, addr)
510 if err != nil {
511 t.Fatal(err)
512 }
513 if len(resp.Answer) != 1 {
514 t.Fatalf("expected 1 answer, got %d", len(resp.Answer))
515 }
516 a := resp.Answer[0].(*dns.A)
517 if !a.A.Equal(net.ParseIP("1.2.3.4")) {
518 t.Errorf("expected upstream IP 1.2.3.4, got %s", a.A)
519 }
520}
521
522// Verify multiple questions are handled (edge case from legacy DNS clients).
523func TestDNSMultipleQuestions(t *testing.T) {
524 proxy := newTestProxyHandler()
525 rs := blocklist.NewRuleSet()
526 rs.AddLine("||blocked.example.com^")
527 proxy.baselineRules.Store(rs)
528
529 upstreamAddr, upstreamShutdown := startMockUpstream(t, "1.2.3.4")
530 defer upstreamShutdown()
531
532 ds := &dnsServer{
533 proxy: proxy,
534 upstream: upstreamAddr,
535 }
536 addr, shutdown := startTestDNSServer(t, ds)
537 defer shutdown()
538
539 // DNS message with multiple questions (unusual but valid)
540 m := new(dns.Msg)
541 m.Id = dns.Id()
542 m.RecursionDesired = true
543 m.Question = []dns.Question{
544 {Name: "blocked.example.com.", Qtype: dns.TypeA, Qclass: dns.ClassINET},
545 {Name: "allowed.example.com.", Qtype: dns.TypeA, Qclass: dns.ClassINET},
546 }
547 resp, err := dns.Exchange(m, addr)
548 if err != nil {
549 // Some implementations reject multi-question; that's acceptable
550 return
551 }
552 // We only process the first question; response is valid as long as no panic
553 _ = resp
554}