diff --git a/internal/tests/integration/javascript_test.go b/internal/tests/integration/javascript_test.go index ae00651973..e796a0e7f8 100644 --- a/internal/tests/integration/javascript_test.go +++ b/internal/tests/integration/javascript_test.go @@ -41,6 +41,7 @@ var jsTestcases = []integrationCase{ {Path: "protocols/javascript/vnc-pass-brute.yaml", TestCase: &javascriptVncPassBrute{}, DisableOn: javascriptDockerDisabled, Serial: true}, {Path: "protocols/javascript/postgres-pass-brute.yaml", TestCase: &javascriptPostgresPassBrute{}, DisableOn: javascriptDockerDisabled, Serial: true}, {Path: "protocols/javascript/mysql-connect.yaml", TestCase: &javascriptMySQLConnect{}, DisableOn: javascriptDockerDisabled, Serial: true}, + {Path: "protocols/javascript/mysql-fingerprint.yaml", TestCase: &javascriptMySQLFingerprint{}, DisableOn: javascriptDockerDisabled, Serial: true}, {Path: "protocols/javascript/multi-ports.yaml", TestCase: &javascriptMultiPortsSSH{}}, {Path: "protocols/javascript/no-port-args.yaml", TestCase: &javascriptNoPortArgs{}}, {Path: "protocols/javascript/telnet-auth-test.yaml", TestCase: &javascriptTelnetAuthTest{}, DisableOn: javascriptDockerDisabled, Serial: true}, @@ -176,6 +177,57 @@ func (j *javascriptMySQLConnect) Execute(filePath string) error { }, javascriptDatabaseReadyTimeout, 0, mysqlReadyCheck("root", "secret"))) } +type javascriptMySQLFingerprint struct{} + +// Execute fingerprints multiple MySQL/MariaDB server versions and asserts the +// extended handshake fields (protocol, version, salt, capabilities, auth plugin). +func (j *javascriptMySQLFingerprint) Execute(filePath string) error { + cases := []struct { + name string + repository string + tag string + env []string + }{ + { + name: "mysql-5.7", + repository: "mysql", + tag: "5.7", + env: []string{"MYSQL_ROOT_PASSWORD=secret"}, + }, + { + name: "mysql-8.0", + repository: "mysql", + tag: "8.0", + env: []string{"MYSQL_ROOT_PASSWORD=secret"}, + }, + { + name: "mysql-8.4", + repository: "mysql", + tag: "8.4", + env: []string{"MYSQL_ROOT_PASSWORD=secret"}, + }, + { + name: "mariadb-11.4", + repository: "mariadb", + tag: "11.4", + env: []string{"MARIADB_ROOT_PASSWORD=secret"}, + }, + } + + var errs []error + for _, tc := range cases { + err := runJavascriptDockerCase(filePath, newJavascriptDockerSpec("3306/tcp", &dockertest.RunOptions{ + Repository: tc.repository, + Tag: tc.tag, + Env: tc.env, + }, javascriptDatabaseReadyTimeout, 0, mysqlReadyCheck("root", "secret"))) + if err != nil { + errs = append(errs, fmt.Errorf("%s: %w", tc.name, err)) + } + } + return multierr.Combine(errs...) +} + type javascriptMultiPortsSSH struct{} func (j *javascriptMultiPortsSSH) Execute(filePath string) error { diff --git a/internal/tests/integration/testdata/protocols/javascript/mysql-fingerprint.yaml b/internal/tests/integration/testdata/protocols/javascript/mysql-fingerprint.yaml new file mode 100644 index 0000000000..c7731052c2 --- /dev/null +++ b/internal/tests/integration/testdata/protocols/javascript/mysql-fingerprint.yaml @@ -0,0 +1,45 @@ +id: mysql-fingerprint + +info: + name: MySQL Fingerprint Test + author: pdteam + severity: info + +javascript: + - pre-condition: | + isPortOpen(Host, Port) + code: | + const mysql = require('nuclei/mysql'); + const client = new mysql.MySQLClient(); + const isMySQL = client.IsMySQL(Host, Port); + const info = client.FingerprintMySQL(Host, Port); + const ok = isMySQL && info && info.ProtocolVersion === 10 && info.Version && info.Version.length > 0 && info.Salt && info.Salt.length > 0 && info.CapabilityFlags > 0 && info.Debug && info.Debug.PacketType === 'handshake'; + Export({ + ok: ok, + isMySQL: isMySQL, + packetType: info.Debug.PacketType, + protocolVersion: info.ProtocolVersion, + version: info.Version, + threadId: info.ThreadID, + authPluginName: info.AuthPluginName, + capabilityFlags: info.CapabilityFlags, + capabilities: info.Capabilities, + status: info.Status, + salt: info.Salt, + }); + args: + Host: "{{Host}}" + Port: "3306" + + matchers: + - type: dsl + dsl: + - "success == true" + - "ok == true" + - "isMySQL == true" + - 'packetType == "handshake"' + - "protocolVersion == 10" + - "len(version) > 0" + - "capabilityFlags > 0" + - "len(salt) > 0" + - "threadId > 0" diff --git a/pkg/js/generated/go/libmysql/mysql.go b/pkg/js/generated/go/libmysql/mysql.go index 48549d7ecf..8b653e54f6 100644 --- a/pkg/js/generated/go/libmysql/mysql.go +++ b/pkg/js/generated/go/libmysql/mysql.go @@ -20,9 +20,10 @@ func init() { // Var and consts // Objects / Classes - "MySQLClient": gojs.GetClassConstructor[lib_mysql.MySQLClient](&lib_mysql.MySQLClient{}), - "MySQLInfo": gojs.GetClassConstructor[lib_mysql.MySQLInfo](&lib_mysql.MySQLInfo{}), - "MySQLOptions": gojs.GetClassConstructor[lib_mysql.MySQLOptions](&lib_mysql.MySQLOptions{}), + "HandshakeInfo": gojs.GetClassConstructor[lib_mysql.HandshakeInfo](&lib_mysql.HandshakeInfo{}), + "MySQLClient": gojs.GetClassConstructor[lib_mysql.MySQLClient](&lib_mysql.MySQLClient{}), + "MySQLInfo": gojs.GetClassConstructor[lib_mysql.MySQLInfo](&lib_mysql.MySQLInfo{}), + "MySQLOptions": gojs.GetClassConstructor[lib_mysql.MySQLOptions](&lib_mysql.MySQLOptions{}), }, ).Register() } diff --git a/pkg/js/generated/ts/mysql.ts b/pkg/js/generated/ts/mysql.ts index 4007f05719..bd9c7fe587 100644 --- a/pkg/js/generated/ts/mysql.ts +++ b/pkg/js/generated/ts/mysql.ts @@ -41,7 +41,7 @@ export class MySQLClient { * const isMySQL = mysql.IsMySQL('acme.com', 3306); * ``` */ - public IsMySQL(host: string, port: number): boolean | null { + public IsMySQL(ctx: any, host: string, port: number): boolean | null { return null; } @@ -58,7 +58,7 @@ export class MySQLClient { * const connected = client.Connect('acme.com', 3306, 'username', 'password'); * ``` */ - public Connect(host: string, port: number, username: string): boolean | null { + public Connect(ctx: any, host: string, port: number, username: string): boolean | null { return null; } @@ -72,7 +72,7 @@ export class MySQLClient { * log(to_json(info)); * ``` */ - public FingerprintMySQL(host: string, port: number): MySQLInfo | null { + public FingerprintMySQL(ctx: any, host: string, port: number): MySQLInfo | null { return null; } @@ -88,7 +88,7 @@ export class MySQLClient { * const connected = client.ConnectWithDSN('username:password@tcp(acme.com:3306)/'); * ``` */ - public ConnectWithDSN(dsn: string): boolean | null { + public ConnectWithDSN(ctx: any, dsn: string): boolean | null { return null; } @@ -106,7 +106,7 @@ export class MySQLClient { * log(to_json(result)); * ``` */ - public ExecuteQueryWithOpts(opts: MySQLOptions, query: string): SQLResult | null | null { + public ExecuteQueryWithOpts(ctx: any, opts: MySQLOptions, query: string): SQLResult | null | null { return null; } @@ -121,7 +121,7 @@ export class MySQLClient { * log(to_json(result)); * ``` */ - public ExecuteQuery(host: string, port: number, username: string): SQLResult | null | null { + public ExecuteQuery(ctx: any, host: string, port: number, username: string): SQLResult | null | null { return null; } @@ -136,7 +136,7 @@ export class MySQLClient { * log(to_json(result)); * ``` */ - public ExecuteQueryOnDB(host: string, port: number, username: string): SQLResult | null | null { + public ExecuteQueryOnDB(ctx: any, host: string, port: number, username: string): SQLResult | null | null { return null; } @@ -145,6 +145,41 @@ export class MySQLClient { +/** + */ +export interface HandshakeInfo { + + PacketType?: string, + + ProtocolVersion?: number, + + Version?: string, + + ThreadID?: number, + + CapabilityFlags?: number, + + Capabilities?: string[], + + CharacterSet?: number, + + StatusFlags?: number, + + Status?: string[], + + AuthPluginDataLen?: number, + + Salt?: string, + + AuthPluginName?: string, + + ErrorMessage?: string, + + ErrorCode?: number, +} + + + /** * MySQLInfo contains information about MySQL server. * this is returned when fingerprint is successful @@ -165,7 +200,25 @@ export interface MySQLInfo { Version?: string, - Debug?: ServiceMySQL, + ProtocolVersion?: number, + + ThreadID?: number, + + CapabilityFlags?: number, + + Capabilities?: string[], + + CharacterSet?: number, + + StatusFlags?: number, + + Status?: string[], + + Salt?: string, + + AuthPluginName?: string, + + Debug?: HandshakeInfo, Raw?: string, } @@ -214,17 +267,3 @@ export interface SQLResult { Columns?: string[], } - - -/** - * ServiceMySQL Interface - */ -export interface ServiceMySQL { - - PacketType?: string, - - ErrorMessage?: string, - - ErrorCode?: number, -} - diff --git a/pkg/js/libs/mysql/fingerprint.go b/pkg/js/libs/mysql/fingerprint.go new file mode 100644 index 0000000000..e72bfcea8a --- /dev/null +++ b/pkg/js/libs/mysql/fingerprint.go @@ -0,0 +1,435 @@ +package mysql + +import ( + "encoding/binary" + "encoding/hex" + "encoding/json" + "fmt" + "io" + "net" + "strings" + "time" +) + +const ( + mysqlProtocolVersion10 = 0x0a + mysqlErrorHeader = 0xff + mysqlFingerprintTimeout = 5 * time.Second + + clientLongPassword = 1 << 0 + clientFoundRows = 1 << 1 + clientLongFlag = 1 << 2 + clientConnectWithDB = 1 << 3 + clientNoSchema = 1 << 4 + clientCompress = 1 << 5 + clientODBC = 1 << 6 + clientLocalFiles = 1 << 7 + clientIgnoreSpace = 1 << 8 + clientProtocol41 = 1 << 9 + clientInteractive = 1 << 10 + clientSSL = 1 << 11 + clientIgnoreSigpipe = 1 << 12 + clientTransactions = 1 << 13 + clientReserved = 1 << 14 + clientSecureConnection = 1 << 15 + clientMultiStatements = 1 << 16 + clientMultiResults = 1 << 17 + clientPSMultiResults = 1 << 18 + clientPluginAuth = 1 << 19 + clientConnectAttrs = 1 << 20 + clientPluginAuthLenEncClientData = 1 << 21 + clientCanHandleExpiredPasswords = 1 << 22 + clientSessionTrack = 1 << 23 + clientDeprecateEOF = 1 << 24 + + serverStatusInTrans = 1 << 0 + serverStatusAutocommit = 1 << 1 + serverMoreResultsExists = 1 << 3 + serverQueryNoGoodIndexUsed = 1 << 4 + serverQueryNoIndexUsed = 1 << 5 + serverStatusCursorExists = 1 << 6 + serverStatusLastRowSent = 1 << 7 + serverStatusDBDropped = 1 << 8 + serverStatusNoBackslashEscapes = 1 << 9 + serverStatusMetadataChanged = 1 << 10 + serverQueryWasSlow = 1 << 11 + serverPSOutParams = 1 << 12 + serverStatusInTransReadonly = 1 << 13 + serverSessionStateChanged = 1 << 14 +) + +var mysqlCapabilityNames = []struct { + flag uint32 + name string +}{ + {clientLongPassword, "LongPassword"}, + {clientFoundRows, "FoundRows"}, + {clientLongFlag, "LongColumnFlag"}, + {clientConnectWithDB, "ConnectWithDatabase"}, + {clientNoSchema, "DontAllowDatabaseTableColumn"}, + {clientCompress, "SupportsCompression"}, + {clientODBC, "ODBCClient"}, + {clientLocalFiles, "SupportsLoadDataLocal"}, + {clientIgnoreSpace, "IgnoreSpaceBeforeParenthesis"}, + {clientProtocol41, "Speaks41ProtocolNew"}, + {clientInteractive, "InteractiveClient"}, + {clientSSL, "SwitchToSSLAfterHandshake"}, + {clientIgnoreSigpipe, "IgnoreSigpipes"}, + {clientTransactions, "SupportsTransactions"}, + {clientReserved, "Speaks41ProtocolOld"}, + {clientSecureConnection, "Support41Auth"}, + {clientMultiStatements, "SupportsMultipleStatements"}, + {clientMultiResults, "SupportsMultipleResults"}, + {clientPSMultiResults, "SupportsPSMultiResults"}, + {clientPluginAuth, "SupportsAuthPlugins"}, + {clientConnectAttrs, "ConnectAttrs"}, + {clientPluginAuthLenEncClientData, "PluginAuthLenEncClientData"}, + {clientCanHandleExpiredPasswords, "CanHandleExpiredPasswords"}, + {clientSessionTrack, "SessionTrack"}, + {clientDeprecateEOF, "DeprecateEOF"}, +} + +var mysqlStatusNames = []struct { + flag uint16 + name string +}{ + {serverStatusInTrans, "InTransaction"}, + {serverStatusAutocommit, "Autocommit"}, + {serverMoreResultsExists, "MoreResultsExists"}, + {serverQueryNoGoodIndexUsed, "NoGoodIndexUsed"}, + {serverQueryNoIndexUsed, "NoIndexUsed"}, + {serverStatusCursorExists, "CursorExists"}, + {serverStatusLastRowSent, "LastRowSent"}, + {serverStatusDBDropped, "DBDropped"}, + {serverStatusNoBackslashEscapes, "NoBackslashEscapes"}, + {serverStatusMetadataChanged, "MetadataChanged"}, + {serverQueryWasSlow, "QueryWasSlow"}, + {serverPSOutParams, "PSOutParams"}, + {serverStatusInTransReadonly, "InTransactionReadonly"}, + {serverSessionStateChanged, "SessionStateChanged"}, +} + +// HandshakeInfo is the extended MySQL initial-handshake / error fingerprint. +type HandshakeInfo struct { + PacketType string `json:"packetType"` + ProtocolVersion int `json:"protocolVersion,omitempty"` + Version string `json:"version,omitempty"` + ThreadID uint32 `json:"threadId,omitempty"` + CapabilityFlags uint32 `json:"capabilityFlags,omitempty"` + Capabilities []string `json:"capabilities,omitempty"` + CharacterSet uint8 `json:"characterSet,omitempty"` + StatusFlags uint16 `json:"statusFlags,omitempty"` + Status []string `json:"status,omitempty"` + AuthPluginDataLen int `json:"authPluginDataLen,omitempty"` + Salt string `json:"salt,omitempty"` + AuthPluginName string `json:"authPluginName,omitempty"` + ErrorMessage string `json:"errorMsg,omitempty"` + ErrorCode int `json:"errorCode,omitempty"` +} + +// fingerprintConn reads the MySQL greeting once and parses an extended fingerprint. +func fingerprintConn(conn net.Conn, timeout time.Duration) (HandshakeInfo, error) { + raw, err := recvMySQLPacket(conn, timeout) + if err != nil { + return HandshakeInfo{}, err + } + if len(raw) == 0 { + return HandshakeInfo{}, fmt.Errorf("empty mysql greeting") + } + return parseMySQLGreeting(raw) +} + +func recvMySQLPacket(conn net.Conn, timeout time.Duration) ([]byte, error) { + if err := conn.SetReadDeadline(time.Now().Add(timeout)); err != nil { + return nil, err + } + header := make([]byte, 4) + if _, err := io.ReadFull(conn, header); err != nil { + return nil, err + } + length := int(uint32(header[0]) | uint32(header[1])<<8 | uint32(header[2])<<16) + if length <= 0 || length > 16*1024*1024 { + return nil, fmt.Errorf("invalid mysql packet length %d", length) + } + payload := make([]byte, length) + if _, err := io.ReadFull(conn, payload); err != nil { + return nil, err + } + out := make([]byte, 0, 4+length) + out = append(out, header...) + out = append(out, payload...) + return out, nil +} + +func parseMySQLGreeting(packet []byte) (HandshakeInfo, error) { + if len(packet) < 5 { + return HandshakeInfo{}, fmt.Errorf("mysql packet too short") + } + if packet[4] == mysqlErrorHeader { + return parseMySQLErrorPacket(packet) + } + return parseMySQLHandshakePacket(packet) +} + +func parseMySQLErrorPacket(packet []byte) (HandshakeInfo, error) { + // Stay compatible with fingerprintx error detection: minimum size and 0xff header. + if len(packet) < 8 { + return HandshakeInfo{}, fmt.Errorf("mysql error packet too short") + } + length := mysqlPacketLength(packet) + if length < 3 || length+4 > len(packet) { + return HandshakeInfo{}, fmt.Errorf("mysql error packet truncated") + } + if packet[4] != mysqlErrorHeader { + return HandshakeInfo{}, fmt.Errorf("mysql error packet has invalid header") + } + + info := HandshakeInfo{ + PacketType: "error", + ErrorCode: int(binary.LittleEndian.Uint16(packet[5:7])), + } + msgStart := 7 + // Protocol 4.1 error packets may include '#' + 5-byte SQLSTATE. + if 4+length > 8 && packet[7] == '#' && 4+length >= 13 { + msgStart = 13 + } + if msgStart < 4+length { + info.ErrorMessage = readPrintableASCII(packet[msgStart : 4+length]) + } + return info, nil +} + +func parseMySQLHandshakePacket(packet []byte) (HandshakeInfo, error) { + // Phase 1: fingerprintx-compatible version detection. This must succeed + // whenever fingerprintx.CheckInitialHandshakePacket would, so Version is + // never lost to stricter extended parsing. + version, versionEnd, err := detectMySQLVersion(packet) + if err != nil { + return HandshakeInfo{}, err + } + + info := HandshakeInfo{ + PacketType: "handshake", + ProtocolVersion: mysqlProtocolVersion10, + Version: version, + } + + // Phase 2: best-effort enrichment. Failures here must not drop Version. + enrichMySQLHandshake(&info, packet, versionEnd) + return info, nil +} + +// detectMySQLVersion mirrors fingerprintx CheckInitialHandshakePacket so we +// accept the same greetings and always surface the server version string. +func detectMySQLVersion(packet []byte) (string, int, error) { + if len(packet) < 35 { + return "", 0, fmt.Errorf("mysql handshake packet too short") + } + + // fingerprintx treats bytes[0:4] as little-endian length (seq usually 0). + // Use the real 3-byte MySQL length for bounds, but keep the same 25..4096 gate. + length := mysqlPacketLength(packet) + if length < 25 || length > 4096 { + return "", 0, fmt.Errorf("mysql handshake packet length out of range") + } + if packet[4] != mysqlProtocolVersion10 { + return "", 0, fmt.Errorf("unsupported mysql protocol version") + } + + version, nullPos, err := readNullTerminatedASCIIString(packet, 5) + if err != nil { + return "", 0, err + } + // nullPos points at the NUL; fingerprintx filler is at nullPos+13. + fillerPos := nullPos + 13 + if fillerPos >= len(packet) { + return "", 0, fmt.Errorf("mysql handshake missing filler byte") + } + if packet[fillerPos] != 0x00 { + return "", 0, fmt.Errorf("mysql handshake filler byte is not zero") + } + return version, nullPos + 1, nil +} + +func enrichMySQLHandshake(info *HandshakeInfo, packet []byte, versionEnd int) { + length := mysqlPacketLength(packet) + if length+4 > len(packet) { + length = len(packet) - 4 + } + if length <= 0 { + return + } + payload := packet[4 : 4+length] + // versionEnd is absolute index of first byte after version NUL in packet. + pos := versionEnd - 4 + if pos < 0 || pos > len(payload) { + return + } + + if pos+4 > len(payload) { + return + } + info.ThreadID = binary.LittleEndian.Uint32(payload[pos : pos+4]) + pos += 4 + + if pos+9 > len(payload) { + return + } + saltPart1 := append([]byte(nil), payload[pos:pos+8]...) + pos += 8 + if payload[pos] != 0x00 { + return + } + pos++ + + if pos+2 > len(payload) { + return + } + capLower := binary.LittleEndian.Uint16(payload[pos : pos+2]) + pos += 2 + + var capUpper uint16 + var status uint16 + var charset uint8 + authDataLen := 0 + if pos < len(payload) { + charset = payload[pos] + pos++ + } + if pos+2 <= len(payload) { + status = binary.LittleEndian.Uint16(payload[pos : pos+2]) + pos += 2 + } + if pos+2 <= len(payload) { + capUpper = binary.LittleEndian.Uint16(payload[pos : pos+2]) + pos += 2 + } + if pos < len(payload) { + authDataLen = int(payload[pos]) + pos++ + } + if pos+10 <= len(payload) { + pos += 10 + } else { + pos = len(payload) + } + + caps := uint32(capLower) | uint32(capUpper)<<16 + info.CapabilityFlags = caps + info.Capabilities = decodeCapabilityFlags(caps) + info.CharacterSet = charset + info.StatusFlags = status + info.Status = decodeStatusFlags(status) + info.AuthPluginDataLen = authDataLen + + part2Len := 13 + if authDataLen > 8 { + part2Len = authDataLen - 8 + } + if part2Len < 0 { + part2Len = 0 + } + if pos+part2Len > len(payload) { + part2Len = len(payload) - pos + } + if part2Len < 0 { + part2Len = 0 + } + saltPart2 := append([]byte(nil), payload[pos:pos+part2Len]...) + pos += part2Len + saltPart2 = bytesTrimRightNull(saltPart2) + info.Salt = formatSalt(append(saltPart1, saltPart2...)) + + if caps&clientPluginAuth != 0 && pos < len(payload) { + if plugin, _, err := readNullTerminatedASCIIString(payload, pos); err == nil { + info.AuthPluginName = plugin + } + } +} + +func mysqlPacketLength(packet []byte) int { + if len(packet) < 3 { + return 0 + } + return int(uint32(packet[0]) | uint32(packet[1])<<8 | uint32(packet[2])<<16) +} + +// readNullTerminatedASCIIString mirrors fingerprintx: printable ASCII only, +// returns the index of the NUL terminator (not the next byte). +func readNullTerminatedASCIIString(buf []byte, start int) (string, int, error) { + if start < 0 || start >= len(buf) { + return "", 0, fmt.Errorf("invalid string offset") + } + var characters []byte + for position := start; position < len(buf); position++ { + c := buf[position] + if c >= 0x20 && c <= 0x7e { + characters = append(characters, c) + continue + } + if c == 0x00 { + return string(characters), position, nil + } + return "", 0, fmt.Errorf("encountered invalid ASCII character") + } + return "", 0, fmt.Errorf("unterminated mysql string") +} + +func readPrintableASCII(buf []byte) string { + var characters []byte + for _, c := range buf { + if c >= 0x20 && c <= 0x7e { + characters = append(characters, c) + } + } + return string(characters) +} + +func bytesTrimRightNull(b []byte) []byte { + for len(b) > 0 && b[len(b)-1] == 0x00 { + b = b[:len(b)-1] + } + return b +} + +func formatSalt(b []byte) string { + var sb strings.Builder + for _, c := range b { + if c >= 0x20 && c <= 0x7e && c != '\\' { + sb.WriteByte(c) + continue + } + sb.WriteByte('\\') + sb.WriteByte('x') + sb.WriteString(hex.EncodeToString([]byte{c})) + } + return sb.String() +} + +func decodeCapabilityFlags(caps uint32) []string { + out := make([]string, 0, len(mysqlCapabilityNames)) + for _, item := range mysqlCapabilityNames { + if caps&item.flag != 0 { + out = append(out, item.name) + } + } + return out +} + +func decodeStatusFlags(status uint16) []string { + out := make([]string, 0, len(mysqlStatusNames)) + for _, item := range mysqlStatusNames { + if status&item.flag != 0 { + out = append(out, item.name) + } + } + return out +} + +func handshakeToJSON(info HandshakeInfo) string { + bin, err := json.Marshal(info) + if err != nil { + return "" + } + return string(bin) +} diff --git a/pkg/js/libs/mysql/fingerprint_test.go b/pkg/js/libs/mysql/fingerprint_test.go new file mode 100644 index 0000000000..e3429ff645 --- /dev/null +++ b/pkg/js/libs/mysql/fingerprint_test.go @@ -0,0 +1,517 @@ +package mysql + +import ( + "encoding/binary" + "encoding/hex" + "net" + "testing" + "time" + + fxmysql "github.com/praetorian-inc/fingerprintx/pkg/plugins/services/mysql" + "github.com/stretchr/testify/require" +) + +func buildPacket(t *testing.T, payload []byte) []byte { + t.Helper() + require.LessOrEqual(t, len(payload), 0xffffff) + header := []byte{ + byte(len(payload)), + byte(len(payload) >> 8), + byte(len(payload) >> 16), + 0x00, + } + return append(header, payload...) +} + +// Golden handshake derived from the MySQL initial-handshake layout documented in +// fingerprintx (version 8.0.28, protocol 10, caching_sha2_password). +func testHandshakePacket(t *testing.T) []byte { + t.Helper() + payload := []byte{0x0a} + payload = append(payload, []byte("8.0.28")...) + payload = append(payload, 0x00) + payload = append(payload, 0x0b, 0x00, 0x00, 0x00) // thread id 11 + payload = append(payload, 0x15, 0x05, 0x6c, 0x51, 0x28, 0x32, 0x48, 0x15) // salt part1 + payload = append(payload, 0x00) // filler + payload = append(payload, 0xff, 0xff) // cap lower + payload = append(payload, 0xff) // charset + payload = append(payload, 0x02, 0x00) // status autocommit + payload = append(payload, 0xff, 0xdf) // cap upper (includes plugin auth) + payload = append(payload, 0x15) // auth plugin data len = 21 + payload = append(payload, make([]byte, 10)...) // reserved + payload = append(payload, 0x26, 0x68, 0x15, 0x1e, 0x2e, 0x7f, 0x69, 0x38, 0x52, 0x6b, 0x6c, 0x5c, 0x00) + payload = append(payload, []byte("caching_sha2_password")...) + payload = append(payload, 0x00) + return buildPacket(t, payload) +} + +func testNativePasswordHandshake(t *testing.T) []byte { + t.Helper() + // MySQL 5.7-style greeting with mysql_native_password. + payload := []byte{0x0a} + payload = append(payload, []byte("5.7.44")...) + payload = append(payload, 0x00) + payload = append(payload, 0x2a, 0x00, 0x00, 0x00) // thread id 42 + payload = append(payload, []byte("12345678")...) // salt part1 + payload = append(payload, 0x00) + payload = append(payload, 0xff, 0xf7) // cap lower + payload = append(payload, 0x08) // charset + payload = append(payload, 0x02, 0x00) // autocommit + payload = append(payload, 0x08, 0x00) // cap upper with CLIENT_PLUGIN_AUTH (bit 19) + payload = append(payload, 0x15) + payload = append(payload, make([]byte, 10)...) + payload = append(payload, []byte("abcdefghijkl")...) + payload = append(payload, 0x00) + payload = append(payload, []byte("mysql_native_password")...) + payload = append(payload, 0x00) + return buildPacket(t, payload) +} + +func TestParseMySQLHandshakePacket(t *testing.T) { + t.Parallel() + + info, err := parseMySQLGreeting(testHandshakePacket(t)) + require.NoError(t, err) + require.Equal(t, "handshake", info.PacketType) + require.Equal(t, 10, info.ProtocolVersion) + require.Equal(t, "8.0.28", info.Version) + require.Equal(t, uint32(11), info.ThreadID) + require.Equal(t, uint8(0xff), info.CharacterSet) + require.Equal(t, uint16(0x0002), info.StatusFlags) + require.Equal(t, []string{"Autocommit"}, info.Status) + require.Equal(t, "caching_sha2_password", info.AuthPluginName) + require.Equal(t, 21, info.AuthPluginDataLen) + require.Equal(t, uint32(0xdfffffff), info.CapabilityFlags) + require.Contains(t, info.Capabilities, "SupportsAuthPlugins") + require.Contains(t, info.Capabilities, "SwitchToSSLAfterHandshake") + require.Contains(t, info.Capabilities, "SupportsTransactions") + require.NotEmpty(t, info.Salt) + + saltRaw, err := hex.DecodeString("15056c51283248152668151e2e7f6938526b6c5c") + require.NoError(t, err) + require.Equal(t, formatSalt(saltRaw), info.Salt) +} + +func TestParseMySQLNativePasswordHandshake(t *testing.T) { + t.Parallel() + + info, err := parseMySQLGreeting(testNativePasswordHandshake(t)) + require.NoError(t, err) + require.Equal(t, "handshake", info.PacketType) + require.Equal(t, "5.7.44", info.Version) + require.Equal(t, uint32(42), info.ThreadID) + require.Equal(t, "mysql_native_password", info.AuthPluginName) + require.Contains(t, info.Capabilities, "SupportsAuthPlugins") + require.Contains(t, info.Salt, "12345678") +} + +func TestParseMySQLHandshakeWithoutPluginAuth(t *testing.T) { + t.Parallel() + + payload := []byte{0x0a} + payload = append(payload, []byte("5.5.5-MariaDB")...) + payload = append(payload, 0x00) + payload = append(payload, 0x01, 0x00, 0x00, 0x00) + payload = append(payload, []byte("abcdefgh")...) + payload = append(payload, 0x00) + payload = append(payload, 0xff, 0xf7) // no plugin-auth in upper caps + payload = append(payload, 0x08) + payload = append(payload, 0x02, 0x00) + payload = append(payload, 0x00, 0x00) // cap upper without plugin auth + payload = append(payload, 0x00) + payload = append(payload, make([]byte, 10)...) + payload = append(payload, []byte("ijklmnopabcd")...) + payload = append(payload, 0x00) + + info, err := parseMySQLGreeting(buildPacket(t, payload)) + require.NoError(t, err) + require.Equal(t, "5.5.5-MariaDB", info.Version) + require.Empty(t, info.AuthPluginName) + require.NotContains(t, info.Capabilities, "SupportsAuthPlugins") +} + +func TestParseMySQLErrorPacket(t *testing.T) { + t.Parallel() + + msg := "Host '1.2.3.4' is not allowed to connect to this MySQL server" + payload := []byte{0xff, 0x6a, 0x04} + payload = append(payload, []byte(msg)...) + + info, err := parseMySQLGreeting(buildPacket(t, payload)) + require.NoError(t, err) + require.Equal(t, "error", info.PacketType) + require.Equal(t, 0x046a, info.ErrorCode) + require.Equal(t, msg, info.ErrorMessage) + require.Empty(t, info.Version) + require.Zero(t, info.ThreadID) +} + +func TestParseMySQLErrorPacketWithSQLState(t *testing.T) { + t.Parallel() + + payload := []byte{0xff, 0x15, 0x04, '#', '2', '8', '0', '0', '0'} + payload = append(payload, []byte("Access denied for user")...) + + info, err := parseMySQLGreeting(buildPacket(t, payload)) + require.NoError(t, err) + require.Equal(t, "error", info.PacketType) + require.Equal(t, 0x0415, info.ErrorCode) + require.Equal(t, "Access denied for user", info.ErrorMessage) +} + +func TestParseMySQLHandshakeRejectsGarbage(t *testing.T) { + t.Parallel() + + _, err := parseMySQLGreeting([]byte{0x01, 0x00, 0x00, 0x00, 0x09}) + require.Error(t, err) + + _, err = parseMySQLGreeting([]byte{0x00}) + require.Error(t, err) + + _, err = parseMySQLGreeting(buildPacket(t, []byte{0x0a, 'x'})) // truncated after version byte + require.Error(t, err) +} + +func TestParseMySQLHandshakeRejectsBadFiller(t *testing.T) { + t.Parallel() + + payload := []byte{0x0a, '8', '.', '0', 0x00} + payload = append(payload, 0x01, 0x00, 0x00, 0x00) + payload = append(payload, []byte("12345678")...) + payload = append(payload, 0x01) // bad filler + payload = append(payload, make([]byte, 20)...) + + _, err := parseMySQLGreeting(buildPacket(t, payload)) + require.Error(t, err) + require.Contains(t, err.Error(), "filler") +} + +func TestFormatSaltEscapesBinary(t *testing.T) { + t.Parallel() + + raw, err := hex.DecodeString("15056c51") + require.NoError(t, err) + got := formatSalt(raw) + require.Equal(t, `\x15\x05lQ`, got) +} + +func TestDecodeCapabilityAndStatusFlags(t *testing.T) { + t.Parallel() + + caps := decodeCapabilityFlags(clientSSL | clientPluginAuth | clientTransactions) + require.Equal(t, []string{"SwitchToSSLAfterHandshake", "SupportsTransactions", "SupportsAuthPlugins"}, caps) + + status := decodeStatusFlags(serverStatusAutocommit | serverStatusInTrans) + require.Equal(t, []string{"InTransaction", "Autocommit"}, status) + + require.Empty(t, decodeCapabilityFlags(0)) + require.Empty(t, decodeStatusFlags(0)) +} + +func TestHandshakeToJSON(t *testing.T) { + t.Parallel() + + info, err := parseMySQLGreeting(testHandshakePacket(t)) + require.NoError(t, err) + raw := handshakeToJSON(info) + require.Contains(t, raw, `"packetType":"handshake"`) + require.Contains(t, raw, `"version":"8.0.28"`) + require.Contains(t, raw, `"authPluginName":"caching_sha2_password"`) +} + +func TestFingerprintConnReadsGreetingOnce(t *testing.T) { + t.Parallel() + + server, client := net.Pipe() + defer func() { _ = server.Close() }() + defer func() { _ = client.Close() }() + + packet := testHandshakePacket(t) + errCh := make(chan error, 1) + go func() { + _, err := server.Write(packet) + errCh <- err + _ = server.Close() + }() + + info, err := fingerprintConn(client, time.Second) + require.NoError(t, err) + require.NoError(t, <-errCh) + require.Equal(t, "8.0.28", info.Version) + require.Equal(t, "caching_sha2_password", info.AuthPluginName) +} + +func TestRecvMySQLPacketRejectsHugeLength(t *testing.T) { + t.Parallel() + + server, client := net.Pipe() + defer func() { _ = server.Close() }() + defer func() { _ = client.Close() }() + + go func() { + // 3-byte length = 0x01000000 (> 16MiB cap) + seq 0 + _, _ = server.Write([]byte{0x00, 0x00, 0x00, 0x01}) + _ = server.Close() + }() + + _, err := recvMySQLPacket(client, time.Second) + require.Error(t, err) + require.Contains(t, err.Error(), "invalid mysql packet length") +} + +func TestCapabilityFlagConstantsMatchWire(t *testing.T) { + t.Parallel() + + buf := make([]byte, 4) + binary.LittleEndian.PutUint32(buf, clientPluginAuth|clientSSL) + require.Equal(t, byte(0x00), buf[0]) + require.Equal(t, byte(0x08), buf[1]) // SSL in lower + require.Equal(t, byte(0x08), buf[2]) // plugin auth in upper (bit 19 -> byte2 bit3) + require.Equal(t, byte(0x00), buf[3]) +} + +func testMariaDBHandshake(t *testing.T) []byte { + t.Helper() + payload := []byte{0x0a} + payload = append(payload, []byte("5.5.5-10.11.8-MariaDB")...) + payload = append(payload, 0x00) + payload = append(payload, 0x01, 0x00, 0x00, 0x00) + payload = append(payload, []byte("abcdefgh")...) + payload = append(payload, 0x00) + payload = append(payload, 0xff, 0xf7) + payload = append(payload, 0x08) + payload = append(payload, 0x02, 0x00) + payload = append(payload, 0x00, 0x00) + payload = append(payload, 0x00) + payload = append(payload, make([]byte, 10)...) + payload = append(payload, []byte("ijklmnopabcd")...) + payload = append(payload, 0x00) + return buildPacket(t, payload) +} + +func testMinimalHandshake(t *testing.T) []byte { + t.Helper() + // Minimal handshake fingerprintx accepts (len>=35) without full salt/plugin tail. + payload := []byte{0x0a} + payload = append(payload, []byte("8.0.28")...) + payload = append(payload, 0x00) + payload = append(payload, 0x0b, 0x00, 0x00, 0x00) + payload = append(payload, 0x15, 0x05, 0x6c, 0x51, 0x28, 0x32, 0x48, 0x15) + payload = append(payload, 0x00) // filler + for len(payload) < 31 { // 4+31 == 35 absolute bytes minimum + payload = append(payload, 0x00) + } + return buildPacket(t, payload) +} + +// TestFingerprintxHandshakeVersionParity ensures every greeting fingerprintx +// accepts yields the same Version from our parser (no version regressions). +func TestFingerprintxHandshakeVersionParity(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + pkt []byte + wantVersion string + wantEnrichment bool + wantAuthPlugin string + }{ + {"mysql80", testHandshakePacket(t), "8.0.28", true, "caching_sha2_password"}, + {"mysql57", testNativePasswordHandshake(t), "5.7.44", true, "mysql_native_password"}, + {"mariadb", testMariaDBHandshake(t), "5.5.5-10.11.8-MariaDB", true, ""}, + {"minimal", testMinimalHandshake(t), "8.0.28", false, ""}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + + fxVersion, fxErr := fxmysql.CheckInitialHandshakePacket(tc.pkt) + require.NoError(t, fxErr, "fingerprintx must accept fixture") + require.Equal(t, tc.wantVersion, fxVersion) + + info, err := parseMySQLGreeting(tc.pkt) + require.NoError(t, err) + require.Equal(t, "handshake", info.PacketType) + require.Equal(t, fxVersion, info.Version, "version must match fingerprintx") + require.NotEmpty(t, info.Version) + require.Equal(t, 10, info.ProtocolVersion) + if tc.wantEnrichment { + require.NotZero(t, info.ThreadID) + require.NotEmpty(t, info.Salt) + require.NotZero(t, info.CapabilityFlags) + } + if tc.wantAuthPlugin != "" { + require.Equal(t, tc.wantAuthPlugin, info.AuthPluginName) + } + }) + } +} + +func TestFingerprintxErrorPacketParity(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + payload []byte + wantMsg string + wantCode int + }{ + { + name: "host_not_allowed", + payload: append([]byte{0xff, 0x6a, 0x04}, + []byte("Host '1.2.3.4' is not allowed to connect to this MySQL server")...), + wantMsg: "Host '1.2.3.4' is not allowed to connect to this MySQL server", + wantCode: 0x046a, + }, + { + name: "access_denied_plain", + payload: append([]byte{0xff, 0x15, 0x04}, + []byte("Access denied for user 'root'@'localhost'")...), + wantMsg: "Access denied for user 'root'@'localhost'", + wantCode: 0x0415, + }, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + pkt := buildPacket(t, tc.payload) + + fxMsg, fxCode, fxErr := fxmysql.CheckErrorMessagePacket(pkt) + require.NoError(t, fxErr) + require.Equal(t, tc.wantMsg, fxMsg) + require.Equal(t, tc.wantCode, fxCode) + + info, err := parseMySQLGreeting(pkt) + require.NoError(t, err) + require.Equal(t, "error", info.PacketType) + require.Equal(t, fxCode, info.ErrorCode) + require.Equal(t, fxMsg, info.ErrorMessage) + require.Empty(t, info.Version) + }) + } +} + +// TestFingerprintxRejectParity: packets fingerprintx rejects as both handshake +// and error must also fail our parser (no false MySQL positives on garbage). +func TestFingerprintxRejectParity(t *testing.T) { + t.Parallel() + + cases := []struct { + name string + pkt []byte + }{ + {"empty", []byte{}}, + {"too_short", []byte{0x01, 0x00, 0x00, 0x00, 0x0a}}, + {"ssh_banner", []byte("SSH-2.0-OpenSSH\r\n")}, + {"http", []byte("HTTP/1.1 200 OK\r\n\r\n")}, + {"wrong_protocol", buildPacket(t, append([]byte{0x09}, make([]byte, 40)...))}, + {"non_ascii_version", func() []byte { + payload := []byte{0x0a, 0x80, 0x00} // non-printable in version + payload = append(payload, make([]byte, 40)...) + return buildPacket(t, payload) + }()}, + {"bad_filler", func() []byte { + payload := []byte{0x0a, '8', '.', '0', 0x00} + payload = append(payload, 0x01, 0x00, 0x00, 0x00) + payload = append(payload, []byte("12345678")...) + payload = append(payload, 0x01) // bad filler + payload = append(payload, make([]byte, 20)...) + return buildPacket(t, payload) + }()}, + } + + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + t.Parallel() + _, hsErr := fxmysql.CheckInitialHandshakePacket(tc.pkt) + _, _, errErr := fxmysql.CheckErrorMessagePacket(tc.pkt) + require.Error(t, hsErr) + require.Error(t, errErr) + + _, err := parseMySQLGreeting(tc.pkt) + require.Error(t, err, "must reject packets fingerprintx rejects") + }) + } +} + +func TestVersionSurvivesTruncatedEnrichment(t *testing.T) { + t.Parallel() + + pkt := testMinimalHandshake(t) + require.GreaterOrEqual(t, len(pkt), 35) + + fxVersion, fxErr := fxmysql.CheckInitialHandshakePacket(pkt) + require.NoError(t, fxErr) + require.Equal(t, "8.0.28", fxVersion) + + info, err := parseMySQLGreeting(pkt) + require.NoError(t, err) + require.Equal(t, "8.0.28", info.Version) + require.Equal(t, "handshake", info.PacketType) + require.Equal(t, uint32(11), info.ThreadID) +} + +func TestVersionAlwaysSetOnValidHandshake(t *testing.T) { + t.Parallel() + + for _, pkt := range [][]byte{ + testHandshakePacket(t), + testNativePasswordHandshake(t), + testMariaDBHandshake(t), + testMinimalHandshake(t), + } { + info, err := parseMySQLGreeting(pkt) + require.NoError(t, err) + require.NotEmpty(t, info.Version) + require.Equal(t, 10, info.ProtocolVersion) + } +} + +func TestFingerprintConnErrorGreeting(t *testing.T) { + t.Parallel() + + msg := "Host '9.9.9.9' is not allowed to connect to this MySQL server" + payload := append([]byte{0xff, 0x6a, 0x04}, []byte(msg)...) + packet := buildPacket(t, payload) + + server, client := net.Pipe() + defer func() { _ = server.Close() }() + defer func() { _ = client.Close() }() + + errCh := make(chan error, 1) + go func() { + _, err := server.Write(packet) + errCh <- err + _ = server.Close() + }() + + info, err := fingerprintConn(client, time.Second) + require.NoError(t, err) + require.NoError(t, <-errCh) + require.Equal(t, "error", info.PacketType) + require.Equal(t, msg, info.ErrorMessage) + require.Empty(t, info.Version) +} + +func TestDebugJSONPreservesFingerprintxFields(t *testing.T) { + t.Parallel() + + hs, err := parseMySQLGreeting(testHandshakePacket(t)) + require.NoError(t, err) + raw := handshakeToJSON(hs) + require.Contains(t, raw, `"packetType":"handshake"`) + require.Contains(t, raw, `"version":"8.0.28"`) + // fingerprintx ServiceMySQL always had these keys on error path; present as omitempty on handshake + require.NotContains(t, raw, `"errorMsg"`) + + msg := "denied" + payload := append([]byte{0xff, 0x6a, 0x04}, []byte(msg)...) + errInfo, err := parseMySQLGreeting(buildPacket(t, payload)) + require.NoError(t, err) + errRaw := handshakeToJSON(errInfo) + require.Contains(t, errRaw, `"packetType":"error"`) + require.Contains(t, errRaw, `"errorMsg":"denied"`) + require.Contains(t, errRaw, `"errorCode":1130`) +} diff --git a/pkg/js/libs/mysql/mysql.go b/pkg/js/libs/mysql/mysql.go index 3f92119d96..48f1b5b866 100644 --- a/pkg/js/libs/mysql/mysql.go +++ b/pkg/js/libs/mysql/mysql.go @@ -6,11 +6,8 @@ import ( "io" "log" "net" - "time" "github.com/go-sql-driver/mysql" - "github.com/praetorian-inc/fingerprintx/pkg/plugins" - mysqlplugin "github.com/praetorian-inc/fingerprintx/pkg/plugins/services/mysql" "github.com/projectdiscovery/nuclei/v3/pkg/js/utils" "github.com/projectdiscovery/nuclei/v3/pkg/protocols/common/protocolstate" ) @@ -42,31 +39,11 @@ func (c *MySQLClient) IsMySQL(ctx context.Context, host string, port int) (bool, // @memo func isMySQL(ctx context.Context, executionId string, host string, port int) (bool, error) { - if !protocolstate.IsHostAllowed(executionId, host) { - // host is not valid according to network policy - return false, protocolstate.ErrHostDenied.Msgf(host) - } - dialer := protocolstate.GetDialersWithId(executionId) - if dialer == nil { - return false, fmt.Errorf("dialers not initialized for %s", executionId) - } - - conn, err := dialer.Fastdialer.Dial(ctx, "tcp", net.JoinHostPort(host, fmt.Sprintf("%d", port))) + // Reuse the memoized fingerprint probe so IsMySQL + FingerprintMySQL share one dial. + _, err := memoizedfingerprintMySQL(ctx, executionId, host, port) if err != nil { return false, err } - defer func() { - _ = conn.Close() - }() - - plugin := &mysqlplugin.MYSQLPlugin{} - service, err := plugin.Run(conn, 5*time.Second, plugins.Target{Host: host}) - if err != nil { - return false, err - } - if service == nil { - return false, nil - } return true, nil } @@ -114,15 +91,24 @@ type ( // MySQLInfo contains information about MySQL server. // this is returned when fingerprint is successful MySQLInfo struct { - Host string `json:"host,omitempty"` - IP string `json:"ip"` - Port int `json:"port"` - Protocol string `json:"protocol"` - TLS bool `json:"tls"` - Transport string `json:"transport"` - Version string `json:"version,omitempty"` - Debug plugins.ServiceMySQL `json:"debug,omitempty"` - Raw string `json:"metadata"` + Host string `json:"host,omitempty"` + IP string `json:"ip"` + Port int `json:"port"` + Protocol string `json:"protocol"` + TLS bool `json:"tls"` + Transport string `json:"transport"` + Version string `json:"version,omitempty"` + ProtocolVersion int `json:"protocolVersion,omitempty"` + ThreadID uint32 `json:"threadId,omitempty"` + CapabilityFlags uint32 `json:"capabilityFlags,omitempty"` + Capabilities []string `json:"capabilities,omitempty"` + CharacterSet uint8 `json:"characterSet,omitempty"` + StatusFlags uint16 `json:"statusFlags,omitempty"` + Status []string `json:"status,omitempty"` + Salt string `json:"salt,omitempty"` + AuthPluginName string `json:"authPluginName,omitempty"` + Debug HandshakeInfo `json:"debug,omitempty"` + Raw string `json:"metadata"` } ) @@ -158,25 +144,33 @@ func fingerprintMySQL(ctx context.Context, executionId string, host string, port _ = conn.Close() }() - plugin := &mysqlplugin.MYSQLPlugin{} - service, err := plugin.Run(conn, 5*time.Second, plugins.Target{Host: host}) + handshake, err := fingerprintConn(conn, mysqlFingerprintTimeout) if err != nil { return info, err } - if service == nil { - return info, fmt.Errorf("something went wrong got null output") + + info.Host = host + info.Port = port + info.Protocol = "mysql" + info.Transport = "tcp" + info.TLS = false + if hostIP := net.ParseIP(host); hostIP != nil { + info.IP = hostIP.String() + } else if remote, ok := conn.RemoteAddr().(*net.TCPAddr); ok && remote.IP != nil { + info.IP = remote.IP.String() } - // fill all fields - info.Host = service.Host - info.IP = service.IP - info.Port = service.Port - info.Protocol = service.Protocol - info.TLS = service.TLS - info.Transport = service.Transport - info.Version = service.Version - info.Debug = service.Metadata().(plugins.ServiceMySQL) - bin, _ := service.Raw.MarshalJSON() - info.Raw = string(bin) + info.Version = handshake.Version + info.ProtocolVersion = handshake.ProtocolVersion + info.ThreadID = handshake.ThreadID + info.CapabilityFlags = handshake.CapabilityFlags + info.Capabilities = handshake.Capabilities + info.CharacterSet = handshake.CharacterSet + info.StatusFlags = handshake.StatusFlags + info.Status = handshake.Status + info.Salt = handshake.Salt + info.AuthPluginName = handshake.AuthPluginName + info.Debug = handshake + info.Raw = handshakeToJSON(handshake) return info, nil } diff --git a/pkg/js/libs/mysql/mysql_fingerprint_client_test.go b/pkg/js/libs/mysql/mysql_fingerprint_client_test.go new file mode 100644 index 0000000000..89267df9df --- /dev/null +++ b/pkg/js/libs/mysql/mysql_fingerprint_client_test.go @@ -0,0 +1,196 @@ +package mysql + +import ( + "context" + "net" + "strings" + "sync/atomic" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/projectdiscovery/nuclei/v3/pkg/protocols/common/protocolstate" + "github.com/projectdiscovery/nuclei/v3/pkg/types" +) + +func initMySQLExec(t *testing.T) (context.Context, string) { + t.Helper() + executionID := "mysql-" + strings.NewReplacer("/", "-", " ", "-").Replace(t.Name()) + require.NoError(t, protocolstate.Init(&types.Options{ExecutionId: executionID})) + t.Cleanup(func() { protocolstate.Close(executionID) }) + ctx := context.WithValue(context.Background(), "executionId", executionID) //nolint:staticcheck + return ctx, executionID +} + +// startGreetingServer serves a fixed MySQL greeting once per accepted connection. +// dialCount increments on each Accept. +func startGreetingServer(t *testing.T, packet []byte, dialCount *atomic.Int32) (host string, port int) { + t.Helper() + ln, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + t.Cleanup(func() { _ = ln.Close() }) + + go func() { + for { + conn, err := ln.Accept() + if err != nil { + return + } + if dialCount != nil { + dialCount.Add(1) + } + go func(c net.Conn) { + defer func() { _ = c.Close() }() + _, _ = c.Write(packet) + }(conn) + } + }() + + addr := ln.Addr().(*net.TCPAddr) + return addr.IP.String(), addr.Port +} + +func TestFingerprintMySQLClientHandshake(t *testing.T) { + t.Parallel() + + ctx, _ := initMySQLExec(t) + host, port := startGreetingServer(t, testHandshakePacket(t), nil) + + client := &MySQLClient{} + info, err := client.FingerprintMySQL(ctx, host, port) + require.NoError(t, err) + + require.Equal(t, host, info.Host) + require.Equal(t, port, info.Port) + require.Equal(t, "mysql", info.Protocol) + require.Equal(t, "tcp", info.Transport) + require.False(t, info.TLS) + require.Equal(t, "8.0.28", info.Version) + require.Equal(t, 10, info.ProtocolVersion) + require.Equal(t, uint32(11), info.ThreadID) + require.Equal(t, "caching_sha2_password", info.AuthPluginName) + require.NotEmpty(t, info.Salt) + require.NotZero(t, info.CapabilityFlags) + require.Contains(t, info.Capabilities, "SupportsAuthPlugins") + require.Equal(t, "handshake", info.Debug.PacketType) + require.Equal(t, info.Version, info.Debug.Version) + require.Contains(t, info.Raw, `"packetType":"handshake"`) + require.Contains(t, info.Raw, `"version":"8.0.28"`) + + ok, err := client.IsMySQL(ctx, host, port) + require.NoError(t, err) + require.True(t, ok) +} + +func TestFingerprintMySQLClientErrorGreeting(t *testing.T) { + t.Parallel() + + msg := "Host '1.2.3.4' is not allowed to connect to this MySQL server" + payload := []byte{0xff, 0x6a, 0x04} + payload = append(payload, []byte(msg)...) + pkt := buildPacket(t, payload) + + ctx, _ := initMySQLExec(t) + host, port := startGreetingServer(t, pkt, nil) + + client := &MySQLClient{} + info, err := client.FingerprintMySQL(ctx, host, port) + require.NoError(t, err) + require.Empty(t, info.Version) + require.Equal(t, "error", info.Debug.PacketType) + require.Equal(t, 0x046a, info.Debug.ErrorCode) + require.Equal(t, msg, info.Debug.ErrorMessage) + require.Contains(t, info.Raw, `"packetType":"error"`) + + ok, err := client.IsMySQL(ctx, host, port) + require.NoError(t, err) + require.True(t, ok, "error greetings are still MySQL, same as fingerprintx") +} + +func TestFingerprintMySQLAndIsMySQLShareOneDial(t *testing.T) { + t.Parallel() + + var dials atomic.Int32 + ctx, _ := initMySQLExec(t) + host, port := startGreetingServer(t, testHandshakePacket(t), &dials) + + client := &MySQLClient{} + ok, err := client.IsMySQL(ctx, host, port) + require.NoError(t, err) + require.True(t, ok) + + info, err := client.FingerprintMySQL(ctx, host, port) + require.NoError(t, err) + require.Equal(t, "8.0.28", info.Version) + require.Equal(t, int32(1), dials.Load(), "memoized probe must dial once") +} + +func TestFingerprintMySQLRejectsNonMySQL(t *testing.T) { + t.Parallel() + + ctx, _ := initMySQLExec(t) + host, port := startGreetingServer(t, []byte("SSH-2.0-OpenSSH_8.0\r\n"), nil) + + client := &MySQLClient{} + _, err := client.FingerprintMySQL(ctx, host, port) + require.Error(t, err) + + ok, err := client.IsMySQL(ctx, host, port) + require.Error(t, err) + require.False(t, ok) +} + +func TestFingerprintMySQLNativePasswordVersion(t *testing.T) { + t.Parallel() + + ctx, _ := initMySQLExec(t) + host, port := startGreetingServer(t, testNativePasswordHandshake(t), nil) + + info, err := (&MySQLClient{}).FingerprintMySQL(ctx, host, port) + require.NoError(t, err) + require.Equal(t, "5.7.44", info.Version) + require.Equal(t, "mysql_native_password", info.AuthPluginName) + require.NotEmpty(t, info.Version, "version must always be detected on valid handshake") +} + +func TestFingerprintMySQLDeniedByNetworkPolicy(t *testing.T) { + t.Parallel() + + executionID := "mysql-deny-" + strings.NewReplacer("/", "-", " ", "-").Replace(t.Name()) + require.NoError(t, protocolstate.Init(&types.Options{ + ExecutionId: executionID, + RestrictLocalNetworkAccess: true, + })) + t.Cleanup(func() { protocolstate.Close(executionID) }) + ctx := context.WithValue(context.Background(), "executionId", executionID) //nolint:staticcheck + + _, err := (&MySQLClient{}).FingerprintMySQL(ctx, "127.0.0.1", 3306) + require.Error(t, err) + require.Contains(t, err.Error(), "127.0.0.1") +} + +func TestFingerprintMySQLVersionNeverEmptyOnHandshakeFixtures(t *testing.T) { + t.Parallel() + + fixtures := []struct { + name string + packet []byte + version string + }{ + {"8.0", testHandshakePacket(t), "8.0.28"}, + {"5.7", testNativePasswordHandshake(t), "5.7.44"}, + } + + for _, fx := range fixtures { + t.Run(fx.name, func(t *testing.T) { + t.Parallel() + ctx, _ := initMySQLExec(t) + host, port := startGreetingServer(t, fx.packet, nil) + info, err := (&MySQLClient{}).FingerprintMySQL(ctx, host, port) + require.NoError(t, err) + require.Equal(t, fx.version, info.Version) + require.NotEmpty(t, info.Version) + require.Equal(t, fx.version, info.Debug.Version) + }) + } +}