Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -244,9 +244,10 @@ class FileDownloader(
apply { initCause(cause) }

private fun executeSegment(source: SourceSet.Source, segment: SegmentChain.Segment, chain: SegmentChain, channel: FileChannel) {
val rangeFrom = segment.position()
//部分 CDN 的鉴权层会拒绝任何携带 Range 的请求,从零开始的段必须以普通 GET 起步
val rangeFrom = segment.position().takeIf { it > 0L } ?: -1L
executeCall(source.url, rangeFrom).use { response ->
val body = checkStatus(response, rangeFrom, rangedRequest = true, source = source, chain = chain)
val body = checkStatus(response, rangeFrom, rangedRequest = rangeFrom >= 0, source = source, chain = chain)
streamInto(body.byteStream(), segment, chain, channel)
}
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ val TIME_OUT = TimeUnit.SECONDS.toMillis(30L)
const val SMALL_FILE_READ_TIMEOUT_MS = 5_000L

const val HOST_CURSEFORGE_API = "api.curseforge.com"
const val HOST_CURSEFORGE_EDGE = "edge.forgecdn.net"
const val CURSEFORGE_CDN_SUFFIX = "forgecdn.net"

const val URL_MCMOD: String = "https://www.mcmod.cn/"
const val URL_MINECRAFT_VERSION_REPOS: String = "https://piston-meta.mojang.com/mc/game/version_manifest_v2.json"
Expand All @@ -67,16 +67,20 @@ const val URL_CLOUD_RENDERER_PLUGINS = "https://www.123865.com/s/YLIUVv-hae0v"
const val URL_CLOUD_DRIVE_DRIVER_PLUGINS = "https://www.123865.com/s/YLIUVv-3ae0v"
const val URL_CLOUD_NATIVE_LIB_PLUGINS = "https://www.123865.com/s/YLIUVv-Hae0v"

private fun isCurseForgeHost(host: String): Boolean =
host == HOST_CURSEFORGE_API ||
host == CURSEFORGE_CDN_SUFFIX ||
host.endsWith(".$CURSEFORGE_CDN_SUFFIX")

/**
* An [Interceptor] for CurseForge API requests.
*
* It automatically injects the `x-api-key` header when the request host matches
* [HOST_CURSEFORGE_API] or [HOST_CURSEFORGE_EDGE], provided the API key is not blank.
* It automatically injects the `x-api-key` header when the request targets a
* CurseForge host, provided the API key is not blank.
*/
private val CURSEFORGE_INTERCEPTOR = Interceptor { chain ->
val request = chain.request()
val host = request.url.host
if (host == HOST_CURSEFORGE_API || host == HOST_CURSEFORGE_EDGE) {
if (isCurseForgeHost(request.url.host)) {
val apiKey = BuildKeys.CURSEFORGE_API
if (apiKey.isNotBlank()) {
val newRequest = request.newBuilder()
Expand Down Expand Up @@ -130,10 +134,7 @@ val GLOBAL_CLIENT = HttpClient(OkHttp) {
}
}.apply {
requestPipeline.intercept(HttpRequestPipeline.State) {
// 检查 host 是否为 CurseForge
// 自动添加 CurseForge 的 api 密钥
val host = context.url.host
if (host == HOST_CURSEFORGE_API || host == HOST_CURSEFORGE_EDGE) {
if (isCurseForgeHost(context.url.host)) {
val apiKey = BuildKeys.CURSEFORGE_API
if (apiKey.isNotBlank()) {
context.header("x-api-key", apiKey)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -88,9 +88,10 @@ private class FakeSource(
}

/**
* 分区守卫源:prefixOnly=true 时仅服务起点位于文件前半段的 Range 请求
* 分区守卫源:prefixOnly=true 时仅服务起点位于文件前半段的请求
* 起点越过中位的一律 500;prefixOnly=false 时拒绝起点为 0 的请求,
* 其余区间正常服务。两者配合可强制一次下载必然发生跨源断点拼接。
* 未携带 Range 的请求按起点 0 处理。
*/
private class GuardedHalfSource(
private val content: ByteArray,
Expand All @@ -102,9 +103,7 @@ private class GuardedHalfSource(
override fun dispatch(request: RecordedRequest): MockResponse {
hits.incrementAndGet()
val range = request.headers["Range"]
?: return MockResponse.Builder().code(500).build()

val start = RANGE.find(range)!!.groupValues[1].toLong().toInt()
val start = range?.let { RANGE.find(it)!!.groupValues[1].toLong().toInt() } ?: 0
val servesPrefixHalf = start < content.size / 2

if (prefixOnly != servesPrefixHalf) {
Expand All @@ -127,6 +126,19 @@ private class GuardedHalfSource(
}
}

/** 模拟拒绝任何 Range 请求的 CDN 鉴权层:携带 Range 一律 404,普通 GET 正常服务 */
private class RangeRejectingSource(private val content: ByteArray) : Dispatcher() {
override fun dispatch(request: RecordedRequest): MockResponse {
if (request.headers["Range"] != null) {
return MockResponse.Builder().code(404).build()
}
return MockResponse.Builder()
.addHeader("Content-Length", content.size.toString())
.body(Buffer().write(content))
.build()
}
}

class FileDownloaderE2ETest {

private lateinit var workDir: File
Expand Down Expand Up @@ -224,6 +236,26 @@ class FileDownloaderE2ETest {
assertArrayEquals(payload, target.readBytes())
}

/** 起点为零的段不得携带 Range:部分 CDN 会直接拒绝任何带 Range 的请求 */
@Test
fun `fresh download must not carry a range header`() = runBlocking {
val source = RangeRejectingSource(payload)
val server = startServer(source)
val (request, target) = engineRequest("no-range.bin", server.url("/file.bin").toString(), sha1 = sha1HexOf(payload))

withTimeout(TEST_TIMEOUT_MS) {
FileDownloader(
request = request,
connections = Semaphore(8),
stats = DownloadStats(),
allowExtraConnection = { false },
client = OkHttpClient()
).download()
}

assertArrayEquals(payload, target.readBytes())
}

@Test
fun `unknown size streams till eof`() = runBlocking {
val server = startServer(FakeSource(payload))
Expand Down