diff --git a/src/client/endpoints.ts b/src/client/endpoints.ts index 9b60f833..333d9ec0 100644 --- a/src/client/endpoints.ts +++ b/src/client/endpoints.ts @@ -27,15 +27,15 @@ export function videoGenerateV2Endpoint(baseUrl: string): string { } export function videoTaskEndpoint(baseUrl: string, taskId: string): string { - return `${baseUrl}/v1/query/video_generation?task_id=${taskId}`; + return `${baseUrl}/v1/query/video_generation?task_id=${encodeURIComponent(taskId)}`; } export function videoTaskV2Endpoint(baseUrl: string, taskId: string): string { - return `${baseUrl}/v2/query/video_generation/${taskId}`; + return `${baseUrl}/v2/query/video_generation/${encodeURIComponent(taskId)}`; } export function fileRetrieveEndpoint(baseUrl: string, fileId: string): string { - return `${baseUrl}/v1/files/retrieve?file_id=${fileId}`; + return `${baseUrl}/v1/files/retrieve?file_id=${encodeURIComponent(fileId)}`; } export function searchEndpoint(baseUrl: string): string { diff --git a/src/commands/speech/synthesize.ts b/src/commands/speech/synthesize.ts index d095201b..4f8f39e2 100644 --- a/src/commands/speech/synthesize.ts +++ b/src/commands/speech/synthesize.ts @@ -9,6 +9,7 @@ import { writeFileSync } from 'fs'; import { readTextFromPathOrStdin } from '../../utils/fs'; import { T2A_FORMATS, formatList, validateAudioFormat, validateT2AStreaming, t2aDefaultSampleRate } from '../../utils/audio-formats'; import { pipeAudioStream } from '../../utils/audio-stream'; +import { validateSafeUrl } from '../../utils/network'; import type { Config } from '../../config/schema'; import type { GlobalFlags } from '../../types/flags'; import type { SpeechRequest, SpeechResponse } from '../../types/api'; @@ -126,6 +127,7 @@ export default defineCommand({ // Download and save subtitle file when --subtitles is requested if (flags.subtitles && response.data.subtitle_file) { try { + await validateSafeUrl(response.data.subtitle_file); // Download the subtitle JSON file from the URL const subtitleRes = await fetch(response.data.subtitle_file); if (!subtitleRes.ok) { diff --git a/src/config/loader.ts b/src/config/loader.ts index af9a7a7a..feb5d728 100644 --- a/src/config/loader.ts +++ b/src/config/loader.ts @@ -1,4 +1,5 @@ import { copyFileSync, existsSync, readFileSync, renameSync, unlinkSync, writeFileSync } from 'fs'; +import { randomBytes } from 'crypto'; import { parseConfigFile, REGIONS, type Config, type ConfigFile, type Region } from './schema'; import { ensureConfigDir, getConfigPath } from './paths'; import { detectOutputFormat, type OutputFormat } from '../output/formatter'; @@ -56,9 +57,15 @@ export async function writeConfigFile( ): Promise { await ensureConfigDir(); const path = getConfigPath(); - const tmp = path + '.tmp'; - writeFileSync(tmp, JSON.stringify(data, null, 2) + '\n', { mode: 0o600 }); - renameWithCrossDeviceFallback(tmp, path, renameOps); + const uniqueSuffix = randomBytes(8).toString('hex'); + const tmp = `${path}.${process.pid}.${Date.now()}.${uniqueSuffix}.tmp`; + writeFileSync(tmp, JSON.stringify(data, null, 2) + '\n', { mode: 0o600, flag: 'wx' }); + try { + renameWithCrossDeviceFallback(tmp, path, renameOps); + } catch (err) { + try { unlinkSync(tmp); } catch { /* best effort */ } + throw err; + } } export function loadConfig(flags: GlobalFlags): Config { diff --git a/src/files/download.ts b/src/files/download.ts index 246b9332..2d1db75f 100644 --- a/src/files/download.ts +++ b/src/files/download.ts @@ -3,6 +3,7 @@ import type { WriteStream } from 'fs'; import { createProgressBar } from '../output/progress'; import { CLIError } from '../errors/base'; import { ExitCode } from '../errors/codes'; +import { validateSafeUrl } from '../utils/network'; const DEFAULT_OVERALL_TIMEOUT_MS = 30 * 60 * 1000; const DEFAULT_IDLE_TIMEOUT_MS = 60 * 1000; @@ -17,6 +18,7 @@ export interface DownloadOpts { idleTimeoutMs?: number; maxBytes?: number; signal?: AbortSignal; + allowPrivate?: boolean; } class RetryableDownloadError extends CLIError { @@ -296,7 +298,7 @@ async function attemptDownload( ? undefined : declaredLength; tmpPath = `${destPath}.tmp-${process.pid}-${Date.now()}-${attempt}-${Math.random().toString(36).slice(2)}`; - writer = createWriteStream(tmpPath); + writer = createWriteStream(tmpPath, { flags: 'wx', mode: 0o600 }); progress = expectedLength && !opts.quiet ? createProgressBar(expectedLength, 'Downloading') : null; @@ -394,6 +396,7 @@ export async function downloadFile( ): Promise<{ size: number }> { // Alibaba Cloud OSS US East blocks HTTP from certain regions. const downloadUrl = url.startsWith('http://') ? url.replace('http://', 'https://') : url; + await validateSafeUrl(downloadUrl, { allowPrivate: opts?.allowPrivate }); const maxRetries = nonNegativeInteger(opts?.retries ?? 3, 'retries'); const baseDelay = nonNegativeNumber(opts?.retryDelayMs ?? 1000, 'retryDelayMs'); const overallTimeoutMs = positiveNumber( diff --git a/src/update/self-update.ts b/src/update/self-update.ts index 1b401ce5..10a94eec 100644 --- a/src/update/self-update.ts +++ b/src/update/self-update.ts @@ -1,4 +1,4 @@ -import { createWriteStream, renameSync, chmodSync, existsSync } from 'fs'; +import { createWriteStream, renameSync, chmodSync, existsSync, mkdtempSync, rmSync } from 'fs'; import { join } from 'path'; import { tmpdir } from 'os'; import { CLIError } from '../errors/base'; @@ -110,7 +110,7 @@ export async function downloadFile(url: string, dest: string, onProgress?: (pct: const total = Number(res.headers.get('content-length') ?? 0); let received = 0; - const writer = createWriteStream(dest); + const writer = createWriteStream(dest, { flags: 'wx', mode: 0o700 }); const reader = res.body!.getReader(); try { @@ -150,7 +150,10 @@ export async function downloadFile(url: string, dest: string, onProgress?: (pct: }); } // Don't leave a half-downloaded binary in /tmp on failure. - try { (await import('fs')).unlinkSync(dest); } catch { /* best-effort — race with concurrent cleanup is fine */ } + // If the error was EEXIST, dest already existed and was not created by this download. + if ((err as { code?: string })?.code !== 'EEXIST') { + try { (await import('fs')).unlinkSync(dest); } catch { /* best-effort — race with concurrent cleanup is fine */ } + } throw err; } finally { // Always release the Web Streams reader lock — the API contract requires @@ -177,40 +180,45 @@ export async function resolveUpdateTarget(channel: Channel): Promise { - const tmp = join(tmpdir(), `mmx-update-${Date.now()}`); - - process.stderr.write(`Downloading ${target.version}...\n`); - let lastPct = -1; - await downloadFile(target.downloadUrl, tmp, (pct) => { - if (pct !== lastPct && pct % 10 === 0) { - process.stderr.write(` ${pct}%\r`); - lastPct = pct; - } - }); - process.stderr.write(' \r'); + const tempDir = mkdtempSync(join(tmpdir(), 'mmx-update-')); + const tmp = join(tempDir, 'mmx-bin'); - process.stderr.write('Verifying checksum...\n'); - await verifySha256(tmp, target.checksum); + try { + process.stderr.write(`Downloading ${target.version}...\n`); + let lastPct = -1; + await downloadFile(target.downloadUrl, tmp, (pct) => { + if (pct !== lastPct && pct % 10 === 0) { + process.stderr.write(` ${pct}%\r`); + lastPct = pct; + } + }); + process.stderr.write(' \r'); - chmodSync(tmp, 0o755); + process.stderr.write('Verifying checksum...\n'); + await verifySha256(tmp, target.checksum); - // Atomic replace: rename works on same filesystem - // If cross-device, fall back to copy+rename - try { - renameSync(tmp, currentBin); - } catch { - const { copyFileSync, unlinkSync } = await import('fs'); - const backup = `${currentBin}.bak`; - copyFileSync(currentBin, backup); + chmodSync(tmp, 0o755); + + // Atomic replace: rename works on same filesystem + // If cross-device, fall back to copy+rename try { - copyFileSync(tmp, currentBin); - chmodSync(currentBin, 0o755); - unlinkSync(tmp); - if (existsSync(backup)) unlinkSync(backup); - } catch (e) { - // Restore backup - if (existsSync(backup)) renameSync(backup, currentBin); - throw e; + renameSync(tmp, currentBin); + } catch { + const { copyFileSync, unlinkSync } = await import('fs'); + const backup = `${currentBin}.bak`; + copyFileSync(currentBin, backup); + try { + copyFileSync(tmp, currentBin); + chmodSync(currentBin, 0o755); + unlinkSync(tmp); + if (existsSync(backup)) unlinkSync(backup); + } catch (e) { + // Restore backup + if (existsSync(backup)) renameSync(backup, currentBin); + throw e; + } } + } finally { + try { rmSync(tempDir, { recursive: true, force: true }); } catch { /* best-effort cleanup */ } } } diff --git a/src/utils/image.ts b/src/utils/image.ts index ab53b632..98ae6ae2 100644 --- a/src/utils/image.ts +++ b/src/utils/image.ts @@ -2,6 +2,7 @@ import { readFileSync, existsSync, statSync } from 'fs'; import { extname } from 'path'; import { CLIError } from '../errors/base'; import { ExitCode } from '../errors/codes'; +import { validateSafeUrl } from './network'; export const IMAGE_MIME_TYPES: Record = { '.jpg': 'image/jpeg', @@ -41,18 +42,46 @@ export async function toDataUri(image: string): Promise { if (image.startsWith('data:')) return image; if (image.startsWith('http://') || image.startsWith('https://')) { + await validateSafeUrl(image); const res = await fetch(image); if (!res.ok) throw new CLIError(`Failed to download image: HTTP ${res.status}`, ExitCode.GENERAL); - const contentType = res.headers.get('content-type') || 'image/jpeg'; - const mime = contentType.split(';')[0]!.trim(); - const buf = await res.arrayBuffer(); - if (buf.byteLength > MAX_IMAGE_SIZE_BYTES) { + + const declaredLength = res.headers.get('content-length'); + if (declaredLength && Number(declaredLength) > MAX_IMAGE_SIZE_BYTES) { throw new CLIError( - `Image too large (${(buf.byteLength / 1024 / 1024).toFixed(1)} MB). Maximum is 50 MB.`, + `Image too large (${(Number(declaredLength) / 1024 / 1024).toFixed(1)} MB). Maximum is 50 MB.`, ExitCode.USAGE, ); } - return `data:${mime};base64,${Buffer.from(buf).toString('base64')}`; + + const contentType = res.headers.get('content-type') || 'image/jpeg'; + const mime = contentType.split(';')[0]!.trim(); + + const reader = res.body?.getReader(); + if (!reader) throw new CLIError('Failed to read image response body.', ExitCode.GENERAL); + + const chunks: Uint8Array[] = []; + let totalBytes = 0; + try { + while (true) { + const { done, value } = await reader.read(); + if (done) break; + totalBytes += value.byteLength; + if (totalBytes > MAX_IMAGE_SIZE_BYTES) { + await reader.cancel(); + throw new CLIError( + `Image too large (${(totalBytes / 1024 / 1024).toFixed(1)} MB). Maximum is 50 MB.`, + ExitCode.USAGE, + ); + } + chunks.push(value); + } + } finally { + reader.releaseLock(); + } + + const buf = Buffer.concat(chunks); + return `data:${mime};base64,${buf.toString('base64')}`; } if (!existsSync(image)) throw new CLIError(`File not found: ${image}`, ExitCode.USAGE); diff --git a/src/utils/network.ts b/src/utils/network.ts new file mode 100644 index 00000000..58582a39 --- /dev/null +++ b/src/utils/network.ts @@ -0,0 +1,202 @@ +import { promises as dns } from 'dns'; +import { isIP } from 'net'; +import { CLIError } from '../errors/base'; +import { ExitCode } from '../errors/codes'; + +/** + * Checks whether an IPv4 address belongs to a private, loopback, link-local, + * CGNAT, or otherwise reserved/non-routable address space. + */ +export function isPrivateOrLoopbackIpv4(ip: string): boolean { + const parts = ip.split('.').map(p => Number(p)); + if (parts.length !== 4 || parts.some(p => !Number.isInteger(p) || p < 0 || p > 255)) { + return true; // Malformed IPv4 is treated as unsafe + } + + const [a, b] = parts as [number, number, number, number]; + + // 0.0.0.0/8 (Current network / "this" network) + if (a === 0) return true; + + // 10.0.0.0/8 (Private-Use - RFC 1918) + if (a === 10) return true; + + // 100.64.0.0/10 (Shared Address Space / CGNAT - RFC 6598) + if (a === 100 && b >= 64 && b <= 127) return true; + + // 127.0.0.0/8 (Loopback - RFC 1122) + if (a === 127) return true; + + // 169.254.0.0/16 (Link Local & Cloud Metadata - RFC 3927) + if (a === 169 && b === 254) return true; + + // 172.16.0.0/12 (Private-Use - RFC 1918) + if (a === 172 && b >= 16 && b <= 31) return true; + + // 192.0.0.0/24 (IETF Protocol Assignments) + if (a === 192 && b === 0 && parts[2] === 0) return true; + + // 192.0.2.0/24 (TEST-NET-1) + if (a === 192 && b === 0 && parts[2] === 2) return true; + + // 192.168.0.0/16 (Private-Use - RFC 1918) + if (a === 192 && b === 168) return true; + + // 198.18.0.0/15 (Benchmarking - RFC 2544) + if (a === 198 && (b === 18 || b === 19)) return true; + + // 198.51.100.0/24 (TEST-NET-2) + if (a === 198 && b === 51 && parts[2] === 100) return true; + + // 203.0.113.0/24 (TEST-NET-3) + if (a === 203 && b === 0 && parts[2] === 113) return true; + + // 224.0.0.0/4 (Multicast - RFC 5771) + if (a >= 224 && a <= 239) return true; + + // 240.0.0.0/4 (Reserved for future use - RFC 1112) + if (a >= 240) return true; + + return false; +} + +/** + * Checks whether an IPv6 address belongs to a private, loopback, link-local, + * unique-local, or otherwise reserved/non-routable address space. + */ +export function isPrivateOrLoopbackIpv6(ip: string): boolean { + const normalized = ip.toLowerCase().trim(); + + // Loopback (::1) + if (normalized === '::1' || normalized === '0:0:0:0:0:0:0:1') return true; + + // Unspecified (::) + if (normalized === '::' || normalized === '0:0:0:0:0:0:0:0') return true; + + // Unique Local Address (fc00::/7 -> fc00... to fdff...) + if (/^f[cd][0-9a-f]{2}:/i.test(normalized) || normalized.startsWith('fc') || normalized.startsWith('fd')) { + return true; + } + + // Link-Local Unicast (fe80::/10 -> fe80... to febf...) + if (/^fe[89ab][0-9a-f]:/i.test(normalized)) { + return true; + } + + // Multicast (ff00::/8) + if (normalized.startsWith('ff')) { + return true; + } + + // IPv4-mapped IPv6 (::ffff:x.x.x.x) + const v4MappedMatch = normalized.match(/^::ffff:(\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})$/); + if (v4MappedMatch && v4MappedMatch[1]) { + return isPrivateOrLoopbackIpv4(v4MappedMatch[1]); + } + + // IPv4-compatible IPv6 (deprecated ::x.x.x.x) + const v4CompatMatch = normalized.match(/^::(\d{1,3}\.\d{1,3}\.\d{1,3}\.\d{1,3})$/); + if (v4CompatMatch && v4CompatMatch[1]) { + return isPrivateOrLoopbackIpv4(v4CompatMatch[1]); + } + + return false; +} + +/** + * Checks whether an IP (v4 or v6) is private or loopback. + */ +export function isPrivateOrLoopbackIp(ip: string): boolean { + const version = isIP(ip); + if (version === 4) return isPrivateOrLoopbackIpv4(ip); + if (version === 6) return isPrivateOrLoopbackIpv6(ip); + return true; // Not a valid IP +} + +/** + * Checks whether a hostname is a known loopback/local/internal domain name. + */ +export function isDisallowedHostname(hostname: string): boolean { + const lower = hostname.toLowerCase().trim(); + + if (lower === 'localhost' || lower.endsWith('.localhost')) return true; + if (lower === 'local' || lower.endsWith('.local')) return true; + if (lower === 'internal' || lower.endsWith('.internal')) return true; + if (lower === 'invalid' || lower.endsWith('.invalid')) return true; + if (lower === 'onion' || lower.endsWith('.onion')) return true; + + return false; +} + +export interface ValidateUrlOptions { + allowPrivate?: boolean; +} + +/** + * Validates that a URL is safe to fetch: + * - Uses HTTP or HTTPS protocol + * - Does not target private, loopback, or cloud-metadata IPs + * - Does not target local/internal hostnames + * + * Resolves to the parsed URL object on success, or throws CLIError. + */ +export async function validateSafeUrl( + urlString: string, + options?: ValidateUrlOptions, +): Promise { + let parsed: URL; + try { + parsed = new URL(urlString); + } catch { + throw new CLIError(`Invalid URL: "${urlString}"`, ExitCode.USAGE); + } + + if (parsed.protocol !== 'http:' && parsed.protocol !== 'https:') { + throw new CLIError( + `Disallowed URL protocol "${parsed.protocol}". Only HTTP and HTTPS are permitted.`, + ExitCode.USAGE, + ); + } + + if (options?.allowPrivate) { + return parsed; + } + + const rawHost = parsed.hostname.replace(/^\[|\]$/g, ''); + + if (isDisallowedHostname(rawHost)) { + throw new CLIError( + `Access to local/internal host "${rawHost}" is disallowed for security reasons (SSRF protection).`, + ExitCode.USAGE, + ); + } + + const ipVersion = isIP(rawHost); + if (ipVersion !== 0) { + if (isPrivateOrLoopbackIp(rawHost)) { + throw new CLIError( + `Access to private/loopback IP address "${rawHost}" is disallowed for security reasons (SSRF protection).`, + ExitCode.USAGE, + ); + } + return parsed; + } + + // If it's a domain name, attempt asynchronous DNS resolution to guard against DNS rebinding + try { + const lookupResult = await dns.lookup(rawHost); + if (lookupResult && isPrivateOrLoopbackIp(lookupResult.address)) { + throw new CLIError( + `Host "${rawHost}" resolved to private/loopback IP address "${lookupResult.address}", which is disallowed (SSRF protection).`, + ExitCode.USAGE, + ); + } + } catch (err) { + // If the error was our CLIError rejection, rethrow it + if (err instanceof CLIError) throw err; + // Otherwise, DNS resolution may fail in offline or mocked testing environments. + // In that case, we let fetch proceed and handle connection/mock behavior naturally. + } + + return parsed; +} diff --git a/test/client/endpoints.test.ts b/test/client/endpoints.test.ts index 79002f4f..48feda44 100644 --- a/test/client/endpoints.test.ts +++ b/test/client/endpoints.test.ts @@ -1,5 +1,12 @@ import { describe, it, expect } from 'bun:test'; -import { fileUploadEndpoint, quotaEndpoint, usageEndpoint } from '../../src/client/endpoints'; +import { + fileRetrieveEndpoint, + fileUploadEndpoint, + quotaEndpoint, + usageEndpoint, + videoTaskEndpoint, + videoTaskV2Endpoint, +} from '../../src/client/endpoints'; describe('quotaEndpoint', () => { it('uses token_plan/remains for global', () => { @@ -30,3 +37,36 @@ describe('usageEndpoint', () => { expect(usageEndpoint('https://api.minimax.io', 'sk-api-abc')).toBe('https://api.minimax.io/account/query_balance'); }); }); + +describe('parameter encoding security in endpoints', () => { + it('encodes query parameters in videoTaskEndpoint to prevent injection', () => { + expect(videoTaskEndpoint('https://api.minimax.io', 'task123')) + .toBe('https://api.minimax.io/v1/query/video_generation?task_id=task123'); + + // Injection attempt with query parameter and fragment + const maliciousTask = '123&admin=true#fragment'; + expect(videoTaskEndpoint('https://api.minimax.io', maliciousTask)) + .toBe('https://api.minimax.io/v1/query/video_generation?task_id=123%26admin%3Dtrue%23fragment'); + }); + + it('encodes path parameters in videoTaskV2Endpoint to prevent path traversal', () => { + expect(videoTaskV2Endpoint('https://api.minimax.io', 'task456')) + .toBe('https://api.minimax.io/v2/query/video_generation/task456'); + + // Traversal or injection attempt + const maliciousTask = '../../admin?extra=1'; + expect(videoTaskV2Endpoint('https://api.minimax.io', maliciousTask)) + .toBe('https://api.minimax.io/v2/query/video_generation/..%2F..%2Fadmin%3Fextra%3D1'); + }); + + it('encodes query parameters in fileRetrieveEndpoint to prevent injection', () => { + expect(fileRetrieveEndpoint('https://api.minimax.io', 'file_xyz')) + .toBe('https://api.minimax.io/v1/files/retrieve?file_id=file_xyz'); + + // Injection attempt + const maliciousFile = 'file1&download=true&token=leak'; + expect(fileRetrieveEndpoint('https://api.minimax.io', maliciousFile)) + .toBe('https://api.minimax.io/v1/files/retrieve?file_id=file1%26download%3Dtrue%26token%3Dleak'); + }); +}); + diff --git a/test/commands/speech/synthesize-security.test.ts b/test/commands/speech/synthesize-security.test.ts new file mode 100644 index 00000000..a09f2d14 --- /dev/null +++ b/test/commands/speech/synthesize-security.test.ts @@ -0,0 +1,91 @@ +import { afterEach, describe, expect, it } from 'bun:test'; +import { mkdtempSync, rmSync } from 'fs'; +import { tmpdir } from 'os'; +import { join } from 'path'; +import { default as synthesizeCommand } from '../../../src/commands/speech/synthesize'; +import type { Config } from '../../../src/config/schema'; + +const originalFetch = globalThis.fetch; +let stderrOutput = ''; +const originalStderrWrite = process.stderr.write; +const tempDirs: string[] = []; + +afterEach(() => { + globalThis.fetch = originalFetch; + process.stderr.write = originalStderrWrite; + stderrOutput = ''; + for (const dir of tempDirs.splice(0)) { + try { rmSync(dir, { recursive: true, force: true }); } catch { /* ignore */ } + } +}); + +describe('speech synthesize security: subtitle SSRF protection', () => { + it('blocks subtitle download when subtitle_file points to private or loopback destination', async () => { + const attemptedFetchUrls: string[] = []; + stderrOutput = ''; + process.stderr.write = ((chunk: unknown) => { + stderrOutput += String(chunk); + return true; + }) as typeof process.stderr.write; + + globalThis.fetch = (async (input: RequestInfo | URL) => { + const urlStr = typeof input === 'string' ? input : (input as URL).href; + attemptedFetchUrls.push(urlStr); + + if (urlStr.includes('/v1/t2a_v2')) { + // Return synthetic API response with malicious subtitle_file URL + return new Response(JSON.stringify({ + data: { + audio: '48656c6c6f', // hex "Hello" + subtitle_file: 'http://169.254.169.254/latest/meta-data', + }, + base_resp: { status_code: 0, status_msg: 'success' }, + }), { + status: 200, + headers: { 'content-type': 'application/json' }, + }); + } + + return new Response('internal data', { status: 200 }); + }) as unknown as typeof fetch; + + const config: Config = { + apiKey: 'test-key', + region: 'global', + baseUrl: 'https://api.minimax.io', + output: 'text', + timeout: 10, + verbose: false, + quiet: false, + noColor: true, + yes: false, + dryRun: false, + nonInteractive: true, + async: false, + }; + + const dir = mkdtempSync(join(tmpdir(), 'mmx-speech-test-')); + tempDirs.push(dir); + const outPath = join(dir, 'test.mp3'); + + await synthesizeCommand.execute(config, { + quiet: false, + verbose: false, + noColor: true, + yes: false, + dryRun: false, + help: false, + nonInteractive: true, + async: false, + text: 'Hello world', + subtitles: true, + out: outPath, + }); + + // Verify fetch was NEVER called for the malicious metadata URL + expect(attemptedFetchUrls).not.toContain('http://169.254.169.254/latest/meta-data'); + // Verify stderr warned about the failure + expect(stderrOutput).toContain('Warning: failed to download subtitles'); + expect(stderrOutput).toContain('SSRF protection'); + }); +}); diff --git a/test/config/loader-security.test.ts b/test/config/loader-security.test.ts new file mode 100644 index 00000000..1cf5491d --- /dev/null +++ b/test/config/loader-security.test.ts @@ -0,0 +1,53 @@ +import { describe, it, expect, beforeEach, afterEach } from 'bun:test'; +import { existsSync, mkdirSync, readFileSync, rmSync, writeFileSync } from 'fs'; +import { join } from 'path'; +import { tmpdir } from 'os'; +import { writeConfigFile } from '../../src/config/loader'; + +describe('writeConfigFile security (unpredictable temp file and symlink resistance)', () => { + const testDir = join(tmpdir(), `mmx-config-sec-test-${Date.now()}`); + const originalConfigDir = process.env.MMX_CONFIG_DIR; + + beforeEach(() => { + process.env.MMX_CONFIG_DIR = testDir; + mkdirSync(testDir, { recursive: true }); + }); + + afterEach(() => { + if (originalConfigDir === undefined) delete process.env.MMX_CONFIG_DIR; + else process.env.MMX_CONFIG_DIR = originalConfigDir; + try { rmSync(testDir, { recursive: true, force: true }); } catch { /* ignore */ } + }); + + it('does not write to a predictable config.json.tmp path', async () => { + const decoyTmp = join(testDir, 'config.json.tmp'); + writeFileSync(decoyTmp, 'DECOY_CONTENT_SHOULD_NOT_BE_TOUCHED'); + + await writeConfigFile({ region: 'cn', output: 'json' }); + + const finalConfig = join(testDir, 'config.json'); + expect(existsSync(finalConfig)).toBe(true); + expect(JSON.parse(readFileSync(finalConfig, 'utf-8')).region).toBe('cn'); + + // The decoy config.json.tmp was NOT overwritten + expect(readFileSync(decoyTmp, 'utf-8')).toBe('DECOY_CONTENT_SHOULD_NOT_BE_TOUCHED'); + }); + + it('cleans up temporary files if rename throws an error', async () => { + const failingRename = { + rename: () => { + const err = new Error('Disk full'); + throw err; + }, + copy: () => {}, + unlink: () => {}, + }; + + await expect(writeConfigFile({ test: 'fail' }, failingRename)).rejects.toThrow('Disk full'); + + // No leftover temporary files in testDir + const files = (await import('fs')).readdirSync(testDir); + const leftoverTmps = files.filter(f => f.endsWith('.tmp')); + expect(leftoverTmps).toEqual([]); + }); +}); diff --git a/test/files/download-ssrf.test.ts b/test/files/download-ssrf.test.ts new file mode 100644 index 00000000..fdbc9623 --- /dev/null +++ b/test/files/download-ssrf.test.ts @@ -0,0 +1,77 @@ +import { afterEach, describe, expect, it } from 'bun:test'; +import { mkdtempSync, rmSync } from 'fs'; +import { tmpdir } from 'os'; +import { join } from 'path'; +import { downloadFile } from '../../src/files/download'; +import { CLIError } from '../../src/errors/base'; + +const originalFetch = globalThis.fetch; +const tempDirs: string[] = []; + +function makeTempDir(): string { + const dir = mkdtempSync(join(tmpdir(), 'mmx-download-ssrf-test-')); + tempDirs.push(dir); + return dir; +} + +afterEach(() => { + globalThis.fetch = originalFetch; + for (const dir of tempDirs.splice(0)) { + rmSync(dir, { recursive: true, force: true }); + } +}); + +describe('media download SSRF and security controls', () => { + it('rejects download URLs targeting loopback IP addresses', async () => { + const dir = makeTempDir(); + const dest = join(dir, 'output.mp4'); + + await expect(downloadFile('http://127.0.0.1:8080/video.mp4', dest, { quiet: true })) + .rejects.toThrow(CLIError); + await expect(downloadFile('http://127.0.0.1:8080/video.mp4', dest, { quiet: true })) + .rejects.toThrow(/SSRF protection/); + }); + + it('rejects download URLs targeting localhost', async () => { + const dir = makeTempDir(); + const dest = join(dir, 'output.mp4'); + + await expect(downloadFile('http://localhost:9000/video.mp4', dest, { quiet: true })) + .rejects.toThrow(/SSRF protection/); + }); + + it('rejects download URLs targeting cloud metadata service (169.254.169.254)', async () => { + const dir = makeTempDir(); + const dest = join(dir, 'output.mp4'); + + await expect(downloadFile('http://169.254.169.254/latest/meta-data', dest, { quiet: true })) + .rejects.toThrow(/SSRF protection/); + }); + + it('rejects download URLs targeting private RFC 1918 networks', async () => { + const dir = makeTempDir(); + const dest = join(dir, 'output.mp4'); + + await expect(downloadFile('http://10.200.1.1/video.mp4', dest, { quiet: true })) + .rejects.toThrow(/SSRF protection/); + await expect(downloadFile('http://192.168.0.10/video.mp4', dest, { quiet: true })) + .rejects.toThrow(/SSRF protection/); + }); + + it('permits private destination when allowPrivate is explicitly true', async () => { + const dir = makeTempDir(); + const dest = join(dir, 'output.mp4'); + + globalThis.fetch = (async () => new Response('video-bytes', { + status: 200, + headers: { 'content-length': '11' }, + })) as unknown as typeof fetch; + + const result = await downloadFile('http://127.0.0.1:8080/video.mp4', dest, { + quiet: true, + allowPrivate: true, + }); + + expect(result.size).toBe(11); + }); +}); diff --git a/test/update/self-update-security.test.ts b/test/update/self-update-security.test.ts new file mode 100644 index 00000000..8acb8a04 --- /dev/null +++ b/test/update/self-update-security.test.ts @@ -0,0 +1,42 @@ +import { afterEach, describe, expect, it } from 'bun:test'; +import { mkdtempSync, readFileSync, rmSync, writeFileSync } from 'fs'; +import { tmpdir } from 'os'; +import { join } from 'path'; +import { downloadFile } from '../../src/update/self-update'; + +const originalFetch = globalThis.fetch; +const tempDirs: string[] = []; + +function makeTempDir(): string { + const dir = mkdtempSync(join(tmpdir(), 'mmx-update-sec-test-')); + tempDirs.push(dir); + return dir; +} + +afterEach(() => { + globalThis.fetch = originalFetch; + for (const dir of tempDirs.splice(0)) { + try { rmSync(dir, { recursive: true, force: true }); } catch { /* ignore */ } + } +}); + +describe('self-update security: safe write semantics and symlink resistance', () => { + it('refuses to overwrite an existing file or symlink using flags: wx', async () => { + const dir = makeTempDir(); + const dest = join(dir, 'mmx-update-target'); + + // Pre-create target (simulating attacker planting a symlink or file) + writeFileSync(dest, 'ORIGINAL_FILE_CONTENT'); + + globalThis.fetch = (async () => new Response('ATTACKER_PAYLOAD', { + status: 200, + headers: { 'content-length': '16' }, + })) as unknown as typeof fetch; + + // Must reject because file already exists (flags: wx) + await expect(downloadFile('https://example.com/binary', dest)).rejects.toThrow(); + + // Original file content must be intact and not overwritten + expect(readFileSync(dest, 'utf-8')).toBe('ORIGINAL_FILE_CONTENT'); + }); +}); diff --git a/test/utils/image-security.test.ts b/test/utils/image-security.test.ts new file mode 100644 index 00000000..878cb74f --- /dev/null +++ b/test/utils/image-security.test.ts @@ -0,0 +1,73 @@ +import { afterEach, describe, expect, it } from 'bun:test'; +import { toDataUri } from '../../src/utils/image'; +import { CLIError } from '../../src/errors/base'; + +const originalFetch = globalThis.fetch; + +afterEach(() => { + globalThis.fetch = originalFetch; +}); + +describe('image security and validation (SSRF and bounded buffering)', () => { + it('rejects remote image URLs pointing to loopback addresses', async () => { + await expect(toDataUri('http://127.0.0.1:8080/avatar.png')).rejects.toThrow(CLIError); + await expect(toDataUri('http://127.0.0.1:8080/avatar.png')).rejects.toThrow(/SSRF protection/); + }); + + it('rejects remote image URLs pointing to localhost', async () => { + await expect(toDataUri('http://localhost:3000/image.jpg')).rejects.toThrow(/SSRF protection/); + }); + + it('rejects remote image URLs pointing to AWS/cloud metadata (169.254.169.254)', async () => { + await expect(toDataUri('http://169.254.169.254/secret.png')).rejects.toThrow(/SSRF protection/); + }); + + it('rejects remote image URLs pointing to private RFC 1918 addresses', async () => { + await expect(toDataUri('http://10.0.0.5/test.png')).rejects.toThrow(/SSRF protection/); + await expect(toDataUri('http://192.168.1.50/pic.jpg')).rejects.toThrow(/SSRF protection/); + }); + + it('rejects oversized image based on Content-Length before downloading full stream', async () => { + globalThis.fetch = (async () => new Response('tiny-body', { + status: 200, + headers: { + 'content-type': 'image/png', + 'content-length': String(51 * 1024 * 1024), // 51 MB + }, + })) as unknown as typeof fetch; + + await expect(toDataUri('https://example.com/huge.png')).rejects.toThrow(/Image too large/); + }); + + it('aborts and rejects oversized response stream when Content-Length is missing or dishonest', async () => { + let canceled = false; + const chunk = new Uint8Array(10 * 1024 * 1024); // 10MB chunk + const stream = new ReadableStream({ + pull(controller) { + controller.enqueue(chunk); + }, + cancel() { + canceled = true; + }, + }); + + globalThis.fetch = (async () => new Response(stream, { + status: 200, + headers: { 'content-type': 'image/png' }, + })) as unknown as typeof fetch; + + await expect(toDataUri('https://example.com/streaming-bomb.png')).rejects.toThrow(/Image too large/); + expect(canceled).toBe(true); + }); + + it('successfully converts valid remote image stream to data URI', async () => { + const imageData = new TextEncoder().encode('fake-image-bytes'); + globalThis.fetch = (async () => new Response(imageData, { + status: 200, + headers: { 'content-type': 'image/png' }, + })) as unknown as typeof fetch; + + const uri = await toDataUri('https://example.com/valid.png'); + expect(uri.startsWith('data:image/png;base64,')).toBe(true); + }); +}); diff --git a/test/utils/network.test.ts b/test/utils/network.test.ts new file mode 100644 index 00000000..7b636d25 --- /dev/null +++ b/test/utils/network.test.ts @@ -0,0 +1,129 @@ +import { describe, it, expect } from 'bun:test'; +import { + isPrivateOrLoopbackIp, + isPrivateOrLoopbackIpv4, + isPrivateOrLoopbackIpv6, + isDisallowedHostname, + validateSafeUrl, +} from '../../src/utils/network'; +import { CLIError } from '../../src/errors/base'; + +describe('network security: IP address validation', () => { + it('detects loopback IPv4 addresses', () => { + expect(isPrivateOrLoopbackIpv4('127.0.0.1')).toBe(true); + expect(isPrivateOrLoopbackIpv4('127.255.255.254')).toBe(true); + expect(isPrivateOrLoopbackIpv4('127.0.0.0')).toBe(true); + }); + + it('detects RFC 1918 private IPv4 addresses', () => { + // 10.0.0.0/8 + expect(isPrivateOrLoopbackIpv4('10.0.0.1')).toBe(true); + expect(isPrivateOrLoopbackIpv4('10.255.255.255')).toBe(true); + + // 172.16.0.0/12 + expect(isPrivateOrLoopbackIpv4('172.16.0.1')).toBe(true); + expect(isPrivateOrLoopbackIpv4('172.31.255.255')).toBe(true); + expect(isPrivateOrLoopbackIpv4('172.15.0.1')).toBe(false); + expect(isPrivateOrLoopbackIpv4('172.32.0.1')).toBe(false); + + // 192.168.0.0/16 + expect(isPrivateOrLoopbackIpv4('192.168.1.1')).toBe(true); + expect(isPrivateOrLoopbackIpv4('192.168.254.254')).toBe(true); + }); + + it('detects link-local and cloud metadata service (169.254.169.254)', () => { + expect(isPrivateOrLoopbackIpv4('169.254.169.254')).toBe(true); + expect(isPrivateOrLoopbackIpv4('169.254.0.1')).toBe(true); + }); + + it('detects current network, CGNAT, multicast, and broadcast', () => { + expect(isPrivateOrLoopbackIpv4('0.0.0.0')).toBe(true); + expect(isPrivateOrLoopbackIpv4('100.64.0.1')).toBe(true); + expect(isPrivateOrLoopbackIpv4('100.127.255.255')).toBe(true); + expect(isPrivateOrLoopbackIpv4('224.0.0.1')).toBe(true); + expect(isPrivateOrLoopbackIpv4('255.255.255.255')).toBe(true); + }); + + it('allows public IPv4 addresses', () => { + expect(isPrivateOrLoopbackIpv4('8.8.8.8')).toBe(false); + expect(isPrivateOrLoopbackIpv4('1.1.1.1')).toBe(false); + expect(isPrivateOrLoopbackIpv4('93.184.216.34')).toBe(false); + }); + + it('detects loopback, link-local, and unique-local IPv6 addresses', () => { + expect(isPrivateOrLoopbackIpv6('::1')).toBe(true); + expect(isPrivateOrLoopbackIpv6('::')).toBe(true); + expect(isPrivateOrLoopbackIpv6('fe80::1')).toBe(true); + expect(isPrivateOrLoopbackIpv6('fc00::1')).toBe(true); + expect(isPrivateOrLoopbackIpv6('fd12:3456:789a::1')).toBe(true); + }); + + it('detects IPv4-mapped IPv6 addresses for private IPv4', () => { + expect(isPrivateOrLoopbackIpv6('::ffff:127.0.0.1')).toBe(true); + expect(isPrivateOrLoopbackIpv6('::ffff:10.0.0.1')).toBe(true); + expect(isPrivateOrLoopbackIpv6('::ffff:169.254.169.254')).toBe(true); + expect(isPrivateOrLoopbackIpv6('::ffff:8.8.8.8')).toBe(false); + }); + + it('correctly handles isPrivateOrLoopbackIp general wrapper', () => { + expect(isPrivateOrLoopbackIp('127.0.0.1')).toBe(true); + expect(isPrivateOrLoopbackIp('::1')).toBe(true); + expect(isPrivateOrLoopbackIp('8.8.8.8')).toBe(false); + expect(isPrivateOrLoopbackIp('2606:4700:4700::1111')).toBe(false); + }); +}); + +describe('network security: hostname validation', () => { + it('detects disallowed hostnames', () => { + expect(isDisallowedHostname('localhost')).toBe(true); + expect(isDisallowedHostname('sub.localhost')).toBe(true); + expect(isDisallowedHostname('myhost.local')).toBe(true); + expect(isDisallowedHostname('server.internal')).toBe(true); + expect(isDisallowedHostname('test.invalid')).toBe(true); + expect(isDisallowedHostname('secret.onion')).toBe(true); + }); + + it('allows normal public hostnames', () => { + expect(isDisallowedHostname('api.minimax.io')).toBe(false); + expect(isDisallowedHostname('example.com')).toBe(false); + expect(isDisallowedHostname('github.com')).toBe(false); + }); +}); + +describe('validateSafeUrl', () => { + it('rejects unsupported protocols', async () => { + await expect(validateSafeUrl('file:///etc/passwd')).rejects.toThrow(CLIError); + await expect(validateSafeUrl('ftp://example.com/file')).rejects.toThrow(/Only HTTP and HTTPS/); + await expect(validateSafeUrl('data:text/plain;base64,abc')).rejects.toThrow(/Only HTTP and HTTPS/); + }); + + it('rejects malformed URLs', async () => { + await expect(validateSafeUrl('not a url')).rejects.toThrow(/Invalid URL/); + }); + + it('rejects loopback and private IP URLs', async () => { + await expect(validateSafeUrl('http://127.0.0.1:8080/admin')).rejects.toThrow(/SSRF protection/); + await expect(validateSafeUrl('http://10.1.2.3/internal')).rejects.toThrow(/SSRF protection/); + await expect(validateSafeUrl('http://192.168.1.1/config')).rejects.toThrow(/SSRF protection/); + await expect(validateSafeUrl('http://172.20.0.1/')).rejects.toThrow(/SSRF protection/); + await expect(validateSafeUrl('http://169.254.169.254/latest/meta-data/')).rejects.toThrow(/SSRF protection/); + await expect(validateSafeUrl('http://[::1]:8080/')).rejects.toThrow(/SSRF protection/); + }); + + it('rejects localhost and local hostnames', async () => { + await expect(validateSafeUrl('http://localhost:3000/')).rejects.toThrow(/SSRF protection/); + await expect(validateSafeUrl('http://api.local/data')).rejects.toThrow(/SSRF protection/); + await expect(validateSafeUrl('http://db.internal:5432/')).rejects.toThrow(/SSRF protection/); + }); + + it('accepts public HTTPS URLs', async () => { + const parsed = await validateSafeUrl('https://api.minimax.io/v1/models'); + expect(parsed.hostname).toBe('api.minimax.io'); + expect(parsed.protocol).toBe('https:'); + }); + + it('allows private URLs when allowPrivate option is enabled', async () => { + const parsed = await validateSafeUrl('http://127.0.0.1:8080/mock', { allowPrivate: true }); + expect(parsed.hostname).toBe('127.0.0.1'); + }); +});