diff --git a/app/src/main/java/io/nekohasekai/sagernet/fmt/olcrtc/OlcrtcFmt.kt b/app/src/main/java/io/nekohasekai/sagernet/fmt/olcrtc/OlcrtcFmt.kt index 89ea351c9..79fd014f9 100644 --- a/app/src/main/java/io/nekohasekai/sagernet/fmt/olcrtc/OlcrtcFmt.kt +++ b/app/src/main/java/io/nekohasekai/sagernet/fmt/olcrtc/OlcrtcFmt.kt @@ -28,9 +28,51 @@ package io.nekohasekai.sagernet.fmt.olcrtc */ private const val SCHEME = "olcrtc://" -private val SUPPORTED_TRANSPORTS = setOf("vp8channel", "datachannel") -private val SUPPORTED_CARRIERS = setOf("jitsi", "telemost", "wbstream") +private const val TRANSPORT_VP8 = "vp8channel" +private const val TRANSPORT_DATA = "datachannel" +private const val DEFAULT_VP8_FPS = 30 +private const val DEFAULT_VP8_BATCH = 8 +private val SUPPORTED_TRANSPORTS_BY_CARRIER = mapOf( + "jitsi" to setOf(TRANSPORT_VP8, TRANSPORT_DATA), + "telemost" to setOf(TRANSPORT_VP8), + "wbstream" to setOf(TRANSPORT_VP8), +) private val DELIMITERS = setOf('<', '>', '&', '=', '@', '#', '$', '?') +private val VP8_FPS_RANGE = 1..120 +private val VP8_BATCH_RANGE = 1..64 + +/** + * Validates fields shared by URI import/export, the profile editor, and runtime args. + * Plain upstream URIs do not carry a client id, so callers opt into that requirement. + */ +fun OlcrtcBean.validateOlcrtcProfile(requireClientId: Boolean = false) { + val carrierName = carrier.orEmpty() + val transportName = transport.orEmpty() + val room = roomId.orEmpty() + val identity = clientId.orEmpty() + val hex = keyHex.orEmpty() + val fps = vp8Fps ?: 0 + val batchSize = vp8BatchSize ?: 0 + val resolver = dnsServer.orEmpty() + + val allowedTransports = SUPPORTED_TRANSPORTS_BY_CARRIER[carrierName] + require(allowedTransports != null) { "olcRTC: unsupported carrier" } + require(transportName in allowedTransports) { "olcRTC: transport is not supported by carrier" } + require(room.isNotBlank()) { "olcRTC: room id / URL is required" } + if (requireClientId) { + require(identity.isNotBlank()) { "olcRTC: client id is required" } + } + require(hex.length == 64 && hex.all { it.isHexDigit() }) { + "olcRTC: encryption key must be 64 hex characters" + } + if (transportName == TRANSPORT_VP8) { + require(fps in VP8_FPS_RANGE) { "olcRTC: VP8 FPS must be between 1 and 120" } + require(batchSize in VP8_BATCH_RANGE) { "olcRTC: VP8 batch size must be between 1 and 64" } + } + require(resolver.isBlank() || resolver.isIpPortLiteral()) { + "olcRTC: DNS resolver must be an IP literal with a valid port" + } +} /** Parses an `olcrtc://` link into an [OlcrtcBean]. Fails fast on malformed input. */ fun parseOlcrtc(url: String): OlcrtcBean { @@ -59,10 +101,7 @@ fun parseOlcrtc(url: String): OlcrtcBean { require(q >= 0) { "invalid olcrtc link: missing '?' before transport" } val carrier = body.substring(0, q).trim() require(carrier.isNotEmpty()) { "invalid olcrtc link: empty carrier" } - require(carrier in SUPPORTED_CARRIERS) { - "olcrtc link unsupported carrier '$carrier' (supported: ${SUPPORTED_CARRIERS.joinToString()})" - } - var afterQ = body.substring(q + 1) + val afterQ = body.substring(q + 1) // roomId after the FIRST '@' that is NOT inside the `<...>` payload block. val payloadEnd = if (afterQ.startsWith("<") || afterQ.contains('<')) afterQ.indexOf('>') else -1 @@ -79,10 +118,14 @@ fun parseOlcrtc(url: String): OlcrtcBean { if (lt >= 0) { val gt = transportPart.indexOf('>', lt) require(gt >= 0) { "invalid olcrtc link: unterminated '<...>' transport payload" } + require(transportPart.substring(gt + 1).isBlank()) { + "invalid olcrtc link: unexpected text after transport payload" + } payload = transportPart.substring(lt + 1, gt) transportPart = transportPart.substring(0, lt) } - val transport = transportPart.trim().ifEmpty { "vp8channel" } + val transport = transportPart.trim().ifEmpty { TRANSPORT_VP8 } + val payloadValues = parsePayload(payload) return OlcrtcBean().apply { // serverAddress/Port are unused by this protocol; keep a stable placeholder. @@ -90,59 +133,69 @@ fun parseOlcrtc(url: String): OlcrtcBean { this.carrier = carrier this.roomId = roomId this.keyHex = keyHex - require(transport in SUPPORTED_TRANSPORTS) { - "olcrtc link unsupported transport '$transport' (supported: ${SUPPORTED_TRANSPORTS.joinToString()})" - } this.transport = transport name = comment + initializeDefaultValues() + + payloadValues.forEach { (key, value) -> + when (key) { + "vp8-fps" -> vp8Fps = value.toIntOrNull() + ?: throw IllegalArgumentException("olcRTC: VP8 FPS must be an integer") + + "vp8-batch" -> vp8BatchSize = value.toIntOrNull() + ?: throw IllegalArgumentException("olcRTC: VP8 batch size must be an integer") - parsePayload(payload).forEach { (k, v) -> - when (k) { - "vp8-fps" -> v.toIntOrNull()?.let { vp8Fps = it } - "vp8-batch" -> v.toIntOrNull()?.let { vp8BatchSize = it } // Our non-standard pairing-token carrier. - "cid", "client-id", "clientid" -> clientId = v + "cid", "client-id", "clientid" -> { + require(value.none { it in DELIMITERS }) { + "olcRTC: client id contains a reserved delimiter" + } + clientId = value + } } } - initializeDefaultValues() - - require(keyHex.isNotBlank()) { "olcrtc link missing encryption key" } - require(keyHex.length == 64 && keyHex.all { it.isHexDigit() }) { - "olcrtc link encryption key must be 64 hex characters" - } + validateOlcrtcProfile() } } /** Serializes an [OlcrtcBean] to a shareable `olcrtc://` link, carrying clientId as `&cid=`. */ fun OlcrtcBean.toUri(): String { - require(carrier.isNotBlank()) { "olcRTC: cannot build share link without a carrier" } - require(roomId.isNotBlank()) { "olcRTC: cannot build share link without a room id" } + validateOlcrtcProfile() + val carrierName = carrier.orEmpty() + val transportName = transport.orEmpty() + val room = roomId.orEmpty() + val shareClientId = clientId.orEmpty() + val hex = keyHex.orEmpty() + val fps = vp8Fps ?: DEFAULT_VP8_FPS + val batchSize = vp8BatchSize ?: DEFAULT_VP8_BATCH + val profileName = name.orEmpty() + // The URI uses bare delimiters with no escaping convention; refuse to emit a link that // would not round-trip rather than silently producing a corrupt one. - require(clientId.none { it in DELIMITERS }) { - "olcRTC: clientId must not contain any of: ${DELIMITERS.joinToString(" ")}" + require(shareClientId.none { it in DELIMITERS }) { + "olcRTC: client id contains a reserved delimiter" } // roomId is emitted raw before '#'; a '$' in it would be mis-parsed as the comment // delimiter on re-import. Refuse rather than emit a link that won't round-trip. - require(roomId.none { it == '$' }) { "olcRTC: room id must not contain '\$'" } - require(name.none { it == '$' }) { "olcRTC: profile name must not contain '\$'" } + require(room.none { it == '$' }) { "olcRTC: room id must not contain '\$'" } + require(profileName.none { it == '$' }) { "olcRTC: profile name must not contain '\$'" } val payloadParts = mutableListOf() - if (transport == "vp8channel") { + if (transportName == TRANSPORT_VP8) { val defaults = OlcrtcBean().apply { initializeDefaultValues() } - if (vp8Fps != defaults.vp8Fps) payloadParts += "vp8-fps=$vp8Fps" - if (vp8BatchSize != defaults.vp8BatchSize) payloadParts += "vp8-batch=$vp8BatchSize" + if (fps != defaults.vp8Fps) payloadParts += "vp8-fps=$fps" + if (batchSize != defaults.vp8BatchSize) payloadParts += "vp8-batch=$batchSize" } - if (clientId.isNotBlank()) payloadParts += "cid=$clientId" + if (shareClientId.isNotBlank()) payloadParts += "cid=$shareClientId" val payload = if (payloadParts.isEmpty()) "" else "<${payloadParts.joinToString("&")}>" val sb = StringBuilder(SCHEME) - sb.append(carrier).append('?').append(transport).append(payload) - sb.append('@').append(roomId) - sb.append('#').append(keyHex) - if (name.isNotBlank()) sb.append('$').append(name) + sb.append(carrierName).append('?').append(transportName).append(payload) + sb.append('@').append(room) + sb.append('#').append(hex) + if (profileName.isNotBlank()) sb.append('$').append(profileName) return sb.toString() } @@ -164,23 +217,36 @@ fun OlcrtcBean.buildOlcrtcArgs( dnsFallback: String, readyTimeoutMs: Long, ): List { - require(!carrier.isNullOrBlank()) { "olcRTC: carrier is required" } - require(!roomId.isNullOrBlank()) { "olcRTC: room id is required" } - val hex = keyHex ?: "" - require(hex.length == 64 && hex.all { it.isHexDigit() }) { - "olcRTC: encryption key must be 64 hex characters" + validateOlcrtcProfile(requireClientId = true) + val carrierName = carrier.orEmpty() + val transportName = transport.orEmpty() + val room = roomId.orEmpty() + val identity = clientId.orEmpty() + val hex = keyHex.orEmpty() + val fps = if (transportName == TRANSPORT_VP8) { + vp8Fps ?: DEFAULT_VP8_FPS + } else { + DEFAULT_VP8_FPS + } + val batchSize = if (transportName == TRANSPORT_VP8) { + vp8BatchSize ?: DEFAULT_VP8_BATCH + } else { + DEFAULT_VP8_BATCH + } + val resolver = dnsServer.orEmpty().ifBlank { dnsFallback } + require(resolver.isIpPortLiteral()) { + "olcRTC: DNS resolver must be an IP literal with a valid port" } - val transportName = if (transport in SUPPORTED_TRANSPORTS) transport else "vp8channel" val args = mutableListOf( - "-carrier", carrier, + "-carrier", carrierName, "-transport", transportName, - "-room", roomId, - "-client-id", clientId ?: "", + "-room", room, + "-client-id", identity, "-key", hex, "-socks-port", port.toString(), - "-dns", (dnsServer ?: "").ifBlank { dnsFallback }, - "-vp8-fps", vp8Fps.toString(), - "-vp8-batch", vp8BatchSize.toString(), + "-dns", resolver, + "-vp8-fps", fps.toString(), + "-vp8-batch", batchSize.toString(), "-protect-path", protectPath, "-ready-timeout-ms", readyTimeoutMs.toString(), ) @@ -199,11 +265,20 @@ fun OlcrtcBean.buildOlcrtcArgs( * may resolve further hosts at runtime; those rely on the sidecar's own protected * resolver. ICE candidates are typically raw IPs, so signaling is the common blocker. */ -fun OlcrtcBean.carrierHost(): String? = when (carrier) { +fun OlcrtcBean.carrierHost(): String? = when (carrier.orEmpty()) { "jitsi" -> { - // roomId is host/room or https://host/room; extract the host. - val s = roomId.substringAfter("://").trimStart('/') - s.substringBefore('/').substringBefore(':').ifBlank { null } + // Mirror upstream's permissive host/room split, but keep bracketed IPv6 intact. + val room = roomId.orEmpty().trim() + val authority = room.substringAfter("://", room).trimStart('/').substringBefore('/').trim() + when { + authority.isBlank() -> null + authority.startsWith('[') -> { + val closingBracket = authority.indexOf(']') + if (closingBracket <= 1) null else authority.substring(1, closingBracket) + } + authority.count { it == ':' } == 1 -> authority.substringBefore(':').ifBlank { null } + else -> authority + } } "telemost" -> "telemost.yandex.ru" "wbstream" -> "stream.wb.ru" @@ -212,10 +287,92 @@ fun OlcrtcBean.carrierHost(): String? = when (carrier) { private fun parsePayload(payload: String): Map { if (payload.isBlank()) return emptyMap() - return payload.split('&').mapNotNull { pair -> - val i = pair.indexOf('=') - if (i <= 0) null else pair.substring(0, i).trim() to pair.substring(i + 1).trim() - }.toMap() + return payload.split('&').associate { pair -> + val separator = pair.indexOf('=') + require(separator > 0) { "invalid olcrtc transport payload" } + val key = pair.substring(0, separator).trim() + require(key.isNotEmpty()) { "invalid olcrtc transport payload" } + key to pair.substring(separator + 1).trim() + } +} + +private fun String.isIpPortLiteral(): Boolean { + val (host, port) = if (startsWith('[')) { + val closingBracket = indexOf(']') + if (closingBracket <= 1 || closingBracket != lastIndexOf(']')) return false + if (closingBracket + 1 >= length || this[closingBracket + 1] != ':') return false + substring(1, closingBracket) to substring(closingBracket + 2) + } else { + val separator = indexOf(':') + if (separator <= 0 || separator != lastIndexOf(':')) return false + substring(0, separator) to substring(separator + 1) + } + if (!port.isValidPort()) return false + return if (startsWith('[')) host.isIpv6Literal() else host.isIpv4Literal() +} + +private fun String.isValidPort(): Boolean { + if (isEmpty() || any { it !in '0'..'9' }) return false + val value = toIntOrNull() ?: return false + return value in 1..65535 } +private fun String.isIpv4Literal(): Boolean { + val octets = split('.') + return octets.size == 4 && octets.all { octet -> + if (octet.isEmpty() || octet.any { it !in '0'..'9' }) return@all false + if (octet.length > 1 && octet.startsWith('0')) return@all false + val value = octet.toIntOrNull() ?: return@all false + value in 0..255 + } +} + +private fun String.isIpv6Literal(): Boolean { + val zoneSeparator = indexOf('%') + val address: String + if (zoneSeparator >= 0) { + if (zoneSeparator == 0 || zoneSeparator != lastIndexOf('%')) return false + val zone = substring(zoneSeparator + 1) + if (zone.isEmpty() || zone.any { !it.isSafeZoneCharacter() }) return false + address = substring(0, zoneSeparator) + } else { + address = this + } + if (':' !in address) return false + + val compression = address.indexOf("::") + if (compression != address.lastIndexOf("::")) return false + val left: List + val right: List + if (compression >= 0) { + left = address.substring(0, compression).ipv6Segments() ?: return false + right = address.substring(compression + 2).ipv6Segments() ?: return false + if (left.any { '.' in it }) return false + } else { + left = address.ipv6Segments() ?: return false + right = emptyList() + } + val segments = left + right + var groups = 0 + segments.forEachIndexed { index, segment -> + if ('.' in segment) { + if (index != segments.lastIndex || !segment.isIpv4Literal()) return false + groups += 2 + } else { + if (segment.length !in 1..4 || !segment.all { it.isHexDigit() }) return false + groups += 1 + } + } + return if (compression >= 0) groups < 8 else groups == 8 +} + +private fun String.ipv6Segments(): List? { + if (isEmpty()) return emptyList() + if (startsWith(':') || endsWith(':')) return null + return split(':').takeIf { segments -> segments.none { it.isEmpty() } } +} + +private fun Char.isSafeZoneCharacter(): Boolean = + this in 'a'..'z' || this in 'A'..'Z' || this in '0'..'9' || this == '_' || this == '-' || this == '.' + private fun Char.isHexDigit(): Boolean = this in '0'..'9' || this in 'a'..'f' || this in 'A'..'F' diff --git a/app/src/main/java/io/nekohasekai/sagernet/ui/profile/OlcrtcSettingsActivity.kt b/app/src/main/java/io/nekohasekai/sagernet/ui/profile/OlcrtcSettingsActivity.kt index 5447f815d..8f2698ec2 100644 --- a/app/src/main/java/io/nekohasekai/sagernet/ui/profile/OlcrtcSettingsActivity.kt +++ b/app/src/main/java/io/nekohasekai/sagernet/ui/profile/OlcrtcSettingsActivity.kt @@ -10,6 +10,7 @@ package io.nekohasekai.sagernet.ui.profile import android.os.Bundle +import android.widget.Toast import androidx.preference.EditTextPreference import androidx.preference.PreferenceFragmentCompat import io.nekohasekai.sagernet.Key @@ -17,7 +18,9 @@ import io.nekohasekai.sagernet.R import io.nekohasekai.sagernet.database.DataStore import io.nekohasekai.sagernet.database.preference.EditTextPreferenceModifiers import io.nekohasekai.sagernet.fmt.olcrtc.OlcrtcBean +import io.nekohasekai.sagernet.fmt.olcrtc.validateOlcrtcProfile import io.nekohasekai.sagernet.ktx.applyDefaultValues +import io.nekohasekai.sagernet.ktx.onMainDispatcher class OlcrtcSettingsActivity : ProfileSettingsActivity() { @@ -51,12 +54,22 @@ class OlcrtcSettingsActivity : ProfileSettingsActivity() { dnsServer = DataStore.olcrtcDnsServer // Fail fast in the editor instead of saving a profile that can only fail at connect. - // DataStore string() may return null if a field was never edited; coalesce to "". - require(!roomId.isNullOrBlank()) { "olcRTC: room id / URL is required" } - require(!clientId.isNullOrBlank()) { "olcRTC: client id is required" } - val hex = keyHex ?: "" - require(hex.length == 64 && hex.all { it in '0'..'9' || it in 'a'..'f' || it in 'A'..'F' }) { - "olcRTC: encryption key must be 64 hex characters" + validateOlcrtcProfile(requireClientId = true) + } + + override suspend fun saveAndExit() { + try { + // Validate a temporary bean before the base path can stop an active profile. + createEntity().apply { serialize() } + super.saveAndExit() + } catch (e: IllegalArgumentException) { + onMainDispatcher { + Toast.makeText( + applicationContext, + e.message ?: "Invalid olcRTC profile", + Toast.LENGTH_LONG, + ).show() + } } } diff --git a/app/src/test/java/io/nekohasekai/sagernet/fmt/OlcrtcFmtTest.kt b/app/src/test/java/io/nekohasekai/sagernet/fmt/OlcrtcFmtTest.kt index e1d267c27..e0923d5fc 100644 --- a/app/src/test/java/io/nekohasekai/sagernet/fmt/OlcrtcFmtTest.kt +++ b/app/src/test/java/io/nekohasekai/sagernet/fmt/OlcrtcFmtTest.kt @@ -2,9 +2,12 @@ package io.nekohasekai.sagernet.fmt import io.nekohasekai.sagernet.fmt.olcrtc.OlcrtcBean import io.nekohasekai.sagernet.fmt.olcrtc.buildOlcrtcArgs +import io.nekohasekai.sagernet.fmt.olcrtc.carrierHost import io.nekohasekai.sagernet.fmt.olcrtc.parseOlcrtc import io.nekohasekai.sagernet.fmt.olcrtc.toUri import org.junit.Assert.assertEquals +import org.junit.Assert.assertFalse +import org.junit.Assert.assertNull import org.junit.Assert.assertThrows import org.junit.Assert.assertTrue import org.junit.Test @@ -43,16 +46,12 @@ class OlcrtcFmtTest { @Test fun roundTrip_preservesFieldsAndCid() { - val bean = OlcrtcBean().apply { - carrier = "telemost" - roomId = "room-7" - clientId = "device-x" - keyHex = key - transport = "vp8channel" - vp8Fps = 60 - vp8BatchSize = 64 - name = "my olc" - }.apply { /* defaults already set explicitly above */ } + val bean = validBean( + carrier = "telemost", + clientId = "device-x", + vp8Fps = 60, + vp8Batch = 64, + ).apply { name = "my olc" } val uri = bean.toUri() assertTrue(uri.startsWith("olcrtc://telemost?vp8channel")) @@ -69,6 +68,14 @@ class OlcrtcFmtTest { assertEquals(bean.name, parsed.name) } + @Test + fun roundTrip_allowsBlankUpstreamClientId() { + val uri = validBean(clientId = "").toUri() + + assertFalse(uri.contains("cid=")) + assertEquals("", parseOlcrtc(uri).clientId) + } + @Test fun parse_rejectsMissingKey() { val link = "olcrtc://jitsi?datachannel@room-01" @@ -88,52 +95,250 @@ class OlcrtcFmtTest { @Test fun parse_rejectsUnsupportedCarrier() { - val link = "olcrtc://zoom?datachannel@room-01#$key" + val link = "olcrtc://unsupported?datachannel@room-01#$key" assertThrows(IllegalArgumentException::class.java) { parseOlcrtc(link) } } @Test - fun buildArgs_unsetDnsUsesFallback() { - // A bean whose dns field was never set (the default) must not crash; the - // fallback dns is used instead. Regression guard for the connect-time NPE. - val bean = OlcrtcBean().apply { - initializeDefaultValues() - carrier = "jitsi" - roomId = "room-01" - clientId = "dev1" - keyHex = key + fun parse_rejectsClientIdReservedDelimiter() { + val error = assertThrows(IllegalArgumentException::class.java) { + parseOlcrtc("olcrtc://jitsi?vp8channel@review-4821#$key") } - val args = bean.buildOlcrtcArgs( - port = 10800, - protectPath = "/tmp/protect", - socksUser = "", - socksPass = "", - verbose = false, - dnsFallback = "9.9.9.9:53", - readyTimeoutMs = 15_000L, + + assertFalse(error.message.orEmpty().contains("device=7")) + } + + @Test + fun parse_enforcesCarrierTransportMatrix() { + val accepted = listOf( + "jitsi" to "vp8channel", + "jitsi" to "datachannel", + "telemost" to "vp8channel", + "wbstream" to "vp8channel", ) + accepted.forEach { (carrier, transport) -> + val parsed = parseOlcrtc("olcrtc://$carrier?$transport@review-4821#$key") + assertEquals(carrier, parsed.carrier) + assertEquals(transport, parsed.transport) + } + + listOf("telemost", "wbstream").forEach { carrier -> + assertThrows(IllegalArgumentException::class.java) { + parseOlcrtc("olcrtc://$carrier?datachannel@review-4821#$key") + } + } + } + + @Test + fun parse_enforcesVp8BoundsAndDefaults() { + val minimum = parseOlcrtc( + "olcrtc://jitsi?vp8channel@review-4821#$key", + ) + assertEquals(1, minimum.vp8Fps) + assertEquals(1, minimum.vp8BatchSize) + + val maximum = parseOlcrtc( + "olcrtc://jitsi?vp8channel@review-4821#$key", + ) + assertEquals(120, maximum.vp8Fps) + assertEquals(64, maximum.vp8BatchSize) + + val defaults = parseOlcrtc("olcrtc://jitsi?vp8channel@review-4821#$key") + assertEquals(30, defaults.vp8Fps) + assertEquals(8, defaults.vp8BatchSize) + } + + @Test + fun datachannel_ignoresUnusedVp8Values() { + val bean = validBean(transport = "datachannel").apply { + vp8Fps = null + vp8BatchSize = 999 + } + + assertFalse(bean.toUri().contains("vp8-")) + val args = buildArgs(bean) + assertEquals("30", args[args.indexOf("-vp8-fps") + 1]) + assertEquals("8", args[args.indexOf("-vp8-batch") + 1]) + } + + @Test + fun parse_rejectsMalformedOrOutOfRangeVp8Values() { + val payloads = listOf( + "vp8-fps=abc", + "vp8-fps=0", + "vp8-fps=121", + "vp8-fps=-1", + "vp8-batch=abc", + "vp8-batch=0", + "vp8-batch=65", + "vp8-batch=-1", + "vp8-fps", + " =ignored", + ) + payloads.forEach { payload -> + assertThrows(IllegalArgumentException::class.java) { + parseOlcrtc("olcrtc://jitsi?vp8channel<$payload>@review-4821#$key") + } + } + } + + @Test + fun parse_rejectsTextAfterTransportPayload() { + assertThrows(IllegalArgumentException::class.java) { + parseOlcrtc("olcrtc://jitsi?vp8channeljunk@review-4821#$key") + } + } + + @Test + fun buildArgs_unsetDnsUsesFallback() { + val args = buildArgs(validBean(dns = "")) val dnsIdx = args.indexOf("-dns") assertTrue(dnsIdx >= 0) assertEquals("9.9.9.9:53", args[dnsIdx + 1]) } + @Test + fun buildArgs_acceptsIpLiteralDnsEndpoints() { + val endpoints = listOf( + "0.0.0.0:1", + "9.9.9.9:53", + "255.255.255.255:65535", + "[2001:4860:4860::8888]:53", + "[::ffff:192.0.2.1]:53", + "[fe80::1%wlan0]:53", + ) + endpoints.forEach { endpoint -> + val args = buildArgs(validBean(dns = endpoint)) + assertEquals(endpoint, args[args.indexOf("-dns") + 1]) + } + } + + @Test + fun buildArgs_rejectsInvalidDnsEndpointsWithoutEchoingThem() { + val endpoints = listOf( + "dns.example:53", + "9.9.9.9", + "2001:4860:4860::8888:53", + ":53", + "[]:53", + "9.9.9.9:", + "9.9.9.9:dns", + "9.9.9.9:0", + "9.9.9.9:65536", + "999.1.1.1:53", + "09.9.9.9:53", + "[2001:db8:::1]:53", + "[192.0.2.1::]:53", + "[fe80::1%]:53", + ) + endpoints.forEach { endpoint -> + val error = assertThrows(IllegalArgumentException::class.java) { + buildArgs(validBean(dns = endpoint)) + } + assertFalse(error.message.orEmpty().contains(endpoint)) + } + } + + @Test + fun buildArgs_rejectsBlankClientAndUnsupportedTransport() { + assertThrows(IllegalArgumentException::class.java) { + buildArgs(validBean(clientId = "")) + } + assertThrows(IllegalArgumentException::class.java) { + buildArgs(validBean(transport = "futurechannel")) + } + } + @Test fun buildArgs_blankRoomThrowsArgumentError() { + assertThrows(IllegalArgumentException::class.java) { + buildArgs(validBean().apply { roomId = "" }) + } + } + + @Test + fun toUri_rejectsUnsupportedCarrierTransportCombination() { + assertThrows(IllegalArgumentException::class.java) { + validBean(carrier = "telemost", transport = "datachannel").toUri() + } + assertThrows(IllegalArgumentException::class.java) { + validBean(transport = "futurechannel").toUri() + } + } + + @Test + fun uninitializedBeanFailsWithArgumentErrorInsteadOfNullPointer() { + val bean = OlcrtcBean() + + assertThrows(IllegalArgumentException::class.java) { bean.toUri() } + assertThrows(IllegalArgumentException::class.java) { buildArgs(bean) } + } + + @Test + fun toUri_treatsNullOptionalJavaFieldsAsBlank() { val bean = OlcrtcBean().apply { - initializeDefaultValues() carrier = "jitsi" + transport = "vp8channel" + roomId = "review-4821" keyHex = key + vp8Fps = 30 + vp8BatchSize = 8 } - assertThrows(IllegalArgumentException::class.java) { - bean.buildOlcrtcArgs( - port = 10800, - protectPath = "/tmp/protect", - socksUser = "", - socksPass = "", - verbose = false, - dnsFallback = "9.9.9.9:53", - readyTimeoutMs = 15_000L, - ) + + assertEquals("olcrtc://jitsi?vp8channel@review-4821#$key", bean.toUri()) + } + + @Test + fun carrierHost_extractsHostnameAndBracketedIpv6() { + val cases = mapOf( + "https://meet.example.org/room" to "meet.example.org", + "https://meet.example.org:8443/room" to "meet.example.org", + "meet.example.org/room" to "meet.example.org", + "https://meet_private.example/room" to "meet_private.example", + "https://[2001:db8::1]/room" to "2001:db8::1", + "https://[2001:db8::1]:8443/room" to "2001:db8::1", + "[2001:db8::1]/room" to "2001:db8::1", + ) + cases.forEach { (room, expectedHost) -> + assertEquals(expectedHost, validBean(roomId = room).carrierHost()) } } + + @Test + fun carrierHost_preservesFixedCarrierHosts() { + assertEquals("telemost.yandex.ru", validBean(carrier = "telemost").carrierHost()) + assertEquals("stream.wb.ru", validBean(carrier = "wbstream").carrierHost()) + assertNull(validBean().apply { carrier = "unsupported" }.carrierHost()) + assertNull(OlcrtcBean().apply { carrier = "jitsi" }.carrierHost()) + } + + private fun validBean( + carrier: String = "jitsi", + transport: String = "vp8channel", + roomId: String = "review-4821", + clientId: String = "device-7", + dns: String = "", + vp8Fps: Int = 30, + vp8Batch: Int = 8, + ) = OlcrtcBean().apply { + initializeDefaultValues() + this.carrier = carrier + this.transport = transport + this.roomId = roomId + this.clientId = clientId + keyHex = key + dnsServer = dns + this.vp8Fps = vp8Fps + vp8BatchSize = vp8Batch + } + + private fun buildArgs(bean: OlcrtcBean) = bean.buildOlcrtcArgs( + port = 10800, + protectPath = "/tmp/protect", + socksUser = "", + socksPass = "", + verbose = false, + dnsFallback = "9.9.9.9:53", + readyTimeoutMs = 15_000L, + ) }