diff --git a/nas/listener.go b/nas/listener.go index b367254..eaeb920 100644 --- a/nas/listener.go +++ b/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 +} diff --git a/nas/tls.go b/nas/tls.go index a624028..c44464a 100644 --- a/nas/tls.go +++ b/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.