Make noise.NewServer take only the private key.

This commit is contained in:
David Fifield
2021-08-01 22:30:49 -06:00
parent 3bb75782f6
commit 6b3e1a32ae
3 changed files with 14 additions and 12 deletions
+8 -9
View File
@@ -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)
}
+5 -2
View File
@@ -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
+1 -1
View File
@@ -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)
}