NAS: Rework TLS proxy to use io.Copy

This commit is contained in:
Palapeli
2026-04-09 05:19:07 -04:00
parent 24d2a40634
commit 00ee95a042
2 changed files with 213 additions and 224 deletions

View File

@@ -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
}

View File

@@ -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.