commit 88787149cf89ee12e50156d2be193b2194e4052f
parent b3a0993e1f1aff1ad84b2e0cd680e5697bfd17dc
Author: Oliver Lowe <o@olowe.co>
Date: Wed, 17 Nov 2021 18:05:31 +1100
Seperate packet handling over UDP and TCP better
Because it's a stream versus just 1 packet so there's some tricky
handling that should be seperate for clarity
Diffstat:
| A | cmd/dohproxy/config.go | | | 47 | +++++++++++++++++++++++++++++++++++++++++++++++ |
| M | dns.go | | | 63 | ++++++++++++++++++++++++++++++++++++++------------------------- |
2 files changed, 85 insertions(+), 25 deletions(-)
diff --git a/cmd/dohproxy/config.go b/cmd/dohproxy/config.go
@@ -0,0 +1,47 @@
+package main
+
+import (
+ "bufio"
+ "fmt"
+ "io"
+ "os"
+ "strings"
+)
+
+type config struct {
+ forwardaddr string
+ listenaddr string
+}
+
+func configFromFile(name string) (config, error) {
+ f, err := os.Open(name)
+ if err != nil {
+ return config{}, err
+ }
+ defer f.Close()
+ return parseConfig(f)
+}
+
+func parseConfig(r io.Reader) (config, error) {
+ sc := bufio.NewScanner(r)
+ var c config
+ for sc.Scan() {
+ line := strings.TrimSpace(sc.Text())
+ if strings.HasPrefix(line, "#") {
+ continue // skip config comments
+ }
+ fields := strings.Fields(line)
+ if len(fields) > 2 {
+ return c, fmt.Errorf("too many values for key %s", fields[0])
+ }
+ switch k := fields[0]; k {
+ case "listen":
+ c.listenaddr = fields[1]
+ case "forward":
+ c.forwardaddr = fields[1]
+ default:
+ return c, fmt.Errorf("unknown key %s", k)
+ }
+ }
+ return c, nil
+}
diff --git a/dns.go b/dns.go
@@ -4,7 +4,6 @@ import (
"crypto/tls"
"encoding/binary"
"fmt"
- "io"
"net"
"golang.org/x/net/dns/dnsmessage"
@@ -42,36 +41,50 @@ func send(msg dnsmessage.Message, conn net.Conn) (dnsmessage.Message, error) {
return dnsmessage.Message{}, err
}
if _, ok := conn.(net.PacketConn); ok {
- if _, err = conn.Write(packed); err != nil {
- return dnsmessage.Message{}, err
+ b, err := dnsPacketExchange(packed, conn)
+ if err != nil {
+ return dnsmessage.Message{}, fmt.Errorf("exchange DNS packet: %v", err)
}
} else {
- // DNS over TCP requires you to prepend the message with a
- // 2-octet length field.
- m := make([]byte, 2+len(packed))
- binary.BigEndian.PutUint16(m, uint16(len(packed)))
- copy(m[2:], packed)
- if _, err = conn.Write(m); err != nil {
- return dnsmessage.Message{}, err
- }
+ b, err := dnsStreamExchange(packed, conn)
+ if err != nil {
+ return dnsmessage.Message{}, fmt.Errorf("exchange DNS TCP stream: %v", err)
+ }
+ var rmsg dnsmessage.Message
+ if err := rmsg.Unpack(b); err != nil {
+ return dnsmessage.Message{}, fmt.Errorf("parse response: %v", err)
+ }
+ return rmsg, nil
+}
+
+func dnsPacketExchange(b []byte, conn net.Conn) ([]byte, error) {
+ if _, err := conn.Write(b); err != nil {
+ return nil, err
+ }
+ buf := make([]byte, 512) // max UDP size per RFC?
+ n, err := conn.Read(buf)
+ if err != nil {
+ return nil, err
+ }
+ return buf[:n], nil
+}
+
+func dnsStreamExchange(b []byte, conn net.Conn) ([]byte, error) {
+ // DNS over TCP requires you to prepend the message with a
+ // 2-octet length field.
+ m := make([]byte, 2+len(b))
+ binary.BigEndian.PutUint16(m, uint16(len(b)))
+ copy(m[2:], b)
+ if _, err := conn.Write(m); err != nil {
+ return nil, err
}
buf := make([]byte, 1024)
n, err := conn.Read(buf)
- if err != nil && err != io.EOF {
- return dnsmessage.Message{}, err
+ if err != nil {
+ return nil, err
}
if n == 0 {
- return dnsmessage.Message{}, fmt.Errorf("empty response")
+ return nil, fmt.Errorf("empty response")
}
- var rmsg dnsmessage.Message
- if _, ok := conn.(net.PacketConn); ok {
- if err := rmsg.Unpack(buf[:n]); err != nil {
- return dnsmessage.Message{}, fmt.Errorf("parse response: %v", err)
- }
- } else {
- if err := rmsg.Unpack(buf[2:n]); err != nil {
- return dnsmessage.Message{}, fmt.Errorf("parse response: %v", err)
- }
- }
- return rmsg, nil
+ return buf[2:n], nil
}