diff --git a/agent/PSProxy.Agent/PSProxy.Agent.cs b/agent/PSProxy.Agent/PSProxy.Agent.cs index c209a1a..e8db9b8 100644 --- a/agent/PSProxy.Agent/PSProxy.Agent.cs +++ b/agent/PSProxy.Agent/PSProxy.Agent.cs @@ -199,6 +199,7 @@ namespace PSProxy.Agent udp.Send(payload, payload.Length); IPEndPoint ep = null; byte[] resp = udp.Receive(ref ep); + if (IsDnsTruncated(resp)) resp = QueryDnsOverTcp(host, dstPort, payload); SendFrame(new Frame(sid, FrameDNSReply, resp)); } } @@ -209,6 +210,35 @@ namespace PSProxy.Agent } } + private static bool IsDnsTruncated(byte[] resp) + { + return resp != null && resp.Length >= 4 && (resp[2] & 0x02) != 0; + } + + private static byte[] QueryDnsOverTcp(string host, int dstPort, byte[] payload) + { + using (var tcp = new TcpClient()) + { + tcp.NoDelay = true; + ConnectWithTimeout(tcp, host, dstPort, 5000); + tcp.ReceiveTimeout = 5000; + tcp.SendTimeout = 5000; + using (NetworkStream stream = tcp.GetStream()) + { + byte[] len = new byte[2]; + WriteU16BE(len, 0, payload.Length); + stream.Write(len, 0, len.Length); + stream.Write(payload, 0, payload.Length); + stream.Flush(); + + byte[] respLen = ReadExact(stream, 2); + int n = ReadU16BE(respLen, 0); + if (n == 0) throw new IOException("empty DNS TCP response"); + return ReadExact(stream, n); + } + } + } + private void PumpTargetToServer(StreamCtx ctx) { byte[] buf = new byte[32768]; @@ -308,12 +338,17 @@ namespace PSProxy.Agent } private byte[] ReadExact(int n) + { + return ReadExact(tls, n); + } + + private static byte[] ReadExact(Stream stream, int n) { byte[] b = new byte[n]; int off = 0; while (off < n) { - int got = tls.Read(b, off, n - off); + int got = stream.Read(b, off, n - off); if (got <= 0) throw new EndOfStreamException(); off += got; } @@ -384,6 +419,8 @@ namespace PSProxy.Agent private static int ReadLen(byte[] b, ref int o) { if (o >= b.Length) throw new CryptographicException("asn1 len"); int v = b[o++]; if ((v & 0x80) == 0) return v; int n = v & 0x7f; if (n < 1 || n > 4 || o + n > b.Length) throw new CryptographicException("asn1 len"); int len = 0; for (int i = 0; i < n; i++) len = (len << 8) | b[o++]; return len; } private static void WriteI32BE(byte[] b, int o, int v) { b[o] = (byte)(v >> 24); b[o + 1] = (byte)(v >> 16); b[o + 2] = (byte)(v >> 8); b[o + 3] = (byte)v; } private static int ReadI32BE(byte[] b, int o) { return ((int)b[o] << 24) | ((int)b[o + 1] << 16) | ((int)b[o + 2] << 8) | (int)b[o + 3]; } + private static void WriteU16BE(byte[] b, int o, int v) { b[o] = (byte)(v >> 8); b[o + 1] = (byte)v; } + private static int ReadU16BE(byte[] b, int o) { return ((int)b[o] << 8) | (int)b[o + 1]; } private static void WriteU64BE(byte[] b, int o, ulong v) { for (int i = 7; i >= 0; i--) { b[o + i] = (byte)v; v >>= 8; } } private static ulong ReadU64BE(byte[] b, int o) { ulong v = 0; for (int i = 0; i < 8; i++) v = (v << 8) | b[o + i]; return v; } } diff --git a/cmd/psproxy-server/main.go b/cmd/psproxy-server/main.go index c81c154..f9c2189 100644 --- a/cmd/psproxy-server/main.go +++ b/cmd/psproxy-server/main.go @@ -10,6 +10,7 @@ import ( "crypto/tls" "crypto/x509" "encoding/base64" + "encoding/binary" "encoding/hex" "encoding/json" "encoding/pem" @@ -447,8 +448,18 @@ func serveDNSRelay(listenAddr string, server *TunnelServer) { if err != nil { log.Fatalf("dns relay listen failed: %v", err) } - log.Printf("dns relay listening on udp://%s", listenAddr) - buf := make([]byte, 4096) + ln, err := net.Listen("tcp", listenAddr) + if err != nil { + _ = conn.Close() + log.Fatalf("dns tcp relay listen failed: %v", err) + } + log.Printf("dns relay listening on udp://%s and tcp://%s", listenAddr, listenAddr) + go serveDNSTCP(ln, server) + serveDNSUDP(conn, server) +} + +func serveDNSUDP(conn *net.UDPConn, server *TunnelServer) { + buf := make([]byte, 65535) for { n, client, err := conn.ReadFromUDP(buf) if err != nil { @@ -467,6 +478,53 @@ func serveDNSRelay(listenAddr string, server *TunnelServer) { } } +func serveDNSTCP(ln net.Listener, server *TunnelServer) { + for { + conn, err := ln.Accept() + if err != nil { + log.Printf("dns tcp relay accept failed: %v", err) + return + } + go handleDNSTCPConn(conn, server) + } +} + +func handleDNSTCPConn(conn net.Conn, server *TunnelServer) { + defer conn.Close() + for { + var lenb [2]byte + if _, err := io.ReadFull(conn, lenb[:]); err != nil { + if !errors.Is(err, io.EOF) && !errors.Is(err, io.ErrUnexpectedEOF) { + log.Printf("dns tcp relay read length failed: %v", err) + } + return + } + n := int(binary.BigEndian.Uint16(lenb[:])) + if n == 0 { + return + } + query := make([]byte, n) + if _, err := io.ReadFull(conn, query); err != nil { + log.Printf("dns tcp relay read query failed: %v", err) + return + } + resp, err := server.QueryDNS(query) + if err != nil { + log.Printf("dns tcp relay query failed: %v", err) + return + } + if len(resp) > 65535 { + log.Printf("dns tcp relay response too large: %d", len(resp)) + return + } + binary.BigEndian.PutUint16(lenb[:], uint16(len(resp))) + if _, err := conn.Write(append(lenb[:], resp...)); err != nil { + log.Printf("dns tcp relay write failed: %v", err) + return + } + } +} + func serveTransparentRelay(listenAddr string, server *TunnelServer) { ln, err := net.Listen("tcp4", listenAddr) if err != nil { diff --git a/cmd/psproxy-server/main_test.go b/cmd/psproxy-server/main_test.go index fee9c52..5947b4a 100644 --- a/cmd/psproxy-server/main_test.go +++ b/cmd/psproxy-server/main_test.go @@ -9,6 +9,8 @@ import ( "crypto/sha256" "crypto/x509" "encoding/base64" + "encoding/binary" + "io" "net" "strings" "testing" @@ -85,6 +87,55 @@ func TestTunnelServerQueryDNS(t *testing.T) { } } +func TestServeDNSTCPForwardsLengthPrefixedQuery(t *testing.T) { + server := NewTunnelServer(staging.NewStore(time.Minute), 2) + agent, peer := net.Pipe() + defer peer.Close() + sess := NewAgentSession(agent, bufio.NewReader(agent), server.maxStreams) + defer sess.Close() + server.SetSession(sess) + go sess.Run() + + ln, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + defer ln.Close() + go serveDNSTCP(ln, server) + + go func() { + f, err := protocol.ReadFrame(peer) + if err != nil { + return + } + if string(f.Payload) != "dns-query" { + return + } + _ = protocol.WriteFrame(peer, protocol.Frame{StreamID: f.StreamID, Type: protocol.FrameDNSReply, Payload: []byte("dns-response")}) + }() + + conn, err := net.Dial("tcp", ln.Addr().String()) + if err != nil { + t.Fatal(err) + } + defer conn.Close() + var lenb [2]byte + binary.BigEndian.PutUint16(lenb[:], uint16(len("dns-query"))) + if _, err := conn.Write(append(lenb[:], []byte("dns-query")...)); err != nil { + t.Fatal(err) + } + if _, err := io.ReadFull(conn, lenb[:]); err != nil { + t.Fatal(err) + } + resp := make([]byte, binary.BigEndian.Uint16(lenb[:])) + if _, err := io.ReadFull(conn, resp); err != nil { + t.Fatal(err) + } + if string(resp) != "dns-response" { + t.Fatalf("unexpected dns tcp response: %q", resp) + } +} + func TestValidateRoutes(t *testing.T) { if err := validateRoutes([]string{"10.0.0.0/24", "192.168.1.10/32"}); err != nil { t.Fatalf("valid routes rejected: %v", err)