From 6b3e1a32aefc4a6d71a4b0ee3e4dcdc2f1bfcc81 Mon Sep 17 00:00:00 2001 From: David Fifield Date: Sun, 1 Aug 2021 22:25:02 -0600 Subject: [PATCH] Make noise.NewServer take only the private key. --- dnstt-server/main.go | 17 ++++++++--------- noise/noise.go | 7 +++++-- noise/noise_test.go | 2 +- 3 files changed, 14 insertions(+), 12 deletions(-) diff --git a/dnstt-server/main.go b/dnstt-server/main.go index fddca11..2da5827 100644 --- a/dnstt-server/main.go +++ b/dnstt-server/main.go @@ -229,9 +229,9 @@ func handleStream(stream *smux.Stream, upstream string, conv uint32) error { // acceptStreams wraps a KCP session in a Noise channel and an smux.Session, // then awaits smux streams. It passes each stream to handleStream. -func acceptStreams(conn *kcp.UDPSession, privkey, pubkey []byte, upstream string) error { +func acceptStreams(conn *kcp.UDPSession, privkey []byte, upstream string) error { // Put a Noise channel on top of the KCP conn. - rw, err := noise.NewServer(conn, privkey, pubkey) + rw, err := noise.NewServer(conn, privkey) if err != nil { return err } @@ -270,7 +270,7 @@ func acceptStreams(conn *kcp.UDPSession, privkey, pubkey []byte, upstream string // acceptSessions listens for incoming KCP connections and passes them to // acceptStreams. -func acceptSessions(ln *kcp.Listener, privkey, pubkey []byte, mtu int, upstream string) error { +func acceptSessions(ln *kcp.Listener, privkey []byte, mtu int, upstream string) error { for { conn, err := ln.AcceptKCP() if err != nil { @@ -298,7 +298,7 @@ func acceptSessions(ln *kcp.Listener, privkey, pubkey []byte, mtu int, upstream log.Printf("end session %08x", conn.GetConv()) conn.Close() }() - err := acceptStreams(conn, privkey, pubkey, upstream) + err := acceptStreams(conn, privkey, upstream) if err != nil { log.Printf("session %08x acceptStreams: %v", conn.GetConv(), err) } @@ -746,10 +746,10 @@ func computeMaxEncodedPayload(limit int) int { return low } -func run(privkey, pubkey []byte, domain dns.Name, upstream string, dnsConn net.PacketConn) error { +func run(privkey []byte, domain dns.Name, upstream string, dnsConn net.PacketConn) error { defer dnsConn.Close() - log.Printf("pubkey %x", pubkey) + log.Printf("pubkey %x", noise.PubkeyFromPrivkey(privkey)) // We have a variable amount of room in which to encode downstream // packets in each response, because each response must contain the @@ -777,7 +777,7 @@ func run(privkey, pubkey []byte, domain dns.Name, upstream string, dnsConn net.P } defer ln.Close() go func() { - err := acceptSessions(ln, privkey, pubkey, mtu, upstream) + err := acceptSessions(ln, privkey, mtu, upstream) if err != nil { log.Printf("acceptSessions: %v", err) } @@ -923,9 +923,8 @@ Example: os.Exit(1) } } - pubkey := noise.PubkeyFromPrivkey(privkey) - err = run(privkey, pubkey, domain, upstream, dnsConn) + err = run(privkey, domain, upstream, dnsConn) if err != nil { log.Fatal(err) } diff --git a/noise/noise.go b/noise/noise.go index 2b7267a..c19b4a9 100644 --- a/noise/noise.go +++ b/noise/noise.go @@ -177,10 +177,13 @@ func NewClient(rwc io.ReadWriteCloser, serverPubkey []byte) (io.ReadWriteCloser, // NewClient wraps an io.ReadWriteCloser in a Noise protocol as a server, and // returns after completing the handshake. It returns a non-nil error if there // is an error during the handshake. -func NewServer(rwc io.ReadWriteCloser, serverPrivkey, serverPubkey []byte) (io.ReadWriteCloser, error) { +func NewServer(rwc io.ReadWriteCloser, serverPrivkey []byte) (io.ReadWriteCloser, error) { config := newConfig() config.Initiator = false - config.StaticKeypair = noise.DHKey{Private: serverPrivkey, Public: serverPubkey} + config.StaticKeypair = noise.DHKey{ + Private: serverPrivkey, + Public: PubkeyFromPrivkey(serverPrivkey), + } handshakeState, err := noise.NewHandshakeState(config) if err != nil { return nil, err diff --git a/noise/noise_test.go b/noise/noise_test.go index fab5df9..b25fc99 100644 --- a/noise/noise_test.go +++ b/noise/noise_test.go @@ -158,7 +158,7 @@ func TestUnexpectedPayload(t *testing.T) { } }() - server, err := NewServer(s, privkey, pubkey) + server, err := NewServer(s, privkey) if err == nil || err.Error() != "unexpected client payload" || server != nil { t.Errorf("NewServer got (%T, %v)", server, err) }