mirror of
https://github.com/WiiLink24/wfc-server.git
synced 2026-09-30 19:16:36 -05:00
NAS: Rework TLS proxy to use io.Copy
This commit is contained in:
116
nas/listener.go
116
nas/listener.go
@@ -29,16 +29,29 @@ type nasIOConn struct {
|
||||
out io.Writer
|
||||
}
|
||||
|
||||
func (c nasInConn) Read(b []byte) (int, error) {
|
||||
return c.in.Read(b)
|
||||
func (c nasInConn) Read(b []byte) (int, error) { return c.in.Read(b) }
|
||||
|
||||
func (c nasInConn) Close() error {
|
||||
err := c.Conn.Close()
|
||||
if closer, ok := c.in.(io.Closer); ok {
|
||||
_ = closer.Close()
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func (c nasIOConn) Read(b []byte) (int, error) {
|
||||
return c.in.Read(b)
|
||||
}
|
||||
func (c nasIOConn) Read(b []byte) (int, error) { return c.in.Read(b) }
|
||||
|
||||
func (c nasIOConn) Write(b []byte) (int, error) {
|
||||
return c.out.Write(b)
|
||||
func (c nasIOConn) Write(b []byte) (int, error) { return c.out.Write(b) }
|
||||
|
||||
func (c nasIOConn) Close() error {
|
||||
// Don't actually close the underlying connection here
|
||||
if closer, ok := c.in.(io.Closer); ok {
|
||||
_ = closer.Close()
|
||||
}
|
||||
if closer, ok := c.out.(io.Closer); ok {
|
||||
_ = closer.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (l *httpListener) Accept() (net.Conn, error) {
|
||||
@@ -50,12 +63,11 @@ func (l *httpListener) Accept() (net.Conn, error) {
|
||||
pr, pw := io.Pipe()
|
||||
go func() {
|
||||
r := bufio.NewReader(conn)
|
||||
_, err := pw.Write(filterDuplicateHost(r))
|
||||
if err != nil {
|
||||
if err := filterDuplicateHost(pw, r); err != nil {
|
||||
_ = pw.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
_, err = io.Copy(pw, r)
|
||||
_, err := io.Copy(pw, r)
|
||||
_ = pw.CloseWithError(err)
|
||||
}()
|
||||
return &nasInConn{
|
||||
@@ -64,40 +76,6 @@ func (l *httpListener) Accept() (net.Conn, error) {
|
||||
}, nil
|
||||
}
|
||||
|
||||
// filterDuplicateHost wraps a net.Conn and filters out duplicate Host headers from the HTTP request,
|
||||
// making the invalid requests sent by DWC acceptable to the standard library's HTTP server.
|
||||
func filterDuplicateHost(r *bufio.Reader) []byte {
|
||||
// Read the first line of the HTTP request
|
||||
line, err := r.ReadString('\n')
|
||||
if err != nil {
|
||||
return []byte(line)
|
||||
}
|
||||
// Is this an HTTP request?
|
||||
if !strings.HasSuffix(line, "HTTP/1.1\r\n") {
|
||||
return []byte(line)
|
||||
}
|
||||
|
||||
// Iterate through the HTTP headers and remove any duplicate Host headers
|
||||
var headers bytes.Buffer
|
||||
hostSeen := false
|
||||
for {
|
||||
headerLine, err := r.ReadString('\n')
|
||||
if err != nil || headerLine == "\r\n" {
|
||||
headers.WriteString(headerLine)
|
||||
break
|
||||
}
|
||||
if strings.HasPrefix(strings.ToLower(headerLine), "host:") {
|
||||
if hostSeen {
|
||||
continue
|
||||
}
|
||||
hostSeen = true
|
||||
}
|
||||
headers.WriteString(headerLine)
|
||||
}
|
||||
|
||||
return []byte(line + headers.String())
|
||||
}
|
||||
|
||||
func listenAndServe() {
|
||||
logging.Notice("NAS", "Starting HTTP server on", aurora.BrightCyan(server.Addr))
|
||||
|
||||
@@ -125,9 +103,21 @@ func (l *tlsListener) Accept() (net.Conn, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// We're gonna need like 3 pipes here, from client -> tls -> filter -> server
|
||||
readFromServer, writeFromServer := io.Pipe()
|
||||
readToServer, writeToServer := io.Pipe()
|
||||
go handleIncomingTLS(conn, readFromServer, writeToServer)
|
||||
readToFilter, writeToFilter := io.Pipe()
|
||||
|
||||
go func() {
|
||||
r := bufio.NewReader(readToFilter)
|
||||
if err := filterDuplicateHost(writeToServer, r); err != nil {
|
||||
_ = writeToServer.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
_, err := io.Copy(writeToServer, r)
|
||||
_ = writeToServer.CloseWithError(err)
|
||||
}()
|
||||
go handleIncomingTLS(conn, readFromServer, writeToFilter)
|
||||
return &nasIOConn{
|
||||
Conn: conn,
|
||||
in: readToServer,
|
||||
@@ -151,3 +141,39 @@ func listenAndServeTLS() {
|
||||
panic(err)
|
||||
}
|
||||
}
|
||||
|
||||
// filterDuplicateHost wraps a net.Conn and filters out duplicate Host headers from the HTTP request,
|
||||
// making the invalid requests sent by DWC acceptable to the standard library's HTTP server.
|
||||
func filterDuplicateHost(w *io.PipeWriter, r *bufio.Reader) error {
|
||||
// Read the first line of the HTTP request
|
||||
line, err := r.ReadString('\n')
|
||||
_, errPipe := w.Write([]byte(line))
|
||||
if err != nil || errPipe != nil {
|
||||
return errPipe
|
||||
}
|
||||
|
||||
// Is this an HTTP request we handle?
|
||||
if !strings.HasSuffix(line, "HTTP/1.1\r\n") {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Iterate through the HTTP headers and remove any duplicate Host headers
|
||||
var headers bytes.Buffer
|
||||
hostSeen := false
|
||||
for {
|
||||
headerLine, err := r.ReadString('\n')
|
||||
_, wErr := w.Write([]byte(headerLine))
|
||||
if err != nil || wErr != nil || headerLine == "\r\n" || headerLine == "\n" {
|
||||
headers.WriteString(headerLine)
|
||||
break
|
||||
}
|
||||
if strings.HasPrefix(strings.ToLower(headerLine), "host:") {
|
||||
if hostSeen {
|
||||
continue
|
||||
}
|
||||
hostSeen = true
|
||||
}
|
||||
headers.WriteString(headerLine)
|
||||
}
|
||||
return errPipe
|
||||
}
|
||||
|
||||
321
nas/tls.go
321
nas/tls.go
@@ -19,7 +19,6 @@ import (
|
||||
"io"
|
||||
"net"
|
||||
"os"
|
||||
"strings"
|
||||
"time"
|
||||
"wwfc/common"
|
||||
"wwfc/logging"
|
||||
@@ -233,18 +232,18 @@ func handleIncomingTLS(rawConn net.Conn, fromServer *io.PipeReader, toServer *io
|
||||
moduleName := "NAS-TLS:" + rawConn.RemoteAddr().String()
|
||||
|
||||
var err error
|
||||
defer func() {
|
||||
_ = rawConn.Close()
|
||||
toServer.CloseWithError(err)
|
||||
}()
|
||||
|
||||
closed := false
|
||||
if rsaKeyWii == nil && rsaKeyDS == nil {
|
||||
// Only handle real TLS requests
|
||||
err = handleRealTLS(moduleName, rawConn, fromServer, toServer)
|
||||
err, closed = proxyRealTLS(moduleName, rawConn, fromServer, toServer)
|
||||
} else {
|
||||
err = handleTLS(moduleName, rawConn, fromServer, toServer)
|
||||
err, closed = handleTLS(moduleName, rawConn, fromServer, toServer)
|
||||
}
|
||||
|
||||
toServer.CloseWithError(err)
|
||||
if !closed {
|
||||
_ = rawConn.Close()
|
||||
}
|
||||
}
|
||||
|
||||
func peekBytes(r *bufio.Reader, expected []byte) (bool, error) {
|
||||
@@ -260,7 +259,7 @@ func peekBytes(r *bufio.Reader, expected []byte) (bool, error) {
|
||||
}
|
||||
|
||||
// handleTLS handles the TLS request from the Wii or the DS. It may call handleRealTLS if the request is from a modern web browser.
|
||||
func handleTLS(moduleName string, rawConn net.Conn, fromServer *io.PipeReader, toServer *io.PipeWriter) error {
|
||||
func handleTLS(moduleName string, rawConn net.Conn, fromServer *io.PipeReader, toServer *io.PipeWriter) (error, bool) {
|
||||
// Recover from panics
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
@@ -275,101 +274,161 @@ func handleTLS(moduleName string, rawConn net.Conn, fromServer *io.PipeReader, t
|
||||
}
|
||||
|
||||
// Read client hello
|
||||
if rsaKeyWii != nil || rsaKeyDS != nil {
|
||||
var peekMatchWii, peekMatchDS bool
|
||||
var err error
|
||||
if rsaKeyWii != nil {
|
||||
peekMatchWii, err = peekBytes(r, []byte{
|
||||
0x80, 0x2B, 0x01, 0x03, 0x01, 0x00, 0x12, 0x00, 0x00, 0x00, 0x10, 0x00,
|
||||
0x00, 0x35, 0x00, 0x00, 0x2F, 0x00, 0x00, 0x0A, 0x00, 0x00, 0x09, 0x00,
|
||||
0x00, 0x05, 0x00, 0x00, 0x04,
|
||||
})
|
||||
}
|
||||
if rsaKeyDS != nil && !peekMatchWii && err == nil {
|
||||
peekMatchDS, err = peekBytes(r, []byte{
|
||||
0x16, 0x03, 0x00, 0x00, 0x2F, 0x01, 0x00, 0x00, 0x2B, 0x03, 0x00,
|
||||
})
|
||||
}
|
||||
|
||||
switch {
|
||||
case err != nil:
|
||||
return err
|
||||
case peekMatchWii:
|
||||
macFn, cipher, clientCipher, err := handleWiiTLSHandshake(moduleName, conn)
|
||||
if err == nil && macFn != nil && cipher != nil && clientCipher != nil {
|
||||
err = proxyConsoleTLS(moduleName, conn, fromServer, toServer, VersionTLS10, macFn, cipher, clientCipher)
|
||||
}
|
||||
return err
|
||||
case peekMatchDS:
|
||||
macFn, cipher, clientCipher, err := handleDSSSLHandshake(moduleName, conn)
|
||||
if err == nil && macFn != nil && cipher != nil && clientCipher != nil {
|
||||
err = proxyConsoleTLS(moduleName, conn, fromServer, toServer, VersionSSL30, macFn, cipher, clientCipher)
|
||||
}
|
||||
return err
|
||||
}
|
||||
var peekMatchWii, peekMatchDS bool
|
||||
var err error
|
||||
if rsaKeyWii != nil {
|
||||
peekMatchWii, err = peekBytes(r, []byte{
|
||||
0x80, 0x2B, 0x01, 0x03, 0x01, 0x00, 0x12, 0x00, 0x00, 0x00, 0x10, 0x00,
|
||||
0x00, 0x35, 0x00, 0x00, 0x2F, 0x00, 0x00, 0x0A, 0x00, 0x00, 0x09, 0x00,
|
||||
0x00, 0x05, 0x00, 0x00, 0x04,
|
||||
})
|
||||
}
|
||||
if rsaKeyDS != nil && !peekMatchWii && err == nil {
|
||||
peekMatchDS, err = peekBytes(r, []byte{
|
||||
0x16, 0x03, 0x00, 0x00, 0x2F, 0x01, 0x00, 0x00, 0x2B, 0x03, 0x00,
|
||||
})
|
||||
}
|
||||
|
||||
_ = conn.SetDeadline(time.Now().UTC().Add(25 * time.Second))
|
||||
var macFn macFunction
|
||||
var cipher, clientCipher *rc4.Cipher
|
||||
var version uint16
|
||||
switch {
|
||||
case err != nil:
|
||||
return err, false
|
||||
|
||||
return handleRealTLS(moduleName, conn, fromServer, toServer)
|
||||
case peekMatchWii:
|
||||
macFn, cipher, clientCipher, err = handleWiiTLSHandshake(moduleName, conn)
|
||||
version = VersionTLS10
|
||||
case peekMatchDS:
|
||||
macFn, cipher, clientCipher, err = handleDSSSLHandshake(moduleName, conn)
|
||||
version = VersionSSL30
|
||||
|
||||
default:
|
||||
return proxyRealTLS(moduleName, conn, fromServer, toServer)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return err, false
|
||||
}
|
||||
if macFn == nil || cipher == nil || clientCipher == nil {
|
||||
return errors.New("invalid TLS handshake result"), false
|
||||
}
|
||||
|
||||
if err == nil && macFn != nil && cipher != nil && clientCipher != nil {
|
||||
conn := &consoleTLSConn{
|
||||
Conn: conn,
|
||||
MacFn: macFn,
|
||||
Cipher: cipher,
|
||||
ClientCipher: clientCipher,
|
||||
Version: version,
|
||||
}
|
||||
return proxyConsoleTLS(moduleName, conn, fromServer, toServer)
|
||||
}
|
||||
return err, false
|
||||
}
|
||||
|
||||
// handleRealTLS handles the TLS request legitimately using crypto/tls
|
||||
func handleRealTLS(moduleName string, rawConn net.Conn, fromServer *io.PipeReader, toServer *io.PipeWriter) error {
|
||||
// Recover from panics
|
||||
defer func() {
|
||||
if r := recover(); r != nil {
|
||||
logging.Error(moduleName, "Panic:", r)
|
||||
if r.(error) != nil {
|
||||
toServer.CloseWithError(r.(error))
|
||||
}
|
||||
}
|
||||
}()
|
||||
// proxyRealTLS handles the TLS request legitimately using crypto/tls
|
||||
func proxyRealTLS(moduleName string, rawConn net.Conn, fromServer *io.PipeReader, toServer *io.PipeWriter) (err error, closed bool) {
|
||||
defer toServer.CloseWithError(err)
|
||||
|
||||
if realTLSConfig == nil {
|
||||
return errors.New("realTLSConfig is not set")
|
||||
return errors.New("realTLSConfig is not set"), false
|
||||
}
|
||||
|
||||
tlsConn := tls.Server(rawConn, realTLSConfig)
|
||||
|
||||
err := tlsConn.Handshake()
|
||||
err = tlsConn.Handshake()
|
||||
if err != nil {
|
||||
return err
|
||||
_ = tlsConn.Close()
|
||||
return err, false
|
||||
}
|
||||
|
||||
// Read bytes from the HTTP server and forward them through the TLS connection
|
||||
go func() {
|
||||
recvBuf := make([]byte, 0x100)
|
||||
|
||||
for {
|
||||
n, err := fromServer.Read(recvBuf)
|
||||
if err != nil {
|
||||
toServer.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
|
||||
_, err = tlsConn.Write(recvBuf[:n])
|
||||
if err != nil {
|
||||
toServer.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
_, err = io.Copy(tlsConn, fromServer)
|
||||
logging.Info(moduleName, "Closed connection from server")
|
||||
_ = tlsConn.Close()
|
||||
}()
|
||||
_, err = io.Copy(toServer, tlsConn)
|
||||
return err, true
|
||||
}
|
||||
|
||||
// Read encrypted content from the client and forward it to the HTTP server
|
||||
buf := make([]byte, 0x1000)
|
||||
for {
|
||||
n, err := tlsConn.Read(buf)
|
||||
if err != nil {
|
||||
return err
|
||||
type consoleTLSConn struct {
|
||||
net.Conn
|
||||
BufferSize int
|
||||
MacFn macFunction
|
||||
Cipher *rc4.Cipher
|
||||
ClientCipher *rc4.Cipher
|
||||
Version uint16
|
||||
seq uint64
|
||||
decodedBuffer bytes.Buffer
|
||||
encodedBuffer bytes.Buffer
|
||||
}
|
||||
|
||||
func (c *consoleTLSConn) Read(b []byte) (n int, err error) {
|
||||
recordLength := uint16(0)
|
||||
for len(b) > c.decodedBuffer.Len() {
|
||||
for c.encodedBuffer.Len() < int(recordLength)+5 {
|
||||
_, err = c.encodedBuffer.ReadFrom(c.Conn)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
}
|
||||
|
||||
_, err = toServer.Write(buf[:n])
|
||||
if err != nil {
|
||||
toServer.CloseWithError(err)
|
||||
return err
|
||||
buf := c.encodedBuffer.Bytes()
|
||||
if buf[0] < 0x15 || buf[0] > 0x17 {
|
||||
return 0, errors.New("invalid record type")
|
||||
}
|
||||
|
||||
if buf[1] != 0x03 || (c.Version == VersionTLS10 && buf[2] != 0x01) || (c.Version == VersionSSL30 && buf[2] != 0x00) {
|
||||
return 0, errors.New("invalid TLS version")
|
||||
}
|
||||
|
||||
recordLength = binary.BigEndian.Uint16(buf[3:5])
|
||||
if recordLength < 17 || (recordLength+5) > 0x1000 {
|
||||
return 0, errors.New("invalid record length")
|
||||
}
|
||||
|
||||
if c.encodedBuffer.Len() < int(recordLength)+5 {
|
||||
continue
|
||||
}
|
||||
|
||||
// Decrypt content
|
||||
c.ClientCipher.XORKeyStream(buf[5:5+recordLength], buf[5:5+recordLength])
|
||||
|
||||
if buf[0] != 0x17 {
|
||||
if buf[0] == 0x15 || buf[5] == 0x01 || buf[6] == 0x00 {
|
||||
// Alert: connection closed
|
||||
if err = c.Close(); err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return 0, io.EOF
|
||||
}
|
||||
return 0, errors.New("non-application data received")
|
||||
}
|
||||
// Write the decrypted content to the buffer
|
||||
c.decodedBuffer.Write(buf[5 : 5+recordLength-16])
|
||||
|
||||
c.encodedBuffer.Next(5 + int(recordLength))
|
||||
}
|
||||
|
||||
_, err = c.decodedBuffer.Read(b)
|
||||
return len(b), err
|
||||
}
|
||||
|
||||
func (c *consoleTLSConn) Write(b []byte) (n int, err error) {
|
||||
var record []byte
|
||||
record, c.seq = encryptTLS(c.MacFn, c.Cipher, b, c.seq, []byte{0x17, 0x03, 0x01, byte(len(b) >> 8), byte(len(b))})
|
||||
return c.Conn.Write(record)
|
||||
}
|
||||
|
||||
func proxyConsoleTLS(moduleName string, conn *consoleTLSConn, fromServer *io.PipeReader, toServer *io.PipeWriter) (err error, closed bool) {
|
||||
go func() {
|
||||
_, err = io.Copy(conn, fromServer)
|
||||
logging.Info(moduleName, "Closed connection from server")
|
||||
_ = conn.Close()
|
||||
}()
|
||||
_, err = io.Copy(toServer, conn)
|
||||
return err, true
|
||||
}
|
||||
|
||||
func handleWiiTLSHandshake(moduleName string, conn nasInConn) (macFn macFunction, cipher *rc4.Cipher, clientCipher *rc4.Cipher, err error) {
|
||||
@@ -729,102 +788,6 @@ func handleDSSSLHandshake(moduleName string, conn nasInConn) (macFn macFunction,
|
||||
return
|
||||
}
|
||||
|
||||
func proxyConsoleTLS(moduleName string, conn nasInConn, fromServer *io.PipeReader, toServer *io.PipeWriter, version uint16, macFn macFunction, cipher *rc4.Cipher, clientCipher *rc4.Cipher) error {
|
||||
// Read bytes from the HTTP server and forward them through the TLS connection
|
||||
go func() {
|
||||
recvBuf := make([]byte, 0x100)
|
||||
|
||||
seq := uint64(1)
|
||||
for {
|
||||
n, err := fromServer.Read(recvBuf)
|
||||
if err != nil {
|
||||
logging.Error(moduleName, "Failed to read from HTTP server:", err)
|
||||
toServer.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
|
||||
// fmt.Printf("Sent:\n% X ", recvBuf[:n])
|
||||
var record []byte
|
||||
record, seq = encryptTLS(macFn, cipher, recvBuf[:n], seq, []byte{0x17, 0x03, 0x01, byte(n >> 8), byte(n)})
|
||||
|
||||
_, err = conn.Write(record)
|
||||
if err != nil {
|
||||
logging.Error(moduleName, "Failed to write to client:", err)
|
||||
toServer.CloseWithError(err)
|
||||
return
|
||||
}
|
||||
}
|
||||
}()
|
||||
|
||||
// Read encrypted content from the client and forward it to the HTTP server
|
||||
index := 0
|
||||
total := 0
|
||||
buf := make([]byte, 0x1000)
|
||||
for {
|
||||
n, err := conn.Read(buf[index:])
|
||||
if err != nil {
|
||||
if errors.Is(err, io.EOF) || strings.Contains(err.Error(), "use of closed network connection") {
|
||||
logging.Info(moduleName, "Connection closed by client after", aurora.BrightCyan(total), "bytes")
|
||||
return err
|
||||
}
|
||||
|
||||
logging.Error(moduleName, "Failed to read from client:", err)
|
||||
return err
|
||||
}
|
||||
|
||||
// fmt.Printf("Received:\n% X ", buf[index:index+n])
|
||||
index += n
|
||||
total += n
|
||||
|
||||
for index >= 5 {
|
||||
if buf[0] < 0x15 || buf[0] > 0x17 {
|
||||
logging.Error(moduleName, "Invalid record type")
|
||||
return errors.New("invalid record type")
|
||||
}
|
||||
|
||||
if buf[1] != 0x03 || (version == VersionTLS10 && buf[2] != 0x01) || (version == VersionSSL30 && buf[2] != 0x00) {
|
||||
logging.Error(moduleName, "Invalid TLS version")
|
||||
return errors.New("invalid TLS version")
|
||||
}
|
||||
|
||||
recordLength := binary.BigEndian.Uint16(buf[3:5])
|
||||
if recordLength < 17 || (recordLength+5) > 0x1000 {
|
||||
logging.Error(moduleName, "Invalid record length")
|
||||
return errors.New("invalid record length")
|
||||
}
|
||||
|
||||
if index < int(recordLength)+5 {
|
||||
break
|
||||
}
|
||||
|
||||
// Decrypt content
|
||||
clientCipher.XORKeyStream(buf[5:5+recordLength], buf[5:5+recordLength])
|
||||
// fmt.Printf("\nDecrypted content:\n% X \n", buf[5:5+recordLength])
|
||||
|
||||
if buf[0] != 0x17 {
|
||||
if buf[0] == 0x15 || buf[5] == 0x01 || buf[6] == 0x00 {
|
||||
logging.Info(moduleName, "Alert connection close by client after", aurora.BrightCyan(total), "bytes")
|
||||
return nil
|
||||
}
|
||||
|
||||
logging.Error(moduleName, "Non-application data received:", aurora.Cyan(fmt.Sprintf("% X ", buf[:5+recordLength])))
|
||||
return errors.New("non-application data received")
|
||||
} else {
|
||||
// Send the decrypted content to the HTTP server
|
||||
_, err = toServer.Write(buf[5 : 5+recordLength-16])
|
||||
if err != nil {
|
||||
logging.Error(moduleName, "Failed to write to HTTP server:", err)
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
buf = buf[5+recordLength:]
|
||||
buf = append(buf, make([]byte, 0x1000-len(buf))...)
|
||||
index -= 5 + int(recordLength)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The following functions are modified from the crypto standard library
|
||||
//
|
||||
// Copyright (c) 2009 The Go Authors. All rights reserved.
|
||||
|
||||
Reference in New Issue
Block a user