diff --git a/src/network/websocketScanner.js b/src/network/websocketScanner.js index 763290e..1945f7b 100644 --- a/src/network/websocketScanner.js +++ b/src/network/websocketScanner.js @@ -1,11 +1,10 @@ import { EventEmitter } from 'events'; import os from 'os'; import { Worker } from 'worker_threads'; -import path from 'path'; -import { fileURLToPath } from 'url'; +import { WebSocketScannerWorker } from './websocketScannerWorker.js'; -const __filename = fileURLToPath(import.meta.url); -const __dirname = path.dirname(__filename); +const isBun = typeof globalThis.Bun !== 'undefined'; +const workerUrl = new URL('./websocketScannerWorker.js', import.meta.url); export class WebSocketScanner extends EventEmitter { constructor(options = {}) { @@ -13,7 +12,11 @@ export class WebSocketScanner extends EventEmitter { this.scanning = false; this.stopped = false; this.workers = []; - this.maxThreads = options.maxThreads || os.cpus().length; + this.inProcessWorker = null; + this.useWorkers = !isBun; + this.maxThreads = this.useWorkers + ? (options.maxThreads || os.cpus().length) + : Math.max(1, Math.min(options.maxThreads || 16, 24)); this.totalIPs = 0; this.scannedIPs = 0; } @@ -59,6 +62,51 @@ export class WebSocketScanner extends EventEmitter { const end = this.ipToNumber(endIP); this.totalIPs = end - start + 1; + + if (this.useWorkers) { + await this.scanWithWorkers(start, end, port, protocol, objectType, foundServices); + } else { + await this.scanInProcess(startIP, endIP, port, protocol, objectType, foundServices); + } + + await this.cleanupWorkers(); + return foundServices; + } + + handleWorkerMessage(message, foundServices) { + if (this.stopped) { + return; + } + + switch (message.type) { + case 'serviceFound': + foundServices.push({ ip: message.ip, hostname: message.hostname }); + this.emit('serviceFound', { ip: message.ip, hostname: message.hostname }); + break; + case 'scanProgress': + this.scannedIPs += message.increment; + this.emit('scanProgress', { + currentIP: message.currentIP, + progress: (this.scannedIPs / this.totalIPs) * 100 + }); + break; + default: + break; + } + } + + async scanInProcess(startIP, endIP, port, protocol, objectType, foundServices) { + const worker = new WebSocketScannerWorker((message) => { + this.handleWorkerMessage(message, foundServices); + }); + this.inProcessWorker = worker; + + await worker.scanRange(startIP, endIP, port, protocol, objectType, { + concurrency: this.maxThreads + }); + } + + async scanWithWorkers(start, end, port, protocol, objectType, foundServices) { const ipsPerWorker = Math.ceil(this.totalIPs / this.maxThreads); const workerPromises = []; @@ -70,15 +118,19 @@ export class WebSocketScanner extends EventEmitter { const workerStart = start + (i * ipsPerWorker); const workerEnd = Math.min(workerStart + ipsPerWorker - 1, end); - if (workerStart > end) break; + if (workerStart > end) { + break; + } - const worker = new Worker(path.join(__dirname, 'websocketScannerWorker.js')); + const worker = new Worker(workerUrl); this.workers.push(worker); const workerPromise = new Promise((resolve) => { let settled = false; const finish = (results) => { - if (settled) return; + if (settled) { + return; + } settled = true; resolve(results || []); }; @@ -89,23 +141,12 @@ export class WebSocketScanner extends EventEmitter { return; } - switch (message.type) { - case 'serviceFound': - foundServices.push({ ip: message.ip, hostname: message.hostname }); - this.emit('serviceFound', { ip: message.ip, hostname: message.hostname }); - break; - case 'scanProgress': - this.scannedIPs += message.increment; - const totalProgress = (this.scannedIPs / this.totalIPs) * 100; - this.emit('scanProgress', { - currentIP: message.currentIP, - progress: totalProgress - }); - break; - case 'scanComplete': - finish(message.results); - break; + if (message.type === 'scanComplete') { + finish(message.results); + return; } + + this.handleWorkerMessage(message, foundServices); }); worker.on('error', () => finish([])); worker.on('exit', () => finish([])); @@ -124,13 +165,23 @@ export class WebSocketScanner extends EventEmitter { } await Promise.all(workerPromises); - await this.cleanupWorkers(); - return foundServices; } async cleanupWorkers() { + if (this.inProcessWorker) { + this.inProcessWorker.stop(); + this.inProcessWorker = null; + } + const workers = this.workers.splice(0, this.workers.length); - await Promise.all(workers.map((worker) => worker.terminate())); + await Promise.all(workers.map((worker) => this.terminateWorker(worker))); + } + + terminateWorker(worker) { + return Promise.race([ + Promise.resolve().then(() => worker.terminate()), + new Promise((resolve) => setTimeout(resolve, 500)) + ]).catch(() => {}); } getLocalNetworkRange() { @@ -140,13 +191,16 @@ export class WebSocketScanner extends EventEmitter { for (const name of Object.keys(interfaces)) { for (const iface of interfaces[name]) { - if (iface.family === 'IPv4' && !iface.internal) { + const isIPv4 = iface.family === 'IPv4' || iface.family === 4; + if (isIPv4 && !iface.internal) { localIP = iface.address; subnetMask = iface.netmask; break; } } - if (localIP) break; + if (localIP) { + break; + } } if (!localIP) { diff --git a/src/network/websocketScannerWorker.js b/src/network/websocketScannerWorker.js index 9ed8a42..4fa2e13 100644 --- a/src/network/websocketScannerWorker.js +++ b/src/network/websocketScannerWorker.js @@ -1,61 +1,129 @@ import { parentPort } from 'worker_threads'; -import WebSocket from 'ws'; import net from 'net'; import http from 'http'; import https from 'https'; -import { loadConfig } from "../config.js"; -import log4js from "log4js"; import dns from 'dns'; +import { loadConfig } from '../config.js'; +import log4js from 'log4js'; const config = loadConfig(); +const isBun = typeof globalThis.Bun !== 'undefined'; -class WebSocketScannerWorker { - constructor() { +let WebSocketImpl; + +async function getWebSocket() { + if (WebSocketImpl) { + return WebSocketImpl; + } + + if (typeof globalThis.WebSocket === 'function') { + WebSocketImpl = globalThis.WebSocket; + return WebSocketImpl; + } + + const mod = await import('ws'); + WebSocketImpl = mod.WebSocket || mod.default; + return WebSocketImpl; +} + +function yieldToEventLoop() { + if (typeof globalThis.Bun?.sleep === 'function') { + return Bun.sleep(0); + } + + return new Promise((resolve) => setImmediate(resolve)); +} + +export class WebSocketScannerWorker { + constructor(reporter = parentPort) { this.randomId = Math.floor(Math.random() * 1000).toString().padStart(3, '0'); this.logger = log4js.getLogger(`WS Scanner #${this.randomId}`); this.logger.level = config.logLevel; this.timeout = 2000; + this.reporter = reporter; + this.stopped = false; } - async scanRange(startIP, endIP, port, protocol, objectType) { + stop() { + this.stopped = true; + } + + report(message) { + if (!this.reporter) { + return; + } + + if (typeof this.reporter.postMessage === 'function') { + this.reporter.postMessage(message); + return; + } + + if (typeof this.reporter === 'function') { + this.reporter(message); + } + } + + async scanRange(startIP, endIP, port, protocol, objectType, { concurrency = 1 } = {}) { const start = this.ipToNumber(startIP); const end = this.ipToNumber(endIP); const foundServices = []; + const total = Math.max(0, end - start + 1); + const probeCount = Math.max(1, Math.min(concurrency, total || 1)); + let next = start; this.logger.info( `Scanning ${startIP} - ${endIP} on port: ${port} using: ${protocol} objectType: ${objectType}` ); - for (let ip = start; ip <= end; ip++) { - const currentIP = this.numberToIP(ip); - - try { - this.logger.debug(`Checking ${currentIP} on port ${port}`); - const isOpen = await this.probe(currentIP, port, protocol, objectType); - if (isOpen) { - const hostname = await this.resolveHostname(currentIP); - foundServices.push({ ip: currentIP, hostname }); - this.logger.info( - `Service found at ${currentIP}${hostname ? ` (${hostname})` : ''}` - ); - parentPort.postMessage({ type: 'serviceFound', ip: currentIP, hostname }); + const runProbe = async () => { + while (!this.stopped) { + const ip = next++; + if (ip > end) { + break; } - } catch (error) { - // Connection failed, continue scanning + + const currentIP = this.numberToIP(ip); + + try { + this.logger.debug(`Checking ${currentIP} on port ${port}`); + const isOpen = await this.probe(currentIP, port, protocol, objectType); + if (this.stopped) { + break; + } + + if (isOpen) { + const hostname = await this.resolveHostname(currentIP); + foundServices.push({ ip: currentIP, hostname }); + this.logger.info( + `Service found at ${currentIP}${hostname ? ` (${hostname})` : ''}` + ); + this.report({ type: 'serviceFound', ip: currentIP, hostname }); + } + } catch (error) { + // Connection failed, continue scanning + } + + this.report({ + type: 'scanProgress', + currentIP, + increment: 1 + }); + + await yieldToEventLoop(); } + }; - parentPort.postMessage({ - type: 'scanProgress', - currentIP, - increment: 1 - }); - } - + await Promise.all(Array.from({ length: probeCount }, () => runProbe())); return foundServices; } async probe(ip, port, protocol, objectType) { - if (protocol === 'serial' || protocol === 'system') { + if (this.stopped || protocol === 'serial' || protocol === 'system') { + return false; + } + + const tcpOpen = await this.checkTcp(ip, port); + if (!tcpOpen || this.stopped) { return false; } @@ -68,7 +136,7 @@ class WebSocketScannerWorker { return this.checkHttp(ip, port, protocol); } - return this.checkTcp(ip, port); + return true; } async resolveHostname(ip) { @@ -83,54 +151,106 @@ class WebSocketScannerWorker { } } - checkWebSocket(url) { + async checkWebSocket(url) { + const WebSocketClient = await getWebSocket(); + return new Promise((resolve) => { - const ws = new WebSocket(url); - const timeout = setTimeout(() => { - ws.terminate(); - resolve(false); - }, this.timeout); + let settled = false; + let ws; - ws.on('open', () => { - clearTimeout(timeout); - ws.close(); - resolve(true); - }); + const finish = (ok) => { + if (settled) { + return; + } + settled = true; + clearTimeout(timer); - ws.on('error', () => { - clearTimeout(timeout); - resolve(false); - }); + if (ws) { + try { + ws.onopen = null; + ws.onerror = null; + ws.onclose = null; + if (typeof ws.removeAllListeners === 'function') { + ws.removeAllListeners(); + } + if (typeof ws.terminate === 'function') { + ws.terminate(); + } else { + ws.close(); + } + } catch { + // Ignore close errors from already-dead sockets + } + } + + resolve(ok); + }; + + const timer = setTimeout(() => finish(false), this.timeout); + + try { + ws = new WebSocketClient(url); + } catch { + finish(false); + return; + } + + if (typeof ws.on === 'function') { + ws.on('open', () => finish(true)); + ws.on('error', () => finish(false)); + ws.on('close', () => finish(false)); + } else { + ws.onopen = () => finish(true); + ws.onerror = () => finish(false); + ws.onclose = () => finish(false); + } }); } checkTcp(ip, port) { return new Promise((resolve) => { const socket = new net.Socket(); - const timeout = setTimeout(() => { - socket.destroy(); - resolve(false); - }, this.timeout); + let settled = false; - socket.once('connect', () => { - clearTimeout(timeout); + const finish = (ok) => { + if (settled) { + return; + } + settled = true; + clearTimeout(timer); + socket.removeAllListeners(); socket.destroy(); - resolve(true); - }); + resolve(ok); + }; - socket.once('error', () => { - clearTimeout(timeout); - socket.destroy(); - resolve(false); - }); + const timer = setTimeout(() => finish(false), this.timeout); - socket.connect(port, ip); + socket.once('connect', () => finish(true)); + socket.once('error', () => finish(false)); + socket.once('timeout', () => finish(false)); + + try { + socket.connect({ port, host: ip, family: 4 }); + } catch { + finish(false); + } }); } checkHttp(ip, port, protocol) { const client = protocol === 'https' ? https : http; return new Promise((resolve) => { + let settled = false; + + const finish = (ok) => { + if (settled) { + return; + } + settled = true; + req.destroy(); + resolve(ok); + }; + const req = client.request( { host: ip, @@ -142,19 +262,12 @@ class WebSocketScannerWorker { }, (res) => { res.resume(); - resolve(true); + finish(true); } ); - req.on('timeout', () => { - req.destroy(); - resolve(false); - }); - - req.on('error', () => { - resolve(false); - }); - + req.on('timeout', () => finish(false)); + req.on('error', () => finish(false)); req.end(); }); } @@ -174,16 +287,21 @@ class WebSocketScannerWorker { } } -parentPort.on('message', async (data) => { - if (data.type === 'scan') { - const scanner = new WebSocketScannerWorker(); +if (parentPort) { + parentPort.on('message', async (data) => { + if (data.type !== 'scan') { + return; + } + + const scanner = new WebSocketScannerWorker(parentPort); const results = await scanner.scanRange( data.startIP, data.endIP, data.port, data.protocol, - data.objectType + data.objectType, + { concurrency: isBun ? Math.max(1, data.concurrency || 1) : 1 } ); parentPort.postMessage({ type: 'scanComplete', results }); - } -}); + }); +}