diff --git a/.env.example b/.env.example index 07721e95..9f061443 100644 --- a/.env.example +++ b/.env.example @@ -5,6 +5,7 @@ DATABASE_PASSWORD=typetype DRAGONFLY_URL=redis://dragonfly:6379 DOWNLOADER_SERVICE_URL=http://typetype-downloader:18093 YOUTUBE_SESSION_ENCRYPTION_KEY=replace-with-at-least-32-random-characters +YOUTUBE_OUTBOUND_PROXY_URL= YOUTUBE_REMOTE_LOGIN_SERVICE_URL=http://typetype-token:8081 YOUTUBE_REMOTE_LOGIN_CALLBACK_BASE_URL=http://typetype-server:8080 YOUTUBE_REMOTE_LOGIN_INTERNAL_TOKEN=replace-with-shared-internal-token diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 00c59028..2eeb0b00 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -19,6 +19,9 @@ jobs: steps: - uses: actions/checkout@v6 + - name: Isolate Gradle user home + run: echo "GRADLE_USER_HOME=${RUNNER_TEMP}/gradle-user-home" >> "$GITHUB_ENV" + - uses: actions/setup-java@v5 with: java-version: "25" diff --git a/.github/workflows/coverage.yml b/.github/workflows/coverage.yml index 4b6978e0..bc6db65a 100644 --- a/.github/workflows/coverage.yml +++ b/.github/workflows/coverage.yml @@ -17,6 +17,9 @@ jobs: steps: - uses: actions/checkout@v6 + - name: Isolate Gradle user home + run: echo "GRADLE_USER_HOME=${RUNNER_TEMP}/gradle-user-home" >> "$GITHUB_ENV" + - name: Set up JDK 25 uses: actions/setup-java@v5 with: diff --git a/.github/workflows/openapi.yml b/.github/workflows/openapi.yml index 89fd674a..890fbddb 100644 --- a/.github/workflows/openapi.yml +++ b/.github/workflows/openapi.yml @@ -18,6 +18,8 @@ jobs: runs-on: ${{ github.event_name == 'pull_request' && 'ubuntu-24.04' || fromJSON('["self-hosted","Linux","X64","arko"]') }} steps: - uses: actions/checkout@v6 + - name: Isolate Gradle user home + run: echo "GRADLE_USER_HOME=${RUNNER_TEMP}/gradle-user-home" >> "$GITHUB_ENV" - uses: actions/setup-java@v5 with: java-version: "25" diff --git a/build.gradle.kts b/build.gradle.kts index 5f8a78cf..8942f4b2 100644 --- a/build.gradle.kts +++ b/build.gradle.kts @@ -36,7 +36,7 @@ dependencies { implementation("io.ktor:ktor-server-call-logging-jvm") implementation("io.ktor:ktor-server-rate-limit-jvm") implementation("ch.qos.logback:logback-classic:1.5.38") - implementation("com.github.Priveetee.PipePipeExtractor:extractor:ec9ba0dd6b0521d83ac3905f1ee980c33930140d") + implementation("com.github.InfinityLoop1308.PipePipeExtractor:extractor:00db0e1d9b553d941f3009eb79e492d95eaf442d") compileOnly("com.github.TeamNewPipe:nanojson:1d9e1aea9049fc9f85e68b43ba39fe7be1c1f751") implementation("org.json:json:20260522") implementation("com.squareup.okhttp3:okhttp:5.4.0") diff --git a/src/main/java/org/schabi/newpipe/extractor/services/youtube/sabr/TypeTypeYoutubeSabrInfoFactory.java b/src/main/java/org/schabi/newpipe/extractor/services/youtube/sabr/TypeTypeYoutubeSabrInfoFactory.java new file mode 100644 index 00000000..614b2a0e --- /dev/null +++ b/src/main/java/org/schabi/newpipe/extractor/services/youtube/sabr/TypeTypeYoutubeSabrInfoFactory.java @@ -0,0 +1,21 @@ +package org.schabi.newpipe.extractor.services.youtube.sabr; + +public final class TypeTypeYoutubeSabrInfoFactory { + private TypeTypeYoutubeSabrInfoFactory() { + } + + public static YoutubeSabrInfo withPlaybackUrlAndClientVersion( + final YoutubeSabrInfo info, + final String serverAbrStreamingUrl, + final String clientVersion) { + return new YoutubeSabrInfo( + info.getProfile(), + info.getVideoId(), + info.getCpn(), + clientVersion, + info.getVisitorData(), + serverAbrStreamingUrl, + info.getVideoPlaybackUstreamerConfig(), + info.getFormats()); + } +} diff --git a/src/main/kotlin/dev/typetype/server/Application.kt b/src/main/kotlin/dev/typetype/server/Application.kt index 438bf76f..498607fb 100644 --- a/src/main/kotlin/dev/typetype/server/Application.kt +++ b/src/main/kotlin/dev/typetype/server/Application.kt @@ -1,6 +1,7 @@ package dev.typetype.server import dev.typetype.server.cache.DragonflyService import dev.typetype.server.db.DatabaseFactory +import dev.typetype.server.downloader.YoutubeProxySelector import dev.typetype.server.services.ActiveSessionService import dev.typetype.server.services.AuthService import dev.typetype.server.services.AdminSettingsService @@ -50,8 +51,16 @@ fun Application.module() { val cacheUrl = System.getenv("DRAGONFLY_URL") ?: "redis://localhost:6379" val cache = DragonflyService(cacheUrl) val subtitleServiceUrl = System.getenv("SUBTITLE_SERVICE_URL") ?: "http://typetype-token:8081" - NewPipeInitializer.init(subtitleServiceUrl) - val svc = ServiceRegistry(cache, subtitleServiceUrl, youtubeSessionEncryptionKey, jwtSecret, adminSettingsService) + val youtubeProxySelector = YoutubeProxySelector.fromUrl(System.getenv("YOUTUBE_OUTBOUND_PROXY_URL")) + NewPipeInitializer.init(subtitleServiceUrl, youtubeProxySelector) + val svc = ServiceRegistry( + cache, + subtitleServiceUrl, + youtubeSessionEncryptionKey, + jwtSecret, + adminSettingsService, + youtubeProxySelector, + ) val youtubeRemoteBrowserConfig = YoutubeRemoteBrowserConfig.fromEnvironment(subtitleServiceUrl) val youtubeRemoteLoginReadinessService = YoutubeRemoteLoginReadinessService( youtubeRemoteBrowserConfig, diff --git a/src/main/kotlin/dev/typetype/server/ExtractionServiceRegistry.kt b/src/main/kotlin/dev/typetype/server/ExtractionServiceRegistry.kt index c9c4cd60..5d0fcc62 100644 --- a/src/main/kotlin/dev/typetype/server/ExtractionServiceRegistry.kt +++ b/src/main/kotlin/dev/typetype/server/ExtractionServiceRegistry.kt @@ -51,6 +51,7 @@ import dev.typetype.server.services.YoutubeSessionStreamService import okhttp3.ConnectionPool import okhttp3.Dispatcher import okhttp3.OkHttpClient +import java.net.ProxySelector import java.util.concurrent.TimeUnit internal class ExtractionServiceRegistry( @@ -58,10 +59,13 @@ internal class ExtractionServiceRegistry( subtitleServiceUrl: String, youtubeSessionEncryptionKey: String?, hlsManifestUrlSigner: ((String) -> String)? = null, + youtubeProxySelector: ProxySelector? = null, ) { private val youtubeSessionSecret = youtubeSessionEncryptionKey?.trim() ?.takeIf { it.length >= MIN_YOUTUBE_SESSION_SECRET_LENGTH } - val httpClient = OkHttpClient() + val httpClient = OkHttpClient.Builder() + .apply { youtubeProxySelector?.let(::proxySelector) } + .build() val proxyHttpClient: OkHttpClient = httpClient.newBuilder() .dispatcher(proxyDispatcher()) .connectionPool(ConnectionPool(64, 5, TimeUnit.MINUTES)) diff --git a/src/main/kotlin/dev/typetype/server/ServiceRegistry.kt b/src/main/kotlin/dev/typetype/server/ServiceRegistry.kt index 16a2f52d..dcc498ef 100644 --- a/src/main/kotlin/dev/typetype/server/ServiceRegistry.kt +++ b/src/main/kotlin/dev/typetype/server/ServiceRegistry.kt @@ -33,12 +33,14 @@ import dev.typetype.server.services.UserVideoMetadataRepairService import dev.typetype.server.services.VideoMetadataResolver import dev.typetype.server.services.WatchLaterService import dev.typetype.server.services.YoutubeTakeoutFactory +import java.net.ProxySelector internal class ServiceRegistry( cache: DragonflyService, subtitleServiceUrl: String, youtubeSessionEncryptionKey: String?, jwtSecret: String, adminSettingsService: AdminSettingsService, + youtubeProxySelector: ProxySelector? = null, ) { init { SubscriptionFeedCacheInvalidation.configure(SubscriptionFeedCacheInvalidator(cache)) @@ -53,6 +55,7 @@ internal class ServiceRegistry( subtitleServiceUrl, youtubeSessionEncryptionKey, publicHlsManifestTokenService::createPath, + youtubeProxySelector, ) val youtubeSessionService = extraction.youtubeSessionService val youtubeSessionStreamService = extraction.youtubeSessionStreamService @@ -103,7 +106,7 @@ internal class ServiceRegistry( val accessControlService = AccessControlService(settingsService, allowedChannelsService, allowedPlaylistsService, adminSettingsService) val blockedService = BlockedService() val bugReportService = BugReportService() - val youtubeTakeoutImportService = YoutubeTakeoutFactory.create(subscriptionsService, playlistService, historyService, favoritesService, watchLaterService, streamService) + val youtubeTakeoutImportService = YoutubeTakeoutFactory.create(subscriptionsService, playlistService, historyService, favoritesService, watchLaterService) val recommendationPoolResolverDependencies = HomeRecommendationPoolResolverDependencies( subscriptionsService = subscriptionsService, subscriptionFeedService = subscriptionFeedService, diff --git a/src/main/kotlin/dev/typetype/server/downloader/OkHttpDownloader.kt b/src/main/kotlin/dev/typetype/server/downloader/OkHttpDownloader.kt index 20f08255..de00fc06 100644 --- a/src/main/kotlin/dev/typetype/server/downloader/OkHttpDownloader.kt +++ b/src/main/kotlin/dev/typetype/server/downloader/OkHttpDownloader.kt @@ -12,6 +12,7 @@ import org.schabi.newpipe.extractor.downloader.Response import org.schabi.newpipe.extractor.downloader.StreamingResponse import org.schabi.newpipe.extractor.exceptions.ReCaptchaException import org.schabi.newpipe.extractor.localization.Localization +import java.net.ProxySelector import java.util.concurrent.TimeUnit class OkHttpDownloader private constructor( @@ -20,13 +21,18 @@ class OkHttpDownloader private constructor( ) : Downloader() { companion object { - fun instance(): OkHttpDownloader = create(STREAMING_READ_TIMEOUT_MS) - - internal fun create(streamingReadTimeoutMs: Long): OkHttpDownloader { - val client = OkHttpClient.Builder() + fun instance(proxySelector: ProxySelector? = null): OkHttpDownloader = + create(STREAMING_READ_TIMEOUT_MS, proxySelector) + + internal fun create( + streamingReadTimeoutMs: Long, + proxySelector: ProxySelector? = null, + ): OkHttpDownloader { + val clientBuilder = OkHttpClient.Builder() .connectTimeout(30, TimeUnit.SECONDS) .readTimeout(30, TimeUnit.SECONDS) - .build() + proxySelector?.let(clientBuilder::proxySelector) + val client = clientBuilder.build() val streamingClient = client.newBuilder() .readTimeout(streamingReadTimeoutMs, TimeUnit.MILLISECONDS) .build() diff --git a/src/main/kotlin/dev/typetype/server/downloader/YoutubeProxySelector.kt b/src/main/kotlin/dev/typetype/server/downloader/YoutubeProxySelector.kt new file mode 100644 index 00000000..24bca222 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/downloader/YoutubeProxySelector.kt @@ -0,0 +1,57 @@ +package dev.typetype.server.downloader + +import java.io.IOException +import java.net.InetSocketAddress +import java.net.Proxy +import java.net.ProxySelector +import java.net.SocketAddress +import java.net.URI + +internal class YoutubeProxySelector private constructor( + private val proxy: Proxy, +) : ProxySelector() { + override fun select(uri: URI): List = + if (isYoutubeHost(uri.host)) listOf(proxy, Proxy.NO_PROXY) else DIRECT + + override fun connectFailed(uri: URI, socketAddress: SocketAddress, exception: IOException) = Unit + + companion object { + fun fromUrl(proxyUrl: String?): ProxySelector? { + val value = proxyUrl?.trim()?.takeIf(String::isNotEmpty) ?: return null + val uri = runCatching { URI(value) } + .getOrElse { throw IllegalArgumentException("Invalid YouTube outbound proxy URL", it) } + require(uri.scheme.equals("http", ignoreCase = true)) { + "YOUTUBE_OUTBOUND_PROXY_URL must use the http scheme" + } + require(!uri.host.isNullOrBlank()) { + "YOUTUBE_OUTBOUND_PROXY_URL must include a host" + } + require(uri.userInfo == null && uri.path.isNullOrEmpty() && uri.query == null && uri.fragment == null) { + "YOUTUBE_OUTBOUND_PROXY_URL must contain only a scheme, host, and port" + } + val port = uri.port.takeIf { it >= 0 } ?: DEFAULT_HTTP_PORT + val address = InetSocketAddress.createUnresolved(uri.host, port) + return YoutubeProxySelector(Proxy(Proxy.Type.HTTP, address)) + } + + private fun isYoutubeHost(host: String?): Boolean { + val normalizedHost = host?.lowercase() ?: return false + return YOUTUBE_DOMAINS.any { domain -> + normalizedHost == domain || normalizedHost.endsWith(".$domain") + } + } + + private val DIRECT = listOf(Proxy.NO_PROXY) + private val YOUTUBE_DOMAINS = setOf( + "youtube.com", + "youtube-nocookie.com", + "youtu.be", + "googlevideo.com", + "ytimg.com", + "ggpht.com", + "googleusercontent.com", + "googleapis.com", + ) + private const val DEFAULT_HTTP_PORT = 80 + } +} diff --git a/src/main/kotlin/dev/typetype/server/models/ExtractionResult.kt b/src/main/kotlin/dev/typetype/server/models/ExtractionResult.kt index fa6bcea5..40e2bac4 100644 --- a/src/main/kotlin/dev/typetype/server/models/ExtractionResult.kt +++ b/src/main/kotlin/dev/typetype/server/models/ExtractionResult.kt @@ -2,6 +2,6 @@ package dev.typetype.server.models sealed class ExtractionResult { data class Success(val data: T) : ExtractionResult() - data class Failure(val message: String) : ExtractionResult() - data class BadRequest(val message: String) : ExtractionResult() + data class Failure(val message: String, val code: String = "error") : ExtractionResult() + data class BadRequest(val message: String, val code: String = "error") : ExtractionResult() } diff --git a/src/main/kotlin/dev/typetype/server/models/VideoItem.kt b/src/main/kotlin/dev/typetype/server/models/VideoItem.kt index 786e9cd4..7b60db8e 100644 --- a/src/main/kotlin/dev/typetype/server/models/VideoItem.kt +++ b/src/main/kotlin/dev/typetype/server/models/VideoItem.kt @@ -23,4 +23,5 @@ data class VideoItem( val isLive: Boolean = false, val isPostLive: Boolean = false, val isLiveContent: Boolean = false, + val requiresMembership: Boolean = false, ) diff --git a/src/main/kotlin/dev/typetype/server/routes/FavoritesRoutes.kt b/src/main/kotlin/dev/typetype/server/routes/FavoritesRoutes.kt index 3161197a..4dca4a02 100644 --- a/src/main/kotlin/dev/typetype/server/routes/FavoritesRoutes.kt +++ b/src/main/kotlin/dev/typetype/server/routes/FavoritesRoutes.kt @@ -14,8 +14,9 @@ import io.ktor.server.routing.post fun Route.favoritesRoutes(favoritesService: FavoritesService, authService: AuthService, metadataRepairService: UserVideoMetadataRepairService? = null) { get("/favorites") { call.withJwtAuth(authService) { userId -> - metadataRepairService?.repairFavorites(userId) - call.respond(favoritesService.getAll(userId)) + val items = favoritesService.getAll(userId) + metadataRepairService?.scheduleFavorites(call.application, userId) + call.respond(items) } } get("/favorites/{videoUrl...}") { diff --git a/src/main/kotlin/dev/typetype/server/routes/PlaylistRoutes.kt b/src/main/kotlin/dev/typetype/server/routes/PlaylistRoutes.kt index ed797566..8bcbece8 100644 --- a/src/main/kotlin/dev/typetype/server/routes/PlaylistRoutes.kt +++ b/src/main/kotlin/dev/typetype/server/routes/PlaylistRoutes.kt @@ -20,7 +20,6 @@ import io.ktor.server.routing.put fun Route.playlistRoutes(playlistService: PlaylistService, authService: AuthService, metadataRepairService: UserVideoMetadataRepairService? = null) { get("/playlists") { call.withJwtAuth(authService) { userId -> - metadataRepairService?.repairPlaylists(userId) call.respond(playlistService.getAll(userId)) } } @@ -35,8 +34,8 @@ fun Route.playlistRoutes(playlistService: PlaylistService, authService: AuthServ get("/playlists/{id}") { call.withJwtAuth(authService) { userId -> val id = call.parameters["id"] ?: return@withJwtAuth call.respond(HttpStatusCode.BadRequest, ErrorResponse("Missing id")) - metadataRepairService?.repairPlaylists(userId) val playlist = playlistService.getById(userId, id) ?: return@withJwtAuth call.respond(HttpStatusCode.NotFound, ErrorResponse("Not found")) + metadataRepairService?.schedulePlaylists(call.application, userId) call.respond(playlist) } } diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackHandler.kt b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackHandler.kt index 2867a478..3e8a7c58 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackHandler.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackHandler.kt @@ -43,6 +43,7 @@ internal class SabrPlaybackHandler( video = video, startTimeMs = startTimeMs, audioOnly = request.audioOnly, + isLive = request.isLive, ) preparation.holder.setActiveTracks(videoActive = !request.audioOnly, audioActive = true) respondPrepared(call, preparation.holder, videoId, preparation.startTimeMs, preparation.ready) @@ -144,6 +145,7 @@ internal class SabrPlaybackHandler( startTimeMs = body?.startTimeMs ?: request.queryParameters["startTimeMs"]?.toLongOrNull(), playerTimeMs = body?.playerTimeMs ?: request.queryParameters["playerTimeMs"]?.toLongOrNull(), audioOnly = body?.audioOnly ?: request.queryParameters["audioOnly"]?.toBooleanStrictOrNull() ?: false, + isLive = body?.isLive ?: request.queryParameters["isLive"]?.toBooleanStrictOrNull() ?: false, ) } diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackModels.kt b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackModels.kt index 5e659675..096749dd 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackModels.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackModels.kt @@ -1,6 +1,7 @@ package dev.typetype.server.routes import kotlinx.serialization.Serializable +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest @Serializable internal data class SabrPlaybackRequest( @@ -10,6 +11,19 @@ internal data class SabrPlaybackRequest( val startTimeMs: Long? = null, val playerTimeMs: Long? = null, val audioOnly: Boolean = false, + val isLive: Boolean = false, +) + +@Serializable +internal data class SabrLivePlaybackResponse( + val active: Boolean, + val postLiveDvr: Boolean, + val headSequence: Long, + val headTimeMs: Long, + val seekableStartMs: Long, + val seekableEndMs: Long, + val atLiveEdge: Boolean, + val targetLatencyMs: Long, ) @Serializable @@ -25,6 +39,7 @@ internal data class SabrPlaybackResponse( val ready: Boolean, val status: String, val retryAfterMs: Long? = null, + val live: SabrLivePlaybackResponse? = null, ) @Serializable @@ -48,6 +63,7 @@ internal data class SabrPlaybackStateResponse( val pendingSegmentDemand: String? = null, val terminalError: String? = null, val diagnosticTrace: String? = null, + val live: SabrLivePlaybackResponse? = null, ) @Serializable @@ -91,6 +107,7 @@ internal data class SabrPlaybackPositionResponse( val readerHeadMs: Long, val readerTailMs: Long, val bufferedEdgeMs: Long, + val live: SabrLivePlaybackResponse? = null, ) @Serializable @@ -113,6 +130,7 @@ internal data class SabrPlaybackPrefetchResponse( val terminalError: String? = null, val recoveryAction: String? = null, val retryVideoItags: List = emptyList(), + val live: SabrLivePlaybackResponse? = null, ) @Serializable @@ -125,6 +143,8 @@ internal data class SabrPlaybackWindowReadyResponse( val endOfStream: Boolean, val audio: SabrPlaybackWindowTrack, val video: SabrPlaybackWindowTrack? = null, + val startTimeMs: Long = 0L, + val live: SabrLivePlaybackResponse? = null, ) @Serializable @@ -159,4 +179,12 @@ internal data class SabrPlaybackWindowPreparingResponse( val terminalError: String? = null, val recoveryAction: String? = null, val retryVideoItags: List = emptyList(), + val live: SabrLivePlaybackResponse? = null, +) + +internal data class SabrPlaybackWindowBuildResult( + val response: SabrPlaybackWindowReadyResponse, + val blockedBy: String?, + val blockedRequests: List, + val isReady: Boolean, ) diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackResponseFactory.kt b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackResponseFactory.kt index f8d20008..dab37d96 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackResponseFactory.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackResponseFactory.kt @@ -1,6 +1,8 @@ package dev.typetype.server.routes import dev.typetype.server.services.SabrSessionHolder +import dev.typetype.server.services.livePlaybackSnapshot +import dev.typetype.server.services.liveRetryAfterMs internal fun SabrSessionHolder.toPlaybackResponse( videoId: String, @@ -19,6 +21,7 @@ internal fun SabrSessionHolder.toPlaybackResponse( ready = ready, status = if (ready) "ready" else playbackState().name.lowercase(), retryAfterMs = if (ready) null else retryAfterMs, + live = livePlaybackSnapshot()?.toResponse(), ) internal fun SabrSessionHolder.toRetryPlaybackResponse(status: String, retryAfterMs: Long): SabrPlaybackResponse = @@ -33,5 +36,18 @@ internal fun SabrSessionHolder.toRetryPlaybackResponse(status: String, retryAfte generation = activeGeneration(), ready = false, status = status, - retryAfterMs = retryAfterMs, + retryAfterMs = if (livePlaybackSnapshot()?.active == true) liveRetryAfterMs() else retryAfterMs, + live = livePlaybackSnapshot()?.toResponse(), + ) + +internal fun dev.typetype.server.services.SabrLivePlaybackSnapshot.toResponse(): SabrLivePlaybackResponse = + SabrLivePlaybackResponse( + active = active, + postLiveDvr = postLiveDvr, + headSequence = headSequence, + headTimeMs = headTimeMs, + seekableStartMs = seekableStartMs, + seekableEndMs = seekableEndMs, + atLiveEdge = atLiveEdge, + targetLatencyMs = targetLatencyMs, ) diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackStateHandler.kt b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackStateHandler.kt index 9fbfd7db..a5b8f07b 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackStateHandler.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackStateHandler.kt @@ -4,6 +4,7 @@ import dev.typetype.server.models.ErrorResponse import dev.typetype.server.services.SabrSessionHolder import dev.typetype.server.services.SabrSessionStore import dev.typetype.server.services.pendingSegmentDemandSummary +import dev.typetype.server.services.livePlaybackSnapshot import io.ktor.http.HttpStatusCode import io.ktor.server.application.ApplicationCall import io.ktor.server.response.respond @@ -36,6 +37,7 @@ internal class SabrPlaybackStateHandler(private val sabrSessionStore: SabrSessio pendingSegmentDemand = pendingSegmentDemandSummary(), terminalError = terminalFailure() ?: networkFailure(), diagnosticTrace = session.diagnosticTrace, + live = livePlaybackSnapshot()?.toResponse(), ) private fun SabrSegmentRequest.summary(): String = "${format.itag}:$sequenceNumber" diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowBuilder.kt b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowBuilder.kt index 5cdf28e4..1a1660db 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowBuilder.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowBuilder.kt @@ -4,7 +4,14 @@ import dev.typetype.server.services.CachedSabrSegment import dev.typetype.server.services.SabrInitializationData import dev.typetype.server.services.SabrSessionHolder import dev.typetype.server.services.SabrSessionStore -import dev.typetype.server.services.findCachedMediaAt +import dev.typetype.server.services.findCachedPlaybackMediaAt +import dev.typetype.server.services.coversPlaybackTime +import dev.typetype.server.services.failLivePlaybackDiscontinuity +import dev.typetype.server.services.livePlaybackSnapshot +import dev.typetype.server.services.playbackContinuationSequence +import dev.typetype.server.services.playbackSegmentDurationMs +import dev.typetype.server.services.playbackSegmentStartMs +import dev.typetype.server.services.resolvePlaybackStartMs import dev.typetype.server.services.playbackStartSequence import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat @@ -14,15 +21,30 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi holder: SabrSessionHolder, request: SabrPlaybackWindowRequest, ): SabrPlaybackWindowBuildResult { + val startTimeMs = holder.resolvePlaybackStartMs(request.playerTimeMs) + val effectiveRequest = if (startTimeMs == request.playerTimeMs) request else request.copy(playerTimeMs = startTimeMs) + val live = holder.livePlaybackSnapshot() SabrInitializationData.ingestRemembered(holder.audioFormat, holder) - if (!request.audioOnly) SabrInitializationData.ingestRemembered(holder.videoFormat, holder) - if (request.audioOnly) return buildAudioOnly(holder, request) - val video = buildTrack(holder, holder.videoFormat, request, request.playerTimeMs) - val decodeStartMs = video.track.segments.firstOrNull()?.startMs ?: request.playerTimeMs - val audio = buildTrack(holder, holder.audioFormat, request, minOf(request.playerTimeMs, decodeStartMs)) + if (!effectiveRequest.audioOnly) SabrInitializationData.ingestRemembered(holder.videoFormat, holder) + if (effectiveRequest.audioOnly) return buildAudioOnly(holder, effectiveRequest, live?.toResponse()) + val video = buildTrack(holder, holder.videoFormat, effectiveRequest, effectiveRequest.playerTimeMs, live?.active == true) + val decodeStartMs = video.track.segments.firstOrNull()?.startMs ?: effectiveRequest.playerTimeMs + val audio = buildTrack( + holder, + holder.audioFormat, + effectiveRequest, + minOf(effectiveRequest.playerTimeMs, decodeStartMs), + live?.active == true, + ) + val playbackStartMs = resolvedPlaybackStartMs( + effectiveRequest.playerTimeMs, + live?.active == true, + video.track, + audio.track, + ) val blocked = blockedTrack(audio, video) - val readyAheadMs = minOf(request.bufferGoalMs.coerceAtLeast(1L), MIN_READY_AHEAD_MS) - val requestedReadyEndMs = request.playerTimeMs.coerceAtLeast(0L) + readyAheadMs + val readyAheadMs = readyAheadMs(effectiveRequest, live?.active == true) + val requestedReadyEndMs = playbackStartMs + readyAheadMs return SabrPlaybackWindowBuildResult( response = SabrPlaybackWindowReadyResponse( sessionId = holder.sessionToken, @@ -30,24 +52,26 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi ready = true, retryAfterMs = null, durationMs = holder.durationMs(), - endOfStream = audio.atEnd && video.atEnd, + endOfStream = live?.active != true && audio.atEnd && video.atEnd, audio = audio.track, video = video.track, + startTimeMs = playbackStartMs, + live = live?.toResponse(), ), blockedBy = blocked?.blockedBy, - blockedRequest = blocked?.blockedRequest, + blockedRequests = listOfNotNull(video.blockedRequest, audio.blockedRequest), isReady = audio.covers(holder.readyEndMs(holder.audioFormat, requestedReadyEndMs)) && video.covers(holder.readyEndMs(holder.videoFormat, requestedReadyEndMs)), ) } - private suspend fun buildAudioOnly( holder: SabrSessionHolder, request: SabrPlaybackWindowRequest, + live: SabrLivePlaybackResponse?, ): SabrPlaybackWindowBuildResult { - val audio = buildTrack(holder, holder.audioFormat, request, request.playerTimeMs) - val readyEndMs = request.playerTimeMs.coerceAtLeast(0L) + - minOf(request.bufferGoalMs.coerceAtLeast(1L), MIN_READY_AHEAD_MS) + val audio = buildTrack(holder, holder.audioFormat, request, request.playerTimeMs, live?.active == true) + val playbackStartMs = resolvedPlaybackStartMs(request.playerTimeMs, live?.active == true, audio.track) + val readyEndMs = playbackStartMs + readyAheadMs(request, live?.active == true) return SabrPlaybackWindowBuildResult( response = SabrPlaybackWindowReadyResponse( sessionId = holder.sessionToken, @@ -55,21 +79,24 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi ready = true, retryAfterMs = null, durationMs = holder.durationMs(), - endOfStream = audio.atEnd, + endOfStream = live?.active != true && audio.atEnd, audio = audio.track, + startTimeMs = playbackStartMs, + live = live, ), blockedBy = audio.blockedBy, - blockedRequest = audio.blockedRequest, + blockedRequests = listOfNotNull(audio.blockedRequest), isReady = audio.covers(holder.readyEndMs(holder.audioFormat, readyEndMs)), ) } + private fun blockedTrack(audio: TrackBuildResult, video: TrackBuildResult): TrackBuildResult? = + sequenceOf(video, audio) + .filter { it.blockedRequest != null } + .minByOrNull { it.coveredEndMs } - private fun blockedTrack(audio: TrackBuildResult, video: TrackBuildResult): TrackBuildResult? = when { - audio.track.segments.isEmpty() && audio.blockedRequest != null -> audio - video.track.segments.isEmpty() && video.blockedRequest != null -> video - audio.blockedRequest != null -> audio - video.blockedRequest != null -> video - else -> null + private fun readyAheadMs(request: SabrPlaybackWindowRequest, activeLive: Boolean): Long { + val minimum = if (activeLive && request.bufferedRanges.isEmpty()) LIVE_STARTUP_READY_AHEAD_MS else MIN_READY_AHEAD_MS + return minOf(request.bufferGoalMs.coerceAtLeast(1L), minimum) } private suspend fun buildTrack( @@ -77,15 +104,18 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi format: YoutubeSabrFormat, request: SabrPlaybackWindowRequest, requestedStartMs: Long, + activeLive: Boolean, ): TrackBuildResult { - val targetMs = request.bufferedEndFor(format, requestedStartMs).coerceAtLeast(0L) - val goalEndMs = request.playerTimeMs.coerceAtLeast(0L) + request.bufferGoalMs.coerceAtLeast(1L) + val continuesServedTrack = !activeLive || holder.lastServedSequence(format) != null + val targetMs = if (continuesServedTrack) request.bufferedEndFor(format, requestedStartMs) else requestedStartMs.coerceAtLeast(0L) + val goalStartMs = if (activeLive) targetMs else request.playerTimeMs.coerceAtLeast(0L) + var goalEndMs = goalStartMs + request.bufferGoalMs.coerceAtLeast(1L) val segments = mutableListOf() var blockedBy: String? = null var blockedRequest: SabrSegmentRequest? = null - var seq = holder.playbackStartSequence(format, targetMs) + var seq = holder.playbackContinuationSequence(format, targetMs, activeLive) var coveredEndMs = targetMs - val endSequence = holder.session.streamState.getEndSegment(format).toInt() + val endSequence = if (activeLive) 0 else holder.session.streamState.getEndSegment(format).toInt() var atEnd = false while (segments.size < MAX_SEGMENTS_PER_TRACK) { if (endSequence > 0 && seq > endSequence) { @@ -93,14 +123,16 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi break } val mediaRequest = SabrSegmentRequest.media(format, seq) - val segment = sabrSessionStore.cachedSegment(holder, mediaRequest) + var segment = sabrSessionStore.cachedSegment(holder, mediaRequest) val expectedStartMs = if (segments.isEmpty()) targetMs else coveredEndMs - if (segment == null || segments.isEmpty() && !segment.covers(targetMs)) { - val authoritative = holder.session.findCachedMediaAt(format, expectedStartMs, seq) + if (segment == null || segments.isEmpty() && !segment.coversPlaybackTime(holder, format, targetMs)) { + val authoritative = sabrSessionStore.findCachedPlaybackMediaAt( + holder, format, expectedStartMs, seq, activeLive && segments.isEmpty(), + ) if (authoritative != null) { - seq = authoritative.header.sequenceNumber + seq = authoritative.sequence holder.session.streamState.jumpBufferedTo(format, seq) - continue + segment = authoritative } } if (segment == null) { @@ -108,12 +140,22 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi blockedRequest = mediaRequest break } - if (segments.isEmpty() && segment.startMs > targetMs && seq > 1) { - seq = previousSequence(seq, segment, targetMs) + if (segments.isEmpty() && holder.failLivePlaybackDiscontinuity( + format, targetMs, segment, continuesServedTrack && request.bufferedRanges.any { it.itag == format.itag }, + ) + ) { + blockedBy = "${format.trackName()}:${format.itag}:$seq discontinuity" + break + } + if (!activeLive && segments.isEmpty() && segment.startMs > targetMs && seq > 1) { + seq = holder.previousPlaybackSequence(format, seq, segment, targetMs) continue } val windowSegment = segment.toWindowSegment(holder, format) segments += windowSegment + if (activeLive && segments.size == 1) { + goalEndMs = maxOf(goalEndMs, windowSegment.startMs + request.bufferGoalMs.coerceAtLeast(1L)) + } coveredEndMs = windowSegment.startMs + windowSegment.durationMs if (endSequence > 0 && seq >= endSequence) { atEnd = true @@ -125,10 +167,12 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi if (blockedBy == null && coveredEndMs < goalEndMs && !atEnd) { blockedBy = "${format.trackName()}:${format.itag}:$seq window capped" } + val mediaBasePath = SabrPlaybackPaths.mediaBasePath(holder.sessionToken) + val defaultInitUrl = "$mediaBasePath/${format.itag}/init?generation=${holder.activeGeneration()}" return TrackBuildResult( track = SabrPlaybackWindowTrack( mime = format.mimeType.orEmpty(), - initUrl = "${SabrPlaybackPaths.mediaBasePath(holder.sessionToken)}/${format.itag}/init?generation=${holder.activeGeneration()}", + initUrl = defaultInitUrl, segments = segments, ), blockedBy = blockedBy, @@ -138,24 +182,14 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi ) } - private fun previousSequence(sequence: Int, segment: CachedSabrSegment, targetMs: Long): Int { - val durationMs = segment.durationMs.coerceAtLeast(1L) - val leadMs = segment.startMs - targetMs - val count = ((leadMs + durationMs - 1L) / durationMs).coerceAtLeast(1L) - return (sequence - count.coerceAtMost((sequence - 1).toLong()).toInt()).coerceAtLeast(1) - } - - private fun CachedSabrSegment.covers(targetMs: Long): Boolean = - startMs >= 0L && durationMs > 0L && targetMs >= startMs && targetMs < startMs + durationMs - private fun CachedSabrSegment.toWindowSegment( holder: SabrSessionHolder, format: YoutubeSabrFormat, ): SabrPlaybackWindowSegment { val startMs = startMs.takeIf { it >= 0L } - ?: holder.session.streamState.getSegmentStartMs(format, sequence).coerceAtLeast(0L) + ?: holder.playbackSegmentStartMs(format, sequence) val durationMs = durationMs.takeIf { it > 0L } - ?: (holder.session.streamState.getSegmentEndMs(format, sequence) - startMs).coerceAtLeast(1L) + ?: holder.playbackSegmentDurationMs(format, sequence) return SabrPlaybackWindowSegment( url = "${SabrPlaybackPaths.mediaBasePath(holder.sessionToken)}/${format.itag}/segment/$sequence?generation=${holder.activeGeneration()}", startMs = startMs, @@ -176,12 +210,6 @@ internal class SabrPlaybackWindowBuilder(private val sabrSessionStore: SabrSessi private companion object { const val MAX_SEGMENTS_PER_TRACK = 12 const val MIN_READY_AHEAD_MS = 1_000L + const val LIVE_STARTUP_READY_AHEAD_MS = 8_000L } } - -internal data class SabrPlaybackWindowBuildResult( - val response: SabrPlaybackWindowReadyResponse, - val blockedBy: String?, - val blockedRequest: SabrSegmentRequest?, - val isReady: Boolean, -) diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowHandler.kt b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowHandler.kt index 744f2837..5e80d282 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowHandler.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowHandler.kt @@ -5,6 +5,9 @@ import dev.typetype.server.services.SabrPlaybackDiagnostics import dev.typetype.server.services.SabrSessionHolder import dev.typetype.server.services.SabrSessionStore import dev.typetype.server.services.pendingSegmentDemandSummary +import dev.typetype.server.services.livePlaybackSnapshot +import dev.typetype.server.services.liveRetryAfterMs +import dev.typetype.server.services.resolvePlaybackStartMs import io.ktor.http.HttpStatusCode import io.ktor.server.application.ApplicationCall import io.ktor.server.request.receive @@ -21,7 +24,7 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi val holder = validatedHolder(call, sessionId, request) ?: return holder.setActiveTracks(videoActive = !request.audioOnly, audioActive = true) - holder.setPlayerTimeMs(request.playerTimeMs) + holder.setPlayerTimeMs(holder.resolvePlaybackStartMs(request.playerTimeMs)) holder.applyClientPreferences() sabrSessionStore.startPump(holder) @@ -29,7 +32,7 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi if (window.isReady) { return call.respond(HttpStatusCode.OK, window.response) } - call.respond(HttpStatusCode.Accepted, holder.preparingResponse(request, window.blockedBy ?: "window pending")) + call.respond(HttpStatusCode.Accepted, holder.preparingResponse(request, window)) } suspend fun position(call: ApplicationCall, sessionId: String) { @@ -37,16 +40,18 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi ?: return call.respond(HttpStatusCode.BadRequest, ErrorResponse("Invalid position request")) val holder = validatedHolder(call, sessionId, request) ?: return holder.setActiveTracks(videoActive = !request.audioOnly, audioActive = true) - holder.setPlayerTimeMs(request.playerTimeMs) + val playerTimeMs = holder.resolvePlaybackStartMs(request.playerTimeMs) + holder.setPlayerTimeMs(playerTimeMs) holder.applyClientPreferences() call.respond( SabrPlaybackPositionResponse( sessionId = holder.sessionToken, generation = holder.activeGeneration(), - playerTimeMs = request.playerTimeMs.coerceAtLeast(0L), + playerTimeMs = playerTimeMs, readerHeadMs = holder.readerHeadMs(), readerTailMs = holder.readerTailMs(), bufferedEdgeMs = holder.session.streamState.getMinBufferedEndMs(), + live = holder.livePlaybackSnapshot()?.toResponse(), ) ) } @@ -56,7 +61,7 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi ?: return call.respond(HttpStatusCode.BadRequest, ErrorResponse("Invalid prefetch request")) val holder = validatedHolder(call, sessionId, request) ?: return holder.setActiveTracks(videoActive = !request.audioOnly, audioActive = true) - holder.setPlayerTimeMs(request.playerTimeMs) + holder.setPlayerTimeMs(holder.resolvePlaybackStartMs(request.playerTimeMs)) holder.applyClientPreferences() sabrSessionStore.startPump(holder) val window = buildWithTargetedPrefetch(holder, request) @@ -72,7 +77,7 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi holder.applyClientPreferences() val window = windowBuilder.build(holder, request) if (window.isReady) return call.respond(HttpStatusCode.OK, window.response) - call.respond(HttpStatusCode.Accepted, holder.preparingResponse(request, window.blockedBy ?: "window pending")) + call.respond(HttpStatusCode.Accepted, holder.preparingResponse(request, window)) } private suspend fun buildWithTargetedPrefetch( @@ -80,19 +85,23 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi request: SabrPlaybackWindowRequest, ): SabrPlaybackWindowBuildResult { val window = windowBuilder.build(holder, request) - val blockedRequest = window.blockedRequest ?: return window - sabrSessionStore.requestSegmentDemand(holder, blockedRequest, request.generation) + window.blockedRequests.forEach { blockedRequest -> + sabrSessionStore.requestSegmentDemand(holder, blockedRequest, request.generation) + } return window } - private suspend fun SabrSessionHolder.preparingResponse(request: SabrPlaybackWindowRequest, blockedBy: String): SabrPlaybackWindowPreparingResponse = + private suspend fun SabrSessionHolder.preparingResponse( + request: SabrPlaybackWindowRequest, + window: SabrPlaybackWindowBuildResult, + ): SabrPlaybackWindowPreparingResponse = SabrPlaybackWindowPreparingResponse( sessionId = sessionToken, generation = activeGeneration(), ready = false, - retryAfterMs = RETRY_AFTER_MS, + retryAfterMs = liveRetryAfterMs(window.blockedRequests), status = playbackState().name.lowercase(), - blockedBy = SabrPlaybackDiagnostics.blocker(this) ?: blockedBy, + blockedBy = SabrPlaybackDiagnostics.blocker(this) ?: window.blockedBy ?: "window pending", playerTimeMs = request.playerTimeMs.coerceAtLeast(0L), readerHeadMs = readerHeadMs(), readerTailMs = readerTailMs(), @@ -103,6 +112,7 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi terminalError = terminalFailure() ?: networkFailure(), recoveryAction = recovery.action(this), retryVideoItags = recovery.retryVideoItags(this), + live = livePlaybackSnapshot()?.toResponse(), ) private suspend fun SabrSessionHolder.prefetchResponse( @@ -112,7 +122,7 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi sessionId = sessionToken, generation = activeGeneration(), ready = window.isReady, - retryAfterMs = if (window.isReady) null else RETRY_AFTER_MS, + retryAfterMs = if (window.isReady) null else liveRetryAfterMs(window.blockedRequests), status = playbackState().name.lowercase(), segmentsUrl = "${SabrPlaybackPaths.mediaBasePath(sessionToken)}/segments", stateUrl = "${SabrPlaybackPaths.mediaBasePath(sessionToken)}/state", @@ -127,6 +137,7 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi terminalError = terminalFailure() ?: networkFailure(), recoveryAction = recovery.action(this), retryVideoItags = recovery.retryVideoItags(this), + live = livePlaybackSnapshot()?.toResponse(), ) private suspend fun validatedHolder( @@ -168,8 +179,4 @@ internal class SabrPlaybackWindowHandler(private val sabrSessionStore: SabrSessi request.videoItag == videoFormat.itag && request.audioItag == audioFormat.itag && request.audioTrackId == audioFormat.audioTrackId private fun SabrSegmentRequest.summary(): String = "${format.itag}:$sequenceNumber" - - private companion object { - const val RETRY_AFTER_MS = 500L - } } diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowTiming.kt b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowTiming.kt index f57e9faf..f3486d6e 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowTiming.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrPlaybackWindowTiming.kt @@ -1,9 +1,15 @@ package dev.typetype.server.routes +import dev.typetype.server.services.CachedSabrSegment import dev.typetype.server.services.SabrSessionHolder +import dev.typetype.server.services.livePlaybackSnapshot +import dev.typetype.server.services.playbackSegmentDurationMs import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat internal fun SabrSessionHolder.durationMs(): Long { + livePlaybackSnapshot()?.let { live -> + if (live.active) return live.seekableEndMs + } val audioEndMs = indexedEndMs(audioFormat) val videoEndMs = indexedEndMs(videoFormat) if (audioEndMs > 0L && videoEndMs > 0L) return maxOf(audioEndMs, videoEndMs) @@ -16,6 +22,9 @@ internal fun SabrSessionHolder.indexedEndMs(format: YoutubeSabrFormat): Long { } internal fun SabrSessionHolder.readyEndMs(format: YoutubeSabrFormat, requestedEndMs: Long): Long { + livePlaybackSnapshot()?.let { live -> + if (live.active) return minOf(live.seekableEndMs, requestedEndMs) + } val durationMs = indexedEndMs(format).takeIf { it > 0L } ?: format.approxDurationMs.coerceAtLeast(0L) return minOf(durationMs, requestedEndMs) } @@ -33,3 +42,28 @@ internal fun SabrPlaybackWindowRequest.bufferedEndFor( ?.coerceAtLeast(startMs) ?: startMs } + +internal fun resolvedPlaybackStartMs( + requestedStartMs: Long, + activeLive: Boolean, + vararg tracks: SabrPlaybackWindowTrack?, +): Long { + if (!activeLive) return requestedStartMs + return tracks.asSequence() + .filterNotNull() + .mapNotNull { it.segments.firstOrNull()?.startMs } + .fold(requestedStartMs.coerceAtLeast(0L), ::maxOf) +} + +internal fun SabrSessionHolder.previousPlaybackSequence( + format: YoutubeSabrFormat, + sequence: Int, + segment: CachedSabrSegment, + targetMs: Long, +): Int { + val durationMs = segment.durationMs.takeIf { it > 0L } + ?: playbackSegmentDurationMs(format, sequence) + val leadMs = segment.startMs - targetMs + val count = ((leadMs + durationMs - 1L) / durationMs).coerceAtLeast(1L) + return (sequence - count.coerceAtMost((sequence - 1).toLong()).toInt()).coerceAtLeast(1) +} diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrSessionStateHandler.kt b/src/main/kotlin/dev/typetype/server/routes/SabrSessionStateHandler.kt index 391a8194..d68583b8 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrSessionStateHandler.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrSessionStateHandler.kt @@ -58,6 +58,7 @@ internal class SabrSessionStateHandler(private val sabrSessionStore: SabrSession val edgeMs = holder.session.streamState.getMinBufferedEndMs() val startMs = holder.session.streamState.getSegmentStartMs(format, sequence).coerceAtLeast(0L) holder.setReaderPosition(format, startMs) + holder.setRequestedSeekTimeMs(holder.playerTimeMs()) if (startMs < edgeMs) { holder.requestRefetch(request) } else if (startMs > edgeMs + FORWARD_SEEK_AHEAD_MS) { diff --git a/src/main/kotlin/dev/typetype/server/routes/SabrStreamContractFilter.kt b/src/main/kotlin/dev/typetype/server/routes/SabrStreamContractFilter.kt index 6d60f856..de4379d2 100644 --- a/src/main/kotlin/dev/typetype/server/routes/SabrStreamContractFilter.kt +++ b/src/main/kotlin/dev/typetype/server/routes/SabrStreamContractFilter.kt @@ -47,7 +47,7 @@ internal fun StreamResponse.withoutSabrStreams(): StreamResponse = copy( ) internal fun StreamResponse.onlySabrStreams(): StreamResponse = copy( - hlsUrl = if (isLive) hlsUrl else "", + hlsUrl = "", dashMpdUrl = "", videoStreams = videoStreams.filter { it.deliveryMethod == SABR_DELIVERY_METHOD }, videoOnlyStreams = videoOnlyStreams.filter { it.deliveryMethod == SABR_DELIVERY_METHOD }, @@ -78,8 +78,8 @@ private fun YoutubeSabrFormat.toVideoStreamItem(videoId: String, info: YoutubeSa height = height.coerceAtLeast(0), fps = 0, contentLength = contentLength.coerceAtLeast(0L), - initStart = initRangeStart.coerceAtLeast(0L), - initEnd = initRangeEnd.coerceAtLeast(0L), + initStart = 0L, + initEnd = 0L, indexStart = 0L, indexEnd = 0L, deliveryMethod = SABR_DELIVERY_METHOD, diff --git a/src/main/kotlin/dev/typetype/server/routes/StreamRoutes.kt b/src/main/kotlin/dev/typetype/server/routes/StreamRoutes.kt index 7fa260a7..4be035bd 100644 --- a/src/main/kotlin/dev/typetype/server/routes/StreamRoutes.kt +++ b/src/main/kotlin/dev/typetype/server/routes/StreamRoutes.kt @@ -126,7 +126,7 @@ private fun Route.streamRoute( call.respond(data) } is ExtractionResult.BadRequest -> call.respond(HttpStatusCode.BadRequest, result.toErrorResponse()) - is ExtractionResult.Failure -> call.respond(HttpStatusCode.UnprocessableEntity, ErrorResponse(result.message)) + is ExtractionResult.Failure -> call.respond(HttpStatusCode.UnprocessableEntity, result.toErrorResponse()) } } } @@ -195,5 +195,7 @@ private fun ExtractionResult.BadRequest.toErrorResponse(): ErrorResponse = if (message == YOUTUBE_SESSION_RECONNECT_ERROR) { ErrorResponse(message, "youtube_session_needs_reconnect") } else { - ErrorResponse(message) + ErrorResponse(message, code) } + +private fun ExtractionResult.Failure.toErrorResponse(): ErrorResponse = ErrorResponse(message, code) diff --git a/src/main/kotlin/dev/typetype/server/routes/WatchLaterRoutes.kt b/src/main/kotlin/dev/typetype/server/routes/WatchLaterRoutes.kt index 3d5d78b2..a84b0af7 100644 --- a/src/main/kotlin/dev/typetype/server/routes/WatchLaterRoutes.kt +++ b/src/main/kotlin/dev/typetype/server/routes/WatchLaterRoutes.kt @@ -16,8 +16,9 @@ import io.ktor.server.routing.post fun Route.watchLaterRoutes(watchLaterService: WatchLaterService, authService: AuthService, metadataRepairService: UserVideoMetadataRepairService? = null) { get("/watch-later") { call.withJwtAuth(authService) { userId -> - metadataRepairService?.repairWatchLater(userId) - call.respond(watchLaterService.getAll(userId)) + val items = watchLaterService.getAll(userId) + metadataRepairService?.scheduleWatchLater(call.application, userId) + call.respond(items) } } post("/watch-later") { diff --git a/src/main/kotlin/dev/typetype/server/services/CachedSabrSegment.kt b/src/main/kotlin/dev/typetype/server/services/CachedSabrSegment.kt index bc77ae0c..205a02a0 100644 --- a/src/main/kotlin/dev/typetype/server/services/CachedSabrSegment.kt +++ b/src/main/kotlin/dev/typetype/server/services/CachedSabrSegment.kt @@ -22,13 +22,16 @@ internal data class CachedSabrSegment( val length: Int get() = byteLength.takeIf { it >= 0 } ?: bytes.size } -internal fun SabrMediaSegment.toCachedSabrSegment(mimeType: String): CachedSabrSegment = CachedSabrSegment( +internal fun SabrMediaSegment.toCachedSabrSegment( + mimeType: String, + bytes: ByteArray = data, +): CachedSabrSegment = CachedSabrSegment( itag = header.itag, sequence = header.sequenceNumber, init = header.isInitSegment, startMs = header.startMs, durationMs = header.durationMs, mimeType = mimeType, - bytesBase64 = Base64.getEncoder().encodeToString(data), - byteLength = length, + bytesBase64 = Base64.getEncoder().encodeToString(bytes), + byteLength = bytes.size, ) diff --git a/src/main/kotlin/dev/typetype/server/services/NewPipeInitializer.kt b/src/main/kotlin/dev/typetype/server/services/NewPipeInitializer.kt index 210be972..1a506bab 100644 --- a/src/main/kotlin/dev/typetype/server/services/NewPipeInitializer.kt +++ b/src/main/kotlin/dev/typetype/server/services/NewPipeInitializer.kt @@ -3,21 +3,27 @@ package dev.typetype.server.services import dev.typetype.server.downloader.OkHttpDownloader import org.schabi.newpipe.extractor.NewPipe import org.schabi.newpipe.extractor.services.youtube.YoutubeApiDecoder +import java.net.ProxySelector object NewPipeInitializer { @Volatile private var initialized = false @Volatile private var decoderServiceUrl: String? = null - fun init(tokenServiceUrl: String? = null): Unit { + fun init( + tokenServiceUrl: String? = null, + youtubeProxySelector: ProxySelector? = null, + ): Unit { NewPipe.setYoutubePlayerClient(YOUTUBE_PLAYER_CLIENT) val normalizedUrl = tokenServiceUrl?.trim()?.takeIf { it.isNotBlank() } if (normalizedUrl != null && normalizedUrl != decoderServiceUrl) { YoutubeApiDecoder.setLocalDecoder(TypetypeTokenYoutubeJavaScriptDecoder(normalizedUrl)) decoderServiceUrl = normalizedUrl } - if (initialized) return - NewPipe.init(OkHttpDownloader.instance()) - initialized = true + if (!initialized) { + NewPipe.init(OkHttpDownloader.instance(youtubeProxySelector)) + initialized = true + } + NewPipe.setYoutubeSessionPoTokenProvider(TypetypeYoutubeSessionPoTokenProvider) } private const val YOUTUBE_PLAYER_CLIENT = "mweb" diff --git a/src/main/kotlin/dev/typetype/server/services/PlaylistService.kt b/src/main/kotlin/dev/typetype/server/services/PlaylistService.kt index 3bdd8fdb..df2e9ad5 100644 --- a/src/main/kotlin/dev/typetype/server/services/PlaylistService.kt +++ b/src/main/kotlin/dev/typetype/server/services/PlaylistService.kt @@ -8,9 +8,11 @@ import dev.typetype.server.models.PlaylistReorderResult import dev.typetype.server.models.PlaylistVideoItem import org.jetbrains.exposed.v1.core.SortOrder import org.jetbrains.exposed.v1.core.and +import org.jetbrains.exposed.v1.core.count import org.jetbrains.exposed.v1.core.eq import org.jetbrains.exposed.v1.jdbc.deleteWhere import org.jetbrains.exposed.v1.jdbc.insert +import org.jetbrains.exposed.v1.jdbc.select import org.jetbrains.exposed.v1.jdbc.selectAll import org.jetbrains.exposed.v1.jdbc.update import java.util.UUID @@ -21,12 +23,14 @@ class PlaylistService { .where { PlaylistsTable.userId eq userId } .orderBy(PlaylistsTable.createdAt to SortOrder.DESC) .toList() + val videoCount = PlaylistVideosTable.id.count() + val videoCounts = PlaylistVideosTable + .select(PlaylistVideosTable.playlistId, videoCount) + .where { PlaylistVideosTable.userId eq userId } + .groupBy(PlaylistVideosTable.playlistId) + .associate { row -> row[PlaylistVideosTable.playlistId] to row[videoCount].toInt() } playlists.map { row -> - val count = PlaylistVideosTable.selectAll() - .where { (PlaylistVideosTable.playlistId eq row[PlaylistsTable.id]) and (PlaylistVideosTable.userId eq userId) } - .count() - .toInt() - row.toPlaylistSummary(count) + row.toPlaylistSummary(videoCounts[row[PlaylistsTable.id]] ?: 0) } } diff --git a/src/main/kotlin/dev/typetype/server/services/PublicCachePolicy.kt b/src/main/kotlin/dev/typetype/server/services/PublicCachePolicy.kt index f9d7fd88..7513ab18 100644 --- a/src/main/kotlin/dev/typetype/server/services/PublicCachePolicy.kt +++ b/src/main/kotlin/dev/typetype/server/services/PublicCachePolicy.kt @@ -23,6 +23,8 @@ internal object PublicCachePolicy { fun channelTtl(url: String, nextpage: String?, sort: String?): Long = when { url.contains("/search", ignoreCase = true) -> 600L + url.contains("/streams", ignoreCase = true) || url.contains("/livestreams", ignoreCase = true) -> + if (nextpage == null) 60L else 300L nextpage != null -> 1_800L url.contains("/shorts", ignoreCase = true) -> 900L sort.equals("latest", ignoreCase = true) -> 900L diff --git a/src/main/kotlin/dev/typetype/server/services/RelatedItemMappers.kt b/src/main/kotlin/dev/typetype/server/services/RelatedItemMappers.kt index 0a61f9de..bb546633 100644 --- a/src/main/kotlin/dev/typetype/server/services/RelatedItemMappers.kt +++ b/src/main/kotlin/dev/typetype/server/services/RelatedItemMappers.kt @@ -54,6 +54,7 @@ internal fun StreamInfoItem.toVideoItem(fallbackAvatarUrl: String = ""): VideoIt isPostLive = apiStreamType.isPostLiveStreamType(), isLiveContent = apiStreamType.isLiveContentType(), isShortFormContent = isShortFormContent, + requiresMembership = requiresMembership(), uploaderVerified = isUploaderVerified, shortDescription = shortDescription?.takeIf { it.isNotBlank() }, publishedAt = PublishedAtMapper.fromUploaded(uploaded), diff --git a/src/main/kotlin/dev/typetype/server/services/SabrCachedSegmentLocator.kt b/src/main/kotlin/dev/typetype/server/services/SabrCachedSegmentLocator.kt index 266b5c78..35871168 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrCachedSegmentLocator.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrCachedSegmentLocator.kt @@ -9,11 +9,16 @@ internal fun YoutubeSabrSession.findCachedMediaAt( format: YoutubeSabrFormat, targetMs: Long, predictedSequence: Int, + allowFollowing: Boolean = false, ): SabrMediaSegment? { for (distance in 1..MAX_SEQUENCE_DISTANCE) { cachedMedia(format, predictedSequence + distance)?.takeIf { it.covers(targetMs) }?.let { return it } cachedMedia(format, predictedSequence - distance)?.takeIf { it.covers(targetMs) }?.let { return it } } + if (!allowFollowing) return null + for (distance in 1..MAX_SEQUENCE_DISTANCE) { + cachedMedia(format, predictedSequence + distance)?.takeIf { it.startsWithinNextSegment(targetMs) }?.let { return it } + } return null } @@ -30,5 +35,10 @@ private fun SabrMediaSegment.covers(targetMs: Long): Boolean { targetMs < startMs + header.durationMs } +private fun SabrMediaSegment.startsWithinNextSegment(targetMs: Long): Boolean { + val leadMs = header.startMs - targetMs + return header.startMs >= 0L && header.durationMs > 0L && leadMs in -TIMING_TOLERANCE_MS..header.durationMs +} + private const val MAX_SEQUENCE_DISTANCE = 24 private const val TIMING_TOLERANCE_MS = 2L diff --git a/src/main/kotlin/dev/typetype/server/services/SabrDemandAttemptFinisher.kt b/src/main/kotlin/dev/typetype/server/services/SabrDemandAttemptFinisher.kt index 84d86c7f..9b8546a9 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrDemandAttemptFinisher.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrDemandAttemptFinisher.kt @@ -4,16 +4,38 @@ import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession internal object SabrDemandAttemptFinisher { + fun interruptCompletedInFlightDemand(holder: SabrSessionHolder, demand: SabrInFlightDemand): Boolean = + synchronized(holder) { + val state = holder.playbackState() + if (state == SabrPlaybackState.TERMINAL || state == SabrPlaybackState.NETWORK_FAILED) return@synchronized false + if (holder.inFlightSegmentDemand()?.identity != demand.identity) return@synchronized false + if (holder.session.getCachedSegment(demand.request) == null) return@synchronized false + holder.setPlaybackState(SabrPlaybackState.IDLE) + true + } + + fun expireStalledInFlightDemand(holder: SabrSessionHolder, demand: SabrInFlightDemand): Boolean = + synchronized(holder) { + val state = holder.playbackState() + if (state == SabrPlaybackState.TERMINAL || state == SabrPlaybackState.NETWORK_FAILED) return@synchronized false + if (holder.inFlightSegmentDemand()?.identity != demand.identity) return@synchronized false + holder.clearSegmentDemands() + val message = "SABR demand stalled for ${demand.request.summary()}" + holder.failTerminal(if (demand.futureLiveRequest) sabrRecoverableFailureMessage(message) else message) + true + } + fun expireStalledDemand( holder: SabrSessionHolder, request: SabrSegmentRequest, identity: String, + recoverable: Boolean = false, ): Boolean = synchronized(holder) { val state = holder.playbackState() if (state == SabrPlaybackState.TERMINAL || state == SabrPlaybackState.NETWORK_FAILED) return@synchronized false val current = holder.nextSegmentDemand() ?: return@synchronized false if (!current.matches(request) || holder.segmentDemandIdentity(current) != identity) return@synchronized false - fail(holder, request, identity) + fail(holder, request, identity, recoverable) } fun finish( @@ -22,6 +44,7 @@ internal object SabrDemandAttemptFinisher { identity: String, result: YoutubeSabrSession.DemandResponseResult, runtime: SabrPumpRuntime, + wasFutureLiveRequest: Boolean, ): Boolean = synchronized(holder) { if (!holder.isSegmentDemandActive(request, identity)) { runtime.finishDemand(identity) @@ -30,13 +53,20 @@ internal object SabrDemandAttemptFinisher { } val resolved = holder.resolveSegmentDemand(request, identity) SabrPumpLogger.finish(holder, "demand", request, result.segmentCount) + if (!resolved && (wasFutureLiveRequest || holder.isFutureLiveRequest(request))) { + runtime.finishDemand(identity) + holder.setPlaybackState(SabrPlaybackState.WAITING_FOR_LIVE) + return@synchronized false + } val action = runtime.demandRecoveryAction( requestKey = identity, targetTrackSegmentCount = result.targetTrackSegmentCount, resolved = resolved, ) val recovering = recover(holder, request, identity, action, runtime) - if (holder.playbackState() != SabrPlaybackState.TERMINAL) { + if (holder.playbackState() != SabrPlaybackState.TERMINAL && + holder.playbackState() != SabrPlaybackState.WAITING_FOR_LIVE + ) { holder.setPlaybackState(SabrPlaybackState.IDLE) } resolved || recovering @@ -67,11 +97,13 @@ internal object SabrDemandAttemptFinisher { holder: SabrSessionHolder, request: SabrSegmentRequest, identity: String, + recoverable: Boolean = false, ): Boolean { val failed = holder.clearSegmentDemand(request, identity) if (failed) { holder.clearSegmentDemands() - holder.failTerminal("SABR demand stalled for ${request.summary()}") + val message = "SABR demand stalled for ${request.summary()}" + holder.failTerminal(if (recoverable) sabrRecoverableFailureMessage(message) else message) } return failed } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrDemandWatchdog.kt b/src/main/kotlin/dev/typetype/server/services/SabrDemandWatchdog.kt index a66bd396..dacd0352 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrDemandWatchdog.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrDemandWatchdog.kt @@ -10,20 +10,42 @@ internal class SabrDemandWatchdog( while (isAlive()) { val state = holder.playbackState() if (state == SabrPlaybackState.TERMINAL || state == SabrPlaybackState.NETWORK_FAILED) return false + val inFlightDemand = holder.inFlightSegmentDemand() + if (inFlightDemand != null) { + if (inFlightDemand.futureLiveRequest) holder.setPlaybackState(SabrPlaybackState.WAITING_FOR_LIVE) + val nowMs = clock() + val lastProgressAtMs = inFlightDemand.observeProgress(holder.session.mediaProgressVersion, nowMs) + val completedIdle = holder.session.getCachedSegment(inFlightDemand.request) != null && + nowMs - lastProgressAtMs >= SabrPumpPolicy.COMPLETED_DEMAND_IDLE_MS + if (completedIdle && SabrDemandAttemptFinisher.interruptCompletedInFlightDemand(holder, inFlightDemand)) { + return true + } + if (nowMs - inFlightDemand.registeredAtMs >= SabrPumpPolicy.DEMAND_TARGET_DEADLINE_MS && + SabrDemandAttemptFinisher.expireStalledInFlightDemand(holder, inFlightDemand) + ) { + return true + } + delay(if (inFlightDemand.futureLiveRequest) maxOf(intervalMs, LIVE_EDGE_POLL_MS) else intervalMs) + continue + } val request = holder.nextSegmentDemand() if (request == null) { delay(intervalMs) continue } + val futureLiveRequest = holder.isFutureLiveRequest(request) + if (futureLiveRequest) { + holder.setPlaybackState(SabrPlaybackState.WAITING_FOR_LIVE) + } val identity = holder.segmentDemandIdentity(request) val registeredAtMs = identity?.let { holder.segmentDemandRegisteredAtMs(request, it) } if (identity != null && registeredAtMs != null && clock() - registeredAtMs >= SabrPumpPolicy.DEMAND_TARGET_DEADLINE_MS && - SabrDemandAttemptFinisher.expireStalledDemand(holder, request, identity) + SabrDemandAttemptFinisher.expireStalledDemand(holder, request, identity, futureLiveRequest) ) { return true } - delay(intervalMs) + delay(if (futureLiveRequest) maxOf(intervalMs, LIVE_EDGE_POLL_MS) else intervalMs) } return false } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrFallbackStreamMapper.kt b/src/main/kotlin/dev/typetype/server/services/SabrFallbackStreamMapper.kt index 5d86870c..5b8f0151 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrFallbackStreamMapper.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrFallbackStreamMapper.kt @@ -77,8 +77,8 @@ private fun YoutubeSabrFormat.toFallbackVideo(videoId: String): VideoStreamItem? height = height.coerceAtLeast(0), fps = 0, contentLength = contentLength.coerceAtLeast(0L), - initStart = initRangeStart.coerceAtLeast(0L), - initEnd = initRangeEnd.coerceAtLeast(0L), + initStart = 0L, + initEnd = 0L, indexStart = 0L, indexEnd = 0L, deliveryMethod = SABR_METHOD, @@ -99,8 +99,8 @@ private fun YoutubeSabrFormat.toFallbackAudio(videoId: String): AudioStreamItem? quality = audioQuality, itag = itag, contentLength = contentLength.coerceAtLeast(0L), - initStart = initRangeStart.coerceAtLeast(0L), - initEnd = initRangeEnd.coerceAtLeast(0L), + initStart = 0L, + initEnd = 0L, indexStart = 0L, indexEnd = 0L, audioTrackId = audioTrackId, diff --git a/src/main/kotlin/dev/typetype/server/services/SabrFallbackStreamService.kt b/src/main/kotlin/dev/typetype/server/services/SabrFallbackStreamService.kt index 5a86dbe3..5be6711a 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrFallbackStreamService.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrFallbackStreamService.kt @@ -23,11 +23,14 @@ internal class SabrFallbackStreamService( return@coroutineScope ExtractionResult.Success(session.toFallbackStreamResponse(videoId)) } val playable = prepared?.await() - if (response.hasPlayableStreams() || videoId == null || playable == null) return@coroutineScope result + if (response.hasSabrStreams() || videoId == null || playable == null) return@coroutineScope result ExtractionResult.Success(response.withSabrFallback(videoId, playable.info)) } } -private fun StreamResponse.hasPlayableStreams(): Boolean = - videoStreams.isNotEmpty() || videoOnlyStreams.isNotEmpty() || audioStreams.isNotEmpty() || - hlsUrl.isNotBlank() || dashMpdUrl.isNotBlank() +private fun StreamResponse.hasSabrStreams(): Boolean = + videoStreams.any { it.deliveryMethod == SABR_METHOD } || + videoOnlyStreams.any { it.deliveryMethod == SABR_METHOD } || + audioStreams.any { it.deliveryMethod == SABR_METHOD } + +private const val SABR_METHOD = "sabr" diff --git a/src/main/kotlin/dev/typetype/server/services/SabrInFlightDemandTracker.kt b/src/main/kotlin/dev/typetype/server/services/SabrInFlightDemandTracker.kt new file mode 100644 index 00000000..c52ac13d --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrInFlightDemandTracker.kt @@ -0,0 +1,64 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import java.util.concurrent.ConcurrentHashMap + +internal class SabrInFlightDemand( + val request: SabrSegmentRequest, + val identity: String, + val registeredAtMs: Long, + val futureLiveRequest: Boolean, +) { + private var lastProgressVersion = Long.MIN_VALUE + private var lastProgressAtMs = registeredAtMs + + fun observeProgress(version: Long, observedAtMs: Long): Long { + if (version != lastProgressVersion) { + lastProgressVersion = version + lastProgressAtMs = observedAtMs + } + return lastProgressAtMs + } +} + +internal object SabrInFlightDemandTracker { + private val demands = ConcurrentHashMap() + + fun begin( + holder: SabrSessionHolder, + request: SabrSegmentRequest, + identity: String, + futureLiveRequest: Boolean, + ): Boolean { + val registeredAtMs = holder.segmentDemandRegisteredAtMs(request, identity) ?: return false + val demand = SabrInFlightDemand( + request, + identity, + registeredAtMs, + futureLiveRequest, + ) + return demands.putIfAbsent(holder.sessionToken, demand) == null + } + + fun current(holder: SabrSessionHolder): SabrInFlightDemand? = demands[holder.sessionToken] + + fun finish(holder: SabrSessionHolder, identity: String): Boolean { + val demand = demands[holder.sessionToken] ?: return false + if (demand.identity != identity) return false + return demands.remove(holder.sessionToken, demand) + } + + fun clearAll(): Unit = demands.clear() +} + +internal fun SabrSessionHolder.beginInFlightSegmentDemand( + request: SabrSegmentRequest, + identity: String, + futureLiveRequest: Boolean, +): Boolean = SabrInFlightDemandTracker.begin(this, request, identity, futureLiveRequest) + +internal fun SabrSessionHolder.inFlightSegmentDemand(): SabrInFlightDemand? = + SabrInFlightDemandTracker.current(this) + +internal fun SabrSessionHolder.finishInFlightSegmentDemand(identity: String): Boolean = + SabrInFlightDemandTracker.finish(this, identity) diff --git a/src/main/kotlin/dev/typetype/server/services/SabrInfoFetcher.kt b/src/main/kotlin/dev/typetype/server/services/SabrInfoFetcher.kt index e6d9bf0b..efaa11c4 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrInfoFetcher.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrInfoFetcher.kt @@ -2,15 +2,11 @@ package dev.typetype.server.services import dev.typetype.server.cache.CacheService import kotlinx.coroutines.Dispatchers -import kotlinx.coroutines.delay import kotlinx.coroutines.withContext import kotlinx.coroutines.withTimeoutOrNull -import org.schabi.newpipe.extractor.localization.ContentCountry -import org.schabi.newpipe.extractor.localization.Localization import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrClientProfile import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo -import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrProbe import org.slf4j.LoggerFactory internal class SabrInfoFetcher( @@ -18,6 +14,7 @@ internal class SabrInfoFetcher( private val sessionClient: TypetypeTokenYoutubeSessionClient? = null, private val infoCache: SabrPreparedInfoCache = SabrPreparedInfoCache(), sharedCache: CacheService? = null, + private val playerInfoProbe: SabrPlayerInfoProbe = PipePipeSabrPlayerInfoProbe, ) { private val repository = SabrInfoRepository(infoCache, sharedCache) @@ -37,14 +34,10 @@ internal class SabrInfoFetcher( return@withContext prepared } } - fetchPlayableWithRetries(videoId, startTimeMs)?.let { + fetchPlayable(videoId, startTimeMs)?.let { logFetch(videoId, startTimeMs, startedAt, "network") return@withContext it } - if (startTimeMs > 0L) fetchPlayableWithRetries(videoId, 0L)?.let { - logFetch(videoId, startTimeMs, startedAt, "network_start0") - return@withContext it - } repository.local(videoId, startTimeMs)?.also { logFetch(videoId, startTimeMs, startedAt, "late_cache") } } @@ -60,7 +53,7 @@ internal class SabrInfoFetcher( suspend fun rememberExtractedInfo(videoId: String, info: YoutubeSabrInfo): Unit { repository.rememberInitialization(videoId, info) - fetchInfoOnce(videoId, startTimeMs = 0L, forceRefresh = false) + fetchInfoOnce(videoId, startTimeMs = 0L) ?.takeIf { it.hasAudioAndVideoFormats() } ?.let { repository.putPrepared(videoId, startTimeMs = 0L, it) } } @@ -70,125 +63,87 @@ internal class SabrInfoFetcher( fun initializationFormat(videoId: String, target: YoutubeSabrFormat): YoutubeSabrFormat? = repository.initializationFormat(videoId, target) - private suspend fun fetchPlayableWithRetries(videoId: String, startTimeMs: Long): SabrPreparedInfo? { - repeat(SabrSessionStoreDefaults.INFO_ATTEMPTS) { attempt -> - fetchInfoOnce(videoId, startTimeMs, forceRefresh = attempt > 0) - ?.let { return repository.putPrepared(videoId, startTimeMs, it) } - if (attempt + 1 < SabrSessionStoreDefaults.INFO_ATTEMPTS) delay(SabrSessionStoreDefaults.INFO_RETRY_DELAY_MS) - } - return null - } + private suspend fun fetchPlayable(videoId: String, startTimeMs: Long): SabrPreparedInfo? = + fetchInfoOnce(videoId, startTimeMs)?.let { repository.putPrepared(videoId, startTimeMs, it) } - private suspend fun fetchInfoOnce(videoId: String, startTimeMs: Long, forceRefresh: Boolean): SabrPreparedInfo? = + private suspend fun fetchInfoOnce(videoId: String, startTimeMs: Long): SabrPreparedInfo? = withTimeoutOrNull(SabrSessionStoreDefaults.INFO_TIMEOUT_MS) { - val refreshedToken = if (forceRefresh) tokenClient.fetch(videoId, forceRefresh = true) else null val tokenSession = sessionClient?.fetchPlaybackSession(videoId) tokenSession?.token ?.takeIf { it.visitorData == tokenSession.info.visitorData } ?.let { sessionToken -> - SabrPreparedInfo(tokenSession.info, sessionToken) + SabrPreparedInfo(tokenSession.info, sessionToken, tokenSession.isLive, tokenSession.isLiveContent) .takeIf { it.hasAudioAndVideoFormats() } ?.let { return@withTimeoutOrNull it } } - val token = refreshedToken ?: tokenClient.fetch(videoId) + val token = tokenClient.fetch(videoId) ?: return@withTimeoutOrNull null.also { logger.warn( - "sabr_probe event=token_missing videoId={} startTimeMs={} forceRefresh={}", + "sabr_probe event=token_missing videoId={} startTimeMs={}", videoId, startTimeMs, - forceRefresh, ) } tokenSession ?.takeIf { it.info.visitorData == token.visitorData } - ?.let { SabrPreparedInfo(it.info, token) } + ?.let { + SabrPreparedInfo(it.info, token, it.isLive, it.isLiveContent) + } ?.takeIf { it.hasAudioAndVideoFormats() } ?.let { return@withTimeoutOrNull it } + val recovery = SabrPlayerContextRecovery(videoId, token, tokenClient, playerInfoProbe) CLIENT_PROFILES.firstNotNullOfOrNull { profile -> - fetchInfoForProfileWithTokenFallback(videoId, startTimeMs, token, profile) + fetchInfoForProfile(videoId, startTimeMs, profile, recovery) ?.takeIf { it.hasAudioAndVideoFormats() } } } - private fun fetchInfoForProfileWithTokenFallback( - videoId: String, - startTimeMs: Long, - token: SabrTokenBundle, - profile: YoutubeSabrClientProfile, - ): SabrPreparedInfo? { - fetchInfoForProfile( - videoId = videoId, - startTimeMs = startTimeMs, - token = token, - profile = profile, - poToken = token.visitorBoundPoToken, - tokenKind = VISITOR_BOUND_TOKEN_KIND, - )?.let { return it } - if (token.visitorBoundPoToken == token.videoBoundPoToken) return null - return fetchInfoForProfile( - videoId = videoId, - startTimeMs = startTimeMs, - token = token, - profile = profile, - poToken = token.videoBoundPoToken, - tokenKind = VIDEO_BOUND_TOKEN_KIND, - ) - } - private fun fetchInfoForProfile( videoId: String, startTimeMs: Long, - token: SabrTokenBundle, profile: YoutubeSabrClientProfile, - poToken: String, - tokenKind: String, + recovery: SabrPlayerContextRecovery, ): SabrPreparedInfo? { - val result = runCatching { - YoutubeSabrProbe.fetchSabrInfo( - videoId, - profile, - LOCALIZATION, - CONTENT_COUNTRY, - poToken, - token.visitorData, - ) - } - result.onFailure { error -> - logger.warn( - "sabr_probe event=fetch_failed videoId={} profile={} tokenKind={} startTimeMs={} errorType={} error={}", - videoId, - profile, - tokenKind, - startTimeMs, - error.javaClass.simpleName, - error.message, - error, - ) - } - return result.getOrNull()?.let { info -> - val prepared = SabrPreparedInfo(info, token) - if (!prepared.hasAudioAndVideoFormats()) { + return when (val result = recovery.fetch(profile)) { + is SabrPlayerProbeResult.Failure -> null.also { logger.warn( - "sabr_probe event=no_av_formats videoId={} profile={} tokenKind={} formatCount={} hasAudio={} hasVideo={} streamingUrl={}", + "sabr_probe event=fetch_failed videoId={} profile={} startTimeMs={} contextRefreshAttempted={} errorType={} error={}", videoId, profile, - tokenKind, - info.formats.size, - info.formats.any { it.isAudio }, - info.formats.any { it.isVideo }, - !info.serverAbrStreamingUrl.isNullOrEmpty(), + startTimeMs, + result.contextRefreshAttempted, + result.error.javaClass.simpleName, + result.error.message, + result.error, ) } - prepared + is SabrPlayerProbeResult.Success -> SabrPreparedInfo(result.info, result.token).also { prepared -> + if (!prepared.hasAudioAndVideoFormats()) logMissingFormats(videoId, profile, result.info) + if (result.contextRefreshed) { + logger.info("sabr_probe event=player_context_refreshed videoId={} profile={}", videoId, profile) + } + } } } + private fun logMissingFormats( + videoId: String, + profile: YoutubeSabrClientProfile, + info: YoutubeSabrInfo, + ): Unit { + logger.warn( + "sabr_probe event=no_av_formats videoId={} profile={} formatCount={} hasAudio={} hasVideo={} streamingUrl={}", + videoId, + profile, + info.formats.size, + info.formats.any { it.isAudio }, + info.formats.any { it.isVideo }, + !info.serverAbrStreamingUrl.isNullOrEmpty(), + ) + } + private companion object { val logger = LoggerFactory.getLogger(SabrInfoFetcher::class.java) - val LOCALIZATION = Localization("en", "US") - val CONTENT_COUNTRY = ContentCountry("US") val CLIENT_PROFILES = listOf(YoutubeSabrClientProfile.WEB, YoutubeSabrClientProfile.MWEB) - const val VIDEO_BOUND_TOKEN_KIND = "video_bound" - const val VISITOR_BOUND_TOKEN_KIND = "visitor_bound" } } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrInitializationData.kt b/src/main/kotlin/dev/typetype/server/services/SabrInitializationData.kt index 7f7f1014..966e84d0 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrInitializationData.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrInitializationData.kt @@ -1,77 +1,79 @@ package dev.typetype.server.services import dev.typetype.server.cache.CacheService -import org.schabi.newpipe.extractor.NewPipe +import kotlinx.coroutines.sync.withLock import org.schabi.newpipe.extractor.localization.Localization +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat import java.security.MessageDigest import java.util.Base64 +import java.util.Collections +import java.util.WeakHashMap import java.util.concurrent.ConcurrentHashMap internal object SabrInitializationData { private const val CACHE_TTL_SECONDS = 21_600L private val memoryCache = ConcurrentHashMap() + private val formatCache = Collections.synchronizedMap(WeakHashMap()) suspend fun ingest( format: YoutubeSabrFormat, holder: SabrSessionHolder, cache: CacheService? = null, ): Boolean { - val data = fetch(format, cache) ?: return false + val data = fetch(holder.key.videoId, format, cache) ?: return false return holder.session.streamState.ingestInitializationData(format, data) } - suspend fun fetch(format: YoutubeSabrFormat, cache: CacheService? = null): ByteArray? { - val key = cacheKey(format) ?: return null - memoryCache[key]?.let { return it } + suspend fun fetch(videoId: String, format: YoutubeSabrFormat, cache: CacheService? = null): ByteArray? { + val key = cacheKey(videoId, format) + formatCache[format]?.let { return it } + memoryCache[key]?.let { + formatCache[format] = it + return it + } cache?.getBytes(key)?.let { bytes -> memoryCache[key] = bytes + formatCache[format] = bytes return bytes } - val bytes = fetchRemote(format) ?: return null - memoryCache[key] = bytes - cache?.setBytes(key, bytes, CACHE_TTL_SECONDS) - return bytes + return null } - fun remember(format: YoutubeSabrFormat, bytes: ByteArray): Unit { - val key = cacheKey(format) ?: return + suspend fun remember( + videoId: String, + format: YoutubeSabrFormat, + bytes: ByteArray, + cache: CacheService? = null, + ): Unit { + val key = cacheKey(videoId, format) memoryCache[key] = bytes + formatCache[format] = bytes + cache?.setBytes(key, bytes, CACHE_TTL_SECONDS) } fun ingestRemembered(format: YoutubeSabrFormat, holder: SabrSessionHolder): Boolean { - val key = cacheKey(format) ?: return false - val bytes = memoryCache[key] ?: return false + val bytes = formatCache[format] ?: return false return holder.session.streamState.ingestInitializationData(format, bytes) } - suspend fun fetchFallback(holder: SabrSessionHolder, format: YoutubeSabrFormat, cache: CacheService? = null): ByteArray? { - val key = cacheKey(format) ?: return null - memoryCache[key]?.let { holder.session.streamState.ingestInitializationData(format, it); return it } - cache?.getBytes(key)?.let { bytes -> - memoryCache[key] = bytes - holder.session.streamState.ingestInitializationData(format, bytes) - return bytes - } - val bytes = runCatchingNonCancellation { - holder.session.fetchInitializationDataFallback(format, Localization("en", "US")) + suspend fun bootstrap( + holder: SabrSessionHolder, + format: YoutubeSabrFormat, + cache: CacheService? = null, + ): ByteArray? { + val initialized = runCatchingNonCancellation { + holder.pumpMutex.withLock { + holder.withPlayerContext { bootstrapInitialization(Localization("en", "US")) } + listOf(holder.audioFormat, holder.videoFormat).mapNotNull { candidate -> + holder.session.getCachedSegment(SabrSegmentRequest.initialization(candidate)) + ?.data + ?.let { candidate to it } + } + } }.getOrNull() ?: return null - memoryCache[key] = bytes - cache?.setBytes(key, bytes, CACHE_TTL_SECONDS) - return bytes - } - - private fun fetchRemote(format: YoutubeSabrFormat): ByteArray? { - val url = format.initializationUrl?.takeUnless { it.isBlank() } ?: return null - val start = format.initRangeStart - val end = format.initRangeEnd - if (start < 0L || end < start) return null - val headers = mapOf("Range" to listOf("bytes=$start-$end")) - return runCatching { NewPipe.getDownloader().get(url, headers) } - .getOrNull() - ?.takeIf { it.responseCode() == 200 || it.responseCode() == 206 } - ?.rawResponseBody() - ?.takeIf { it.isNotEmpty() } + initialized.forEach { (candidate, bytes) -> remember(holder.key.videoId, candidate, bytes, cache) } + return initialized.firstOrNull { (candidate) -> candidate.matches(format) }?.second } private suspend fun CacheService.getBytes(key: String): ByteArray? = runCatching { @@ -82,16 +84,22 @@ internal object SabrInitializationData { runCatching { set(key, Base64.getEncoder().encodeToString(bytes), ttlSeconds) } } - private fun cacheKey(format: YoutubeSabrFormat): String? { - val url = format.initializationUrl?.takeUnless { it.isBlank() } ?: return null - val start = format.initRangeStart - val end = format.initRangeEnd - if (start < 0L || end < start) return null - val raw = listOf(format.itag, format.lastModified, format.xtags.orEmpty(), start, end) + private fun cacheKey(videoId: String, format: YoutubeSabrFormat): String { + val raw = listOf( + videoId, + format.itag, + format.lastModified, + format.xtags.orEmpty(), + format.mimeType.orEmpty(), + format.audioTrackId.orEmpty(), + ) .joinToString("|") - return "sabr:init:v2:${sha256(raw)}" + return "sabr:init:v3:${sha256(raw)}" } + private fun YoutubeSabrFormat.matches(other: YoutubeSabrFormat): Boolean = + itag == other.itag && audioTrackId == other.audioTrackId && xtags == other.xtags + private fun sha256(value: String): String = MessageDigest.getInstance("SHA-256") .digest(value.toByteArray()) .joinToString("") { byte -> "%02x".format(byte) } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrInitializationSegmentFetcher.kt b/src/main/kotlin/dev/typetype/server/services/SabrInitializationSegmentFetcher.kt index 49a4a357..337677b4 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrInitializationSegmentFetcher.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrInitializationSegmentFetcher.kt @@ -21,7 +21,7 @@ internal suspend fun fetchSabrInitializationSegment( if (holder.session.isBeyondEnd(request)) return@withLock null val result = runCatchingNonCancellation { holder.session.prepareForInitialization(request.format) - holder.session.fetchSegment(request, localization) + holder.withPlayerContext { fetchSegment(request, localization) } } result.onFailure { error -> logger.warn( diff --git a/src/main/kotlin/dev/typetype/server/services/SabrLiveMediaNormalizer.kt b/src/main/kotlin/dev/typetype/server/services/SabrLiveMediaNormalizer.kt new file mode 100644 index 00000000..580dcfba --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrLiveMediaNormalizer.kt @@ -0,0 +1,110 @@ +package dev.typetype.server.services + +internal data class SabrLiveMediaParts( + val initialization: ByteArray, + val media: ByteArray, +) + +internal object SabrLiveMediaNormalizer { + fun split(mimeType: String, data: ByteArray): SabrLiveMediaParts? = when { + mimeType.substringBefore(';').trim().lowercase().endsWith("/mp4") -> splitMp4(data) + mimeType.substringBefore(';').trim().lowercase().endsWith("/webm") -> splitWebM(data) + else -> null + } + + private fun splitMp4(data: ByteArray): SabrLiveMediaParts? { + var offset = 0 + var initializationEnd = -1 + while (offset + MP4_HEADER_SIZE <= data.size) { + val size32 = data.readUnsignedInt(offset) + val type = String(data, offset + 4, 4, Charsets.US_ASCII) + val headerSize = if (size32 == 1L) MP4_EXTENDED_HEADER_SIZE else MP4_HEADER_SIZE + if (offset + headerSize > data.size) return null + val size = when (size32) { + 0L -> data.size.toLong() - offset + 1L -> data.readUnsignedLong(offset + MP4_HEADER_SIZE) ?: return null + else -> size32 + } + if (size < headerSize || size > data.size.toLong() - offset) return null + val end = offset + size.toInt() + if (type == "moov") initializationEnd = end + if (type == "moof" && initializationEnd > 0) { + return SabrLiveMediaParts( + initialization = data.copyOfRange(0, initializationEnd), + media = data.copyOfRange(initializationEnd, data.size), + ) + } + offset = end + } + return null + } + + private fun splitWebM(data: ByteArray): SabrLiveMediaParts? { + val ebml = data.readEbmlElement(0) ?: return null + if (ebml.id != EBML_HEADER_ID || ebml.size == null) return null + val segmentOffset = ebml.payloadOffset + ebml.size.toInt() + val segment = data.readEbmlElement(segmentOffset) ?: return null + if (segment.id != WEBM_SEGMENT_ID) return null + var offset = segment.payloadOffset + while (offset < data.size) { + val element = data.readEbmlElement(offset) ?: return null + if (element.id == WEBM_CLUSTER_ID) { + if (offset <= segment.payloadOffset) return null + return SabrLiveMediaParts( + initialization = data.copyOfRange(0, offset), + media = data.copyOfRange(offset, data.size), + ) + } + val size = element.size ?: return null + val next = element.payloadOffset.toLong() + size + if (next <= offset || next > data.size) return null + offset = next.toInt() + } + return null + } + + private fun ByteArray.readEbmlElement(offset: Int): EbmlElement? { + val idLength = vintLength(offset, 4) ?: return null + val id = readRawValue(offset, idLength) ?: return null + val sizeOffset = offset + idLength + val sizeLength = vintLength(sizeOffset, 8) ?: return null + val rawSize = readRawValue(sizeOffset, sizeLength) ?: return null + val marker = 1L shl (7 * sizeLength) + val size = (rawSize and (marker - 1L)).takeUnless { it == marker - 1L } + return EbmlElement(id, size, sizeOffset + sizeLength) + } + + private fun ByteArray.vintLength(offset: Int, maximum: Int): Int? { + if (offset !in indices) return null + val first = this[offset].toInt() and 0xff + if (first == 0) return null + val length = Integer.numberOfLeadingZeros(first) - 23 + return length.takeIf { it in 1..maximum && offset + it <= size } + } + + private fun ByteArray.readRawValue(offset: Int, length: Int): Long? { + if (length !in 1..8 || offset < 0 || offset + length > size) return null + var value = 0L + repeat(length) { index -> value = (value shl 8) or (this[offset + index].toLong() and 0xff) } + return value + } + + private fun ByteArray.readUnsignedInt(offset: Int): Long = + ((this[offset].toLong() and 0xff) shl 24) or + ((this[offset + 1].toLong() and 0xff) shl 16) or + ((this[offset + 2].toLong() and 0xff) shl 8) or + (this[offset + 3].toLong() and 0xff) + + private fun ByteArray.readUnsignedLong(offset: Int): Long? { + val value = readRawValue(offset, 8) ?: return null + return value.takeIf { it >= 0L } + } + + private data class EbmlElement(val id: Long, val size: Long?, val payloadOffset: Int) + + private const val MP4_HEADER_SIZE = 8 + private const val MP4_EXTENDED_HEADER_SIZE = 16 + private const val EBML_HEADER_ID = 0x1A45DFA3L + private const val WEBM_SEGMENT_ID = 0x18538067L + private const val WEBM_CLUSTER_ID = 0x1F43B675L +} diff --git a/src/main/kotlin/dev/typetype/server/services/SabrLivePlayback.kt b/src/main/kotlin/dev/typetype/server/services/SabrLivePlayback.kt new file mode 100644 index 00000000..ffbf4b3a --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrLivePlayback.kt @@ -0,0 +1,121 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat + +internal data class SabrLivePlaybackSnapshot( + val active: Boolean, + val postLiveDvr: Boolean, + val headSequence: Long, + val headTimeMs: Long, + val seekableStartMs: Long, + val seekableEndMs: Long, + val atLiveEdge: Boolean, + val targetLatencyMs: Long, +) + +internal fun SabrSessionHolder.livePlaybackSnapshot(): SabrLivePlaybackSnapshot? { + val state = session.streamState + val postLiveDvr = runCatching { state.isPostLiveDvr }.getOrDefault(false) + val sessionLive = runCatching { session.isLive }.getOrDefault(false) + val stateLive = runCatching { state.isLive }.getOrDefault(false) + val detectedLive = expectsLive() || sessionLive || stateLive || postLiveDvr + if (!detectedLive) return null + + val observedEndMs = maxOf( + state.observedEndMs(audioFormat), + state.observedEndMs(videoFormat), + state.getBufferedEndMs(audioFormat), + state.getBufferedEndMs(videoFormat), + state.getMinBufferedEndMs(), + 0L, + ) + val reportedHeadTimeMs = runCatching { state.liveHeadTimeMs }.getOrDefault(0L) + val headTimeMs = reportedHeadTimeMs.takeIf { it > 0L } ?: observedEndMs + val active = !postLiveDvr && (expectsLive() || sessionLive || stateLive) + val seekableEndMs = if (active) headTimeMs else observedEndMs + val seekableStartMs = if (active) { + (seekableEndMs - LIVE_DVR_WINDOW_MS).coerceAtLeast(0L) + } else { + 0L + } + val extractorAtLiveEdge = runCatching { session.isAtLiveEdge }.getOrDefault(false) + val readerAtLiveEdge = headTimeMs - maxOf(playerTimeMs(), readerHeadMs()) <= LIVE_EDGE_TOLERANCE_MS + return SabrLivePlaybackSnapshot( + active = active, + postLiveDvr = postLiveDvr, + headSequence = maxOf( + runCatching { session.liveHeadSequenceNumber }.getOrDefault(0L), + runCatching { state.liveHeadSequenceNumber }.getOrDefault(0L), + 0L, + ), + headTimeMs = headTimeMs, + seekableStartMs = seekableStartMs, + seekableEndMs = seekableEndMs, + atLiveEdge = active && (extractorAtLiveEdge || readerAtLiveEdge), + targetLatencyMs = LIVE_TARGET_LATENCY_MS, + ) +} + +internal fun SabrSessionHolder.resolvePlaybackStartMs(requestedStartMs: Long): Long { + val requested = requestedStartMs.coerceAtLeast(0L) + val live = livePlaybackSnapshot() ?: return requested + if (!live.active) return requested.coerceAtMost(live.seekableEndMs.takeIf { it > 0L } ?: requested) + if (requested > 0L) return requested.coerceIn(live.seekableStartMs, live.seekableEndMs) + val targetStartMs = (live.seekableEndMs - live.targetLatencyMs).coerceAtLeast(live.seekableStartMs) + return maxOf(targetStartMs, availableLiveMediaStartMs() ?: targetStartMs) +} + +private fun SabrSessionHolder.availableLiveMediaStartMs(): Long? { + val audioStartMs = observedMediaSegment(audioFormat)?.header?.startMs?.takeIf { it >= 0L } + ?: return null + if (!isVideoActive()) return audioStartMs + val videoStartMs = observedMediaSegment(videoFormat)?.header?.startMs?.takeIf { it >= 0L } + ?: return null + return maxOf(audioStartMs, videoStartMs) +} + +internal fun SabrSessionHolder.isFutureLiveRequest(request: SabrSegmentRequest): Boolean { + if (request.isInitializationSegment) return false + livePlaybackSnapshot()?.takeIf { it.active } ?: return false + if (session.getCachedSegment(request) == null && session.getReadableSegment(request) != null) return true + val state = session.streamState + observedMediaSegment(request.format)?.let { observed -> + val distanceFromComplete = request.sequenceNumber.toLong() - observed.header.sequenceNumber.toLong() + return distanceFromComplete in 1L..LIVE_FUTURE_SEGMENT_TOLERANCE.toLong() + } + val trackHeadSequence = runCatching { state.getMaxSegment(request.format) }.getOrDefault(0) + if (trackHeadSequence <= 0) return false + if (request.sequenceNumber > trackHeadSequence) { + return request.sequenceNumber <= trackHeadSequence + LIVE_FUTURE_SEGMENT_TOLERANCE + } + if (request.sequenceNumber != trackHeadSequence) return false + val requestStartMs = playbackSegmentStartMs(request.format, request.sequenceNumber) + val completeEndMs = runCatching { state.getBufferedEndMs(request.format) }.getOrDefault(0L) + return requestStartMs > 0L && requestStartMs >= completeEndMs +} + +internal fun SabrSessionHolder.isHistoricalLiveRequest(request: SabrSegmentRequest): Boolean { + if (request.isInitializationSegment || livePlaybackSnapshot()?.active != true) return false + val observed = observedMediaSegment(request.format) ?: return false + return request.sequenceNumber < observed.header.sequenceNumber +} + +internal fun SabrSessionHolder.liveRetryAfterMs(blockedRequests: List = emptyList()): Long = + if (livePlaybackSnapshot()?.active == true && + (blockedRequests.isEmpty() || blockedRequests.all(::isFutureLiveRequest)) + ) LIVE_EDGE_POLL_MS else DEFAULT_PLAYBACK_RETRY_MS + +private fun org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState.observedEndMs( + format: YoutubeSabrFormat, +): Long { + val sequence = runCatching { getMaxSegment(format) }.getOrDefault(0) + return if (sequence > 0) runCatching { getSegmentEndMs(format, sequence) }.getOrDefault(0L).coerceAtLeast(0L) else 0L +} + +internal const val LIVE_EDGE_POLL_MS = 2_000L +internal const val DEFAULT_PLAYBACK_RETRY_MS = 500L +private const val LIVE_TARGET_LATENCY_MS = 20_000L +private const val LIVE_EDGE_TOLERANCE_MS = 15_000L +private const val LIVE_DVR_WINDOW_MS = 12L * 60L * 60L * 1_000L +private const val LIVE_FUTURE_SEGMENT_TOLERANCE = 2 diff --git a/src/main/kotlin/dev/typetype/server/services/SabrLivePlaybackDiscontinuity.kt b/src/main/kotlin/dev/typetype/server/services/SabrLivePlaybackDiscontinuity.kt new file mode 100644 index 00000000..fd443749 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrLivePlaybackDiscontinuity.kt @@ -0,0 +1,17 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat + +internal fun SabrSessionHolder.failLivePlaybackDiscontinuity( + format: YoutubeSabrFormat, + targetMs: Long, + segment: CachedSabrSegment, + hasBufferedMedia: Boolean, +): Boolean { + if (!hasBufferedMedia || livePlaybackSnapshot()?.active != true) return false + if (segment.startMs <= targetMs + MAX_CONTIGUOUS_LIVE_GAP_MS) return false + failTerminal(sabrRecoverableFailureMessage("live ${format.itag} media discontinuity")) + return true +} + +private const val MAX_CONTIGUOUS_LIVE_GAP_MS = 500L diff --git a/src/main/kotlin/dev/typetype/server/services/SabrLiveWarmupRequest.kt b/src/main/kotlin/dev/typetype/server/services/SabrLiveWarmupRequest.kt new file mode 100644 index 00000000..b9780886 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrLiveWarmupRequest.kt @@ -0,0 +1,63 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrBufferedRange +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat + +internal data class SabrLiveWarmupTarget(val sequence: Int, val timeMs: Long) + +internal fun SabrSessionHolder.liveWarmupTarget(): SabrLiveWarmupTarget? { + val live = livePlaybackSnapshot() + ?.takeIf { it.active && it.headSequence > 1L && it.headTimeMs > 0L } + ?: return null + val segmentDurationMs = (live.headTimeMs / live.headSequence) + .coerceIn(MIN_LIVE_SEGMENT_DURATION_MS, MAX_LIVE_SEGMENT_DURATION_MS) + val segmentsBehind = Math.floorDiv( + live.targetLatencyMs + segmentDurationMs - 1L, + segmentDurationMs, + ) + val targetSequence = (live.headSequence - segmentsBehind) + .coerceIn(1L, Int.MAX_VALUE.toLong()) + .toInt() + val targetTimeMs = (live.headTimeMs - live.targetLatencyMs).coerceAtLeast(0L) + return SabrLiveWarmupTarget(targetSequence, targetTimeMs) +} + +internal inline fun withLiveWarmupRequestShape( + holder: SabrSessionHolder, + target: SabrLiveWarmupTarget?, + block: () -> T, +): T { + target ?: return block() + val state = holder.session.streamState + val ranges = buildList { + if (holder.isAudioActive()) add(holder.audioFormat.liveWarmupRange(target.sequence, target.timeMs)) + if (holder.isVideoActive()) add(holder.videoFormat.liveWarmupRange(target.sequence, target.timeMs)) + } + state.setPlayerTimeMs(target.timeMs) + state.setActiveTrackTypes(holder.isVideoActive(), holder.isAudioActive()) + state.setBufferedRangesOverride(ranges) + return try { + block() + } finally { + state.setBufferedRangesOverride(null) + state.setActiveTrackTypes(holder.isVideoActive(), holder.isAudioActive()) + } +} + +private fun YoutubeSabrFormat.liveWarmupRange(targetSequence: Int, targetTimeMs: Long): SabrBufferedRange { + val bufferedSequence = (targetSequence - 1).coerceAtLeast(0) + return SabrBufferedRange( + itag, + lastModified, + xtags, + 0L, + targetTimeMs.coerceAtLeast(1L), + if (bufferedSequence > 0) 1 else 0, + bufferedSequence, + TIMESCALE, + ) +} + +private const val MIN_LIVE_SEGMENT_DURATION_MS = 100L +private const val MAX_LIVE_SEGMENT_DURATION_MS = 30_000L +private const val TIMESCALE = 1_000 diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPendingSeek.kt b/src/main/kotlin/dev/typetype/server/services/SabrPendingSeek.kt new file mode 100644 index 00000000..3cc1c668 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrPendingSeek.kt @@ -0,0 +1,42 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest + +internal fun SabrSessionHolder.consumeMatchingSeek(request: SabrSegmentRequest): Boolean { + pendingRefetchRequest()?.takeIf { it.matches(request) }?.let { + consumeRefetch() + setPlaybackState(SabrPlaybackState.REPOSITIONING) + prepareForExplicitRewind(request) + return true + } + pendingForwardSeekRequest()?.takeIf { it.matches(request) }?.let { + consumeForwardSeek() + setPlaybackState(SabrPlaybackState.REPOSITIONING) + prepareForExplicitForwardJump(request) + return true + } + return false +} + +internal fun SabrSessionHolder.prepareForExplicitRewind(request: SabrSegmentRequest): Unit = + session.prepareForRewind(request, explicitSeekPositionMs()) + +internal fun SabrSessionHolder.prepareForExplicitForwardJump(request: SabrSegmentRequest): Unit = + session.prepareForForwardJump(request, explicitSeekPositionMs()) + +internal fun SabrSessionHolder.prepareForHistoricalLiveRewind(request: SabrSegmentRequest): Unit { + val seekPositionMs = requestedSeekTimeMs()?.takeIf { request.contains(this, it) } + if (seekPositionMs == null) session.prepareForRewind(request) else session.prepareForRewind(request, seekPositionMs) +} + +private fun SabrSessionHolder.explicitSeekPositionMs(): Long = requestedSeekTimeMs() ?: playerTimeMs() + +private fun SabrSegmentRequest.contains(holder: SabrSessionHolder, positionMs: Long): Boolean { + val startMs = holder.playbackSegmentStartMs(format, sequenceNumber) + val endMs = holder.playbackSegmentEndMs(format, sequenceNumber) + return positionMs >= startMs && (endMs <= startMs || positionMs < endMs) +} + +private fun SabrSegmentRequest.matches(other: SabrSegmentRequest): Boolean = + format.itag == other.format.itag && sequenceNumber == other.sequenceNumber && + isInitializationSegment == other.isInitializationSegment diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPlaybackCachedSegmentLocator.kt b/src/main/kotlin/dev/typetype/server/services/SabrPlaybackCachedSegmentLocator.kt new file mode 100644 index 00000000..0ab8cb26 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrPlaybackCachedSegmentLocator.kt @@ -0,0 +1,63 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat + +internal suspend fun SabrSessionStore.findCachedPlaybackMediaAt( + holder: SabrSessionHolder, + format: YoutubeSabrFormat, + targetMs: Long, + predictedSequence: Int, + allowFollowing: Boolean = false, +): CachedSabrSegment? { + for (distance in 0..MAX_SEQUENCE_DISTANCE) { + cachedMedia(holder, format, predictedSequence + distance)?.let { + if (it.coversPlaybackTime(holder, format, targetMs)) return it + } + if (distance > 0) { + cachedMedia(holder, format, predictedSequence - distance)?.let { + if (it.coversPlaybackTime(holder, format, targetMs)) return it + } + } + } + if (!allowFollowing) return null + for (distance in 0..MAX_SEQUENCE_DISTANCE) { + cachedMedia(holder, format, predictedSequence + distance)?.let { + if (it.startsWithinNextSegment(holder, format, targetMs)) return it + } + } + return null +} + +private suspend fun SabrSessionStore.cachedMedia( + holder: SabrSessionHolder, + format: YoutubeSabrFormat, + sequence: Int, +): CachedSabrSegment? { + if (sequence < 1) return null + return cachedSegment(holder, SabrSegmentRequest.media(format, sequence)) +} + +internal fun CachedSabrSegment.coversPlaybackTime( + holder: SabrSessionHolder, + format: YoutubeSabrFormat, + targetMs: Long, +): Boolean { + val effectiveStartMs = startMs.takeIf { it >= 0L } ?: holder.playbackSegmentStartMs(format, sequence) + val effectiveDurationMs = durationMs.takeIf { it > 0L } ?: holder.playbackSegmentDurationMs(format, sequence) + return targetMs >= effectiveStartMs - TIMING_TOLERANCE_MS && targetMs < effectiveStartMs + effectiveDurationMs +} + +private fun CachedSabrSegment.startsWithinNextSegment( + holder: SabrSessionHolder, + format: YoutubeSabrFormat, + targetMs: Long, +): Boolean { + val effectiveStartMs = startMs.takeIf { it >= 0L } ?: holder.playbackSegmentStartMs(format, sequence) + val effectiveDurationMs = durationMs.takeIf { it > 0L } ?: holder.playbackSegmentDurationMs(format, sequence) + val leadMs = effectiveStartMs - targetMs + return leadMs in -TIMING_TOLERANCE_MS..effectiveDurationMs +} + +private const val MAX_SEQUENCE_DISTANCE = 24 +private const val TIMING_TOLERANCE_MS = 2L diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPlaybackSegmentSelection.kt b/src/main/kotlin/dev/typetype/server/services/SabrPlaybackSegmentSelection.kt index ae2739c5..43b344dc 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrPlaybackSegmentSelection.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrPlaybackSegmentSelection.kt @@ -3,6 +3,71 @@ package dev.typetype.server.services import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat internal fun SabrSessionHolder.playbackStartSequence(format: YoutubeSabrFormat, playerTimeMs: Long): Int { + liveSequenceAt(format, playerTimeMs)?.let { return it } return session.streamState.getSegmentNumberAtOrAfterTimeMs(format, playerTimeMs.coerceAtLeast(0L)) .coerceAtLeast(1) } + +internal fun SabrSessionHolder.playbackContinuationSequence( + format: YoutubeSabrFormat, + playerTimeMs: Long, + continueAfterLastServed: Boolean, +): Int = lastServedSequence(format) + ?.takeIf { continueAfterLastServed && it < Int.MAX_VALUE } + ?.plus(1) + ?: playbackStartSequence(format, playerTimeMs) + +internal fun SabrSessionHolder.playbackSegmentStartMs(format: YoutubeSabrFormat, sequence: Int): Long { + val observed = observedMediaSegment(format) + val observedStartMs = observed?.header?.startMs?.takeIf { it >= 0L } + if (observed != null && observedStartMs != null) { + val durationMs = playbackSegmentDurationMs(format, sequence) + val sequenceDelta = sequence.toLong() - observed.header.sequenceNumber + return (observedStartMs + sequenceDelta * durationMs).coerceAtLeast(0L) + } + return session.streamState.getSegmentStartMs(format, sequence).coerceAtLeast(0L) +} + +internal fun SabrSessionHolder.playbackSegmentDurationMs(format: YoutubeSabrFormat, sequence: Int): Long { + val observed = observedMediaSegment(format) + val observedStartMs = observed?.header?.startMs?.takeIf { it >= 0L } + val observedDurationMs = observed?.header?.durationMs?.takeIf { it > 0L } + if (observedDurationMs != null) return observedDurationMs + if (observed != null && observedStartMs != null) { + livePlaybackSnapshot() + ?.takeIf { it.active } + ?.segmentDurationFrom(observed.header.sequenceNumber, observedStartMs) + ?.takeIf { it in MIN_LIVE_SEGMENT_DURATION_MS..MAX_LIVE_SEGMENT_DURATION_MS } + ?.let { return it } + } + val startMs = session.streamState.getSegmentStartMs(format, sequence).coerceAtLeast(0L) + return (session.streamState.getSegmentEndMs(format, sequence) - startMs).coerceAtLeast(1L) +} + +internal fun SabrSessionHolder.playbackSegmentEndMs(format: YoutubeSabrFormat, sequence: Int): Long = + playbackSegmentStartMs(format, sequence) + playbackSegmentDurationMs(format, sequence) + +private fun SabrSessionHolder.liveSequenceAt(format: YoutubeSabrFormat, playerTimeMs: Long): Int? { + livePlaybackSnapshot()?.takeIf { it.active } ?: return null + val observed = observedMediaSegment(format) ?: return null + val observedStartMs = observed.header.startMs.takeIf { it >= 0L } + ?: return null + val segmentDurationMs = playbackSegmentDurationMs(format, observed.header.sequenceNumber) + if (segmentDurationMs !in MIN_LIVE_SEGMENT_DURATION_MS..MAX_LIVE_SEGMENT_DURATION_MS) return null + val offset = Math.floorDiv(playerTimeMs.coerceAtLeast(0L) - observedStartMs, segmentDurationMs) + val sequence = (observed.header.sequenceNumber.toLong() + offset) + .coerceIn(1L, Int.MAX_VALUE.toLong()) + .toInt() + val availableHead = runCatching { session.streamState.getMaxSegment(format) }.getOrDefault(0) + return if (availableHead > 0) sequence.coerceAtMost(availableHead) else sequence +} + +private fun SabrLivePlaybackSnapshot.segmentDurationFrom(observedSequence: Int, observedStartMs: Long): Long? { + val sequenceDelta = headSequence - observedSequence + val timeDeltaMs = headTimeMs - observedStartMs + if (sequenceDelta <= 0L || timeDeltaMs <= 0L) return null + return Math.floorDiv(timeDeltaMs + sequenceDelta / 2L, sequenceDelta).takeIf { it > 0L } +} + +private const val MIN_LIVE_SEGMENT_DURATION_MS = 100L +private const val MAX_LIVE_SEGMENT_DURATION_MS = 30_000L diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPlaybackSessionService.kt b/src/main/kotlin/dev/typetype/server/services/SabrPlaybackSessionService.kt index 09fda027..4d5ea812 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrPlaybackSessionService.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrPlaybackSessionService.kt @@ -15,6 +15,7 @@ internal class SabrPlaybackSessionService(private val sessionStore: SabrSessionS video: YoutubeSabrFormat, startTimeMs: Long, audioOnly: Boolean = false, + isLive: Boolean = false, ): SabrPlaybackPreparation { val holder = sessionStore.getOrCreate( videoId = videoId, @@ -28,8 +29,15 @@ internal class SabrPlaybackSessionService(private val sessionStore: SabrSessionS purpose = SabrSessionPurpose.PLAYBACK, audioOnly = audioOnly, ) - SabrPlaybackInitializationPreloader.preload(sessionStore, holder, INITIALIZATION_PRELOAD_TIMEOUT_MS) - return prepareHolder(holder, startTimeMs, audioOnly) + if (isLive || prepared.isLive || prepared.isLiveContent) holder.markExpectedLive() + if (holder.expectsLive()) { + holder.setActiveTracks(videoActive = !audioOnly, audioActive = true) + holder.session.streamState.setSelectVideoFormatBeforeAudio(!audioOnly) + sessionStore.ensureWarmed(holder, LIVE_INITIAL_PUMPS) + } else { + SabrPlaybackInitializationPreloader.preload(sessionStore, holder, INITIALIZATION_PRELOAD_TIMEOUT_MS) + } + return prepareHolder(holder, holder.resolvePlaybackStartMs(startTimeMs), audioOnly) } suspend fun seek( @@ -49,6 +57,7 @@ internal class SabrPlaybackSessionService(private val sessionStore: SabrSessionS video = video, startTimeMs = playerTimeMs, audioOnly = audioOnly, + isLive = source.expectsLive() || prepared.isLive || prepared.isLiveContent, ) } @@ -123,6 +132,15 @@ internal class SabrPlaybackSessionService(private val sessionStore: SabrSessionS while (segment == null && holder.terminalFailure() == null && holder.networkFailure() == null) { delay(SEGMENT_WAIT_MS) segment = sessionStore.cachedSegment(holder, request) + if (segment == null && holder.livePlaybackSnapshot()?.active == true) { + segment = sessionStore.findCachedPlaybackMediaAt( + holder = holder, + format = request.format, + targetMs = holder.playbackSegmentStartMs(request.format, request.sequenceNumber), + predictedSequence = request.sequenceNumber, + allowFollowing = true, + ) + } } segment } @@ -143,6 +161,7 @@ internal class SabrPlaybackSessionService(private val sessionStore: SabrSessionS holder.session.streamState.setSelectVideoFormatBeforeAudio(startTimeMs > SEEK_FORMAT_ORDER_MS) if (startTimeMs > SEEK_FORMAT_ORDER_MS) holder.anchorReaderPositions(startTimeMs) if (startTimeMs > 0L) { + holder.setRequestedSeekTimeMs(startTimeMs) holder.requestReposition(startTimeMs, holder.activeGeneration()) } sessionStore.startPump(holder) @@ -170,16 +189,15 @@ internal class SabrPlaybackSessionService(private val sessionStore: SabrSessionS } val targets = listOfNotNull(request, companion) targets.forEach { target -> - val targetStartMs = session.streamState - .getSegmentStartMs(target.format, target.sequenceNumber) - .coerceAtLeast(0L) + val targetStartMs = playbackSegmentStartMs(target.format, target.sequenceNumber) setReaderPosition(target.format, targetStartMs, generation) } val missing = targets.filter { session.getCachedSegment(it) == null } if (missing.isEmpty()) return - val anchor = missing.first() - val startMs = session.streamState.getSegmentStartMs(anchor.format, anchor.sequenceNumber).coerceAtLeast(0L) missing.forEach { requestSegmentDemand(it, generation) } + if (livePlaybackSnapshot()?.active == true) return + val anchor = missing.first() + val startMs = playbackSegmentStartMs(anchor.format, anchor.sequenceNumber) if (startMs < session.streamState.getMinBufferedEndMs()) { requestRefetch(anchor) } else { @@ -194,5 +212,6 @@ internal class SabrPlaybackSessionService(private val sessionStore: SabrSessionS const val INITIALIZATION_PRELOAD_TIMEOUT_MS = 6_000L const val SEEK_FORMAT_ORDER_MS = 1_000L const val SEGMENT_WAIT_MS = 250L + const val LIVE_INITIAL_PUMPS = 8 } } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPlaybackState.kt b/src/main/kotlin/dev/typetype/server/services/SabrPlaybackState.kt index 281c4cfe..3739acba 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrPlaybackState.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrPlaybackState.kt @@ -5,6 +5,7 @@ internal enum class SabrPlaybackState { PREPARING, REQUESTING, REPOSITIONING, + WAITING_FOR_LIVE, THROTTLED, NETWORK_FAILED, TERMINAL, diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPlayerContextRecovery.kt b/src/main/kotlin/dev/typetype/server/services/SabrPlayerContextRecovery.kt new file mode 100644 index 00000000..76ddc7c6 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrPlayerContextRecovery.kt @@ -0,0 +1,60 @@ +package dev.typetype.server.services + +import kotlinx.coroutines.CancellationException +import org.schabi.newpipe.extractor.exceptions.AntiBotException +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrProtocolException +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrClientProfile +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo + +internal class SabrPlayerContextRecovery( + private val videoId: String, + initialToken: SabrTokenBundle, + private val tokenClient: TypetypeTokenSabrTokenClient, + private val probe: SabrPlayerInfoProbe, +) { + private var activeToken = initialToken + private var refreshAttempted = false + + fun fetch(profile: YoutubeSabrClientProfile): SabrPlayerProbeResult { + val initialToken = activeToken + val initialFailure = try { + return SabrPlayerProbeResult.Success(probe.fetch(videoId, profile, initialToken), initialToken, false) + } catch (error: CancellationException) { + throw error + } catch (error: Exception) { + error + } + if (refreshAttempted || !initialFailure.isPlayerAdmissionFailure(profile)) { + return SabrPlayerProbeResult.Failure(initialFailure, refreshAttempted) + } + refreshAttempted = true + val refreshedToken = tokenClient.fetch(videoId, forceRefresh = true) + ?: return SabrPlayerProbeResult.Failure(initialFailure, true) + activeToken = refreshedToken + return try { + SabrPlayerProbeResult.Success(probe.fetch(videoId, profile, refreshedToken), refreshedToken, true) + } catch (error: CancellationException) { + throw error + } catch (refreshedFailure: Exception) { + refreshedFailure.addSuppressed(initialFailure) + SabrPlayerProbeResult.Failure(refreshedFailure, true) + } + } + + private fun Exception.isPlayerAdmissionFailure(profile: YoutubeSabrClientProfile): Boolean = + this is AntiBotException || + this is SabrProtocolException && message == "Player response has no streamingData for $profile" +} + +internal sealed interface SabrPlayerProbeResult { + data class Success( + val info: YoutubeSabrInfo, + val token: SabrTokenBundle, + val contextRefreshed: Boolean, + ) : SabrPlayerProbeResult + + data class Failure( + val error: Exception, + val contextRefreshAttempted: Boolean, + ) : SabrPlayerProbeResult +} diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPlayerInfoProbe.kt b/src/main/kotlin/dev/typetype/server/services/SabrPlayerInfoProbe.kt new file mode 100644 index 00000000..44ecba53 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrPlayerInfoProbe.kt @@ -0,0 +1,28 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.localization.ContentCountry +import org.schabi.newpipe.extractor.localization.Localization +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrClientProfile +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrProbe + +internal fun interface SabrPlayerInfoProbe { + fun fetch( + videoId: String, + profile: YoutubeSabrClientProfile, + token: SabrTokenBundle, + ): YoutubeSabrInfo +} + +internal object PipePipeSabrPlayerInfoProbe : SabrPlayerInfoProbe { + private val localization = Localization("en", "US") + private val contentCountry = ContentCountry("US") + + override fun fetch( + videoId: String, + profile: YoutubeSabrClientProfile, + token: SabrTokenBundle, + ): YoutubeSabrInfo = TypetypeYoutubeSessionPoTokenProvider.withToken(token) { + YoutubeSabrProbe.fetchSabrInfo(videoId, profile, localization, contentCountry) + } +} diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPreparedInfo.kt b/src/main/kotlin/dev/typetype/server/services/SabrPreparedInfo.kt index ea52839b..98de8cd2 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrPreparedInfo.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrPreparedInfo.kt @@ -5,6 +5,8 @@ import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo internal class SabrPreparedInfo( val info: YoutubeSabrInfo, val initialToken: SabrTokenBundle?, + val isLive: Boolean = false, + val isLiveContent: Boolean = false, ) internal fun SabrPreparedInfo.hasAudioAndVideoFormats(): Boolean = diff --git a/src/main/kotlin/dev/typetype/server/services/SabrPumpPolicy.kt b/src/main/kotlin/dev/typetype/server/services/SabrPumpPolicy.kt index 91619124..cf2c4dca 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrPumpPolicy.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrPumpPolicy.kt @@ -13,12 +13,20 @@ internal object SabrPumpPolicy { const val MIN_SERVER_READAHEAD_CUSHION_MS = 3_000L const val SERVER_AHEAD_MARGIN_MS = 16_000L const val DEMAND_TARGET_DEADLINE_MS = 15_000L + const val COMPLETED_DEMAND_IDLE_MS = 1_000L const val MAX_DEMAND_RESPONSES_WITHOUT_TARGET = 3 const val MAX_AHEAD_BYTES = 24L * 1024L * 1024L const val BACK_BUFFER_MS = 12_000L const val MIN_BACK_BUFFER_MS = 2_000L const val BACK_BUFFER_BYTES = 4L * 1024L * 1024L + fun demandDelayMs(intervalMs: Long, backoffMs: Long, activeLive: Boolean, futureDemand: Boolean?): Long = + maxOf( + intervalMs, + backoffMs, + LIVE_EDGE_POLL_MS.takeIf { futureDemand ?: activeLive } ?: 0L, + ) + fun backBufferMs(holder: SabrSessionHolder): Long { val cachedBytes = runCatchingNonCancellation { holder.session.cachedBytes }.getOrDefault(0L) if (cachedBytes > MAX_AHEAD_BYTES) return MIN_BACK_BUFFER_MS diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSegmentCache.kt b/src/main/kotlin/dev/typetype/server/services/SabrSegmentCache.kt index 9d7ecb1d..e4de84b8 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrSegmentCache.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrSegmentCache.kt @@ -10,7 +10,11 @@ internal class SabrSegmentCache { fun put(holder: SabrSessionHolder, segment: SabrMediaSegment): Unit { val format = holder.formatForItag(segment.header.itag) ?: return - val cached = segment.toCachedSabrSegment(format.mimeType.orEmpty()) + val mimeType = format.mimeType.orEmpty() + val mediaParts = segment.takeIf { holder.expectsLive() && !it.header.isInitSegment } + ?.let { SabrLiveMediaNormalizer.split(mimeType, it.data) } + if (mediaParts != null) holder.rememberLiveInitialization(format.itag, mediaParts.initialization) + val cached = segment.toCachedSabrSegment(mimeType, mediaParts?.media ?: segment.data) val key = key(holder, format, cached.init, cached.sequence) holder.putCachedSegment(key, cached) } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSegmentDemandResolution.kt b/src/main/kotlin/dev/typetype/server/services/SabrSegmentDemandResolution.kt index 9a83f2e3..4c871728 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrSegmentDemandResolution.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrSegmentDemandResolution.kt @@ -3,14 +3,16 @@ package dev.typetype.server.services import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest internal fun SabrSessionHolder.resolveSegmentDemand(request: SabrSegmentRequest, identity: String): Boolean { - val requestedCached = session.getCachedSegment(request) != null - val rebased = if (requestedCached) null else session.findCachedMediaAt( + val requested = session.getCachedSegment(request) + val rebased = if (requested != null) null else session.findCachedMediaAt( format = request.format, - targetMs = session.streamState.getBufferedEndMs(request.format), + targetMs = playbackSegmentStartMs(request.format, request.sequenceNumber), predictedSequence = request.sequenceNumber, + allowFollowing = livePlaybackSnapshot()?.active == true, ) - val resolved = requestedCached || rebased != null - if (!resolved || !clearSegmentDemand(request, identity)) return false + val resolved = requested ?: rebased ?: return false + if (!clearSegmentDemand(request, identity)) return false + observeMediaSegment(resolved) if (rebased != null) session.streamState.jumpBufferedTo(request.format, rebased.header.sequenceNumber) return true } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSegmentDemandTracker.kt b/src/main/kotlin/dev/typetype/server/services/SabrSegmentDemandTracker.kt index 9b0163b0..cefd7153 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrSegmentDemandTracker.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrSegmentDemandTracker.kt @@ -11,7 +11,7 @@ internal object SabrSegmentDemandTracker { fun request(holder: SabrSessionHolder, request: SabrSegmentRequest, registeredAtMs: Long): Unit { if (request.isInitializationSegment) return if (holder.session.getCachedSegment(request) != null) return clear(holder, request) - if (holder.session.isBeyondEnd(request)) return clear(holder, request) + if (holder.hasFiniteEndBefore(request)) return clear(holder, request) val requestKey = key(holder, request) demands.putIfAbsent(requestKey, SegmentDemand(request, order.incrementAndGet(), registeredAtMs)) } @@ -37,13 +37,15 @@ internal object SabrSegmentDemandTracker { for ((key, value) in demands) { if (!key.startsWith(prefix)) continue val request = value.request - if (holder.session.getCachedSegment(request) != null || holder.session.isBeyondEnd(request)) { + if (holder.session.getCachedSegment(request) != null || holder.hasFiniteEndBefore(request)) { demands.remove(key, value) continue } - val startMs = holder.session.streamState - .getSegmentStartMs(request.format, request.sequenceNumber) - .coerceAtLeast(0L) + val startMs = if (holder.isFutureLiveRequest(request)) { + holder.livePlaybackSnapshot()?.seekableEndMs ?: Long.MAX_VALUE + } else { + holder.playbackSegmentStartMs(request.format, request.sequenceNumber) + } if (startMs < selectedStartMs || startMs == selectedStartMs && value.order < selectedOrder) { selected = request selectedStartMs = startMs @@ -88,6 +90,9 @@ internal object SabrSegmentDemandTracker { private fun identity(requestKey: String, demand: SegmentDemand): String = "$requestKey:${demand.order}" + private fun SabrSessionHolder.hasFiniteEndBefore(request: SabrSegmentRequest): Boolean = + session.isBeyondEnd(request) && livePlaybackSnapshot()?.active != true + private data class SegmentDemand( val request: SabrSegmentRequest, val order: Long, diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSessionHolder.kt b/src/main/kotlin/dev/typetype/server/services/SabrSessionHolder.kt index f81c79ad..c1cd9fa4 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrSessionHolder.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrSessionHolder.kt @@ -1,6 +1,7 @@ package dev.typetype.server.services import kotlinx.coroutines.sync.Mutex +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaSegment import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo @@ -19,15 +20,19 @@ internal class SabrSessionHolder( val sessionToken: String, val key: SabrSessionKey, @Volatile var lastRequestAt: Instant, + @Volatile var playerContextToken: SabrTokenBundle? = null, val pumpMutex: Mutex = Mutex(), ) { private val readerPositions = ConcurrentHashMap() private val lastServedSequences = ConcurrentHashMap() private val activeItags: MutableSet = ConcurrentHashMap.newKeySet() + private val observedMediaAnchors = ConcurrentHashMap() + private val liveInitializationData = ConcurrentHashMap() private val pendingRefetch = AtomicReference() private val pendingForwardSeek = AtomicReference() private val pumpStarted = AtomicBoolean(false) private val unauthorizedRefreshAttempted = AtomicBoolean(false) + private val expectedLive = AtomicBoolean(false) private val activeGeneration = AtomicLong(0L) private val segmentMemory = SabrMemorySegmentCache(MAX_SEGMENT_MEMORY_BYTES) private val playbackStatus = SabrPlaybackStatus() @@ -69,6 +74,30 @@ internal class SabrSessionHolder( fun markUnauthorizedRefreshAttempted(): Boolean = unauthorizedRefreshAttempted.compareAndSet(false, true) + fun markExpectedLive(): Unit { + expectedLive.set(true) + } + + fun expectsLive(): Boolean = expectedLive.get() + + fun observeMediaSegment(segment: SabrMediaSegment): Unit { + val header = segment.header + if (header.isInitSegment || header.sequenceNumber <= 0 || + (header.itag != audioFormat.itag && header.itag != videoFormat.itag) + ) return + observedMediaAnchors.compute(header.itag) { _, current -> + if (current == null || header.startMs >= current.header.startMs) segment else current + } + } + + fun observedMediaSegment(format: YoutubeSabrFormat): SabrMediaSegment? = observedMediaAnchors[format.itag] + + fun rememberLiveInitialization(itag: Int, data: ByteArray): Unit { + liveInitializationData.putIfAbsent(itag, data) + } + + fun liveInitialization(format: YoutubeSabrFormat): ByteArray? = liveInitializationData[format.itag] + fun setLastServedSequence(itag: Int, sequence: Int): Unit = setLastServedSequence(itag, sequence, activeGeneration()) diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSessionPlayerContext.kt b/src/main/kotlin/dev/typetype/server/services/SabrSessionPlayerContext.kt new file mode 100644 index 00000000..2d0154d2 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/SabrSessionPlayerContext.kt @@ -0,0 +1,8 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession + +internal inline fun SabrSessionHolder.withPlayerContext(crossinline block: YoutubeSabrSession.() -> T): T { + val token = playerContextToken ?: return session.block() + return TypetypeYoutubeSessionPoTokenProvider.withToken(token) { session.block() } +} diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSessionProgress.kt b/src/main/kotlin/dev/typetype/server/services/SabrSessionProgress.kt index eb156823..cc2cd770 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrSessionProgress.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrSessionProgress.kt @@ -9,8 +9,9 @@ internal fun SabrSessionHolder.markServed(segment: SabrMediaSegment): Unit { internal fun SabrSessionHolder.markServed(segment: SabrMediaSegment, generation: Long): Unit { if (!segment.header.isInitSegment) { + observeMediaSegment(segment) val format = if (audioFormat.itag == segment.header.itag) audioFormat else videoFormat - setReaderPosition(format, segment.header.startMs + segment.header.durationMs, generation) + setReaderPosition(format, playbackSegmentEndMs(format, segment.header.sequenceNumber), generation) setLastServedSequence(segment.header.itag, segment.header.sequenceNumber, generation) evictCachedSegmentsBefore(readerTailMs() - SabrPumpPolicy.backBufferMs(this)) } @@ -18,8 +19,9 @@ internal fun SabrSessionHolder.markServed(segment: SabrMediaSegment, generation: internal fun SabrSessionHolder.markPrepared(segment: SabrMediaSegment): Unit { if (!segment.header.isInitSegment) { + observeMediaSegment(segment) val format = if (audioFormat.itag == segment.header.itag) audioFormat else videoFormat - setReaderPosition(format, segment.header.startMs + segment.header.durationMs) + setReaderPosition(format, playbackSegmentEndMs(format, segment.header.sequenceNumber)) } } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSessionPump.kt b/src/main/kotlin/dev/typetype/server/services/SabrSessionPump.kt index 4d42f1f6..ae9283ee 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrSessionPump.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrSessionPump.kt @@ -16,11 +16,26 @@ internal class SabrSessionPump( suspend fun ensureWarmed(holder: SabrSessionHolder, maxPumps: Int) { val localization = Localization("en", "US") var pumps = 0 + var liveWarmupTarget: SabrLiveWarmupTarget? = null holder.setPlaybackState(SabrPlaybackState.PREPARING) - while (pumps < maxPumps && !isWarmEnough(holder) && !holder.session.isComplete) { + while (pumps < maxPumps && !isWarmEnough(holder) && (!holder.session.isComplete || holder.expectsLive())) { holder.pumpMutex.withLock { holder.setPlaybackState(SabrPlaybackState.REQUESTING) - runCatchingNonCancellation { holder.session.pumpOnceStreaming(localization) } + if (holder.expectsLive()) { + runCatchingNonCancellation { + withLiveWarmupRequestShape(holder, liveWarmupTarget) { + holder.withPlayerContext { pumpOnce(localization) } + } + } + .getOrDefault(emptyList()) + .forEach { segment -> + segmentCache?.put(holder, segment) + holder.observeMediaSegment(segment) + } + if (liveWarmupTarget == null) liveWarmupTarget = holder.liveWarmupTarget() + } else { + runCatchingNonCancellation { holder.withPlayerContext { pumpOnceStreaming(localization) } } + } } pumps++ } @@ -31,10 +46,14 @@ internal class SabrSessionPump( holder.setPlaybackState(SabrPlaybackState.IDLE) } - private fun isWarmEnough(holder: SabrSessionHolder): Boolean = - bothFormatsKnown(holder) || + private fun isWarmEnough(holder: SabrSessionHolder): Boolean { + val audioObserved = holder.observedMediaSegment(holder.audioFormat) != null + val videoObserved = !holder.isVideoActive() || holder.observedMediaSegment(holder.videoFormat) != null + if (holder.expectsLive()) return audioObserved && videoObserved + return bothFormatsKnown(holder) || holder.session.streamState.getMaxSegment(holder.audioFormat) > 0 && holder.session.streamState.getMaxSegment(holder.videoFormat) > 0 + } suspend fun fetchSegment( holder: SabrSessionHolder, @@ -53,7 +72,7 @@ internal class SabrSessionPump( segmentCache?.put(holder, segment) return segment } - if (holder.session.isBeyondEnd(request)) return null + if (holder.session.isBeyondEnd(request) && !holder.isFutureLiveRequest(request)) return null val localization = Localization("en", "US") if (request.isInitializationSegment) { return fetchSabrInitializationSegment(holder, request, localization, segmentCache, markServed) @@ -97,28 +116,28 @@ internal class SabrSessionPump( if (markServed) holder.markServed(cached) return@withLock cached } - if (holder.session.isBeyondEnd(request)) return@withLock null - holder.consumeMatchingSeek(request) + if (holder.session.isBeyondEnd(request) && !holder.isFutureLiveRequest(request)) return@withLock null + val exactSeekPrepared = holder.consumeMatchingSeek(request) val edgeMs = holder.session.streamState.getMinBufferedEndMs() val startMs = holder.session.streamState.getSegmentStartMs(request.format, request.sequenceNumber) if (startMs < edgeMs - REWIND_GAP_MS) { holder.setReaderPosition(request.format, startMs.coerceAtLeast(0L)) - holder.session.prepareForRewind(request) - runCatchingNonCancellation { holder.session.pumpOnceStreamingUntilCached(localization, request) } + if (!exactSeekPrepared) holder.session.prepareForRewind(request) + runCatchingNonCancellation { holder.withPlayerContext { pumpOnceStreamingUntilCached(localization, request) } } return@withLock holder.session.getCachedSegment(request)?.also { if (markServed) holder.markServed(it) } } else { holder.session.streamState.setPlayerTimeMs(maxOf(edgeMs, startMs + PLAYER_TIME_OFFSET_MS)) } - val fetched = runCatchingNonCancellation { holder.session.fetchSegment(request, localization) }.getOrNull() + val fetched = runCatchingNonCancellation { holder.withPlayerContext { fetchSegment(request, localization) } }.getOrNull() fetched?.also { if (markServed) holder.markServed(it) } } if (segment != null) { holder.clearSegmentDemand(request) segmentCache?.put(holder, segment) } - if (segment != null || holder.session.isBeyondEnd(request)) return segment + if (segment != null || holder.session.isBeyondEnd(request) && !holder.isFutureLiveRequest(request)) return segment pumps++ delay(FETCH_RETRY_DELAY_MS) } @@ -138,10 +157,12 @@ internal class SabrSessionPump( if (markServed) holder.markServed(cached) return@withLock cached } - if (holder.session.isBeyondEnd(request)) return@withLock null + if (holder.session.isBeyondEnd(request) && !holder.isFutureLiveRequest(request)) return@withLock null holder.consumeMatchingSeek(request) withTargetedRequestShape(holder, request) { - holder.session.fetchTargetedSegment(holder, request, localization, targetPlayerTime(holder, request)) + holder.withPlayerContext { + fetchTargetedSegment(holder, request, localization, targetPlayerTime(holder, request)) + } ?: SabrWindowSegmentFetcher.fetch(holder, request, localization, segmentCache) } } @@ -150,7 +171,7 @@ internal class SabrSessionPump( holder.clearSegmentDemand(request) segmentCache?.put(holder, segment) } - if (segment != null || holder.session.isBeyondEnd(request)) return segment + if (segment != null || holder.session.isBeyondEnd(request) && !holder.isFutureLiveRequest(request)) return segment attempt++ delay(FETCH_RETRY_DELAY_MS) } @@ -174,19 +195,4 @@ internal class SabrSessionPump( return playerTimeMs.takeIf { it < endMs } } - private fun SabrSessionHolder.consumeMatchingSeek(request: SabrSegmentRequest): Unit { - pendingRefetchRequest()?.takeIf { it.matches(request) }?.let { - consumeRefetch() - setPlaybackState(SabrPlaybackState.REPOSITIONING) - session.prepareForRewind(request) - } - pendingForwardSeekRequest()?.takeIf { it.matches(request) }?.let { - consumeForwardSeek() - setPlaybackState(SabrPlaybackState.REPOSITIONING) - session.prepareForForwardJump(request) - } - } - - private fun SabrSegmentRequest.matches(other: SabrSegmentRequest): Boolean = - format.itag == other.format.itag && sequenceNumber == other.sequenceNumber && isInitializationSegment == other.isInitializationSegment } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSessionPumpLoop.kt b/src/main/kotlin/dev/typetype/server/services/SabrSessionPumpLoop.kt index c38f95c0..8032c4be 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrSessionPumpLoop.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrSessionPumpLoop.kt @@ -62,13 +62,13 @@ internal class SabrSessionPumpLoop( ): Boolean { prepareEviction(holder) holder.consumeRefetch()?.let { request -> - if (holder.session.isBeyondEnd(request)) { + if (holder.session.isBeyondEnd(request) && !holder.isFutureLiveRequest(request)) { holder.clearSegmentDemand(request) return true } runtime.activateSeekMode() holder.setPlaybackState(SabrPlaybackState.REPOSITIONING) - holder.session.prepareForRewind(request) + holder.prepareForExplicitRewind(request) SabrPumpLogger.start(holder, "refetch", request) val pumped = pumpOnce(holder, localization, runtime) SabrPumpLogger.finish(holder, "refetch", request, pumped) @@ -76,13 +76,13 @@ internal class SabrSessionPumpLoop( return true } holder.consumeForwardSeek()?.let { request -> - if (holder.session.isBeyondEnd(request)) { + if (holder.session.isBeyondEnd(request) && !holder.isFutureLiveRequest(request)) { holder.clearSegmentDemand(request) return true } runtime.activateSeekMode() holder.setPlaybackState(SabrPlaybackState.REPOSITIONING) - holder.session.prepareForForwardJump(request) + holder.prepareForExplicitForwardJump(request) SabrPumpLogger.start(holder, "forward_seek", request) val pumped = pumpUntilCached(holder, localization, request, runtime) SabrPumpLogger.finish(holder, "forward_seek", request, pumped.segmentCount) @@ -91,13 +91,30 @@ internal class SabrSessionPumpLoop( } holder.nextSegmentDemand()?.let { request -> val demandIdentity = holder.segmentDemandIdentity(request) ?: return true - SabrPumpLogger.start(holder, "demand", request) - runtime.beginDemand(demandIdentity) - val result = pumpDemand(holder, localization, request, runtime) - return SabrDemandAttemptFinisher.finish(holder, request, demandIdentity, result, runtime) + val wasFutureLiveRequest = holder.isFutureLiveRequest(request) + if (!holder.beginInFlightSegmentDemand(request, demandIdentity, wasFutureLiveRequest)) return true + try { + SabrPumpLogger.start(holder, "demand", request) + runtime.beginDemand(demandIdentity) + val result = pumpDemand(holder, localization, request, runtime) + return SabrDemandAttemptFinisher.finish( + holder, + request, + demandIdentity, + result, + runtime, + wasFutureLiveRequest, + ) + } finally { + holder.finishInFlightSegmentDemand(demandIdentity) + } + } + if (holder.livePlaybackSnapshot()?.active == true) { + holder.setPlaybackState(SabrPlaybackState.IDLE) + return false } - if (holder.session.requestNumber == 0) return false - if (holder.session.isComplete && !holder.hasPendingSeek()) return false + if (holder.session.requestNumber == 0 && !holder.expectsLive()) return false + if (holder.session.isComplete && holder.livePlaybackSnapshot()?.active != true && !holder.hasPendingSeek()) return false if (runtime.isThrottled(holder)) { holder.setPlaybackState(SabrPlaybackState.THROTTLED) return false @@ -111,7 +128,7 @@ internal class SabrSessionPumpLoop( } private suspend fun pumpOnce(holder: SabrSessionHolder, localization: Localization, runtime: SabrPumpRuntime): Int { return try { - runInterruptible(Dispatchers.IO) { holder.session.pumpOnceStreaming(localization) } + runInterruptible(Dispatchers.IO) { holder.withPlayerContext { pumpOnceStreaming(localization) } } .also { unauthorizedRecovery.verify(holder) } } finally { runtime.recordRequest() @@ -125,7 +142,7 @@ internal class SabrSessionPumpLoop( runtime: SabrPumpRuntime, ): YoutubeSabrSession.DemandResponseResult { return try { - runInterruptible(Dispatchers.IO) { holder.session.pumpOnceStreamingForDemand(localization, request) } + runInterruptible(Dispatchers.IO) { holder.withPlayerContext { pumpOnceStreamingForDemand(localization, request) } } .also { unauthorizedRecovery.verify(holder) } } finally { runtime.recordRequest() @@ -139,8 +156,21 @@ internal class SabrSessionPumpLoop( runtime: SabrPumpRuntime, ): YoutubeSabrSession.DemandResponseResult { val edgeMs = holder.session.streamState.getMinBufferedEndMs() - val startMs = holder.session.streamState.getSegmentStartMs(request.format, request.sequenceNumber).coerceAtLeast(0L) + val startMs = holder.playbackSegmentStartMs(request.format, request.sequenceNumber) holder.setReaderPosition(request.format, startMs) + if (holder.livePlaybackSnapshot()?.active == true) { + if (holder.isHistoricalLiveRequest(request)) { + holder.setPlaybackState(SabrPlaybackState.REPOSITIONING) + holder.prepareForHistoricalLiveRewind(request) + return withTargetedRequestShape(holder, request) { + pumpUntilCached(holder, localization, request, runtime) + } + } + holder.setPlaybackState(SabrPlaybackState.REQUESTING) + return withTargetedRequestShape(holder, request) { + pumpUntilCached(holder, localization, request, runtime) + } + } return when { startMs < edgeMs -> { holder.setPlaybackState(SabrPlaybackState.REPOSITIONING) @@ -165,7 +195,13 @@ internal class SabrSessionPumpLoop( holder.session.evictPlayed() } - private fun demandDelayMs(holder: SabrSessionHolder, intervalMs: Long): Long = - if (holder.pendingSegmentDemandSummary() == null) intervalMs - else maxOf(intervalMs, holder.session.demandBackoffRemainingMs) + private fun demandDelayMs(holder: SabrSessionHolder, intervalMs: Long): Long { + val demand = holder.nextSegmentDemand() + return SabrPumpPolicy.demandDelayMs( + intervalMs, + holder.session.demandBackoffRemainingMs.takeIf { demand != null } ?: 0L, + holder.livePlaybackSnapshot()?.active == true, + demand?.let(holder::isFutureLiveRequest), + ) + } } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrSessionStore.kt b/src/main/kotlin/dev/typetype/server/services/SabrSessionStore.kt index f8341b41..ff7bce9a 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrSessionStore.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrSessionStore.kt @@ -72,7 +72,16 @@ internal class SabrSessionStore( runCatching { provider.getPoToken(info, session.streamState) } .getOrNull() ?.let { session.streamState.setPoToken(it) } - val holder = SabrSessionHolder(session, info, audioFormat, videoFormat, playbackToken ?: SabrSessionTokenGenerator.newToken(), key, Instant.now()) + val holder = SabrSessionHolder( + session, + info, + audioFormat, + videoFormat, + playbackToken ?: SabrSessionTokenGenerator.newToken(), + key, + Instant.now(), + initialToken, + ) holder.setPlayerTimeMs(normalizedStartTimeMs) registry.put(key, holder) if (startPump) startPump(holder) @@ -118,7 +127,7 @@ internal class SabrSessionStore( holder.session.getCachedSegment(request)?.let { segmentCache.put(holder, it) holder.clearSegmentDemand(request) - return it.toCachedSabrSegment(request.format.mimeType.orEmpty()) + return segmentCache.get(holder, request) } return segmentCache.get(holder, request)?.also { holder.clearSegmentDemand(request) } } @@ -165,21 +174,22 @@ internal class SabrSessionStore( holder: SabrSessionHolder, format: YoutubeSabrFormat, ): ByteArray? { + holder.liveInitialization(format)?.let { return it } val request = SabrSegmentRequest.initialization(format) holder.session.getCachedSegment(request)?.let { segmentCache.put(holder, it); return it.data } - val sourceFormat = infoFetcher.initializationFormat(holder.key.videoId, format) ?: format - SabrInitializationData.fetch(sourceFormat, initCache)?.let { + SabrInitializationData.fetch(holder.key.videoId, format, initCache)?.let { holder.session.streamState.ingestInitializationData(format, it) return it } - SabrInitializationData.fetchFallback(holder, format, initCache)?.let { return it } - return pump.fetchSegment(holder, request)?.also { segmentCache.put(holder, it) }?.data + SabrInitializationData.bootstrap(holder, format, initCache)?.let { return it } + val segment = pump.fetchSegment(holder, request) ?: return null + segmentCache.put(holder, segment) + SabrInitializationData.remember(holder.key.videoId, format, segment.data, initCache) + return segment.data } private suspend fun fetchDirectInitialization(holder: SabrSessionHolder, format: YoutubeSabrFormat): Unit { - val source = infoFetcher.initializationFormat(holder.key.videoId, format) ?: format - val data = SabrInitializationData.fetch(source, initCache) ?: return - SabrInitializationData.remember(format, data) + val data = SabrInitializationData.fetch(holder.key.videoId, format, initCache) ?: return holder.pumpMutex.withLock { holder.session.streamState.ingestInitializationData(format, data) } } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrTargetRequest.kt b/src/main/kotlin/dev/typetype/server/services/SabrTargetRequest.kt index 6737282c..7b49e7a4 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrTargetRequest.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrTargetRequest.kt @@ -15,7 +15,7 @@ internal fun YoutubeSabrSession.fetchTargetedSegment( val targetPlayerTimeMs = when { playerTimeMs != null -> playerTimeMs request.isInitializationSegment -> 0L - else -> targetTimeInsideSegment(request) + else -> holder.targetTimeInsideSegment(request) } val result = runCatchingNonCancellation { prepareForTarget(request, targetPlayerTimeMs) @@ -55,9 +55,9 @@ private fun SabrMediaSegment.matches(request: SabrSegmentRequest): Boolean { } } -private fun YoutubeSabrSession.targetTimeInsideSegment(request: SabrSegmentRequest): Long { - val startMs = streamState.getSegmentStartMs(request.format, request.sequenceNumber).coerceAtLeast(0L) - val nextStartMs = streamState.getSegmentStartMs(request.format, request.sequenceNumber + 1).coerceAtLeast(0L) +private fun SabrSessionHolder.targetTimeInsideSegment(request: SabrSegmentRequest): Long { + val startMs = playbackSegmentStartMs(request.format, request.sequenceNumber) + val nextStartMs = playbackSegmentStartMs(request.format, request.sequenceNumber + 1) if (nextStartMs <= startMs + 1L) return startMs return minOf(startMs + TARGET_SEGMENT_OFFSET_MS, nextStartMs - 1L) } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrTargetRequestShape.kt b/src/main/kotlin/dev/typetype/server/services/SabrTargetRequestShape.kt index 3b701a8d..0b14da86 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrTargetRequestShape.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrTargetRequestShape.kt @@ -13,7 +13,7 @@ internal inline fun withTargetedRequestShape( ): T { val companion = holder.companionFormat(request.format) val state = holder.session.streamState - val requestStartMs = state.getSegmentStartMs(request.format, request.sequenceNumber).coerceAtLeast(0L) + val requestStartMs = holder.playbackSegmentStartMs(request.format, request.sequenceNumber) val targetPlayerTimeMs = request.targetPlayerTimeMs(holder, requestStartMs) val ranges = listOf(request.targetRange(holder), companion.targetCompanionRange(holder, targetPlayerTimeMs)) holder.session.prepareForMediaSegment(request) @@ -53,14 +53,14 @@ private fun YoutubeSabrFormat.targetCompanionRange(holder: SabrSessionHolder, ta private fun SabrSegmentRequest.targetPlayerTimeMs(holder: SabrSessionHolder, startMs: Long): Long { val playerTimeMs = holder.playerTimeMs() - val endMs = holder.session.streamState.getSegmentEndMs(format, sequenceNumber).coerceAtLeast(0L) + val endMs = holder.playbackSegmentEndMs(format, sequenceNumber) if (endMs > startMs && playerTimeMs >= startMs && playerTimeMs < endMs) return playerTimeMs if (endMs > startMs + 1L) return minOf(startMs + TARGET_SEGMENT_OFFSET_MS, endMs - 1L) return startMs } private fun YoutubeSabrFormat.bufferedRange(holder: SabrSessionHolder, bufferedSequence: Int): SabrBufferedRange { - val endMs = holder.session.streamState.getSegmentEndMs(this, bufferedSequence) + val endMs = holder.playbackSegmentEndMs(this, bufferedSequence) val durationMs = endMs.takeIf { it > 0L } ?: 1L return SabrBufferedRange( itag, diff --git a/src/main/kotlin/dev/typetype/server/services/SabrUnauthorizedResponseRecovery.kt b/src/main/kotlin/dev/typetype/server/services/SabrUnauthorizedResponseRecovery.kt index 57dae5f7..6f8d83de 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrUnauthorizedResponseRecovery.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrUnauthorizedResponseRecovery.kt @@ -9,10 +9,11 @@ internal class SabrUnauthorizedResponseRecovery( val response = holder.session.diagnosticTrace.substringAfterLast("response n=", missingDelimiterValue = "") if (!response.contains(" http=403 ")) return if (!holder.markUnauthorizedRefreshAttempted()) throw unauthorized() - val token = refreshPoToken(holder.key.videoId) - ?.streamingPoTokenBytesFor(holder.info) + val refreshed = refreshPoToken(holder.key.videoId) ?: throw unauthorized() + val token = refreshed.streamingPoTokenBytesFor(holder.info) ?.takeUnless { holder.session.streamState.poToken?.contentEquals(it) == true } ?: throw unauthorized() + holder.playerContextToken = refreshed holder.session.streamState.setPoToken(token) holder.session.addDiagnosticEvent("upstream 403 refreshed TypeType PO token") } diff --git a/src/main/kotlin/dev/typetype/server/services/SabrWindowSegmentFetcher.kt b/src/main/kotlin/dev/typetype/server/services/SabrWindowSegmentFetcher.kt index 841baf7c..9969c0f4 100644 --- a/src/main/kotlin/dev/typetype/server/services/SabrWindowSegmentFetcher.kt +++ b/src/main/kotlin/dev/typetype/server/services/SabrWindowSegmentFetcher.kt @@ -13,7 +13,7 @@ internal object SabrWindowSegmentFetcher { segmentCache: SabrSegmentCache?, ): SabrMediaSegment? { val result = runCatchingNonCancellation { - holder.session.pumpOnceStreamingUntilCached(localization, request) + holder.withPlayerContext { pumpOnceStreamingUntilCached(localization, request) } holder.session.getCachedSegment(request) } result.onFailure { error -> diff --git a/src/main/kotlin/dev/typetype/server/services/StreamExtractionErrorMapper.kt b/src/main/kotlin/dev/typetype/server/services/StreamExtractionErrorMapper.kt index 0ced431e..4f7956b6 100644 --- a/src/main/kotlin/dev/typetype/server/services/StreamExtractionErrorMapper.kt +++ b/src/main/kotlin/dev/typetype/server/services/StreamExtractionErrorMapper.kt @@ -6,27 +6,44 @@ import org.schabi.newpipe.extractor.exceptions.GeographicRestrictionException import org.schabi.newpipe.extractor.exceptions.NeedLoginException import org.schabi.newpipe.extractor.exceptions.PaidContentException import org.schabi.newpipe.extractor.exceptions.PrivateContentException +import org.schabi.newpipe.extractor.exceptions.VideoNotReleaseException import org.schabi.newpipe.extractor.exceptions.YoutubeMusicPremiumContentException internal object StreamExtractionErrorMapper { const val MEMBERS_ONLY_FALLBACK = "This video is only available for members" + const val PAID_CONTENT_FALLBACK = "This video is a paid video" - fun map(error: Throwable, sourceUrl: String? = null, fallback: String = "Extraction failed"): ExtractionResult = when { - error is NeedLoginException || - error is PaidContentException || - error is YoutubeMusicPremiumContentException -> ExtractionResult.BadRequest(sanitize(error.message) ?: MEMBERS_ONLY_FALLBACK) - else -> mapByType(error, fallback) - } + fun map(error: Throwable, sourceUrl: String? = null, fallback: String = "Extraction failed"): ExtractionResult = + mapByType(error, fallback) private fun mapByType(error: Throwable, fallback: String): ExtractionResult = when (error) { - is NeedLoginException, - is PaidContentException, - is YoutubeMusicPremiumContentException -> ExtractionResult.BadRequest(sanitize(error.message) ?: MEMBERS_ONLY_FALLBACK) + is NeedLoginException -> ExtractionResult.BadRequest( + sanitize(error.message) ?: MEMBERS_ONLY_FALLBACK, + "members_only", + ) + is PaidContentException -> paidContent(error.message) + is YoutubeMusicPremiumContentException -> ExtractionResult.BadRequest( + sanitize(error.message) ?: PAID_CONTENT_FALLBACK, + "paid_content", + ) + is VideoNotReleaseException -> ExtractionResult.Failure( + sanitize(error.message) ?: "This premiere has not started yet", + "scheduled_premiere", + ) is GeographicRestrictionException, is AgeRestrictedContentException, is PrivateContentException -> ExtractionResult.BadRequest(sanitize(error.message) ?: "Content not available") else -> ExtractionResult.Failure(sanitize(error.message) ?: fallback) } + private fun paidContent(message: String?): ExtractionResult.BadRequest { + val sanitized = sanitize(message) + return if (sanitized == MEMBERS_ONLY_FALLBACK) { + ExtractionResult.BadRequest(sanitized, "members_only") + } else { + ExtractionResult.BadRequest(sanitized ?: PAID_CONTENT_FALLBACK, "paid_content") + } + } + private fun sanitize(message: String?): String? = ExtractionErrorSanitizer.sanitize(message) } diff --git a/src/main/kotlin/dev/typetype/server/services/SubscriptionFeedService.kt b/src/main/kotlin/dev/typetype/server/services/SubscriptionFeedService.kt index a97ee610..496cdb00 100644 --- a/src/main/kotlin/dev/typetype/server/services/SubscriptionFeedService.kt +++ b/src/main/kotlin/dev/typetype/server/services/SubscriptionFeedService.kt @@ -11,6 +11,7 @@ import kotlinx.coroutines.sync.Semaphore import kotlinx.coroutines.sync.withPermit import kotlinx.coroutines.withTimeout import kotlinx.serialization.builtins.ListSerializer +import java.net.URI import java.util.Base64 class SubscriptionFeedService( @@ -58,30 +59,63 @@ class SubscriptionFeedService( val videos = coroutineScope { subs.map { sub -> async { - semaphore.withPermit { - runCatching { - withTimeout(CHANNEL_TIMEOUT_MS) { - (channelService.getChannel(sub.channelUrl, null) as? ExtractionResult.Success) - ?.data?.videos.orEmpty() - } - }.getOrElse { emptyList() } - } + runCatching { fetchForSubscription(sub.channelUrl) }.getOrElse { emptyList() } } }.map { it.await() }.flatten() } val sorted = videos.sortedWith( - compareByDescending { v: VideoItem -> if (v.uploaded == -1L) Long.MIN_VALUE else v.uploaded } + compareByDescending { it.isLive } + .thenByDescending { if (it.uploaded == -1L) Long.MIN_VALUE else it.uploaded } ) runCatching { cache.set(key, CacheJson.encodeToString(ListSerializer(VideoItem.serializer()), sorted), FEED_TTL_SECONDS) } return sorted } + private suspend fun fetchForSubscription(channelUrl: String): List = coroutineScope { + val channel = async { fetchVideos(channelUrl) } + val livestreams = if (isYoutubeUrl(channelUrl)) { + async { fetchVideos(channelUrl.toLivestreamsTabUrl()) } + } else null + mergeVideos(channel.await(), livestreams?.await().orEmpty()) + } + + private suspend fun fetchVideos(url: String): List = semaphore.withPermit { + runCatching { + withTimeout(CHANNEL_TIMEOUT_MS) { + (channelService.getChannel(url, null) as? ExtractionResult.Success)?.data?.videos.orEmpty() + } + }.getOrElse { + emptyList() + } + } + + private fun mergeVideos(channel: List, livestreams: List): List = + buildMap { + channel.forEach { put(it.feedKey(), it) } + livestreams.forEach { put(it.feedKey(), it) } + }.values.toList() + + private fun VideoItem.feedKey(): String = url.ifBlank { id.ifBlank { "$uploaderUrl|$title" } } + + private fun String.toLivestreamsTabUrl(): String { + val uri = URI(this) + val path = uri.path.trimEnd('/') + val segments = path.split('/').filter(String::isNotBlank) + val basePath = if (segments.size >= 2 && segments.last() in YOUTUBE_CHANNEL_TABS) { + path.substringBeforeLast('/') + } else { + path + } + return URI(uri.scheme, uri.userInfo, uri.host, uri.port, "$basePath/streams", null, null).toString() + } + private fun encodeNextPage(page: Int): String = Base64.getEncoder().encodeToString("""{"page":$page}""".toByteArray()) companion object { - private const val FEED_TTL_SECONDS = 300L + private const val FEED_TTL_SECONDS = 60L private const val MAX_CONCURRENT_FETCHES = 20 private const val CHANNEL_TIMEOUT_MS = 15_000L + private val YOUTUBE_CHANNEL_TABS = setOf("featured", "videos", "shorts", "streams", "playlists", "community", "about") } } diff --git a/src/main/kotlin/dev/typetype/server/services/TypetypeTokenYoutubeSessionClient.kt b/src/main/kotlin/dev/typetype/server/services/TypetypeTokenYoutubeSessionClient.kt index 1eb2ca2d..7f19cd47 100644 --- a/src/main/kotlin/dev/typetype/server/services/TypetypeTokenYoutubeSessionClient.kt +++ b/src/main/kotlin/dev/typetype/server/services/TypetypeTokenYoutubeSessionClient.kt @@ -4,11 +4,13 @@ import com.grack.nanojson.JsonParser import kotlinx.coroutines.suspendCancellableCoroutine import okhttp3.Call import okhttp3.Callback +import okhttp3.HttpUrl.Companion.toHttpUrlOrNull import okhttp3.OkHttpClient import okhttp3.Request import okhttp3.Response import org.json.JSONObject import org.schabi.newpipe.extractor.services.youtube.YoutubeParsingHelper +import org.schabi.newpipe.extractor.services.youtube.sabr.TypeTypeYoutubeSabrInfoFactory import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrClientProfile import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrProbe @@ -96,11 +98,21 @@ internal class TypetypeTokenYoutubeSessionClient( ), ), ) - YoutubeSabrProbe.fromPlayerResponse( + val info = YoutubeSabrProbe.fromPlayerResponse( videoId, YoutubeSabrClientProfile.MWEB, YoutubeParsingHelper.generateContentPlaybackNonce(), JsonParser.`object`().from(playerResponse.toString()), ) + val playbackUrl = optString("serverAbrStreamingUrl").takeIf { it.isNotBlank() } + val clientVersion = playbackUrl + ?.toHttpUrlOrNull() + ?.queryParameter("cver") + ?.takeIf { it.isNotBlank() } + if (playbackUrl != null && clientVersion != null) { + TypeTypeYoutubeSabrInfoFactory.withPlaybackUrlAndClientVersion(info, playbackUrl, clientVersion) + } else { + info + } }.getOrNull() } diff --git a/src/main/kotlin/dev/typetype/server/services/TypetypeYoutubeSessionPoTokenProvider.kt b/src/main/kotlin/dev/typetype/server/services/TypetypeYoutubeSessionPoTokenProvider.kt new file mode 100644 index 00000000..f4f28006 --- /dev/null +++ b/src/main/kotlin/dev/typetype/server/services/TypetypeYoutubeSessionPoTokenProvider.kt @@ -0,0 +1,28 @@ +package dev.typetype.server.services + +import org.schabi.newpipe.extractor.localization.ContentCountry +import org.schabi.newpipe.extractor.localization.Localization +import org.schabi.newpipe.extractor.services.youtube.YoutubeSessionPoToken +import org.schabi.newpipe.extractor.services.youtube.YoutubeSessionPoTokenProvider + +internal object TypetypeYoutubeSessionPoTokenProvider : YoutubeSessionPoTokenProvider { + private val scopedToken = ThreadLocal() + + fun withToken(token: SabrTokenBundle, block: () -> T): T { + val previous = scopedToken.get() + val sessionToken = YoutubeSessionPoToken(token.visitorData, token.visitorBoundPoToken) + scopedToken.set(sessionToken) + return try { + block() + } finally { + if (previous == null) scopedToken.remove() else scopedToken.set(previous) + } + } + + override fun getSessionPoToken( + clientName: String, + localization: Localization, + contentCountry: ContentCountry, + loggedIn: Boolean, + ): YoutubeSessionPoToken? = scopedToken.get() +} diff --git a/src/main/kotlin/dev/typetype/server/services/UserVideoMetadataRepairService.kt b/src/main/kotlin/dev/typetype/server/services/UserVideoMetadataRepairService.kt index d36f0f96..2f1757ec 100644 --- a/src/main/kotlin/dev/typetype/server/services/UserVideoMetadataRepairService.kt +++ b/src/main/kotlin/dev/typetype/server/services/UserVideoMetadataRepairService.kt @@ -4,6 +4,10 @@ import dev.typetype.server.db.DatabaseFactory import dev.typetype.server.db.tables.FavoritesTable import dev.typetype.server.db.tables.PlaylistVideosTable import dev.typetype.server.db.tables.WatchLaterTable +import kotlinx.coroutines.CancellationException +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.Dispatchers +import kotlinx.coroutines.launch import org.jetbrains.exposed.v1.core.and import org.jetbrains.exposed.v1.core.eq import org.jetbrains.exposed.v1.core.lessEq @@ -11,14 +15,51 @@ import org.jetbrains.exposed.v1.core.like import org.jetbrains.exposed.v1.core.or import org.jetbrains.exposed.v1.jdbc.selectAll import org.jetbrains.exposed.v1.jdbc.update +import org.slf4j.LoggerFactory +import java.util.concurrent.ConcurrentHashMap class UserVideoMetadataRepairService(private val resolver: VideoMetadataResolver) { + private val logger = LoggerFactory.getLogger(UserVideoMetadataRepairService::class.java) + private val inFlight = ConcurrentHashMap.newKeySet() + private val lastScheduledAt = ConcurrentHashMap() + + fun schedulePlaylists(scope: CoroutineScope, userId: String): Unit = schedule(scope, "playlists:$userId") { + repairPlaylists(userId) + } + + fun scheduleWatchLater(scope: CoroutineScope, userId: String): Unit = schedule(scope, "watch-later:$userId") { + repairWatchLater(userId) + } + + fun scheduleFavorites(scope: CoroutineScope, userId: String): Unit = schedule(scope, "favorites:$userId") { + repairFavorites(userId) + } + suspend fun repairPlaylists(userId: String): Int = repair(userId, ::playlistCandidateUrls) suspend fun repairWatchLater(userId: String): Int = repair(userId, ::watchLaterCandidateUrls) suspend fun repairFavorites(userId: String): Int = repair(userId, ::favoriteCandidateUrls) + private fun schedule(scope: CoroutineScope, key: String, repair: suspend () -> Unit) { + val now = System.currentTimeMillis() + val lastScheduled = lastScheduledAt[key] + if (lastScheduled != null && now - lastScheduled < REPAIR_COOLDOWN_MS) return + if (!inFlight.add(key)) return + lastScheduledAt[key] = now + scope.launch(Dispatchers.IO) { + try { + repair() + } catch (error: CancellationException) { + throw error + } catch (error: Exception) { + logger.warn("Background video metadata repair failed", error) + } finally { + inFlight.remove(key) + } + } + } + private suspend fun repair(userId: String, candidates: suspend (String) -> List): Int { val urls = candidates(userId).take(MAX_REPAIR_PER_REQUEST) if (urls.isEmpty()) return 0 @@ -104,5 +145,6 @@ class UserVideoMetadataRepairService(private val resolver: VideoMetadataResolver const val FALLBACK_TITLE_PATTERN = "YouTube video %" const val YOUTUBE_THUMB_PATTERN = "https://i.ytimg.com/vi/%" const val MAX_REPAIR_PER_REQUEST = 25 + const val REPAIR_COOLDOWN_MS = 5 * 60 * 1000L } } diff --git a/src/main/kotlin/dev/typetype/server/services/YoutubeTakeoutFactory.kt b/src/main/kotlin/dev/typetype/server/services/YoutubeTakeoutFactory.kt index 0c846b53..480922df 100644 --- a/src/main/kotlin/dev/typetype/server/services/YoutubeTakeoutFactory.kt +++ b/src/main/kotlin/dev/typetype/server/services/YoutubeTakeoutFactory.kt @@ -7,15 +7,13 @@ object YoutubeTakeoutFactory { historyService: HistoryService, favoritesService: FavoritesService, watchLaterService: WatchLaterService, - streamService: StreamService? = null, ): YoutubeTakeoutImportJobService { val previewLookup = YoutubeTakeoutPreviewLookupService(historyService, favoritesService, watchLaterService) val signalImport = YoutubeTakeoutSignalImportService(favoritesService, watchLaterService, historyService) - val metadataResolver = streamService?.let(::VideoMetadataResolver) return YoutubeTakeoutImportJobService( parser = YoutubeTakeoutParserService(), previewService = YoutubeTakeoutPreviewService(subscriptionsService, playlistService, previewLookup), - importerService = YoutubeTakeoutImporterService(subscriptionsService, playlistService, signalImport, YoutubeTakeoutPlaylistKeyService(), metadataResolver), + importerService = YoutubeTakeoutImporterService(subscriptionsService, playlistService, signalImport, YoutubeTakeoutPlaylistKeyService()), store = YoutubeTakeoutImportJobStore(), statusStore = YoutubeTakeoutImportJobStatusStore(), archiveStore = YoutubeTakeoutImportJobArchiveStore(), diff --git a/src/main/kotlin/dev/typetype/server/services/YoutubeTakeoutImporterService.kt b/src/main/kotlin/dev/typetype/server/services/YoutubeTakeoutImporterService.kt index e359a49a..e1beb779 100644 --- a/src/main/kotlin/dev/typetype/server/services/YoutubeTakeoutImporterService.kt +++ b/src/main/kotlin/dev/typetype/server/services/YoutubeTakeoutImporterService.kt @@ -15,7 +15,6 @@ class YoutubeTakeoutImporterService( private val playlistService: PlaylistService, private val signalImportService: YoutubeTakeoutSignalImportService, private val playlistKeyService: YoutubeTakeoutPlaylistKeyService = YoutubeTakeoutPlaylistKeyService(), - private val metadataResolver: VideoMetadataResolver? = null, ) { suspend fun commit(userId: String, parsed: YoutubeTakeoutParsedData, plan: YoutubeTakeoutCommitPlan): YoutubeTakeoutImportReportItem = coroutineScope { val (issues, issueSummary) = YoutubeTakeoutIssueService.build(parsed.warnings, parsed.errors, stage = "commit") @@ -43,9 +42,7 @@ class YoutubeTakeoutImporterService( var itemImported = 0 var itemSkipped = 0 val createdBySource = mutableMapOf() - val playlistItems = metadataResolver?.enrichPlaylistItems(parsed.playlistItems) ?: parsed.playlistItems - val watchLater = metadataResolver?.enrichPlaylistVideos(parsed.watchLater) ?: parsed.watchLater - val favorites = metadataResolver?.enrichFavorites(parsed.favorites) ?: parsed.favorites.map { it.withYoutubeFallbackTitle() } + val favorites = parsed.favorites.map { it.withYoutubeFallbackTitle() } if (plan.importPlaylists) { parsed.playlists.forEach { item -> if (YoutubeTakeoutSystemPlaylist.canonicalKey(item.name) != null || YoutubeTakeoutSystemPlaylist.canonicalKey(item.id) != null) { @@ -74,7 +71,7 @@ class YoutubeTakeoutImporterService( } } if (plan.importPlaylistItems) { - playlistItems.forEach { (playlistKey, videos) -> + parsed.playlistItems.forEach { (playlistKey, videos) -> val normalizedKey = playlistKey.lowercase() if (YoutubeTakeoutSystemPlaylist.canonicalKey(normalizedKey) != null) { itemSkipped += videos.size @@ -94,7 +91,7 @@ class YoutubeTakeoutImporterService( } } val favoriteDeferred = if (plan.importFavorites) async { signalImportService.importFavorites(userId, favorites) } else null - val watchLaterDeferred = if (plan.importWatchLater) async { signalImportService.importWatchLater(userId, watchLater) } else null + val watchLaterDeferred = if (plan.importWatchLater) async { signalImportService.importWatchLater(userId, parsed.watchLater) } else null val historyDeferred = if (plan.importHistory) async { signalImportService.importHistory(userId, parsed.history) } else null val emptyStats = YoutubeTakeoutImportStats(0, 0, 0) val favoriteStats = favoriteDeferred?.await() ?: emptyStats diff --git a/src/test/kotlin/dev/typetype/server/EncodedVideoUrlDeleteRoutesTest.kt b/src/test/kotlin/dev/typetype/server/EncodedVideoUrlDeleteRoutesTest.kt index f1b1dbc2..066ccf15 100644 --- a/src/test/kotlin/dev/typetype/server/EncodedVideoUrlDeleteRoutesTest.kt +++ b/src/test/kotlin/dev/typetype/server/EncodedVideoUrlDeleteRoutesTest.kt @@ -1,11 +1,14 @@ package dev.typetype.server import dev.typetype.server.models.PlaylistItem +import dev.typetype.server.models.WatchLaterItem import dev.typetype.server.routes.favoritesRoutes import dev.typetype.server.routes.playlistRoutes +import dev.typetype.server.routes.watchLaterRoutes import dev.typetype.server.services.AuthService import dev.typetype.server.services.FavoritesService import dev.typetype.server.services.PlaylistService +import dev.typetype.server.services.WatchLaterService import io.ktor.client.request.delete import io.ktor.client.request.get import io.ktor.client.request.headers @@ -33,6 +36,7 @@ class EncodedVideoUrlDeleteRoutesTest { private val auth = AuthService.fixed(TEST_USER_ID) private val favoritesService = FavoritesService() private val playlistService = PlaylistService() + private val watchLaterService = WatchLaterService() companion object { private const val VIDEO_URL = "https://www.youtube.com/watch?v=ZZZZZZZZZZZ" @@ -52,6 +56,7 @@ class EncodedVideoUrlDeleteRoutesTest { routing { favoritesRoutes(favoritesService, auth) playlistRoutes(playlistService, auth) + watchLaterRoutes(watchLaterService, auth) } } block() @@ -91,4 +96,19 @@ class EncodedVideoUrlDeleteRoutesTest { }.bodyAsText() assertFalse(afterDelete.contains(VIDEO_URL)) } + + @Test + fun `DELETE watch later removes encoded youtube watch url`() = withApp { + watchLaterService.add(TEST_USER_ID, WatchLaterItem(url = VIDEO_URL, title = "test", thumbnail = "", duration = 0L)) + + val delete = client.delete("/watch-later/$ENCODED_VIDEO_URL") { + headers.append(HttpHeaders.Authorization, "Bearer test-jwt") + } + assertEquals(HttpStatusCode.NoContent, delete.status) + + val afterDelete = client.get("/watch-later") { + headers.append(HttpHeaders.Authorization, "Bearer test-jwt") + }.bodyAsText() + assertFalse(afterDelete.contains(VIDEO_URL)) + } } diff --git a/src/test/kotlin/dev/typetype/server/PublicCachePolicyTest.kt b/src/test/kotlin/dev/typetype/server/PublicCachePolicyTest.kt index ef549f53..03d49275 100644 --- a/src/test/kotlin/dev/typetype/server/PublicCachePolicyTest.kt +++ b/src/test/kotlin/dev/typetype/server/PublicCachePolicyTest.kt @@ -41,6 +41,13 @@ class PublicCachePolicyTest { assertEquals(900L, PublicCachePolicy.channelTtl("https://www.youtube.com/channel/id", null, "latest")) } + @Test + fun `channel ttl keeps livestream state fresh`() { + assertEquals(60L, PublicCachePolicy.channelTtl("https://www.youtube.com/channel/id/streams", null, null)) + assertEquals(60L, PublicCachePolicy.channelTtl("https://www.youtube.com/channel/id/livestreams", null, null)) + assertEquals(300L, PublicCachePolicy.channelTtl("https://www.youtube.com/channel/id/streams", "cursor", null)) + } + @Test fun `comments ttl is shortest on first youtube page`() { assertEquals(180L, PublicCachePolicy.commentsTtl("https://youtube.com/watch?v=id", null)) diff --git a/src/test/kotlin/dev/typetype/server/SabrPlaybackGranularRoutesTest.kt b/src/test/kotlin/dev/typetype/server/SabrPlaybackGranularRoutesTest.kt index 4e425247..b9d216cf 100644 --- a/src/test/kotlin/dev/typetype/server/SabrPlaybackGranularRoutesTest.kt +++ b/src/test/kotlin/dev/typetype/server/SabrPlaybackGranularRoutesTest.kt @@ -5,6 +5,7 @@ import dev.typetype.server.services.CachedSabrSegment import dev.typetype.server.services.SabrSessionHolder import dev.typetype.server.services.SabrSessionKey import dev.typetype.server.services.SabrSessionStore +import dev.typetype.server.services.SabrSegmentDemandTracker import io.ktor.client.request.post import io.ktor.client.request.setBody import io.ktor.client.statement.bodyAsText @@ -26,6 +27,7 @@ import io.mockk.mockk import io.mockk.verify import org.junit.jupiter.api.Assertions.assertEquals import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.AfterEach import org.junit.jupiter.api.Test import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat @@ -35,6 +37,9 @@ import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState import java.time.Instant class SabrPlaybackGranularRoutesTest { + @AfterEach + fun clearDemands(): Unit = SabrSegmentDemandTracker.clearAll() + @Test fun `position updates player time without prefetching`() = testApplication { val store = mockk(relaxed = true) @@ -108,7 +113,7 @@ class SabrPlaybackGranularRoutesTest { assertTrue(response.bodyAsText().contains("segment/1")) verify(exactly = 1) { store.startPump(holder) } verify(exactly = 0) { store.warmPlaybackAsync(holder) } - verify(exactly = 1) { store.requestSegmentDemand(holder, any(), holder.activeGeneration()) } + verify(exactly = 2) { store.requestSegmentDemand(holder, any(), holder.activeGeneration()) } } @Test @@ -129,7 +134,7 @@ class SabrPlaybackGranularRoutesTest { } @Test - fun `window queues first blocker without direct fetch after nonzero start`() = testApplication { + fun `window queues video first when blocked tracks have equal coverage`() = testApplication { val store = emptyStore() val holder = holder() every { store.lookupByToken("session-token") } returns holder @@ -142,8 +147,8 @@ class SabrPlaybackGranularRoutesTest { val body = response.bodyAsText() assertEquals(HttpStatusCode.Accepted, response.status) - assertTrue(body.contains("audio:140:1 pending")) - verify(exactly = 1) { store.requestSegmentDemand(holder, any(), holder.activeGeneration()) } + assertTrue(body.contains("video:136:1 pending")) + verify(exactly = 2) { store.requestSegmentDemand(holder, any(), holder.activeGeneration()) } coVerify(exactly = 0) { store.fetchInitializationData(holder, any()) } } @@ -225,6 +230,7 @@ class SabrPlaybackGranularRoutesTest { val state = mockk() every { session.streamState } returns state every { session.getCachedSegment(any()) } returns null + every { session.isBeyondEnd(any()) } returns false every { state.setActiveTrackTypes(true, true) } returns Unit every { state.getSegmentNumberAtOrAfterTimeMs(any(), any()) } returns 1 every { state.getEndSegment(any()) } returns 0L @@ -256,7 +262,6 @@ class SabrPlaybackGranularRoutesTest { every { format.bitrate } returns if (isAudio) 128_000 else 2_000_000 every { format.height } returns if (isAudio) 0 else 720 every { format.approxDurationMs } returns 900_000L - every { format.initializationUrl } returns null return format } diff --git a/src/test/kotlin/dev/typetype/server/SabrStreamContractFilterTest.kt b/src/test/kotlin/dev/typetype/server/SabrStreamContractFilterTest.kt index 6244cb9b..1a502840 100644 --- a/src/test/kotlin/dev/typetype/server/SabrStreamContractFilterTest.kt +++ b/src/test/kotlin/dev/typetype/server/SabrStreamContractFilterTest.kt @@ -122,8 +122,6 @@ class SabrStreamContractFilterTest { every { format.height } returns if (isAudio) 0 else 1080 every { format.qualityLabel } returns if (isAudio) null else "1080p" every { format.contentLength } returns 1_000_000L - every { format.initRangeStart } returns 0L - every { format.initRangeEnd } returns 0L return format } diff --git a/src/test/kotlin/dev/typetype/server/StreamExtractionErrorMapperTest.kt b/src/test/kotlin/dev/typetype/server/StreamExtractionErrorMapperTest.kt index 5ee5977b..b88fa204 100644 --- a/src/test/kotlin/dev/typetype/server/StreamExtractionErrorMapperTest.kt +++ b/src/test/kotlin/dev/typetype/server/StreamExtractionErrorMapperTest.kt @@ -8,6 +8,7 @@ import org.junit.jupiter.api.Test import org.schabi.newpipe.extractor.exceptions.NeedLoginException import org.schabi.newpipe.extractor.exceptions.PaidContentException import org.schabi.newpipe.extractor.exceptions.PrivateContentException +import org.schabi.newpipe.extractor.exceptions.VideoNotReleaseException import org.schabi.newpipe.extractor.exceptions.YoutubeMusicPremiumContentException class StreamExtractionErrorMapperTest { @@ -15,16 +16,34 @@ class StreamExtractionErrorMapperTest { fun `maps membership restrictions using extractor messages when available`() { val login = StreamExtractionErrorMapper.map(NeedLoginException("This video is only available for members")) val paid = StreamExtractionErrorMapper.map(PaidContentException("This video is only available for members")) - assertEquals(ExtractionResult.BadRequest("This video is only available for members"), login) - assertEquals(ExtractionResult.BadRequest("This video is only available for members"), paid) + assertEquals(ExtractionResult.BadRequest("This video is only available for members", "members_only"), login) + assertEquals(ExtractionResult.BadRequest("This video is only available for members", "members_only"), paid) } @Test fun `maps membership restrictions to fallback when extractor message is blank`() { val login = StreamExtractionErrorMapper.map(NeedLoginException("")) val paid = StreamExtractionErrorMapper.map(PaidContentException("")) - assertEquals(ExtractionResult.BadRequest(StreamExtractionErrorMapper.MEMBERS_ONLY_FALLBACK), login) - assertEquals(ExtractionResult.BadRequest(StreamExtractionErrorMapper.MEMBERS_ONLY_FALLBACK), paid) + assertEquals( + ExtractionResult.BadRequest(StreamExtractionErrorMapper.MEMBERS_ONLY_FALLBACK, "members_only"), + login, + ) + assertEquals( + ExtractionResult.BadRequest(StreamExtractionErrorMapper.PAID_CONTENT_FALLBACK, "paid_content"), + paid, + ) + } + + @Test + fun `keeps paid videos distinct from members-only videos`() { + val result = StreamExtractionErrorMapper.map(PaidContentException("This video is a paid video")) + assertEquals(ExtractionResult.BadRequest("This video is a paid video", "paid_content"), result) + } + + @Test + fun `maps upcoming premieres to a stable availability code`() { + val result = StreamExtractionErrorMapper.map(VideoNotReleaseException("Premieres in 200 days")) + assertEquals(ExtractionResult.Failure("Premieres in 200 days", "scheduled_premiere"), result) } @Test @@ -35,9 +54,12 @@ class StreamExtractionErrorMapperTest { } @Test - fun `maps youtube music premium exception to member-only style message`() { + fun `maps youtube music premium exception to paid content`() { val mapped = StreamExtractionErrorMapper.map(YoutubeMusicPremiumContentException()) - assertTrue(mapped is ExtractionResult.BadRequest) + assertEquals( + ExtractionResult.BadRequest("This video is a YouTube Music Premium video", "paid_content"), + mapped, + ) } @Test diff --git a/src/test/kotlin/dev/typetype/server/StreamRoutesDeliveryModeTest.kt b/src/test/kotlin/dev/typetype/server/StreamRoutesDeliveryModeTest.kt index e2e1de5f..179c8109 100644 --- a/src/test/kotlin/dev/typetype/server/StreamRoutesDeliveryModeTest.kt +++ b/src/test/kotlin/dev/typetype/server/StreamRoutesDeliveryModeTest.kt @@ -136,8 +136,8 @@ class StreamRoutesDeliveryModeTest { } @Test - fun `sabr endpoint preserves hls for live streams`() = testApplication { - val live = testStreamResponse(videoOnlyStreams = emptyList(), audioStreams = emptyList()).copy( + fun `sabr endpoint removes hls for live streams`() = testApplication { + val live = sabrResponse().copy( hlsUrl = "/streams/hls-manifest?url=live", isLive = true, isLiveContent = true, @@ -149,7 +149,8 @@ class StreamRoutesDeliveryModeTest { val response = client.get("/streams/youtube/sabr?url=$VIDEO_URL") assertEquals(HttpStatusCode.OK, response.status) - assertTrue(response.bodyAsText().contains("\"hlsUrl\":\"/streams/hls-manifest?url=live\"")) + assertTrue(response.bodyAsText().contains("\"hlsUrl\":\"\"")) + assertFalse(response.bodyAsText().contains("hls-manifest")) } @Test diff --git a/src/test/kotlin/dev/typetype/server/StreamRoutesSignedHlsTest.kt b/src/test/kotlin/dev/typetype/server/StreamRoutesSignedHlsTest.kt index e769e90c..b3d97e3a 100644 --- a/src/test/kotlin/dev/typetype/server/StreamRoutesSignedHlsTest.kt +++ b/src/test/kotlin/dev/typetype/server/StreamRoutesSignedHlsTest.kt @@ -63,9 +63,19 @@ class StreamRoutesSignedHlsTest { } @Test - fun `anonymous sabr signs live hls url`() = testApplication { + fun `anonymous sabr removes live hls url`() = testApplication { coEvery { streamService.getStreamInfo(any()) } returns ExtractionResult.Success( - publicHlsStream().copy(isLive = true, isLiveContent = true, hasLiveManifest = true), + publicHlsStream().copy( + videoOnlyStreams = listOf( + testVideoStream().copy(url = "", deliveryMethod = "sabr", sabrSessionUrl = "/sabr/session/test"), + ), + audioStreams = listOf( + testAudioStream(url = "", deliveryMethod = "sabr", sabrSessionUrl = "/sabr/session/test"), + ), + isLive = true, + isLiveContent = true, + hasLiveManifest = true, + ), ) installApp() @@ -73,10 +83,9 @@ class StreamRoutesSignedHlsTest { val body = response.bodyAsText() assertEquals(HttpStatusCode.OK, response.status) - assertTrue(body.contains("\"hlsUrl\":\"/streams/hls-manifest?token=")) + assertTrue(body.contains("\"hlsUrl\":\"\"")) + assertFalse(body.contains("hls-manifest")) assertFalse(body.contains(MANIFEST_URL)) - val token = body.substringAfter("/streams/hls-manifest?token=").substringBefore('"') - assertTrue(tokenService.verify(token) is PublicHlsManifestTokenResult.Valid) } @Test diff --git a/src/test/kotlin/dev/typetype/server/StreamRoutesTest.kt b/src/test/kotlin/dev/typetype/server/StreamRoutesTest.kt index 00135a9f..411da2a0 100644 --- a/src/test/kotlin/dev/typetype/server/StreamRoutesTest.kt +++ b/src/test/kotlin/dev/typetype/server/StreamRoutesTest.kt @@ -48,25 +48,27 @@ class StreamRoutesTest { @Test fun `GET streams returns 422 on Failure`() = testApplication { coEvery { streamService.getStreamInfo(any()) } returns - ExtractionResult.Failure("Extraction failed") + ExtractionResult.Failure("Premieres in 200 days", "scheduled_premiere") application { install(ContentNegotiation) { json() } routing { streamRoutes(streamService) } } val response = client.get("/streams?url=https://youtube.com/watch?v=bad") assertEquals(HttpStatusCode.UnprocessableEntity, response.status) + assertTrue(response.bodyAsText().contains("\"code\":\"scheduled_premiere\"")) } @Test fun `GET streams returns 400 on BadRequest`() = testApplication { coEvery { streamService.getStreamInfo(any()) } returns - ExtractionResult.BadRequest("Unsupported URL") + ExtractionResult.BadRequest("This video is a paid video", "paid_content") application { install(ContentNegotiation) { json() } routing { streamRoutes(streamService) } } val response = client.get("/streams?url=https://unsupported.com/video") assertEquals(HttpStatusCode.BadRequest, response.status) + assertTrue(response.bodyAsText().contains("\"code\":\"paid_content\"")) } @Test diff --git a/src/test/kotlin/dev/typetype/server/SubscriptionFeedRoutesTest.kt b/src/test/kotlin/dev/typetype/server/SubscriptionFeedRoutesTest.kt index 877ac24d..44bf6238 100644 --- a/src/test/kotlin/dev/typetype/server/SubscriptionFeedRoutesTest.kt +++ b/src/test/kotlin/dev/typetype/server/SubscriptionFeedRoutesTest.kt @@ -2,6 +2,7 @@ package dev.typetype.server import dev.typetype.server.cache.CacheService import dev.typetype.server.models.ExtractionResult +import dev.typetype.server.models.SubscriptionFeedResponse import dev.typetype.server.SubscriptionFeedTestFixtures.channel import dev.typetype.server.SubscriptionFeedTestFixtures.subscription import dev.typetype.server.SubscriptionFeedTestFixtures.video @@ -22,7 +23,9 @@ import io.ktor.server.routing.routing import io.ktor.server.testing.ApplicationTestBuilder import io.ktor.server.testing.testApplication import io.mockk.coEvery +import io.mockk.coVerify import io.mockk.mockk +import kotlinx.serialization.json.Json import org.junit.jupiter.api.Assertions.assertEquals import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.BeforeAll @@ -77,6 +80,63 @@ class SubscriptionFeedRoutesTest { assertTrue(body.indexOf("2000") < body.indexOf("1000")) } + @Test + fun `GET subscriptions feed merges livestream tab and promotes active live`() = withApp { + val channelUrl = "https://www.youtube.com/channel/UC1" + val liveUrl = "https://www.youtube.com/watch?v=live" + subscriptionsService.add(TEST_USER_ID, subscription(channelUrl, "Live channel")) + coEvery { channelService.getChannel(channelUrl, null) } returns channel( + video(3000L, url = "https://www.youtube.com/watch?v=normal"), + video(2000L, url = liveUrl), + ) + coEvery { channelService.getChannel("$channelUrl/streams", null) } returns channel( + video(-1L, url = liveUrl, live = true), + video(5000L, url = "https://www.youtube.com/watch?v=upcoming"), + ) + + val body = client.get("/subscriptions/feed") { + headers.append(HttpHeaders.Authorization, "Bearer test-jwt") + }.bodyAsText() + val feed = Json.decodeFromString(body) + + assertEquals(liveUrl, feed.videos.first().url) + assertTrue(feed.videos.first().isLive) + assertEquals(1, feed.videos.count { it.url == liveUrl }) + assertTrue(feed.videos.any { it.url.endsWith("v=upcoming") }) + coVerify { cacheService.set(any(), any(), 60L) } + } + + @Test + fun `GET subscriptions feed keeps channel videos when livestream tab fails`() = withApp { + val channelUrl = "https://www.youtube.com/channel/UC2" + subscriptionsService.add(TEST_USER_ID, subscription(channelUrl, "Channel")) + coEvery { channelService.getChannel(channelUrl, null) } returns channel(video(1000L)) + coEvery { channelService.getChannel("$channelUrl/streams", null) } returns ExtractionResult.Failure("err") + + val body = client.get("/subscriptions/feed") { + headers.append(HttpHeaders.Authorization, "Bearer test-jwt") + }.bodyAsText() + + assertTrue(body.contains("V1000")) + } + + @Test + fun `GET subscriptions feed replaces existing channel tab with livestream tab`() = withApp { + val channelUrl = "https://www.youtube.com/channel/UC3/videos" + val livestreamUrl = "https://www.youtube.com/channel/UC3/streams" + subscriptionsService.add(TEST_USER_ID, subscription(channelUrl, "Imported channel")) + coEvery { channelService.getChannel(channelUrl, null) } returns channel(video(1000L)) + coEvery { channelService.getChannel(livestreamUrl, null) } returns channel( + video(-1L, url = "https://www.youtube.com/watch?v=live", live = true), + ) + + val body = client.get("/subscriptions/feed") { + headers.append(HttpHeaders.Authorization, "Bearer test-jwt") + }.bodyAsText() + + assertTrue(body.contains("v=live")) + } + @Test fun `GET subscriptions feed pagination works`() = withApp { subscriptionsService.add(TEST_USER_ID, subscription(1)) diff --git a/src/test/kotlin/dev/typetype/server/SubscriptionFeedTestFixtures.kt b/src/test/kotlin/dev/typetype/server/SubscriptionFeedTestFixtures.kt index d1eb80dd..a44151fb 100644 --- a/src/test/kotlin/dev/typetype/server/SubscriptionFeedTestFixtures.kt +++ b/src/test/kotlin/dev/typetype/server/SubscriptionFeedTestFixtures.kt @@ -11,6 +11,7 @@ internal object SubscriptionFeedTestFixtures { channel: String = "Ch", url: String = "u/$uploaded", short: Boolean = false, + live: Boolean = false, ): VideoItem = VideoItem( id = "id-$uploaded-$channel", title = "V$uploaded", @@ -23,10 +24,12 @@ internal object SubscriptionFeedTestFixtures { viewCount = 0L, uploadDate = "", uploaded = uploaded, - streamType = "video_stream", + streamType = if (live) "live_stream" else "video_stream", isShortFormContent = short, uploaderVerified = false, shortDescription = null, + isLive = live, + isLiveContent = live, ) fun channel(vararg videos: VideoItem): ExtractionResult = ExtractionResult.Success( diff --git a/src/test/kotlin/dev/typetype/server/UserVideoMetadataRepairRoutesTest.kt b/src/test/kotlin/dev/typetype/server/UserVideoMetadataRepairRoutesTest.kt index ffae1b3d..2af7e2b3 100644 --- a/src/test/kotlin/dev/typetype/server/UserVideoMetadataRepairRoutesTest.kt +++ b/src/test/kotlin/dev/typetype/server/UserVideoMetadataRepairRoutesTest.kt @@ -24,17 +24,27 @@ import io.ktor.server.application.install import io.ktor.server.plugins.contentnegotiation.ContentNegotiation import io.ktor.server.routing.routing import io.ktor.server.testing.testApplication +import kotlinx.coroutines.CompletableDeferred +import kotlinx.coroutines.CoroutineScope +import kotlinx.coroutines.SupervisorJob +import kotlinx.coroutines.cancel +import kotlinx.coroutines.delay +import kotlinx.coroutines.job +import kotlinx.coroutines.runBlocking +import kotlinx.coroutines.withTimeout +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertFalse import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.BeforeAll import org.junit.jupiter.api.BeforeEach import org.junit.jupiter.api.Test +import java.util.concurrent.atomic.AtomicInteger class UserVideoMetadataRepairRoutesTest { private val playlists = PlaylistService() private val watchLater = WatchLaterService() private val favorites = FavoritesService() private val auth = AuthService.fixed(TEST_USER_ID) - private val repair = UserVideoMetadataRepairService(VideoMetadataResolver(fakeStreamService())) companion object { private const val VIDEO_URL = "https://www.youtube.com/watch?v=abc123" @@ -48,7 +58,10 @@ class UserVideoMetadataRepairRoutesTest { fun clean() = TestDatabase.truncateAll() @Test - fun `list routes repair fallback video metadata`() = testApplication { + fun `collection routes return before fallback metadata repair completes`() = testApplication { + val repairStarted = CompletableDeferred() + val releaseRepair = CompletableDeferred() + val repair = UserVideoMetadataRepairService(VideoMetadataResolver(fakeStreamService(repairStarted, releaseRepair))) val playlist = playlists.create(TEST_USER_ID, PlaylistItem(name = "Imported", description = "")) playlists.addVideo(TEST_USER_ID, playlist.id, fallbackVideo(VIDEO_URL)) watchLater.add(TEST_USER_ID, WatchLaterItem(url = VIDEO_URL, title = "YouTube video abc123", thumbnail = "https://i.ytimg.com/vi/abc123/hqdefault.jpg", duration = 0L)) @@ -62,18 +75,73 @@ class UserVideoMetadataRepairRoutesTest { } } - val playlistBody = client.get("/playlists/${playlist.id}") { header(HttpHeaders.Authorization, "Bearer test-jwt") }.bodyAsText() - val watchLaterBody = client.get("/watch-later") { header(HttpHeaders.Authorization, "Bearer test-jwt") }.bodyAsText() - val favoritesBody = client.get("/favorites") { header(HttpHeaders.Authorization, "Bearer test-jwt") }.bodyAsText() - - assertTrue(playlistBody.contains("Resolved abc123")) - assertTrue(watchLaterBody.contains("Resolved abc123")) - assertTrue(favoritesBody.contains("Resolved abc123")) - assertTrue(watchLaterBody.contains("\"channelName\":\"Channel\"")) - assertTrue(watchLaterBody.contains("\"viewCount\":10")) - assertTrue(favoritesBody.contains("\"thumbnail\":\"https://thumb.test/abc123.jpg\"")) - assertTrue(favoritesBody.contains("\"viewCount\":10")) - assertTrue(favoritesBody.contains("\"publishedAt\":1")) + val playlistsBody = client.get("/playlists") { header(HttpHeaders.Authorization, "Bearer test-jwt") }.bodyAsText() + assertTrue(playlistsBody.contains("\"videoCount\":1")) + assertFalse(repairStarted.isCompleted) + + val (playlistBody, watchLaterBody, favoritesBody) = withTimeout(1_000) { + listOf( + client.get("/playlists/${playlist.id}") { header(HttpHeaders.Authorization, "Bearer test-jwt") }.bodyAsText(), + client.get("/watch-later") { header(HttpHeaders.Authorization, "Bearer test-jwt") }.bodyAsText(), + client.get("/favorites") { header(HttpHeaders.Authorization, "Bearer test-jwt") }.bodyAsText(), + ) + } + + assertTrue(playlistBody.contains("YouTube video abc123")) + assertTrue(watchLaterBody.contains("YouTube video abc123")) + assertFalse(playlistBody.contains("Resolved abc123")) + assertFalse(watchLaterBody.contains("Resolved abc123")) + assertFalse(favoritesBody.contains("Resolved abc123")) + + withTimeout(5_000) { repairStarted.await() } + releaseRepair.complete(Unit) + withTimeout(5_000) { + while ( + playlists.getById(TEST_USER_ID, playlist.id)?.videos?.single()?.title != "Resolved abc123" || + watchLater.getAll(TEST_USER_ID).single().title != "Resolved abc123" || + favorites.getAll(TEST_USER_ID).single().title != "Resolved abc123" + ) { + delay(10) + } + } + } + + @Test + fun `repeated requests coalesce failed metadata repairs during cooldown`() = runBlocking { + val attempts = AtomicInteger() + val repairStarted = CompletableDeferred() + val releaseRepair = CompletableDeferred() + val streamService = object : StreamService { + override suspend fun getStreamInfo(url: String): ExtractionResult { + attempts.incrementAndGet() + repairStarted.complete(Unit) + releaseRepair.await() + return ExtractionResult.Failure("Unavailable") + } + } + val repair = UserVideoMetadataRepairService(VideoMetadataResolver(streamService)) + val playlist = playlists.create(TEST_USER_ID, PlaylistItem(name = "Imported", description = "")) + playlists.addVideo(TEST_USER_ID, playlist.id, fallbackVideo(VIDEO_URL)) + val scope = CoroutineScope(SupervisorJob()) + + try { + repair.schedulePlaylists(scope, TEST_USER_ID) + repair.schedulePlaylists(scope, TEST_USER_ID) + + withTimeout(5_000) { repairStarted.await() } + assertEquals(1, attempts.get()) + releaseRepair.complete(Unit) + withTimeout(5_000) { + while (scope.coroutineContext.job.children.any()) delay(10) + } + + repair.schedulePlaylists(scope, TEST_USER_ID) + delay(100) + + assertEquals(1, attempts.get()) + } finally { + scope.cancel() + } } private fun fallbackVideo(url: String): PlaylistVideoItem = PlaylistVideoItem( @@ -83,8 +151,12 @@ class UserVideoMetadataRepairRoutesTest { duration = 0L, ) - private fun fakeStreamService(): StreamService = object : StreamService { - override suspend fun getStreamInfo(url: String): ExtractionResult = ExtractionResult.Success(stream(url)) + private fun fakeStreamService(started: CompletableDeferred, release: CompletableDeferred): StreamService = object : StreamService { + override suspend fun getStreamInfo(url: String): ExtractionResult { + started.complete(Unit) + release.await() + return ExtractionResult.Success(stream(url)) + } } private fun stream(url: String): StreamResponse = StreamResponse( diff --git a/src/test/kotlin/dev/typetype/server/YoutubeProxySelectorTest.kt b/src/test/kotlin/dev/typetype/server/YoutubeProxySelectorTest.kt new file mode 100644 index 00000000..e41bd942 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/YoutubeProxySelectorTest.kt @@ -0,0 +1,69 @@ +package dev.typetype.server + +import dev.typetype.server.downloader.YoutubeProxySelector +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Assertions.assertThrows +import org.junit.jupiter.api.Test +import java.net.InetSocketAddress +import java.net.Proxy +import java.net.URI + +class YoutubeProxySelectorTest { + @Test + fun `routes youtube extraction and media hosts through proxy`() { + val selector = YoutubeProxySelector.fromUrl("http://proxy.internal:8080")!! + + listOf( + "https://www.youtube.com/youtubei/v1/player", + "https://youtu.be/video-id", + "https://youtubei.googleapis.com/youtubei/v1/player", + "https://rr1---sn.example.googlevideo.com/videoplayback", + "https://i.ytimg.com/vi/id/hqdefault.jpg", + "https://yt3.googleusercontent.com/avatar", + ).forEach { url -> + val proxies = selector.select(URI(url)) + val proxy = proxies.first() + assertEquals(Proxy.Type.HTTP, proxy.type()) + assertEquals(InetSocketAddress.createUnresolved("proxy.internal", 8080), proxy.address()) + assertEquals(Proxy.NO_PROXY, proxies.last()) + } + } + + @Test + fun `keeps unrelated and deceptive hosts direct`() { + val selector = YoutubeProxySelector.fromUrl("http://proxy.internal:8080")!! + + listOf( + "http://typetype-token:8081/youtube/sabr/session", + "https://www.nicovideo.jp/watch/id", + "https://evilgooglevideo.com/video", + ).forEach { url -> + assertEquals(Proxy.NO_PROXY, selector.select(URI(url)).single()) + } + } + + @Test + fun `accepts an absent proxy and defaults the http port`() { + assertNull(YoutubeProxySelector.fromUrl(null)) + assertNull(YoutubeProxySelector.fromUrl(" ")) + + val selector = YoutubeProxySelector.fromUrl("http://proxy.internal")!! + val address = selector.select(URI("https://www.youtube.com")).first().address() + assertEquals(InetSocketAddress.createUnresolved("proxy.internal", 80), address) + } + + @Test + fun `rejects unsupported or ambiguous proxy urls`() { + listOf( + "https://proxy.internal:8080", + "http:///missing-host", + "http://user:pass@proxy.internal:8080", + "http://proxy.internal:8080/path", + ).forEach { value -> + assertThrows(IllegalArgumentException::class.java) { + YoutubeProxySelector.fromUrl(value) + } + } + } +} diff --git a/src/test/kotlin/dev/typetype/server/YoutubeTakeoutImportRoutesTest.kt b/src/test/kotlin/dev/typetype/server/YoutubeTakeoutImportRoutesTest.kt index d528e11f..a3fa6052 100644 --- a/src/test/kotlin/dev/typetype/server/YoutubeTakeoutImportRoutesTest.kt +++ b/src/test/kotlin/dev/typetype/server/YoutubeTakeoutImportRoutesTest.kt @@ -1,15 +1,11 @@ package dev.typetype.server import dev.typetype.server.routes.youtubeTakeoutImportRoutes -import dev.typetype.server.models.ExtractionResult -import dev.typetype.server.models.StreamResponse import dev.typetype.server.services.AuthService import dev.typetype.server.services.FavoritesService import dev.typetype.server.services.HistoryService import dev.typetype.server.services.PlaylistService -import dev.typetype.server.services.StreamService import dev.typetype.server.services.SubscriptionsService -import dev.typetype.server.services.VideoMetadataResolver import dev.typetype.server.services.WatchLaterService import dev.typetype.server.services.YoutubeTakeoutImportJobService import dev.typetype.server.services.YoutubeTakeoutImporterService @@ -106,9 +102,11 @@ class YoutubeTakeoutImportRoutesTest { val zip = createTakeoutZip() val upload = uploadArchive(zip) val jobId = Json.parseToJsonElement(upload.bodyAsText()).jsonObject["jobId"]!!.jsonPrimitive.content + val preview = client.get("/imports/youtube-takeout/$jobId/preview") { header(HttpHeaders.Authorization, "Bearer test-jwt") } val commit = client.post("/imports/youtube-takeout/$jobId/commit") { header(HttpHeaders.Authorization, "Bearer test-jwt") } + assertEquals(HttpStatusCode.OK, preview.status) assertEquals(HttpStatusCode.Accepted, commit.status) val body = Json.parseToJsonElement(commit.bodyAsText()).jsonObject assertEquals("running", body["status"]!!.jsonPrimitive.content) @@ -116,33 +114,6 @@ class YoutubeTakeoutImportRoutesTest { Files.deleteIfExists(zip) } - @Test - fun `metadata extraction failures do not block import completion`() = testApplication { - val service = YoutubeTakeoutImportJobService( - YoutubeTakeoutParserService(), - YoutubeTakeoutPreviewService(subscriptions, playlists, previewLookup), - YoutubeTakeoutImporterService( - subscriptions, - playlists, - signalImport, - metadataResolver = VideoMetadataResolver(ThrowingStreamService()), - ), - ) - application { - install(ContentNegotiation) { json() } - routing { youtubeTakeoutImportRoutes(service, auth) } - } - val zip = createTakeoutZip() - val upload = uploadArchive(zip) - val jobId = Json.parseToJsonElement(upload.bodyAsText()).jsonObject["jobId"]!!.jsonPrimitive.content - - val commit = client.post("/imports/youtube-takeout/$jobId/commit") { header(HttpHeaders.Authorization, "Bearer test-jwt") } - - assertEquals(HttpStatusCode.Accepted, commit.status) - assertEventuallyCompleted(jobId) - Files.deleteIfExists(zip) - } - @Test fun `preview failure leaves terminal failed status`() = testApplication { application { @@ -183,11 +154,6 @@ class YoutubeTakeoutImportRoutesTest { })) } - private class ThrowingStreamService : StreamService { - override suspend fun getStreamInfo(url: String): ExtractionResult = - throw IllegalStateException("Sign in to confirm you're not a bot") - } - private fun createTakeoutZip(): Path { val zip = Files.createTempFile("yt-takeout-", ".zip") ZipOutputStream(Files.newOutputStream(zip)).use { out -> diff --git a/src/test/kotlin/dev/typetype/server/YoutubeTakeoutImporterMetadataTest.kt b/src/test/kotlin/dev/typetype/server/YoutubeTakeoutImporterMetadataTest.kt index 7a1bcd20..49ed624e 100644 --- a/src/test/kotlin/dev/typetype/server/YoutubeTakeoutImporterMetadataTest.kt +++ b/src/test/kotlin/dev/typetype/server/YoutubeTakeoutImporterMetadataTest.kt @@ -1,18 +1,14 @@ package dev.typetype.server -import dev.typetype.server.models.ExtractionResult import dev.typetype.server.models.FavoriteItem import dev.typetype.server.models.PlaylistItem import dev.typetype.server.models.PlaylistVideoItem -import dev.typetype.server.models.StreamResponse import dev.typetype.server.models.YoutubeTakeoutCommitPlan import dev.typetype.server.models.YoutubeTakeoutParsedData import dev.typetype.server.services.FavoritesService import dev.typetype.server.services.HistoryService import dev.typetype.server.services.PlaylistService -import dev.typetype.server.services.StreamService import dev.typetype.server.services.SubscriptionsService -import dev.typetype.server.services.VideoMetadataResolver import dev.typetype.server.services.WatchLaterService import dev.typetype.server.services.YoutubeTakeoutImporterService import dev.typetype.server.services.YoutubeTakeoutSignalImportService @@ -28,8 +24,11 @@ class YoutubeTakeoutImporterMetadataTest { private val favorites = FavoritesService() private val watchLater = WatchLaterService() private val history = HistoryService() - private val signalImport = YoutubeTakeoutSignalImportService(favorites, watchLater, history) - private val importer = YoutubeTakeoutImporterService(subscriptions, playlists, signalImport, metadataResolver = VideoMetadataResolver(fakeStreamService())) + private val importer = YoutubeTakeoutImporterService( + subscriptions, + playlists, + YoutubeTakeoutSignalImportService(favorites, watchLater, history), + ) companion object { private const val VIDEO_URL = "https://www.youtube.com/watch?v=abc123" @@ -43,46 +42,30 @@ class YoutubeTakeoutImporterMetadataTest { fun clean() = TestDatabase.truncateAll() @Test - fun `commit enriches takeout list videos and favorites before persisting`() = runBlocking { + fun `commit persists parsed takeout metadata without stream extraction`() = runBlocking { importer.commit(TEST_USER_ID, parsed(), YoutubeTakeoutCommitPlan(true, true, true, true, true, false)) val playlistId = playlists.getAll(TEST_USER_ID).single().id - val playlistVideo = playlists.getById(TEST_USER_ID, playlistId)?.videos?.single() - val watchLaterVideo = watchLater.getAll(TEST_USER_ID).single() - val favorite = favorites.getAll(TEST_USER_ID).single() - assertEquals("Resolved abc123", playlistVideo?.title) - assertEquals("Resolved abc123", watchLaterVideo.title) - assertEquals("Resolved abc123", favorite.title) - assertEquals("https://thumb.test/abc123.jpg", favorite.thumbnail) + assertEquals("Takeout title", playlists.getById(TEST_USER_ID, playlistId)?.videos?.single()?.title) + assertEquals("Takeout title", watchLater.getAll(TEST_USER_ID).single().title) + assertEquals("YouTube video abc123", favorites.getAll(TEST_USER_ID).single().title) } - private fun parsed(): YoutubeTakeoutParsedData = YoutubeTakeoutParsedData( + private fun parsed() = YoutubeTakeoutParsedData( subscriptions = emptyList(), playlists = listOf(PlaylistItem(id = "PL1", name = "Imported")), - playlistItems = mapOf("PL1" to listOf(fallbackVideo())), + playlistItems = mapOf("PL1" to listOf(video())), favorites = listOf(FavoriteItem(videoUrl = VIDEO_URL)), - watchLater = listOf(fallbackVideo()), + watchLater = listOf(video()), history = emptyList(), warnings = emptyList(), errors = emptyList(), ) - private fun fallbackVideo(): PlaylistVideoItem = PlaylistVideoItem( + private fun video() = PlaylistVideoItem( url = VIDEO_URL, - title = "YouTube video abc123", + title = "Takeout title", thumbnail = "https://i.ytimg.com/vi/abc123/hqdefault.jpg", duration = 0L, ) - - private fun fakeStreamService(): StreamService = object : StreamService { - override suspend fun getStreamInfo(url: String): ExtractionResult = ExtractionResult.Success(stream()) - } - - private fun stream(): StreamResponse = StreamResponse( - id = "abc123", title = "Resolved abc123", uploaderName = "Channel", uploaderUrl = "https://www.youtube.com/channel/UC1", uploaderAvatarUrl = "https://avatar.test/uc1.jpg", - thumbnailUrl = "https://thumb.test/abc123.jpg", description = "", duration = 120L, viewCount = 10L, likeCount = 0L, dislikeCount = 0L, uploadDate = "", uploaded = 1L, - uploaderSubscriberCount = 0L, uploaderVerified = false, category = "", license = "", visibility = "", tags = emptyList(), streamType = "video", isShortFormContent = false, - requiresMembership = false, startPosition = 0L, streamSegments = emptyList(), hlsUrl = "", dashMpdUrl = "", videoStreams = emptyList(), audioStreams = emptyList(), - originalAudioTrackId = null, preferredDefaultAudioTrackId = null, videoOnlyStreams = emptyList(), subtitles = emptyList(), previewFrames = emptyList(), sponsorBlockSegments = emptyList(), relatedStreams = emptyList(), - ) } diff --git a/src/test/kotlin/dev/typetype/server/YoutubeTakeoutParserServiceTest.kt b/src/test/kotlin/dev/typetype/server/YoutubeTakeoutParserServiceTest.kt index a89429d9..a2644aff 100644 --- a/src/test/kotlin/dev/typetype/server/YoutubeTakeoutParserServiceTest.kt +++ b/src/test/kotlin/dev/typetype/server/YoutubeTakeoutParserServiceTest.kt @@ -70,6 +70,38 @@ class YoutubeTakeoutParserServiceTest { Files.deleteIfExists(zip) } + @Test + fun `parse preserves takeout playlist order and added dates`() { + val zip = Files.createTempFile("yt-takeout-playlist-order-", ".zip") + ZipOutputStream(Files.newOutputStream(zip)).use { out -> + out.putNextEntry(ZipEntry("Takeout/YouTube et YouTube Music/playlists/playlists.csv")) + out.write("ID de la playlist,Titre (d'origine) de la playlist\nPL123456,Imported\n".toByteArray()) + out.closeEntry() + out.putNextEntry(ZipEntry("Takeout/YouTube et YouTube Music/playlists/Imported.csv")) + out.write( + """ + ID vidéo,Code temporel de création de la vidéo de la playlist + newer000001,2026-01-02T00:00:00+00:00 + older000001,2026-01-01T00:00:00+00:00 + """.trimIndent().plus("\n").toByteArray(), + ) + out.closeEntry() + } + + val parsed = YoutubeTakeoutParserService().parse(zip) + val imported = parsed.playlistItems.values.single() + + assertEquals( + listOf( + "https://www.youtube.com/watch?v=newer000001", + "https://www.youtube.com/watch?v=older000001", + ), + imported.map { it.url }, + ) + assertEquals(listOf(1_767_312_000_000L, 1_767_225_600_000L), imported.map { it.addedAt }) + Files.deleteIfExists(zip) + } + private fun createZip(): Path { val zip = Files.createTempFile("yt-takeout-parser-", ".zip") ZipOutputStream(Files.newOutputStream(zip)).use { out -> diff --git a/src/test/kotlin/dev/typetype/server/YoutubeTakeoutSystemPlaylistImportTest.kt b/src/test/kotlin/dev/typetype/server/YoutubeTakeoutSystemPlaylistImportTest.kt index 15849050..1d1ecfcf 100644 --- a/src/test/kotlin/dev/typetype/server/YoutubeTakeoutSystemPlaylistImportTest.kt +++ b/src/test/kotlin/dev/typetype/server/YoutubeTakeoutSystemPlaylistImportTest.kt @@ -45,6 +45,29 @@ class YoutubeTakeoutSystemPlaylistImportTest { assertEquals(1_700_000_100_000L, favorites.getAll(TEST_USER_ID).single().favoritedAt) } + @Test + fun `commit preserves imported playlist order and added dates`() = runBlocking { + val newer = video("newer", 1_700_000_100_000L) + val older = video("older", 1_700_000_000_000L) + val parsed = YoutubeTakeoutParsedData( + subscriptions = emptyList(), + playlists = listOf(PlaylistItem(id = "PL123456", name = "Imported")), + playlistItems = mapOf("PL123456" to listOf(newer, older)), + favorites = emptyList(), + watchLater = emptyList(), + history = emptyList(), + warnings = emptyList(), + errors = emptyList(), + ) + + importer.commit(TEST_USER_ID, parsed, YoutubeTakeoutCommitPlan(false, true, true, false, false, false)) + + val playlistId = playlists.getAll(TEST_USER_ID).single().id + val imported = playlists.getById(TEST_USER_ID, playlistId)?.videos.orEmpty() + assertEquals(listOf(newer.url, older.url), imported.map { it.url }) + assertEquals(listOf(newer.addedAt, older.addedAt), imported.map { it.addedAt }) + } + private fun parsed(): YoutubeTakeoutParsedData = YoutubeTakeoutParsedData( subscriptions = emptyList(), playlists = listOf(PlaylistItem(id = "WL", name = "Watch later"), PlaylistItem(id = "LL", name = "Liked videos")), diff --git a/src/test/kotlin/dev/typetype/server/routes/SabrLiveGapWindowTest.kt b/src/test/kotlin/dev/typetype/server/routes/SabrLiveGapWindowTest.kt new file mode 100644 index 00000000..ee962c7f --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/routes/SabrLiveGapWindowTest.kt @@ -0,0 +1,133 @@ +package dev.typetype.server.routes + +import dev.typetype.server.services.CachedSabrSegment +import dev.typetype.server.services.SabrSessionHolder +import dev.typetype.server.services.SabrSessionKey +import dev.typetype.server.services.SabrSessionStore +import io.mockk.coEvery +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertFalse +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant + +class SabrLiveGapWindowTest { + @Test + fun `live gap requests a fresh session before playback reaches it`() = runTest { + val audio = format(140, isAudio = true) + val video = format(299, isAudio = false) + val session = mockk(relaxed = true) + val state = mockk(relaxed = true) + every { session.streamState } returns state + every { session.isLive } returns true + every { state.isLive } returns true + every { state.liveHeadTimeMs } returns 510_000L + every { state.getSegmentStartMs(video, 92) } returns 475_000L + val holder = holder(session, audio, video) + holder.setLastServedSequence(video.itag, 91) + holder.setLastServedSequence(audio.itag, 48) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + val request = secondArg() + when { + request.format.itag == 299 && request.sequenceNumber == 93 -> cached(299, 93, 480_000L, 5_000L) + request.format.itag == 299 && request.sequenceNumber == 94 -> cached(299, 94, 485_000L, 5_000L) + request.format.itag == 140 && request.sequenceNumber == 49 -> cached(140, 49, 479_259L, 9_985L) + else -> null + } + } + val ranges = listOf( + SabrPlaybackBufferedRange(video.itag, 450_000L, 475_000L), + SabrPlaybackBufferedRange(audio.itag, 450_000L, 479_259L), + ) + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest(0L, 470_000L, 299, 140, bufferGoalMs = 8_000L, bufferedRanges = ranges), + ) + + assertFalse(result.isReady) + assertEquals("SABR recoverable failure: live 299 media discontinuity", holder.terminalFailure()) + assertNull(result.blockedRequests.firstOrNull { it.format.itag == video.itag }) + } + + @Test + fun `live gap requests a fresh session when playhead remains behind`() = runTest { + val audio = format(140, isAudio = true) + val video = format(299, isAudio = false) + val session = mockk(relaxed = true) + val state = mockk(relaxed = true) + every { session.streamState } returns state + every { session.isLive } returns true + every { state.isLive } returns true + every { state.liveHeadTimeMs } returns 510_000L + every { state.getSegmentNumberAtOrAfterTimeMs(video, any()) } returns 92 + val holder = holder(session, audio, video) + holder.setLastServedSequence(video.itag, 93) + holder.setLastServedSequence(audio.itag, 49) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + val request = secondArg() + when { + request.format.itag == 299 && request.sequenceNumber in 94..96 -> + cached(299, request.sequenceNumber, 485_000L + (request.sequenceNumber - 94) * 5_000L, 5_000L) + request.format.itag == 140 && request.sequenceNumber in 50..51 -> + cached(140, request.sequenceNumber, 489_244L + (request.sequenceNumber - 50) * 9_985L, 9_985L) + else -> null + } + } + val ranges = listOf( + SabrPlaybackBufferedRange(video.itag, 450_000L, 477_000L), + SabrPlaybackBufferedRange(video.itag, 480_000L, 485_000L), + SabrPlaybackBufferedRange(audio.itag, 450_000L, 489_244L), + ) + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest(0L, 477_942L, 299, 140, bufferGoalMs = 8_000L, bufferedRanges = ranges), + ) + + assertFalse(result.isReady) + assertFalse(result.blockedRequests.any { it.sequenceNumber == 92 }) + assertEquals("SABR recoverable failure: live 299 media discontinuity", holder.terminalFailure()) + } + + private fun holder( + session: YoutubeSabrSession, + audio: YoutubeSabrFormat, + video: YoutubeSabrFormat, + ): SabrSessionHolder = SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session", + key = SabrSessionKey("video", "user", audio.itag, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + ) + + private fun format(itag: Int, isAudio: Boolean): YoutubeSabrFormat = mockk(relaxed = true) { + every { this@mockk.itag } returns itag + every { this@mockk.isAudio } returns isAudio + every { mimeType } returns if (isAudio) "audio/mp4" else "video/mp4" + } + + private fun cached(itag: Int, sequence: Int, startMs: Long, durationMs: Long): CachedSabrSegment = CachedSabrSegment( + itag = itag, + sequence = sequence, + init = false, + startMs = startMs, + durationMs = durationMs, + mimeType = if (itag == 140) "audio/mp4" else "video/mp4", + bytesBase64 = "AA==", + byteLength = 1, + ) +} diff --git a/src/test/kotlin/dev/typetype/server/routes/SabrLivePlaybackWindowBuilderTest.kt b/src/test/kotlin/dev/typetype/server/routes/SabrLivePlaybackWindowBuilderTest.kt new file mode 100644 index 00000000..1c6b69ac --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/routes/SabrLivePlaybackWindowBuilderTest.kt @@ -0,0 +1,237 @@ +package dev.typetype.server.routes + +import dev.typetype.server.services.CachedSabrSegment +import dev.typetype.server.services.SabrSessionHolder +import dev.typetype.server.services.SabrSessionKey +import dev.typetype.server.services.SabrSessionStore +import io.mockk.coEvery +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertFalse +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant + +class SabrLivePlaybackWindowBuilderTest { + @Test + fun `active live window never reports end of stream at the current head`() = runTest { + val audio = format(itag = 140, isAudio = true) + val video = format(itag = 299, isAudio = false) + val session = mockk(relaxed = true) + val streamState = mockk(relaxed = true) + every { session.streamState } returns streamState + every { session.isLive } returns true + every { session.isAtLiveEdge } returns true + every { streamState.isLive } returns true + every { streamState.getMaxSegment(audio) } returns 90 + every { streamState.getMaxSegment(video) } returns 100 + every { streamState.getSegmentEndMs(audio, 90) } returns 898_000L + every { streamState.getSegmentEndMs(video, 100) } returns 897_000L + every { streamState.liveHeadTimeMs } returns 898_000L + every { streamState.getSegmentNumberAtOrAfterTimeMs(audio, any()) } returns 90 + every { streamState.getSegmentNumberAtOrAfterTimeMs(video, any()) } returns 100 + val holder = holder(session, audio, video) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + val request = secondArg() + when (request.format.itag) { + 140 -> cached(140, 90, 888_000L, 10_000L) + else -> cached(299, 100, 892_000L, 5_000L) + } + } + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest(0L, 896_000L, 299, 140, bufferGoalMs = 1_000L), + ) + + assertTrue(result.isReady) + assertTrue(result.response.live?.active == true) + assertEquals(898_000L, result.response.durationMs) + assertEquals(false, result.response.endOfStream) + assertEquals("/api/sabr/playback/session/140/init?generation=0", result.response.audio.initUrl) + val videoTrack = requireNotNull(result.response.video) + assertEquals("/api/sabr/playback/session/299/init?generation=0", videoTrack.initUrl) + } + + @Test + fun `live window requests every track missing from the target window`() = runTest { + val audio = format(itag = 140, isAudio = true) + val video = format(itag = 248, isAudio = false) + val session = mockk(relaxed = true) + val streamState = mockk(relaxed = true) + every { session.streamState } returns streamState + every { session.isLive } returns true + every { streamState.isLive } returns true + every { streamState.liveHeadTimeMs } returns 130_000L + every { streamState.getSegmentNumberAtOrAfterTimeMs(video, any()) } returns 100 + every { streamState.getSegmentNumberAtOrAfterTimeMs(audio, any()) } returns 100 + val holder = holder(session, audio, video) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + val request = secondArg() + when { + request.format.itag == 248 && request.sequenceNumber == 100 -> + cached(248, 100, 100_000L, 2_000L) + request.format.itag == 140 && request.sequenceNumber in 100..105 -> + cached(140, request.sequenceNumber, 100_000L + (request.sequenceNumber - 100) * 2_000L, 2_000L) + else -> null + } + } + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest(0L, 100_000L, 248, 140, bufferGoalMs = 30_000L), + ) + + assertEquals(listOf(140, 248), result.blockedRequests.map { it.format.itag }.sorted()) + assertEquals(listOf(101, 106), result.blockedRequests.map { it.sequenceNumber }) + } + + @Test + fun `active live window starts at the shared audio and video timestamp`() = runTest { + val audio = format(itag = 140, isAudio = true) + val video = format(itag = 137, isAudio = false) + val session = mockk(relaxed = true) + val streamState = mockk(relaxed = true) + every { session.streamState } returns streamState + every { session.isLive } returns true + every { streamState.isLive } returns true + every { streamState.liveHeadTimeMs } returns 112_000L + every { streamState.getSegmentNumberAtOrAfterTimeMs(audio, any()) } returns 100 + every { streamState.getSegmentNumberAtOrAfterTimeMs(video, any()) } returns 100 + val holder = holder(session, audio, video) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + val request = secondArg() + when { + request.format.itag == 137 && request.sequenceNumber == 100 -> cached(137, 100, 100_002L, 5_000L) + request.format.itag == 140 && request.sequenceNumber == 100 -> cached(140, 100, 102_010L, 5_000L) + else -> null + } + } + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest(0L, 100_000L, 137, 140, bufferGoalMs = 1_000L), + ) + + assertTrue(result.isReady) + assertEquals(102_010L, result.response.startTimeMs) + assertEquals(102_010L, result.response.audio.segments.single().startMs) + } + + @Test + fun `active live startup waits beyond five seconds of shared media`() = runTest { + val audio = format(itag = 140, isAudio = true) + val video = format(itag = 299, isAudio = false) + val session = mockk(relaxed = true) + val streamState = mockk(relaxed = true) + every { session.streamState } returns streamState + every { session.isLive } returns true + every { streamState.isLive } returns true + every { streamState.liveHeadTimeMs } returns 120_000L + every { streamState.getSegmentNumberAtOrAfterTimeMs(any(), any()) } returns 50 + val holder = holder(session, audio, video) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + val request = secondArg() + if (request.sequenceNumber in 50..52) { + cached(request.format.itag, request.sequenceNumber, 100_000L + (request.sequenceNumber - 50) * 2_000L, 2_000L) + } else { + null + } + } + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest(0L, 100_000L, 299, 140, bufferGoalMs = 8_000L), + ) + + assertFalse(result.isReady) + assertEquals(listOf(53, 53), result.blockedRequests.map { it.sequenceNumber }) + } + + @Test + fun `asymmetric live buffers keep requesting the shorter track`() = runTest { + val audio = format(itag = 140, isAudio = true) + val video = format(itag = 299, isAudio = false) + val session = mockk(relaxed = true) + val streamState = mockk(relaxed = true) + every { session.streamState } returns streamState + every { session.isLive } returns true + every { streamState.isLive } returns true + every { streamState.liveHeadTimeMs } returns 70_000L + every { streamState.getSegmentNumberAtOrAfterTimeMs(audio, any()) } returns 6 + every { streamState.getSegmentNumberAtOrAfterTimeMs(video, any()) } returns 8 + val holder = holder(session, audio, video) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + val request = secondArg() + when { + request.format.itag == 140 && request.sequenceNumber == 6 -> cached(140, 6, 49_923L, 9_985L) + request.format.itag == 299 && request.sequenceNumber == 8 -> cached(299, 8, 42_000L, 6_000L) + else -> null + } + } + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest( + 0L, + 34_954L, + 299, + 140, + bufferGoalMs = 8_000L, + bufferedRanges = listOf( + SabrPlaybackBufferedRange(140, 9_984L, 49_923L), + SabrPlaybackBufferedRange(299, 24_000L, 42_000L), + ), + ), + ) + + assertFalse(result.isReady) + assertEquals(listOf(299 to 9), result.blockedRequests.map { it.format.itag to it.sequenceNumber }) + } + + private fun holder( + session: YoutubeSabrSession, + audio: YoutubeSabrFormat, + video: YoutubeSabrFormat, + ): SabrSessionHolder = SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session", + key = SabrSessionKey("video", "user", 140, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + ) + + private fun format(itag: Int, isAudio: Boolean): YoutubeSabrFormat { + val format = mockk() + every { format.itag } returns itag + every { format.isAudio } returns isAudio + every { format.mimeType } returns if (isAudio) "audio/mp4" else "video/mp4" + every { format.approxDurationMs } returns 900_000L + return format + } + + private fun cached(itag: Int, sequence: Int, startMs: Long, durationMs: Long): CachedSabrSegment = CachedSabrSegment( + itag = itag, + sequence = sequence, + init = false, + startMs = startMs, + durationMs = durationMs, + mimeType = if (itag == 140) "audio/mp4" else "video/mp4", + bytesBase64 = "AA==", + byteLength = 1, + ) +} diff --git a/src/test/kotlin/dev/typetype/server/routes/SabrLiveRoundedBoundaryWindowTest.kt b/src/test/kotlin/dev/typetype/server/routes/SabrLiveRoundedBoundaryWindowTest.kt new file mode 100644 index 00000000..3c488ebc --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/routes/SabrLiveRoundedBoundaryWindowTest.kt @@ -0,0 +1,99 @@ +package dev.typetype.server.routes + +import dev.typetype.server.services.CachedSabrSegment +import dev.typetype.server.services.SabrSessionHolder +import dev.typetype.server.services.SabrSessionKey +import dev.typetype.server.services.SabrSessionStore +import io.mockk.coEvery +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaHeader +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaSegment +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant + +class SabrLiveRoundedBoundaryWindowTest { + @Test + fun `live startup advances past rounded target boundary`() = runTest { + val audio = format(140, true) + val video = format(299, false) + val session = mockk(relaxed = true) + val state = mockk(relaxed = true) + every { session.streamState } returns state + every { session.isLive } returns true + every { state.isLive } returns true + every { state.isPostLiveDvr } returns false + every { state.liveHeadTimeMs } returns 4_594_000L + every { state.liveHeadSequenceNumber } returns 4_594L + val holder = holder(session, audio, video) + holder.observeMediaSegment(mediaSegment(audio.itag, 4_573, 4_574_000L)) + holder.observeMediaSegment(mediaSegment(video.itag, 4_573, 4_574_000L)) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + val request = secondArg() + request.takeIf { it.sequenceNumber == 4_573 }?.let { cached(it.format.itag) } + } + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest(0L, 4_573_966L, video.itag, audio.itag, bufferGoalMs = 1L), + ) + + assertTrue(result.isReady) + assertTrue(result.blockedRequests.isEmpty()) + assertEquals(4_574_000L, result.response.startTimeMs) + assertEquals(4_573, result.response.video?.segments?.single()?.url?.sequenceFromUrl()) + assertEquals(4_573, result.response.audio.segments.single().url.sequenceFromUrl()) + } + + private fun holder( + session: YoutubeSabrSession, + audio: YoutubeSabrFormat, + video: YoutubeSabrFormat, + ): SabrSessionHolder = SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session", + key = SabrSessionKey("video", "user", audio.itag, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + ) + + private fun format(itag: Int, isAudio: Boolean): YoutubeSabrFormat = mockk { + every { this@mockk.itag } returns itag + every { this@mockk.isAudio } returns isAudio + every { mimeType } returns if (isAudio) "audio/mp4" else "video/mp4" + every { approxDurationMs } returns 4_594_000L + } + + private fun mediaSegment(itag: Int, sequence: Int, startMs: Long): SabrMediaSegment { + val header = mockk(relaxed = true) + every { header.itag } returns itag + every { header.sequenceNumber } returns sequence + every { header.startMs } returns startMs + every { header.durationMs } returns 1_000L + return mockk { every { this@mockk.header } returns header } + } + + private fun cached(itag: Int): CachedSabrSegment = CachedSabrSegment( + itag = itag, + sequence = 4_573, + init = false, + startMs = 4_574_000L, + durationMs = 1_000L, + mimeType = if (itag == 140) "audio/mp4" else "video/mp4", + bytesBase64 = "AA==", + byteLength = 1, + ) + + private fun String.sequenceFromUrl(): Int = substringAfter("/segment/").substringBefore('?').toInt() +} diff --git a/src/test/kotlin/dev/typetype/server/routes/SabrPlaybackAudioOnlyWindowTest.kt b/src/test/kotlin/dev/typetype/server/routes/SabrPlaybackAudioOnlyWindowTest.kt index 788c7fbe..40599f0a 100644 --- a/src/test/kotlin/dev/typetype/server/routes/SabrPlaybackAudioOnlyWindowTest.kt +++ b/src/test/kotlin/dev/typetype/server/routes/SabrPlaybackAudioOnlyWindowTest.kt @@ -76,7 +76,6 @@ class SabrPlaybackAudioOnlyWindowTest { every { format.isAudio } returns isAudio every { format.mimeType } returns if (isAudio) "audio/mp4" else "video/mp4" every { format.approxDurationMs } returns 420_000L - every { format.initializationUrl } returns null return format } diff --git a/src/test/kotlin/dev/typetype/server/routes/SabrPlaybackWindowBuilderTest.kt b/src/test/kotlin/dev/typetype/server/routes/SabrPlaybackWindowBuilderTest.kt index 828831cb..e813a8a5 100644 --- a/src/test/kotlin/dev/typetype/server/routes/SabrPlaybackWindowBuilderTest.kt +++ b/src/test/kotlin/dev/typetype/server/routes/SabrPlaybackWindowBuilderTest.kt @@ -61,12 +61,15 @@ class SabrPlaybackWindowBuilderTest { } @Test - fun `window starts from initial media without waiting for buffer goal`() = runTest { + fun `live window ignores client ranges before serving its own media`() = runTest { val audio = format(itag = 140, isAudio = true) val video = format(itag = 401, isAudio = false) val session = mockk(relaxed = true) val streamState = mockk(relaxed = true) every { session.streamState } returns streamState + every { session.isLive } returns true + every { streamState.isLive } returns true + every { streamState.liveHeadTimeMs } returns 330_000L every { streamState.getSegmentNumberAtOrAfterTimeMs(video, 300_000L) } returns 60 every { streamState.getSegmentNumberAtOrAfterTimeMs(audio, 296_600L) } returns 30 val holder = holder(session, audio, video) @@ -80,7 +83,8 @@ class SabrPlaybackWindowBuilderTest { else -> null } } - val request = SabrPlaybackWindowRequest(0L, 300_000L, 401, 140, bufferGoalMs = 30_000L) + val ranges = listOf(SabrPlaybackBufferedRange(401, 290_000L, 320_000L)) + val request = SabrPlaybackWindowRequest(0L, 300_000L, 401, 140, bufferGoalMs = 30_000L, bufferedRanges = ranges) assertTrue(SabrPlaybackWindowBuilder(store).build(holder, request).isReady) } @@ -166,6 +170,39 @@ class SabrPlaybackWindowBuilderTest { } } + @Test + fun `live window continues after the last segment served on each track`() = runTest { + val audio = format(itag = 140, isAudio = true) + val video = format(itag = 299, isAudio = false) + val session = mockk(relaxed = true) + val streamState = mockk(relaxed = true) + every { session.streamState } returns streamState + every { session.isLive } returns true + every { session.liveHeadSequenceNumber } returns 45L + every { streamState.isLive } returns true + every { streamState.liveHeadSequenceNumber } returns 45L + every { streamState.liveHeadTimeMs } returns 220_000L + every { streamState.getSegmentNumberAtOrAfterTimeMs(video, any()) } returns 43 + every { streamState.getSegmentNumberAtOrAfterTimeMs(audio, any()) } returns 23 + val holder = holder(session, audio, video) + holder.setLastServedSequence(video.itag, 41) + holder.setLastServedSequence(audio.itag, 21) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } returns null + val ranges = listOf( + SabrPlaybackBufferedRange(video.itag, 190_000L, 208_000L), + SabrPlaybackBufferedRange(audio.itag, 190_000L, 208_000L), + ) + + val result = SabrPlaybackWindowBuilder(store).build( + holder, + SabrPlaybackWindowRequest(0L, 200_000L, video.itag, audio.itag, bufferGoalMs = 8_000L, bufferedRanges = ranges), + ) + + assertEquals(42, result.blockedRequests.first { !it.format.isAudio }.sequenceNumber) + assertEquals(22, result.blockedRequests.first { it.format.isAudio }.sequenceNumber) + } + @Test fun `final indexed segments complete the window without requesting beyond end`() = runTest { val audio = format(itag = 140, isAudio = true) @@ -219,7 +256,6 @@ class SabrPlaybackWindowBuilderTest { every { format.isAudio } returns isAudio every { format.mimeType } returns if (isAudio) "audio/mp4" else "video/mp4" every { format.approxDurationMs } returns 900_000L - every { format.initializationUrl } returns null return format } diff --git a/src/test/kotlin/dev/typetype/server/routes/SabrStreamContractFilterTest.kt b/src/test/kotlin/dev/typetype/server/routes/SabrStreamContractFilterTest.kt new file mode 100644 index 00000000..f0078556 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/routes/SabrStreamContractFilterTest.kt @@ -0,0 +1,105 @@ +package dev.typetype.server.routes + +import dev.typetype.server.testAudioStream +import dev.typetype.server.testStreamResponse +import dev.typetype.server.testVideoStream +import dev.typetype.server.services.SabrPreparedInfo +import dev.typetype.server.services.SabrSessionStore +import io.mockk.coEvery +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo + +class SabrStreamContractFilterTest { + @Test + fun `sabr contract never exposes hls`() { + val response = testStreamResponse( + videoOnlyStreams = listOf( + testVideoStream().copy( + url = "", + deliveryMethod = "sabr", + sabrSessionUrl = "/sabr/session/video-id?videoItag=137", + ), + ), + audioStreams = listOf( + testAudioStream( + url = "", + deliveryMethod = "sabr", + sabrSessionUrl = "/sabr/session/video-id?audioItag=140", + ), + ), + hlsUrl = "https://example.com/live.m3u8", + ).copy(isLive = true, isLiveContent = true, hasLiveManifest = true) + + val sabr = response.onlySabrStreams() + + assertEquals("", sabr.hlsUrl) + assertTrue(sabr.videoOnlyStreams.isNotEmpty()) + assertTrue(sabr.audioStreams.isNotEmpty()) + } + + @Test + fun `sabr contract exposes every available video codec`() = runTest { + val h264 = videoFormat(137, "video/mp4; codecs=\"avc1.4d4028\"") + val vp9 = videoFormat(248, "video/webm; codecs=\"vp9\"") + val av1 = videoFormat(399, "video/mp4; codecs=\"av01.0.08M.08\"") + val audio = mockk(relaxed = true) { + every { isAudio } returns true + every { isVideo } returns false + every { itag } returns 140 + every { mimeType } returns "audio/mp4; codecs=\"mp4a.40.2\"" + } + val info = mockk { + every { formats } returns listOf(h264, vp9, av1, audio) + every { findFormatByItag(any()) } answers { + formats.firstOrNull { it.itag == firstArg() } + } + } + val store = mockk() + coEvery { store.fetchInfo(VIDEO_ID, cachedFirst = true) } returns SabrPreparedInfo(info, null) + val response = testStreamResponse( + videoOnlyStreams = listOf( + testVideoStream().copy( + itag = 137, + codec = "avc1.4d4028", + deliveryMethod = "sabr", + sabrSessionUrl = "/sabr/session/$VIDEO_ID?videoItag=137", + ), + ), + audioStreams = listOf( + testAudioStream( + itag = 140, + codec = "mp4a.40.2", + deliveryMethod = "sabr", + sabrSessionUrl = "/sabr/session/$VIDEO_ID?audioItag=140", + ), + ), + ) + + val sabr = response.withPlayableSabrStreams(YOUTUBE_URL, store).onlySabrStreams() + + assertEquals(setOf(137, 248, 399), sabr.videoOnlyStreams.map { it.itag }.toSet()) + assertEquals(setOf("avc1.4d4028", "vp9", "av01.0.08M.08"), sabr.videoOnlyStreams.map { it.codec }.toSet()) + } + + private fun videoFormat(itag: Int, mime: String): YoutubeSabrFormat = + mockk(relaxed = true) { + every { isAudio } returns false + every { isVideo } returns true + every { this@mockk.itag } returns itag + every { mimeType } returns mime + every { height } returns 1080 + every { width } returns 1920 + every { qualityLabel } returns "1080p" + } + + private companion object { + const val VIDEO_ID = "X4VbdwhkE10" + const val YOUTUBE_URL = "https://www.youtube.com/watch?v=$VIDEO_ID" + } +} diff --git a/src/test/kotlin/dev/typetype/server/services/RelatedItemMappersTest.kt b/src/test/kotlin/dev/typetype/server/services/RelatedItemMappersTest.kt new file mode 100644 index 00000000..55ca5660 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/RelatedItemMappersTest.kt @@ -0,0 +1,22 @@ +package dev.typetype.server.services + +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.stream.StreamInfoItem +import org.schabi.newpipe.extractor.stream.StreamType + +class RelatedItemMappersTest { + @Test + fun `video item preserves membership requirement`() { + val item = StreamInfoItem( + 0, + "https://www.youtube.com/watch?v=member", + "Members-only video", + StreamType.VIDEO_STREAM, + ).apply { + setRequiresMembership(true) + } + + assertTrue(item.toVideoItem().requiresMembership) + } +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrDemandWatchdogLifecycleTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrDemandWatchdogLifecycleTest.kt index fd8e300d..e84a32b2 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrDemandWatchdogLifecycleTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrDemandWatchdogLifecycleTest.kt @@ -83,12 +83,107 @@ class SabrDemandWatchdogLifecycleTest { } } + @Test + fun `future live demand expires as recoverable`() = runTest { + withTracker { holder -> + val request = SabrSegmentRequest.media(holder.videoFormat, 50) + every { holder.session.isLive } returns true + every { holder.session.streamState.isLive } returns true + every { holder.session.streamState.getMaxSegment(holder.videoFormat) } returns 49 + holder.requestSegmentDemand(request, registeredAtMs = 0L) + var expired = false + val job = launch { + expired = SabrDemandWatchdog( + clock = { testScheduler.currentTime }, + intervalMs = 100L, + ).monitor({ true }, holder) + } + runCurrent() + + advanceTimeBy(SabrPumpPolicy.DEMAND_TARGET_DEADLINE_MS) + runCurrent() + advanceTimeBy(LIVE_EDGE_POLL_MS) + runCurrent() + + assertTrue(job.isCompleted) + assertTrue(expired) + assertEquals( + "$SABR_RECOVERABLE_FAILURE_PREFIX SABR demand stalled for 299:50", + holder.terminalFailure(), + ) + } + } + + @Test + fun `completed in flight demand interrupts without terminal failure`() = runTest { + withTracker { holder -> + val request = SabrSegmentRequest.media(holder.videoFormat, 50) + var cached = false + var progressVersion = 0L + every { holder.session.getCachedSegment(request) } answers { + if (cached) mockk(relaxed = true) else null + } + every { holder.session.mediaProgressVersion } answers { progressVersion } + holder.requestSegmentDemand(request, registeredAtMs = 0L) + val identity = requireNotNull(holder.segmentDemandIdentity(request)) + assertTrue(holder.beginInFlightSegmentDemand(request, identity, futureLiveRequest = false)) + cached = true + progressVersion = 1L + holder.clearSegmentDemand(request) + var interrupted = false + val job = launch { + interrupted = SabrDemandWatchdog( + clock = { testScheduler.currentTime }, + intervalMs = 100L, + ).monitor({ true }, holder) + } + runCurrent() + + advanceTimeBy(SabrPumpPolicy.COMPLETED_DEMAND_IDLE_MS) + runCurrent() + + assertTrue(job.isCompleted) + assertTrue(interrupted) + assertFalse(holder.playbackState() == SabrPlaybackState.TERMINAL) + assertEquals(null, holder.terminalFailure()) + } + } + + @Test + fun `missing in flight demand expires at its original deadline`() = runTest { + withTracker { holder -> + val request = SabrSegmentRequest.media(holder.videoFormat, 50) + holder.requestSegmentDemand(request, registeredAtMs = 0L) + val identity = requireNotNull(holder.segmentDemandIdentity(request)) + assertTrue(holder.beginInFlightSegmentDemand(request, identity, futureLiveRequest = false)) + holder.clearSegmentDemand(request) + var expired = false + val job = launch { + expired = SabrDemandWatchdog( + clock = { testScheduler.currentTime }, + intervalMs = 100L, + ).monitor({ true }, holder) + } + runCurrent() + + advanceTimeBy(SabrPumpPolicy.DEMAND_TARGET_DEADLINE_MS) + runCurrent() + + assertTrue(job.isCompleted) + assertTrue(expired) + assertEquals(SabrPlaybackState.TERMINAL, holder.playbackState()) + assertEquals("SABR demand stalled for 299:50", holder.terminalFailure()) + } + } + private suspend fun withTracker(block: suspend (SabrSessionHolder) -> Unit) { SabrSegmentDemandTracker.clearAll() + SabrInFlightDemandTracker.clearAll() try { block(holder()) } finally { SabrSegmentDemandTracker.clearAll() + SabrInFlightDemandTracker.clearAll() } } diff --git a/src/test/kotlin/dev/typetype/server/services/SabrFallbackStreamServiceTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrFallbackStreamServiceTest.kt index 109fcf0c..2f556726 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrFallbackStreamServiceTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrFallbackStreamServiceTest.kt @@ -36,7 +36,7 @@ class SabrFallbackStreamServiceTest { } @Test - fun `prepares sabr while keeping already playable extraction`() = runTest { + fun `replaces classic extraction with prepared sabr formats`() = runTest { val delegate = mockk() val sessionStore = mockk() val tokenSessionClient = mockk() @@ -47,11 +47,37 @@ class SabrFallbackStreamServiceTest { val result = service.getStreamInfo(YOUTUBE_URL) - assertEquals(ExtractionResult.Success(response), result) + val enriched = (result as ExtractionResult.Success).data + assertEquals(listOf(137), enriched.videoOnlyStreams.map { it.itag }) + assertEquals(listOf(140), enriched.audioStreams.map { it.itag }) + assertEquals("sabr", enriched.videoOnlyStreams.single().deliveryMethod) + assertEquals("sabr", enriched.audioStreams.single().deliveryMethod) coVerify(exactly = 1) { sessionStore.fetchInfo(VIDEO_ID, cachedFirst = true) } coVerify(exactly = 0) { tokenSessionClient.fetchPlaybackSession(any()) } } + @Test + fun `enriches a live hls extraction with prepared sabr formats`() = runTest { + val delegate = mockk() + val sessionStore = mockk() + val tokenSessionClient = mockk() + val response = testStreamResponse( + videoOnlyStreams = emptyList(), + audioStreams = emptyList(), + hlsUrl = LIVE_HLS_URL, + ).copy(isLive = true, isLiveContent = true, hasLiveManifest = true, streamType = "live_stream") + coEvery { delegate.getStreamInfo(YOUTUBE_URL) } returns ExtractionResult.Success(response) + coEvery { sessionStore.fetchInfo(VIDEO_ID, cachedFirst = true) } returns preparedInfo() + val service = SabrFallbackStreamService(delegate, sessionStore, tokenSessionClient) + + val result = service.getStreamInfo(YOUTUBE_URL) + + val enriched = (result as ExtractionResult.Success).data + assertEquals(LIVE_HLS_URL, enriched.hlsUrl) + assertEquals(listOf(137), enriched.videoOnlyStreams.map { it.itag }) + assertEquals(listOf(140), enriched.audioStreams.map { it.itag }) + } + @Test fun `recovers youtube extraction failure with token playback session`() = runTest { val delegate = mockk() diff --git a/src/test/kotlin/dev/typetype/server/services/SabrLiveMediaNormalizerTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrLiveMediaNormalizerTest.kt new file mode 100644 index 00000000..8f060e27 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrLiveMediaNormalizerTest.kt @@ -0,0 +1,69 @@ +package dev.typetype.server.services + +import org.junit.jupiter.api.Assertions.assertArrayEquals +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Test + +class SabrLiveMediaNormalizerTest { + @Test + fun `splits fragmented mp4 after movie metadata`() { + val ftyp = mp4Box("ftyp", byteArrayOf(1, 2)) + val moov = mp4Box("moov", byteArrayOf(3, 4, 5)) + val emsg = mp4Box("emsg", byteArrayOf(6)) + val moof = mp4Box("moof", byteArrayOf(7, 8)) + val mdat = mp4Box("mdat", byteArrayOf(9, 10)) + + val parts = requireNotNull(SabrLiveMediaNormalizer.split("video/mp4; codecs=avc1", ftyp + moov + emsg + moof + mdat)) + + assertArrayEquals(ftyp + moov, parts.initialization) + assertArrayEquals(emsg + moof + mdat, parts.media) + } + + @Test + fun `does not parse mp4 payload after first media fragment`() { + val ftyp = mp4Box("ftyp", byteArrayOf(1)) + val moov = mp4Box("moov", byteArrayOf(2)) + val moof = mp4Box("moof", byteArrayOf(3)) + val opaqueMedia = byteArrayOf(0, 0, 0, 1, 0, 0, 0, 0) + + val parts = requireNotNull(SabrLiveMediaNormalizer.split("audio/mp4", ftyp + moov + moof + opaqueMedia)) + + assertArrayEquals(ftyp + moov, parts.initialization) + assertArrayEquals(moof + opaqueMedia, parts.media) + } + + @Test + fun `splits webm before first cluster`() { + val ebml = ebmlElement(byteArrayOf(0x1A, 0x45, 0xDF.toByte(), 0xA3.toByte()), byteArrayOf(1)) + val info = ebmlElement(byteArrayOf(0x15, 0x49, 0xA9.toByte(), 0x66), byteArrayOf(2)) + val tracks = ebmlElement(byteArrayOf(0x16, 0x54, 0xAE.toByte(), 0x6B), byteArrayOf(3)) + val cluster = ebmlElement(byteArrayOf(0x1F, 0x43, 0xB6.toByte(), 0x75), byteArrayOf(4, 5)) + val segment = byteArrayOf(0x18, 0x53, 0x80.toByte(), 0x67, 0xFF.toByte()) + info + tracks + cluster + + val parts = requireNotNull(SabrLiveMediaNormalizer.split("video/webm; codecs=vp9", ebml + segment)) + + assertArrayEquals(ebml + segment.copyOfRange(0, segment.size - cluster.size), parts.initialization) + assertArrayEquals(cluster, parts.media) + } + + @Test + fun `keeps malformed and regular media untouched`() { + assertNull(SabrLiveMediaNormalizer.split("video/mp4", mp4Box("mdat", byteArrayOf(1)))) + assertNull(SabrLiveMediaNormalizer.split("video/unknown", byteArrayOf(1, 2, 3))) + } + + private fun mp4Box(type: String, payload: ByteArray): ByteArray { + val size = payload.size + 8 + return byteArrayOf( + (size ushr 24).toByte(), + (size ushr 16).toByte(), + (size ushr 8).toByte(), + size.toByte(), + ) + type.toByteArray(Charsets.US_ASCII) + payload + } + + private fun ebmlElement(id: ByteArray, payload: ByteArray): ByteArray { + require(payload.size < 127) + return id + byteArrayOf((0x80 or payload.size).toByte()) + payload + } +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrLivePlaybackSessionServiceTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrLivePlaybackSessionServiceTest.kt new file mode 100644 index 00000000..0dc0adca --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrLivePlaybackSessionServiceTest.kt @@ -0,0 +1,139 @@ +package dev.typetype.server.services + +import io.mockk.coEvery +import io.mockk.coVerify +import io.mockk.every +import io.mockk.mockk +import io.mockk.verify +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaHeader +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaSegment +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant + +class SabrLivePlaybackSessionServiceTest { + @Test + fun `live start uses the warmed media pair when it is closer to the head`() { + val audio = format(140, isAudio = true) + val video = format(137, isAudio = false) + val holder = holder(audio, video) + val state = holder.session.streamState + every { holder.session.isLive } returns true + every { state.isLive } returns true + every { state.isPostLiveDvr } returns false + every { state.liveHeadTimeMs } returns 1_005_000L + holder.markExpectedLive() + holder.observeMediaSegment(mediaSegment(audio.itag, 995_002L)) + holder.observeMediaSegment(mediaSegment(video.itag, 995_000L)) + + assertEquals(995_002L, holder.resolvePlaybackStartMs(0L)) + } + + @Test + fun `live prepare warms metadata and starts behind the live head`() = runTest { + val audio = format(140, isAudio = true) + val video = format(137, isAudio = false) + val info = mockk() + val prepared = SabrPreparedInfo(info, token(), isLive = true, isLiveContent = true) + val holder = holder(audio, video) + val state = holder.session.streamState + every { holder.session.isLive } returns true + every { holder.session.liveHeadSequenceNumber } returns 200L + every { state.isLive } returns true + every { state.isPostLiveDvr } returns false + every { state.liveHeadSequenceNumber } returns 200L + every { state.liveHeadTimeMs } returns 1_005_000L + every { state.getMaxSegment(audio) } returns 100 + every { state.getMaxSegment(video) } returns 200 + every { state.getSegmentEndMs(audio, 100) } returns 1_000_000L + every { state.getSegmentEndMs(video, 200) } returns 1_002_000L + every { state.getBufferedEndMs(audio) } returns 1_000_000L + every { state.getBufferedEndMs(video) } returns 1_002_000L + every { state.getMinBufferedEndMs() } returns 1_000_000L + every { state.getSegmentNumberAtOrAfterTimeMs(video, 985_000L) } returns 198 + every { state.getSegmentNumberAtOrAfterTimeMs(audio, 985_000L) } returns 99 + every { state.getSegmentStartMs(video, 198) } returns 990_000L + every { state.getSegmentStartMs(audio, 99) } returns 990_000L + val store = mockk() + every { + store.getOrCreate( + "video", + "user", + info, + audio, + video, + prepared.initialToken, + 0L, + false, + SabrSessionPurpose.PLAYBACK, + false, + ) + } returns holder + coEvery { store.ensureWarmed(holder, 8) } returns Unit + every { store.startPump(holder) } returns Unit + + val result = SabrPlaybackSessionService(store).prepare("video", "user", prepared, audio, video, 0L) + + assertEquals(985_000L, result.startTimeMs) + assertEquals(985_000L, holder.playerTimeMs()) + assertTrue(holder.expectsLive()) + coVerify(exactly = 1) { store.ensureWarmed(holder, 8) } + coVerify(exactly = 0) { store.fetchInitializationData(any(), any()) } + verify(exactly = 1) { store.startPump(holder) } + } + + private fun holder(audio: YoutubeSabrFormat, video: YoutubeSabrFormat): SabrSessionHolder { + val session = mockk() + val state = mockk(relaxed = true) + every { session.streamState } returns state + every { session.getCachedSegment(any()) } returns null + every { session.isBeyondEnd(any()) } returns false + every { session.prepareForInitialization(any()) } returns Unit + every { state.setActiveTrackTypes(any(), any()) } returns Unit + every { state.setSelectVideoFormatBeforeAudio(any()) } returns Unit + return SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session-token", + key = SabrSessionKey("video", "user", audio.itag, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + ) + } + + private fun format(itag: Int, isAudio: Boolean): YoutubeSabrFormat { + val format = mockk() + every { format.itag } returns itag + every { format.isAudio } returns isAudio + every { format.audioTrackId } returns null + every { format.mimeType } returns if (isAudio) "audio/mp4" else "video/mp4" + every { format.bitrate } returns if (isAudio) 128_000 else 2_000_000 + return format + } + + private fun mediaSegment(itag: Int, startMs: Long): SabrMediaSegment { + val header = mockk(relaxed = true) + every { header.itag } returns itag + every { header.sequenceNumber } returns 200 + every { header.startMs } returns startMs + every { header.durationMs } returns 5_000L + every { header.isInitSegment } returns false + return mockk { every { this@mockk.header } returns header } + } + + private fun token(): SabrTokenBundle = SabrTokenBundle( + videoId = "video", + visitorBoundPoToken = "visitor-token", + visitorBoundPoTokenBytes = byteArrayOf(1), + visitorData = "visitor-data", + videoBoundPoToken = "video-token", + videoBoundPoTokenBytes = byteArrayOf(2), + ) +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrLivePlaybackTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrLivePlaybackTest.kt new file mode 100644 index 00000000..48c68125 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrLivePlaybackTest.kt @@ -0,0 +1,258 @@ +package dev.typetype.server.services + +import io.mockk.every +import io.mockk.mockk +import org.junit.jupiter.api.AfterEach +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertFalse +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaHeader +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaSegment +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant + +class SabrLivePlaybackTest { + @AfterEach + fun clearDemands(): Unit = SabrSegmentDemandTracker.clearAll() + + @Test + fun `active live resolves zero start behind the observed head`() { + val fixture = fixture() + + val live = requireNotNull(fixture.holder.livePlaybackSnapshot()) + + assertTrue(live.active) + assertFalse(live.postLiveDvr) + assertEquals(1_005_000L, live.seekableEndMs) + assertEquals(985_000L, fixture.holder.resolvePlaybackStartMs(0L)) + assertEquals(1_005_000L, fixture.holder.resolvePlaybackStartMs(1_100_000L)) + } + + @Test + fun `reported live head wins over sequence based duration estimates`() { + val fixture = fixture() + every { fixture.state.getSegmentEndMs(fixture.video, 200) } returns 9_011_868_000L + every { fixture.state.getBufferedEndMs(fixture.video) } returns 9_011_868_000L + + val live = requireNotNull(fixture.holder.livePlaybackSnapshot()) + + assertEquals(1_005_000L, live.headTimeMs) + assertEquals(1_005_000L, live.seekableEndMs) + } + + @Test + fun `active live maps time from an observed sabr segment for every codec`() { + val fixture = fixture() + val header = mockk { + every { isInitSegment } returns false + every { itag } returns fixture.video.itag + every { sequenceNumber } returns 180 + every { startMs } returns 965_000L + every { durationMs } returns 2_000L + } + val segment = mockk { + every { this@mockk.header } returns header + } + fixture.holder.observeMediaSegment(segment) + + assertEquals(200L, requireNotNull(fixture.holder.livePlaybackSnapshot()).headSequence) + assertEquals(segment, fixture.holder.observedMediaSegment(fixture.video)) + assertEquals(965_000L, segment.header.startMs) + assertEquals(195, fixture.holder.playbackStartSequence(fixture.video, 995_000L)) + assertEquals(195, fixture.holder.playbackStartSequence(fixture.video, 995_001L)) + assertEquals(180, fixture.holder.playbackStartSequence(fixture.video, 965_000L)) + assertEquals(179, fixture.holder.playbackStartSequence(fixture.video, 964_999L)) + assertEquals(200, fixture.holder.playbackStartSequence(fixture.video, 1_006_000L)) + } + + @Test + fun `active live derives missing segment duration from the sabr head`() { + val fixture = fixture() + val header = mockk { + every { isInitSegment } returns false + every { itag } returns fixture.video.itag + every { sequenceNumber } returns 180 + every { startMs } returns 965_000L + every { durationMs } returns -1L + } + val segment = mockk { + every { this@mockk.header } returns header + } + fixture.holder.observeMediaSegment(segment) + + assertEquals(195, fixture.holder.playbackStartSequence(fixture.video, 995_000L)) + assertEquals(180, fixture.holder.playbackStartSequence(fixture.video, 965_000L)) + } + + @Test + fun `live duration tolerates millisecond drift between media and head timestamps`() { + val fixture = fixture() + every { fixture.state.liveHeadTimeMs } returns 1_004_999L + val header = mockk { + every { isInitSegment } returns false + every { itag } returns fixture.video.itag + every { sequenceNumber } returns 180 + every { startMs } returns 965_000L + every { durationMs } returns -1L + } + fixture.holder.observeMediaSegment(mockk { every { this@mockk.header } returns header }) + + assertEquals(2_000L, fixture.holder.playbackSegmentDurationMs(fixture.video, 180)) + assertEquals(195, fixture.holder.playbackStartSequence(fixture.video, 995_000L)) + } + + @Test + fun `only the next live media segments wait for production`() { + val fixture = fixture() + + assertTrue(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.video, 201))) + assertTrue(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.video, 202))) + assertTrue(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.audio, 101))) + assertTrue(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.audio, 102))) + assertFalse(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.video, 203))) + assertFalse(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.video, 200))) + assertFalse(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.audio, 103))) + assertFalse(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.audio, 100))) + } + + @Test + fun `advertised live head waits until its media is complete`() { + val fixture = fixture() + every { fixture.state.getSegmentStartMs(fixture.video, 200) } returns 1_002_000L + + assertTrue(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.video, 200))) + } + + @Test + fun `in flight live media waits beyond the reported head`() { + val fixture = fixture() + val request = SabrSegmentRequest.media(fixture.audio, 103) + every { fixture.session.getReadableSegment(request) } returns mockk() + + assertTrue(fixture.holder.isFutureLiveRequest(request)) + } + + @Test + fun `live head can advance beyond the last complete media segment`() { + val fixture = fixture() + val header = mockk { + every { isInitSegment } returns false + every { itag } returns fixture.video.itag + every { sequenceNumber } returns 198 + every { startMs } returns 998_000L + } + fixture.holder.observeMediaSegment(mockk { every { this@mockk.header } returns header }) + + assertTrue(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.video, 200))) + assertFalse(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.video, 197))) + } + + @Test + fun `live retries immediately behind the head and paces future media`() { + val fixture = fixture() + val available = SabrSegmentRequest.media(fixture.video, 200) + val future = SabrSegmentRequest.media(fixture.video, 201) + + assertEquals(DEFAULT_PLAYBACK_RETRY_MS, fixture.holder.liveRetryAfterMs(listOf(available))) + assertEquals(LIVE_EDGE_POLL_MS, fixture.holder.liveRetryAfterMs(listOf(future))) + } + + @Test + fun `future live demand remains retryable after repeated responses`() { + val fixture = fixture() + val request = SabrSegmentRequest.media(fixture.video, 201) + fixture.holder.requestSegmentDemand(request, registeredAtMs = 0L) + val identity = requireNotNull(fixture.holder.segmentDemandIdentity(request)) + val result = mockk { + every { segmentCount } returns 2 + every { targetTrackSegmentCount } returns 1 + } + val runtime = SabrPumpRuntime { 20_000L } + val wasFutureLiveRequest = fixture.holder.isFutureLiveRequest(request) + every { fixture.state.getMaxSegment(fixture.video) } returns 201 + + repeat(4) { + assertFalse( + SabrDemandAttemptFinisher.finish( + fixture.holder, + request, + identity, + result, + runtime, + wasFutureLiveRequest, + ), + ) + } + + assertEquals(SabrPlaybackState.WAITING_FOR_LIVE, fixture.holder.playbackState()) + assertNull(fixture.holder.terminalFailure()) + assertEquals("299:201", fixture.holder.pendingSegmentDemandSummary()) + } + + @Test + fun `post live dvr is finite instead of an active live edge`() { + val fixture = fixture(postLiveDvr = true) + + val live = requireNotNull(fixture.holder.livePlaybackSnapshot()) + + assertFalse(live.active) + assertTrue(live.postLiveDvr) + assertFalse(fixture.holder.isFutureLiveRequest(SabrSegmentRequest.media(fixture.video, 201))) + } + + private fun fixture(postLiveDvr: Boolean = false): Fixture { + val audio = format(140, true) + val video = format(299, false) + val state = mockk(relaxed = true) + val session = mockk(relaxed = true) + every { session.streamState } returns state + every { session.isLive } returns !postLiveDvr + every { session.isAtLiveEdge } returns !postLiveDvr + every { session.liveHeadSequenceNumber } returns 200L + every { session.getCachedSegment(any()) } returns null + every { session.getReadableSegment(any()) } returns null + every { state.isLive } returns !postLiveDvr + every { state.isPostLiveDvr } returns postLiveDvr + every { state.liveHeadSequenceNumber } returns 200L + every { state.liveHeadTimeMs } returns 1_005_000L + every { state.getMaxSegment(audio) } returns 100 + every { state.getMaxSegment(video) } returns 200 + every { state.getSegmentEndMs(audio, 100) } returns 1_000_000L + every { state.getSegmentEndMs(video, 200) } returns 1_002_000L + every { state.getBufferedEndMs(audio) } returns 1_000_000L + every { state.getBufferedEndMs(video) } returns 1_002_000L + every { state.getMinBufferedEndMs() } returns 1_000_000L + val holder = SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session", + key = SabrSessionKey("video", "user", audio.itag, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + ) + return Fixture(holder, session, state, audio, video) + } + + private fun format(itag: Int, isAudio: Boolean): YoutubeSabrFormat { + val format = mockk() + every { format.itag } returns itag + every { format.isAudio } returns isAudio + every { format.bitrate } returns if (isAudio) 128_000 else 2_000_000 + return format + } + + private data class Fixture( + val holder: SabrSessionHolder, + val session: YoutubeSabrSession, + val state: YoutubeSabrStreamState, + val audio: YoutubeSabrFormat, + val video: YoutubeSabrFormat, + ) +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrLiveProtocolProbeTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrLiveProtocolProbeTest.kt new file mode 100644 index 00000000..df511f17 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrLiveProtocolProbeTest.kt @@ -0,0 +1,170 @@ +package dev.typetype.server.services + +import kotlinx.coroutines.runBlocking +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Tag +import org.junit.jupiter.api.Test +import org.junit.jupiter.api.condition.EnabledIfSystemProperty +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import java.security.MessageDigest + +@EnabledIfSystemProperty(named = "sabr.probe", matches = "true") +@Tag("network") +class SabrLiveProtocolProbeTest { + @Test + fun `retrieves every requested live sabr codec`(): Unit = runBlocking { + val videoId = sabrProbeVideoId() + val tokenServiceUrl = sabrProbeTokenServiceUrl() + NewPipeInitializer.init(tokenServiceUrl) + val store = SabrSessionStore(tokenServiceUrl = tokenServiceUrl) + try { + val tokenClient = TypetypeTokenSabrTokenClient(tokenServiceUrl) + val prepared = SabrInfoFetcher( + tokenClient, + TypetypeTokenYoutubeSessionClient(tokenServiceUrl), + ).fetchInfo(videoId) ?: error("Missing SABR player metadata") + println( + "profile=${prepared.info.profile} clientVersion=${prepared.info.clientVersion} " + + "visitorMatches=${prepared.info.visitorData == prepared.initialToken?.visitorData}" + ) + val audio = prepared.info.formats.first { it.itag == sabrProbeAudioItag() && it.isAudio } + val requestedVideoItags = sabrProbeVideoItags() + val videos = requestedVideoItags.map { itag -> + prepared.info.formats.first { it.itag == itag && it.isVideo } + } + assertEquals(requestedVideoItags, videos.map { it.itag }) + videos.forEach { video -> + probe(store, videoId, prepared, audio, video) + } + } finally { + store.release() + } + } + + private suspend fun probe( + store: SabrSessionStore, + videoId: String, + prepared: SabrPreparedInfo, + audio: YoutubeSabrFormat, + video: YoutubeSabrFormat, + ) { + val label = "${video.codecFamily()}-${video.itag}" + val holder = store.getOrCreate( + videoId = videoId, + userId = "live-protocol-probe-$label", + info = prepared.info, + audioFormat = audio, + videoFormat = video, + initialToken = prepared.initialToken, + startPump = false, + ) + holder.markExpectedLive() + store.ensureWarmed(holder, maxPumps = 8) + val live = requireNotNull(holder.livePlaybackSnapshot()) + val segments = listOfNotNull( + holder.observedMediaSegment(audio), + holder.observedMediaSegment(video), + ) + println( + "$label requests=${holder.session.requestNumber} segments=${segments.size} " + + "headSeq=${live.headSequence} headMs=${live.headTimeMs} startMs=${holder.resolvePlaybackStartMs(0L)}" + ) + println("$label trace=${holder.session.diagnosticTrace}") + assertEquals(setOf(audio.itag, video.itag), segments.map { it.header.itag }.toSet()) + segments.forEach { segment -> + val header = segment.header + val format = if (header.itag == audio.itag) audio else video + val parts = requireNotNull(SabrLiveMediaNormalizer.split(format.mimeType.orEmpty(), segment.data)) { + "$label could not split live ${format.mimeType}" + } + assertTrue(header.sequenceNumber > 0, "$label returned bootstrap media as playable media") + assertTrue(parts.initialization.isNotEmpty(), "$label returned an empty initialization") + assertTrue(parts.media.isNotEmpty(), "$label returned empty media") + assertTrue(holder.liveInitialization(format)?.isNotEmpty() == true, "$label did not retain initialization") + println( + " itag=${header.itag} seq=${header.sequenceNumber} startMs=${header.startMs} " + + "durationMs=${header.durationMs} bytes=${segment.length} sha256=${fingerprint(segment.data)} " + + "boxes=${mp4BoxNames(segment.data)} tfdt=${mp4DecodeTimes(segment.data, header.timeRangeTimescale)}" + ) + } + } + + private fun YoutubeSabrFormat.codecFamily(): String { + val mime = mimeType.orEmpty().lowercase() + return when { + "avc1" in mime -> "H.264" + "vp09" in mime || "vp9" in mime -> "VP9" + "av01" in mime -> "AV1" + else -> error("Unsupported probe codec for itag $itag: $mime") + } + } + + private fun mp4BoxNames(data: ByteArray): String { + val names = mutableListOf() + var offset = 0 + while (offset + 8 <= data.size && names.size < 12) { + val size = ((data[offset].toLong() and 0xff) shl 24) or + ((data[offset + 1].toLong() and 0xff) shl 16) or + ((data[offset + 2].toLong() and 0xff) shl 8) or + (data[offset + 3].toLong() and 0xff) + val type = String(data, offset + 4, 4, Charsets.US_ASCII) + if (size < 8L || size > data.size - offset) break + names += "$type:$size" + offset += size.toInt() + } + return names.joinToString(",") + } + + private fun fingerprint(data: ByteArray): String = MessageDigest.getInstance("SHA-256") + .digest(data) + .take(6) + .joinToString("") { "%02x".format(it) } + + private fun mp4DecodeTimes(data: ByteArray, timescale: Int): String { + if (timescale <= 0) return "unavailable" + val decodeTimes = mutableListOf() + collectDecodeTimes(data, 0, data.size, decodeTimes) + if (decodeTimes.isEmpty()) return "none" + val firstMs = decodeTimes.first() * 1_000L / timescale + val lastMs = decodeTimes.last() * 1_000L / timescale + return "count=${decodeTimes.size},firstMs=$firstMs,lastMs=$lastMs" + } + + private fun collectDecodeTimes(data: ByteArray, start: Int, end: Int, output: MutableList) { + var offset = start + while (offset + 8 <= end) { + val size = data.readUnsignedInt(offset) + if (size < 8L || size > end - offset) return + val type = String(data, offset + 4, 4, Charsets.US_ASCII) + val payloadStart = offset + 8 + val boxEnd = offset + size.toInt() + when (type) { + "moof", "traf" -> collectDecodeTimes(data, payloadStart, boxEnd, output) + "tfdt" -> data.readTfdt(payloadStart, boxEnd)?.let(output::add) + } + offset = boxEnd + } + } + + private fun ByteArray.readTfdt(offset: Int, end: Int): Long? { + if (offset + 8 > end) return null + return if (this[offset].toInt() == 1) { + if (offset + 12 > end) null else readUnsignedLong(offset + 4) + } else { + readUnsignedInt(offset + 4) + } + } + + private fun ByteArray.readUnsignedInt(offset: Int): Long = + ((this[offset].toLong() and 0xff) shl 24) or + ((this[offset + 1].toLong() and 0xff) shl 16) or + ((this[offset + 2].toLong() and 0xff) shl 8) or + (this[offset + 3].toLong() and 0xff) + + private fun ByteArray.readUnsignedLong(offset: Int): Long { + var value = 0L + repeat(8) { index -> value = (value shl 8) or (this[offset + index].toLong() and 0xff) } + return value + } +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrLiveSessionWarmupTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrLiveSessionWarmupTest.kt new file mode 100644 index 00000000..7e6b100b --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrLiveSessionWarmupTest.kt @@ -0,0 +1,142 @@ +package dev.typetype.server.services + +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertArrayEquals +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrBufferedRange +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaHeader +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaSegment +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant + +class SabrLiveSessionWarmupTest { + @Test + fun `warmup keeps bootstrap initialization and requests real live media`() = runTest { + val audio = format(140, audio = true, "audio/mp4") + val video = format(299, audio = false, "video/mp4") + val streamState = mockk(relaxed = true) + val session = mockk() + val rangeOverrides = mutableListOf?>() + var pumps = 0 + every { session.streamState } returns streamState + every { session.isComplete } returns false + every { session.isLive } answers { pumps > 0 } + every { session.isAtLiveEdge } returns false + every { session.liveHeadSequenceNumber } answers { if (pumps > 0) liveHeadSequence(pumps) else -1L } + every { streamState.isPostLiveDvr } returns false + every { streamState.isLive } answers { pumps > 0 } + every { streamState.liveHeadSequenceNumber } answers { if (pumps > 0) liveHeadSequence(pumps) else -1L } + every { streamState.liveHeadTimeMs } answers { if (pumps > 0) liveHeadTimeMs(pumps) else -1L } + every { streamState.getBufferedEndMs(audio) } returns 0L + every { streamState.getBufferedEndMs(video) } returns 0L + every { streamState.getMinBufferedEndMs() } returns 0L + every { streamState.setBufferedRangesOverride(any()) } answers { + rangeOverrides += firstArg?>() + } + val audioInit = mp4Box("ftyp", byteArrayOf(1)) + mp4Box("moov", byteArrayOf(2)) + val videoInit = mp4Box("ftyp", byteArrayOf(3)) + mp4Box("moov", byteArrayOf(4)) + val bootstrap = listOf( + segment(140, 0, 0L, -1L, audioInit + mediaFragment(5)), + segment(299, 0, 0L, -1L, videoInit + mediaFragment(6)), + ) + val liveMedia = listOf( + segment(140, TARGET_SEQUENCE, TARGET_TIME_MS, 2_000L, audioInit + mediaFragment(7)), + segment(299, TARGET_SEQUENCE, TARGET_TIME_MS, 2_000L, videoInit + mediaFragment(8)), + ) + every { session.pumpOnce(any()) } answers { + pumps++ + when (pumps) { + 1 -> bootstrap + 2 -> emptyList() + else -> liveMedia + } + } + val holder = holder(session, audio, video) + holder.markExpectedLive() + + SabrSessionPump(SabrSegmentCache()).ensureWarmed(holder, maxPumps = 8) + + assertEquals(3, pumps) + assertEquals(TARGET_SEQUENCE, holder.observedMediaSegment(audio)?.header?.sequenceNumber) + assertEquals(TARGET_SEQUENCE, holder.observedMediaSegment(video)?.header?.sequenceNumber) + assertArrayEquals(audioInit, holder.liveInitialization(audio)) + assertArrayEquals(videoInit, holder.liveInitialization(video)) + assertEquals(liveHeadTimeMs(pumps) - 20_000L, holder.resolvePlaybackStartMs(0L)) + val targetedRanges = rangeOverrides.filterNotNull().map { ranges -> ranges.map(SabrBufferedRange::summarize) } + val expectedRanges = listOf( + "itag=140:seq=1-3319:time=0+6640000:timescale=1000", + "itag=299:seq=1-3319:time=0+6640000:timescale=1000", + ) + assertEquals(listOf(expectedRanges, expectedRanges), targetedRanges) + assertNull(rangeOverrides.last()) + } + + private fun holder( + session: YoutubeSabrSession, + audio: YoutubeSabrFormat, + video: YoutubeSabrFormat, + ): SabrSessionHolder = SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session-token", + key = SabrSessionKey("video", "user", audio.itag, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + ) + + private fun format(itag: Int, audio: Boolean, mimeType: String): YoutubeSabrFormat = + mockk(relaxed = true) { + every { this@mockk.itag } returns itag + every { isAudio } returns audio + every { isVideo } returns !audio + every { this@mockk.mimeType } returns mimeType + every { lastModified } returns 1L + } + + private fun segment( + itag: Int, + sequence: Int, + startMs: Long, + durationMs: Long, + data: ByteArray, + ): SabrMediaSegment { + val header = mockk { + every { isInitSegment } returns false + every { this@mockk.itag } returns itag + every { sequenceNumber } returns sequence + every { this@mockk.startMs } returns startMs + every { this@mockk.durationMs } returns durationMs + } + return mockk { + every { this@mockk.header } returns header + every { this@mockk.data } returns data + } + } + + private fun mediaFragment(value: Byte): ByteArray = + mp4Box("moof", byteArrayOf(value)) + mp4Box("mdat", byteArrayOf(value)) + + private fun liveHeadSequence(pumps: Int): Long = LIVE_HEAD_SEQUENCE + (pumps - 1) * 3L + + private fun liveHeadTimeMs(pumps: Int): Long = LIVE_HEAD_TIME_MS + (pumps - 1) * 6_000L + + private fun mp4Box(type: String, payload: ByteArray): ByteArray { + val size = payload.size + 8 + return byteArrayOf(0, 0, 0, size.toByte()) + type.toByteArray(Charsets.US_ASCII) + payload + } + + private companion object { + const val LIVE_HEAD_SEQUENCE = 3_330L + const val LIVE_HEAD_TIME_MS = 6_660_000L + const val TARGET_SEQUENCE = 3_320 + const val TARGET_TIME_MS = 6_640_000L + } +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrPlaybackCachedSegmentLocatorTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrPlaybackCachedSegmentLocatorTest.kt new file mode 100644 index 00000000..de9c2f6c --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrPlaybackCachedSegmentLocatorTest.kt @@ -0,0 +1,57 @@ +package dev.typetype.server.services + +import io.mockk.coEvery +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState + +class SabrPlaybackCachedSegmentLocatorTest { + @Test + fun `finds the cached live segment covering a rounded audio boundary`() = runTest { + val store = mockk() + val holder = mockk() + val format = mockk() + val session = mockk() + val state = mockk() + every { format.itag } returns 140 + every { holder.observedMediaSegment(format) } returns null + every { holder.session } returns session + every { session.streamState } returns state + every { state.getSegmentStartMs(format, any()) } returns 0L + every { state.getSegmentEndMs(format, any()) } returns 2_000L + coEvery { store.cachedSegment(holder, any()) } answers { + val sequence = secondArg().sequenceNumber + when (sequence) { + 1841446 -> cached(sequence, 3_682_888_446L) + 1841447 -> cached(sequence, 3_682_890_443L) + else -> null + } + } + + val segment = store.findCachedPlaybackMediaAt( + holder, + format, + targetMs = 3_682_888_446L, + predictedSequence = 1841447, + ) + + assertEquals(1841446, segment?.sequence) + } + + private fun cached(sequence: Int, startMs: Long): CachedSabrSegment = CachedSabrSegment( + itag = 140, + sequence = sequence, + init = false, + startMs = startMs, + durationMs = -1L, + mimeType = "audio/mp4", + bytesBase64 = "AA==", + byteLength = 1, + ) +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrPlaybackLiveGapServiceTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrPlaybackLiveGapServiceTest.kt new file mode 100644 index 00000000..0339b356 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrPlaybackLiveGapServiceTest.kt @@ -0,0 +1,86 @@ +package dev.typetype.server.services + +import io.mockk.coEvery +import io.mockk.every +import io.mockk.mockk +import kotlinx.coroutines.test.runTest +import org.junit.jupiter.api.Assertions.assertArrayEquals +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant +import java.util.Base64 + +class SabrPlaybackLiveGapServiceTest { + @Test + fun `serves the next live segment when YouTube skips a sequence`() = runTest { + val video = format(299, isAudio = false) + val audio = format(140, isAudio = true) + val session = mockk(relaxed = true) + val state = mockk(relaxed = true) + every { session.streamState } returns state + every { session.isLive } returns true + every { state.isLive } returns true + every { state.getSegmentStartMs(video, 92) } returns 475_000L + val holder = holder(session, audio, video) + val replacement = cachedSegment(299, 93, 480_000L, byteArrayOf(1, 2, 3)) + val store = mockk() + coEvery { store.cachedSegment(holder, any()) } answers { + secondArg().takeIf { it.sequenceNumber == 93 }?.let { replacement } + } + every { store.requestSegmentDemand(holder, any(), 0L) } returns Unit + + val result = SabrPlaybackSessionService(store).fetchMedia( + holder = holder, + format = video, + sequence = 92, + timeoutMs = 1_000L, + generation = 0L, + ) as SabrPlaybackSegmentResult.Ready + + assertArrayEquals(byteArrayOf(1, 2, 3), result.bytes) + assertEquals(93, holder.lastServedSequence(video)) + } + + private fun holder( + session: YoutubeSabrSession, + audio: YoutubeSabrFormat, + video: YoutubeSabrFormat, + ): SabrSessionHolder = SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session-token", + key = SabrSessionKey("video", "user", audio.itag, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + ) + + private fun format(itag: Int, isAudio: Boolean): YoutubeSabrFormat { + val format = mockk(relaxed = true) + every { format.itag } returns itag + every { format.isAudio } returns isAudio + every { format.mimeType } returns if (isAudio) "audio/mp4" else "video/mp4" + return format + } + + private fun cachedSegment( + itag: Int, + sequence: Int, + startMs: Long, + bytes: ByteArray, + ): CachedSabrSegment = CachedSabrSegment( + itag = itag, + sequence = sequence, + init = false, + startMs = startMs, + durationMs = 5_000L, + mimeType = "video/mp4", + bytesBase64 = Base64.getEncoder().encodeToString(bytes), + byteLength = bytes.size, + ) +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrPlayerContextRecoveryTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrPlayerContextRecoveryTest.kt new file mode 100644 index 00000000..e1f25a8f --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrPlayerContextRecoveryTest.kt @@ -0,0 +1,131 @@ +package dev.typetype.server.services + +import io.mockk.every +import io.mockk.mockk +import io.mockk.verify +import kotlinx.coroutines.CancellationException +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertInstanceOf +import org.junit.jupiter.api.Assertions.assertSame +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Assertions.assertThrows +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.exceptions.AntiBotException +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrProtocolException +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrClientProfile +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import java.io.IOException + +class SabrPlayerContextRecoveryTest { + @Test + fun `refreshes rejected player context once`() { + val tokenClient = mockk() + val initial = token("old-visitor") + val refreshed = token("fresh-visitor") + val info = mockk() + every { tokenClient.fetch("video", forceRefresh = true, refreshVideo = false) } returns refreshed + val visitors = mutableListOf() + val probe = SabrPlayerInfoProbe { _, profile, token -> + visitors += token.visitorData + if (token === initial) throw missingStreamingData(profile) + info + } + + val result = SabrPlayerContextRecovery("video", initial, tokenClient, probe) + .fetch(YoutubeSabrClientProfile.WEB) + + val success = assertInstanceOf(SabrPlayerProbeResult.Success::class.java, result) + assertSame(info, success.info) + assertSame(refreshed, success.token) + assertTrue(success.contextRefreshed) + assertEquals(listOf("old-visitor", "fresh-visitor"), visitors) + verify(exactly = 1) { tokenClient.fetch("video", forceRefresh = true, refreshVideo = false) } + } + + @Test + fun `refreshes typed anti bot rejection`() { + val tokenClient = mockk() + val refreshed = token("fresh-visitor") + every { tokenClient.fetch("video", forceRefresh = true, refreshVideo = false) } returns refreshed + var attempts = 0 + val probe = SabrPlayerInfoProbe { _, _, _ -> + attempts++ + if (attempts == 1) throw AntiBotException("Sign in to confirm you're not a bot") + mockk() + } + + val result = SabrPlayerContextRecovery("video", token("old-visitor"), tokenClient, probe) + .fetch(YoutubeSabrClientProfile.MWEB) + + assertInstanceOf(SabrPlayerProbeResult.Success::class.java, result) + assertEquals(2, attempts) + } + + @Test + fun `does not refresh network failure`() { + val tokenClient = mockk(relaxed = true) + val failure = IOException("network unavailable") + val probe = SabrPlayerInfoProbe { _, _, _ -> throw failure } + + val result = SabrPlayerContextRecovery("video", token("visitor"), tokenClient, probe) + .fetch(YoutubeSabrClientProfile.WEB) + + val failed = assertInstanceOf(SabrPlayerProbeResult.Failure::class.java, result) + assertSame(failure, failed.error) + verify(exactly = 0) { tokenClient.fetch(any(), any(), any()) } + } + + @Test + fun `does not refresh unrelated protocol failure`() { + val tokenClient = mockk(relaxed = true) + val failure = SabrProtocolException("Missing serverAbrStreamingUrl") + val probe = SabrPlayerInfoProbe { _, _, _ -> throw failure } + + val result = SabrPlayerContextRecovery("video", token("visitor"), tokenClient, probe) + .fetch(YoutubeSabrClientProfile.WEB) + + val failed = assertInstanceOf(SabrPlayerProbeResult.Failure::class.java, result) + assertSame(failure, failed.error) + verify(exactly = 0) { tokenClient.fetch(any(), any(), any()) } + } + + @Test + fun `propagates cancellation without refreshing`() { + val tokenClient = mockk(relaxed = true) + val probe = SabrPlayerInfoProbe { _, _, _ -> throw CancellationException("cancelled") } + val recovery = SabrPlayerContextRecovery("video", token("visitor"), tokenClient, probe) + + assertThrows(CancellationException::class.java) { + recovery.fetch(YoutubeSabrClientProfile.WEB) + } + verify(exactly = 0) { tokenClient.fetch(any(), any(), any()) } + } + + @Test + fun `shares one refresh budget across profiles`() { + val tokenClient = mockk() + every { tokenClient.fetch("video", forceRefresh = true, refreshVideo = false) } returns token("fresh") + val probe = SabrPlayerInfoProbe { _, profile, _ -> throw missingStreamingData(profile) } + val recovery = SabrPlayerContextRecovery("video", token("old"), tokenClient, probe) + + val web = recovery.fetch(YoutubeSabrClientProfile.WEB) + val mweb = recovery.fetch(YoutubeSabrClientProfile.MWEB) + + val failedWeb = assertInstanceOf(SabrPlayerProbeResult.Failure::class.java, web) + assertTrue(failedWeb.contextRefreshAttempted) + assertInstanceOf(SabrPlayerProbeResult.Failure::class.java, mweb) + verify(exactly = 1) { tokenClient.fetch("video", forceRefresh = true, refreshVideo = false) } + } + + private fun missingStreamingData(profile: YoutubeSabrClientProfile): SabrProtocolException = + SabrProtocolException("Player response has no streamingData for $profile") + + private fun token(visitorData: String): SabrTokenBundle = SabrTokenBundle( + videoId = "video", + visitorBoundPoToken = "player-$visitorData", + visitorBoundPoTokenBytes = byteArrayOf(1), + visitorData = visitorData, + videoBoundPoToken = "video-$visitorData", + videoBoundPoTokenBytes = byteArrayOf(2), + ) +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrPreparedInfoCacheTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrPreparedInfoCacheTest.kt index 2761a6ad..84e17fd8 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrPreparedInfoCacheTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrPreparedInfoCacheTest.kt @@ -7,9 +7,11 @@ import io.mockk.mockk import io.mockk.verify import kotlinx.coroutines.test.runTest import org.junit.jupiter.api.Assertions.assertFalse +import org.junit.jupiter.api.Assertions.assertEquals import org.junit.jupiter.api.Assertions.assertSame import org.junit.jupiter.api.Assertions.assertTrue import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrClientProfile import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo @@ -76,7 +78,7 @@ class SabrPreparedInfoCacheTest { val result = fetcher.fetchInfo("video") - assertSame(info, result?.info) + assertEquals(info.formats, result?.info?.formats) assertSame(token, result?.initialToken) coVerify(exactly = 1) { sessionClient.fetchPlaybackSession("video") } verify(exactly = 0) { tokenClient.fetch(any(), any(), any()) } @@ -94,7 +96,7 @@ class SabrPreparedInfoCacheTest { val result = fetcher.fetchInfo("video") - assertSame(info, result?.info) + assertEquals(info.formats, result?.info?.formats) assertSame(token, result?.initialToken) verify(exactly = 1) { tokenClient.fetch("video", forceRefresh = false, refreshVideo = false) } } @@ -112,19 +114,53 @@ class SabrPreparedInfoCacheTest { val result = fetcher.fetchInfo("video") - assertSame(info, result?.info) + assertEquals(info.formats, result?.info?.formats) assertSame(fallback, result?.initialToken) verify(exactly = 1) { tokenClient.fetch("video", forceRefresh = false, refreshVideo = false) } } + @Test + fun `info fetcher keeps refreshed player context and metadata together`() = runTest { + val tokenClient = mockk() + val initial = token(visitorData = "old-visitor") + val refreshed = token(visitorData = "fresh-visitor") + val info = info(listOf(format(isAudio = true), format(isAudio = false)), "fresh-visitor") + every { tokenClient.fetch("video", forceRefresh = false, refreshVideo = false) } returns initial + every { tokenClient.fetch("video", forceRefresh = true, refreshVideo = false) } returns refreshed + val probe = SabrPlayerInfoProbe { _, profile, token -> + if (token === initial) { + throw org.schabi.newpipe.extractor.services.youtube.sabr.SabrProtocolException( + "Player response has no streamingData for $profile", + ) + } + info + } + val fetcher = SabrInfoFetcher(tokenClient, playerInfoProbe = probe) + + val result = fetcher.fetchInfo("video") + + assertSame(info, result?.info) + assertSame(refreshed, result?.initialToken) + verify(exactly = 1) { tokenClient.fetch("video", forceRefresh = true, refreshVideo = false) } + } + private fun preparedInfo(formats: List): SabrPreparedInfo { return SabrPreparedInfo(info(formats), token()) } - private fun info(formats: List): YoutubeSabrInfo { + private fun info( + formats: List, + visitorData: String = "visitor-data", + ): YoutubeSabrInfo { val info = mockk() every { info.formats } returns formats - every { info.visitorData } returns "visitor-data" + every { info.visitorData } returns visitorData + every { info.profile } returns YoutubeSabrClientProfile.MWEB + every { info.videoId } returns "video" + every { info.cpn } returns "cpn" + every { info.clientVersion } returns "2.20260718.00.00" + every { info.serverAbrStreamingUrl } returns "https://example.com/sabr" + every { info.videoPlaybackUstreamerConfig } returns "config" return info } diff --git a/src/test/kotlin/dev/typetype/server/services/SabrProbeDiagnostics.kt b/src/test/kotlin/dev/typetype/server/services/SabrProbeDiagnostics.kt index fbe4fc88..4caf767c 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrProbeDiagnostics.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrProbeDiagnostics.kt @@ -10,8 +10,7 @@ internal fun printSabrProbeFormat(label: String, format: YoutubeSabrFormat): Uni "size=${format.width}x${format.height} bitrate=${format.bitrate} mime=${format.mimeType} " + "quality=${format.qualityLabel} audioTrack=${format.audioTrackId} " + "xtags=${format.xtags} drc=${format.isDrc} original=${format.isOriginalAudio} " + - "approxDurationMs=${format.approxDurationMs} init=${!format.initializationUrl.isNullOrBlank()} " + - "initRange=${format.initRangeStart}-${format.initRangeEnd}" + "approxDurationMs=${format.approxDurationMs}" ) } diff --git a/src/test/kotlin/dev/typetype/server/services/SabrProbeTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrProbeTest.kt index 3e626713..e7dc7e05 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrProbeTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrProbeTest.kt @@ -38,14 +38,9 @@ class SabrProbeTest { println("\n========== SABR probe: $videoId ==========") try { val token = tokenClient.fetch(videoId) ?: error("No SABR token") - val info = YoutubeSabrProbe.fetchSabrInfo( - videoId, - profile, - loc, - country, - token.visitorPoToken, - token.visitorData, - ) + val info = TypetypeYoutubeSessionPoTokenProvider.withToken(token) { + YoutubeSabrProbe.fetchSabrInfo(videoId, profile, loc, country) + } println("serverAbrStreamingUrl present: ${!info.serverAbrStreamingUrl.isNullOrEmpty()}") println("videoPlaybackUstreamerConfig present: ${!info.videoPlaybackUstreamerConfig.isNullOrEmpty()}") println("--- formats (itag | A/V | height×width | bitrate | mime | audioTrack | approxDurMs) ---") diff --git a/src/test/kotlin/dev/typetype/server/services/SabrPumpPolicyTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrPumpPolicyTest.kt new file mode 100644 index 00000000..cdb17d1f --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrPumpPolicyTest.kt @@ -0,0 +1,13 @@ +package dev.typetype.server.services + +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Test + +class SabrPumpPolicyTest { + @Test + fun `live demand only waits at the actual edge`() { + assertEquals(100L, SabrPumpPolicy.demandDelayMs(100L, 0L, activeLive = true, futureDemand = false)) + assertEquals(LIVE_EDGE_POLL_MS, SabrPumpPolicy.demandDelayMs(100L, 0L, activeLive = true, futureDemand = true)) + assertEquals(LIVE_EDGE_POLL_MS, SabrPumpPolicy.demandDelayMs(100L, 0L, activeLive = true, futureDemand = null)) + } +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrSeekRepositionPumpTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrSeekRepositionPumpTest.kt index e838a2d1..5125f12b 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrSeekRepositionPumpTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrSeekRepositionPumpTest.kt @@ -6,6 +6,7 @@ import io.mockk.verify import kotlinx.coroutines.test.runTest import org.junit.jupiter.api.Assertions.assertEquals import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaHeader import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaSegment import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat @@ -60,13 +61,15 @@ class SabrSeekRepositionPumpTest { every { session.getCachedSegment(any()) } returns null every { session.pumpOnceStreamingForDemand(any(), request) } returns mockk(relaxed = true) val holder = holder(session, audio, video) + holder.setRequestedSeekTimeMs(120_321L) holder.requestSegmentDemand(request) holder.requestForwardSeek(request) var rounds = 0 SabrSessionPumpLoop().run({ rounds++ < 1 }, holder, intervalMs = 0L) - verify(exactly = 1) { session.prepareForForwardJump(request) } + verify(exactly = 1) { session.prepareForForwardJump(request, 120_321L) } + verify(exactly = 0) { session.prepareForForwardJump(request) } verify(exactly = 0) { session.pumpOnceStreaming(any()) } verify(exactly = 1) { session.pumpOnceStreamingForDemand(any(), request) } } finally { @@ -102,6 +105,68 @@ class SabrSeekRepositionPumpTest { } } + @Test + fun `historical live demand rewinds from observed media edge`() = runTest { + SabrSegmentDemandTracker.clearAll() + try { + val audio = format(140, true) + val video = format(299, false) + val request = SabrSegmentRequest.media(audio, 3_073) + val session = mockk(relaxed = true) + val state = mockk(relaxed = true) + every { session.streamState } returns state + every { session.isLive } returns true + every { state.isLive } returns true + every { state.isPostLiveDvr } returns false + every { session.getCachedSegment(any()) } returns null + every { session.pumpOnceStreamingForDemand(any(), request) } returns mockk(relaxed = true) + val holder = holder(session, audio, video) + holder.observeMediaSegment(mediaSegment(audio.itag, sequence = 3_076)) + holder.requestSegmentDemand(request) + var rounds = 0 + + SabrSessionPumpLoop().run({ rounds++ < 1 }, holder, intervalMs = 0L) + + verify(exactly = 1) { session.prepareForRewind(request) } + verify(exactly = 1) { session.prepareForMediaSegment(request) } + verify(exactly = 1) { state.setBufferedRangesOverride(null) } + verify(exactly = 1) { session.pumpOnceStreamingForDemand(any(), request) } + } finally { + SabrSegmentDemandTracker.clearAll() + } + } + + @Test + fun `historical live seek preserves its exact player position`() = runTest { + SabrSegmentDemandTracker.clearAll() + try { + val audio = format(140, true) + val video = format(299, false) + val request = SabrSegmentRequest.media(audio, 3_073) + val session = mockk(relaxed = true) + val state = mockk(relaxed = true) + every { session.streamState } returns state + every { session.isLive } returns true + every { state.isLive } returns true + every { state.isPostLiveDvr } returns false + every { session.getCachedSegment(any()) } returns null + every { session.pumpOnceStreamingForDemand(any(), request) } returns mockk(relaxed = true) + val holder = holder(session, audio, video) + holder.observeMediaSegment(mediaSegment(audio.itag, sequence = 3_076)) + holder.setRequestedSeekTimeMs(15_365_500L) + holder.requestSegmentDemand(request) + var rounds = 0 + + SabrSessionPumpLoop().run({ rounds++ < 1 }, holder, intervalMs = 0L) + + verify(exactly = 1) { session.prepareForRewind(request, 15_365_500L) } + verify(exactly = 0) { session.prepareForRewind(request) } + verify(exactly = 1) { session.pumpOnceStreamingForDemand(any(), request) } + } finally { + SabrSegmentDemandTracker.clearAll() + } + } + private fun holder( session: YoutubeSabrSession, audio: YoutubeSabrFormat, @@ -122,6 +187,18 @@ class SabrSeekRepositionPumpTest { every { format.isAudio } returns isAudio every { format.audioTrackId } returns null every { format.bitrate } returns if (isAudio) 128_000 else 2_000_000 + every { format.lastModified } returns 0L + every { format.xtags } returns "" return format } + + private fun mediaSegment(itag: Int, sequence: Int): SabrMediaSegment { + val header = mockk(relaxed = true) + every { header.itag } returns itag + every { header.sequenceNumber } returns sequence + every { header.startMs } returns sequence * 5_000L + every { header.durationMs } returns 5_000L + every { header.isInitSegment } returns false + return mockk { every { this@mockk.header } returns header } + } } diff --git a/src/test/kotlin/dev/typetype/server/services/SabrSegmentCacheTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrSegmentCacheTest.kt index e823c70e..51484ca6 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrSegmentCacheTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrSegmentCacheTest.kt @@ -36,6 +36,22 @@ class SabrSegmentCacheTest { assertArrayEquals(byteArrayOf(1, 2, 3), cached.bytes) } + @Test + fun `live cache separates initialization from mp4 media`() { + val segmentCache = SabrSegmentCache() + val audio = format(140, isAudio = true) + val video = format(137, isAudio = false) + val holder = holder(audio, video).also { it.markExpectedLive() } + val initialization = mp4Box("ftyp", 1) + mp4Box("moov", 2) + val media = mp4Box("moof", 3) + mp4Box("mdat", 4) + val request = SabrSegmentRequest.media(video, 7) + + segmentCache.put(holder, segment(137, 7, initialization + media)) + + assertArrayEquals(initialization, holder.liveInitialization(video)) + assertArrayEquals(media, requireNotNull(segmentCache.get(holder, request)).bytes) + } + private fun holder(audio: YoutubeSabrFormat, video: YoutubeSabrFormat): SabrSessionHolder { val session = mockk() val state = mockk() @@ -76,4 +92,7 @@ class SabrSegmentCacheTest { every { segment.length } returns bytes.size return segment } + + private fun mp4Box(type: String, value: Byte): ByteArray = + byteArrayOf(0, 0, 0, 9) + type.toByteArray(Charsets.US_ASCII) + byteArrayOf(value) } diff --git a/src/test/kotlin/dev/typetype/server/services/SabrSegmentDemandResolutionTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrSegmentDemandResolutionTest.kt new file mode 100644 index 00000000..9454dba1 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrSegmentDemandResolutionTest.kt @@ -0,0 +1,123 @@ +package dev.typetype.server.services + +import io.mockk.every +import io.mockk.mockk +import io.mockk.verify +import org.junit.jupiter.api.AfterEach +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Assertions.assertSame +import org.junit.jupiter.api.Assertions.assertTrue +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaHeader +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaSegment +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant + +class SabrSegmentDemandResolutionTest { + @AfterEach + fun clearDemands(): Unit = SabrSegmentDemandTracker.clearAll() + + @Test + fun `resolves skipped live sequence against requested segment time`() { + val format = mockk() + val session = mockk(relaxed = true) + val state = mockk(relaxed = true) + val request = SabrSegmentRequest.media(format, 1_904) + val replacement = segment(sequence = 1_905, startMs = 9_520_445L, durationMs = 5_000L) + every { format.itag } returns 140 + every { format.isAudio } returns true + every { session.streamState } returns state + every { state.getSegmentStartMs(format, 1_904) } returns 9_520_445L + every { session.getCachedSegment(any()) } answers { + firstArg().takeIf { it.sequenceNumber == 1_905 }?.let { replacement } + } + val holder = holder(session, format) + holder.requestSegmentDemand(request) + val identity = requireNotNull(holder.segmentDemandIdentity(request)) + + assertTrue(holder.resolveSegmentDemand(request, identity)) + + assertNull(holder.pendingSegmentDemandSummary()) + assertSame(replacement, holder.observedMediaSegment(format)) + verify(exactly = 1) { state.jumpBufferedTo(format, 1_905) } + } + + @Test + fun `resolves live sequence gap from the next available segment`() { + val format = mockk() + val session = mockk(relaxed = true) + val state = mockk(relaxed = true) + val request = SabrSegmentRequest.media(format, 92) + val replacement = segment(itag = 299, sequence = 93, startMs = 480_000L, durationMs = 5_000L) + every { format.itag } returns 299 + every { format.isAudio } returns false + every { session.streamState } returns state + every { session.isLive } returns true + every { state.isLive } returns true + every { state.getSegmentStartMs(format, 92) } returns 475_000L + every { session.getCachedSegment(any()) } answers { + firstArg().takeIf { it.sequenceNumber == 93 }?.let { replacement } + } + val holder = holder(session, format) + holder.requestSegmentDemand(request) + val identity = requireNotNull(holder.segmentDemandIdentity(request)) + + assertTrue(holder.resolveSegmentDemand(request, identity)) + + assertNull(holder.pendingSegmentDemandSummary()) + assertSame(replacement, holder.observedMediaSegment(format)) + verify(exactly = 1) { state.jumpBufferedTo(format, 93) } + } + + @Test + fun `resolved requested segment advances complete media anchor`() { + val format = mockk() + val session = mockk(relaxed = true) + val request = SabrSegmentRequest.media(format, 2_752) + val cached = segment(sequence = 2_752, startMs = 13_757_066L, durationMs = 4_997L) + var cacheReady = false + every { format.itag } returns 140 + every { format.isAudio } returns true + every { session.getCachedSegment(request) } answers { cached.takeIf { cacheReady } } + val holder = holder(session, format) + holder.requestSegmentDemand(request) + val identity = requireNotNull(holder.segmentDemandIdentity(request)) + cacheReady = true + + assertTrue(holder.resolveSegmentDemand(request, identity)) + + assertSame(cached, holder.observedMediaSegment(format)) + assertNull(holder.pendingSegmentDemandSummary()) + } + + private fun holder(session: YoutubeSabrSession, audio: YoutubeSabrFormat): SabrSessionHolder { + val video = mockk() + every { video.itag } returns 299 + every { video.isAudio } returns false + return SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session-token", + key = SabrSessionKey("video", "user", audio.itag, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + ) + } + + private fun segment(sequence: Int, startMs: Long, durationMs: Long, itag: Int = 140): SabrMediaSegment { + val header = mockk() + every { header.sequenceNumber } returns sequence + every { header.startMs } returns startMs + every { header.durationMs } returns durationMs + every { header.itag } returns itag + every { header.isInitSegment } returns false + val segment = mockk() + every { segment.header } returns header + return segment + } +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrSegmentDemandTrackerTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrSegmentDemandTrackerTest.kt index 8cdd8089..d9a1194d 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrSegmentDemandTrackerTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrSegmentDemandTrackerTest.kt @@ -56,6 +56,20 @@ class SabrSegmentDemandTrackerTest { } } + @Test + fun `active live keeps a missing track demand behind the current head`() { + withTracker { holder -> + val request = SabrSegmentRequest.media(holder.videoFormat, 85) + holder.markExpectedLive() + every { holder.session.isBeyondEnd(request) } returns true + + holder.requestSegmentDemand(request) + + assertEquals("299:85", holder.pendingSegmentDemandSummary()) + assertEquals(request, holder.nextSegmentDemand()) + } + } + @Test fun `duplicate registration preserves exact demand identity`() { withTracker { holder -> diff --git a/src/test/kotlin/dev/typetype/server/services/SabrSessionPlayerContextTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrSessionPlayerContextTest.kt new file mode 100644 index 00000000..3f7d2f86 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/SabrSessionPlayerContextTest.kt @@ -0,0 +1,73 @@ +package dev.typetype.server.services + +import io.mockk.every +import io.mockk.mockk +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.localization.ContentCountry +import org.schabi.newpipe.extractor.localization.Localization +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrInfo +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrSession +import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrStreamState +import java.time.Instant + +class SabrSessionPlayerContextTest { + @Test + fun `holder exposes only its own token during extractor calls`() { + val first = holder(token("first", "first-token")) + val second = holder(token("second", "second-token")) + + first.withPlayerContext { + assertEquals("first", currentToken()?.visitorData) + second.withPlayerContext { + assertEquals("second", currentToken()?.visitorData) + } + assertEquals("first", currentToken()?.visitorData) + } + + assertNull(currentToken()) + } + + private fun holder(token: SabrTokenBundle): SabrSessionHolder { + val audio = format(140, audio = true) + val video = format(137, audio = false) + val streamState = mockk(relaxed = true) + val session = mockk { + every { this@mockk.streamState } returns streamState + } + return SabrSessionHolder( + session = session, + info = mockk(), + audioFormat = audio, + videoFormat = video, + sessionToken = "session-token", + key = SabrSessionKey("video", "user", audio.itag, null, video.itag, 0L), + lastRequestAt = Instant.EPOCH, + playerContextToken = token, + ) + } + + private fun format(itag: Int, audio: Boolean): YoutubeSabrFormat = mockk(relaxed = true) { + every { this@mockk.itag } returns itag + every { isAudio } returns audio + every { isVideo } returns !audio + } + + private fun currentToken() = TypetypeYoutubeSessionPoTokenProvider.getSessionPoToken( + "MWEB", + Localization("en", "US"), + ContentCountry("US"), + false, + ) + + private fun token(visitorData: String, playerToken: String) = SabrTokenBundle( + videoId = "video", + visitorBoundPoToken = playerToken, + visitorBoundPoTokenBytes = byteArrayOf(1), + visitorData = visitorData, + videoBoundPoToken = "video-token", + videoBoundPoTokenBytes = byteArrayOf(2), + ) +} diff --git a/src/test/kotlin/dev/typetype/server/services/SabrSessionPumpLoopTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrSessionPumpLoopTest.kt index 10c4fdd2..687491ae 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrSessionPumpLoopTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrSessionPumpLoopTest.kt @@ -90,6 +90,30 @@ class SabrSessionPumpLoopTest { assertEquals(100L, testScheduler.currentTime) } + @Test + fun `active live pump waits for a reader demand after warmup`() = runTest { + val audio = format(140, isAudio = true) + val video = format(299, isAudio = false) + val session = mockk(relaxed = true) + val streamState = mockk(relaxed = true) + every { session.requestNumber } returns 42 + every { session.streamState } returns streamState + every { session.isLive } returns true + every { streamState.isPostLiveDvr } returns false + every { streamState.isLive } returns true + val holder = holder(session, audio, video) + holder.setReaderPosition(audio, 1_000_000L) + holder.setReaderPosition(video, 1_000_000L) + var rounds = 0 + + SabrSessionPump().pumpLoop({ rounds++ == 0 }, holder, intervalMs = 100L) + + verify(exactly = 0) { session.pumpOnceStreaming(any()) } + verify(exactly = 1) { session.setPlayHeadMs(match { it in 988_000L..998_000L }) } + verify(exactly = 1) { session.evictPlayed() } + assertEquals(LIVE_EDGE_POLL_MS, testScheduler.currentTime) + } + @Test fun `non target media response keeps demand loop paced`() = runTest { SabrSegmentDemandTracker.clearAll() @@ -113,7 +137,7 @@ class SabrSessionPumpLoopTest { } every { streamState.getMinBufferedEndMs() } returns 487_134L every { streamState.getBufferedEndMs(video) } returns 491_203L - every { streamState.getSegmentStartMs(video, 98) } returns 487_134L + every { streamState.getSegmentStartMs(video, 98) } returns 491_203L every { streamState.setPlayerTimeMs(any()) } returns Unit every { streamState.jumpBufferedTo(video, 101) } returns Unit every { session.pumpOnceStreamingForDemand(any(), request) } returns demandResult(5, 0) @@ -232,6 +256,8 @@ class SabrSessionPumpLoopTest { every { header.sequenceNumber } returns sequence every { header.startMs } returns startMs every { header.durationMs } returns durationMs + every { header.itag } returns 299 + every { header.isInitSegment } returns false val segment = mockk() every { segment.header } returns header return segment diff --git a/src/test/kotlin/dev/typetype/server/services/SabrTransientDemandFailureTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrTransientDemandFailureTest.kt index 33ca9a4d..b9613340 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrTransientDemandFailureTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrTransientDemandFailureTest.kt @@ -8,6 +8,7 @@ import kotlinx.coroutines.test.runTest import org.junit.jupiter.api.Assertions.assertEquals import org.junit.jupiter.api.Assertions.assertNull import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaHeader import org.schabi.newpipe.extractor.services.youtube.sabr.SabrMediaSegment import org.schabi.newpipe.extractor.services.youtube.sabr.SabrSegmentRequest import org.schabi.newpipe.extractor.services.youtube.sabr.YoutubeSabrFormat @@ -26,7 +27,7 @@ class SabrTransientDemandFailureTest { val audio = format(140, true) val video = format(137, false) val request = SabrSegmentRequest.media(audio, 50) - val segment = mockk() + val segment = mediaSegment(140, 50, 499_414L) val session = mockk(relaxed = true) val streamState = mockk(relaxed = true) var cached = false @@ -64,7 +65,7 @@ class SabrTransientDemandFailureTest { val audio = format(140, true) val video = format(137, false) val request = SabrSegmentRequest.media(audio, 39) - val segment = mockk() + val segment = mediaSegment(140, 39, 379_414L) val session = mockk(relaxed = true) val streamState = mockk(relaxed = true) val result = mockk() @@ -237,4 +238,15 @@ class SabrTransientDemandFailureTest { every { result.targetTrackSegmentCount } returns targetTrackSegmentCount return result } + + private fun mediaSegment(itag: Int, sequence: Int, startMs: Long): SabrMediaSegment { + val header = mockk() + every { header.itag } returns itag + every { header.sequenceNumber } returns sequence + every { header.startMs } returns startMs + every { header.isInitSegment } returns false + val segment = mockk() + every { segment.header } returns header + return segment + } } diff --git a/src/test/kotlin/dev/typetype/server/services/SabrUnauthorizedResponseRecoveryTest.kt b/src/test/kotlin/dev/typetype/server/services/SabrUnauthorizedResponseRecoveryTest.kt index 71a73239..8e87c440 100644 --- a/src/test/kotlin/dev/typetype/server/services/SabrUnauthorizedResponseRecoveryTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/SabrUnauthorizedResponseRecoveryTest.kt @@ -35,12 +35,14 @@ class SabrUnauthorizedResponseRecoveryTest { @Test fun appliesFreshVideoTokenFromSameVisitorSession(): Unit { val freshToken = byteArrayOf(2) + val refreshed = bundle("visitor-a", freshToken) val (holder, state) = holder(visitorData = "visitor-a", currentToken = byteArrayOf(1)) - val recovery = SabrUnauthorizedResponseRecovery { bundle("visitor-a", freshToken) } + val recovery = SabrUnauthorizedResponseRecovery { refreshed } recovery.verify(holder) verify(exactly = 1) { state.setPoToken(match { it.contentEquals(freshToken) }) } + verify(exactly = 1) { holder.playerContextToken = refreshed } } private fun holder( @@ -50,7 +52,7 @@ class SabrUnauthorizedResponseRecoveryTest { val state = mockk(relaxed = true) val session = mockk(relaxed = true) val info = mockk() - val holder = mockk() + val holder = mockk(relaxed = true) every { state.poToken } returns currentToken every { session.streamState } returns state every { session.diagnosticTrace } returns "response n=4 http=403 segments=count=0" diff --git a/src/test/kotlin/dev/typetype/server/services/TypetypeTokenYoutubeSessionClientTest.kt b/src/test/kotlin/dev/typetype/server/services/TypetypeTokenYoutubeSessionClientTest.kt index 14e7ddbd..ff5366b5 100644 --- a/src/test/kotlin/dev/typetype/server/services/TypetypeTokenYoutubeSessionClientTest.kt +++ b/src/test/kotlin/dev/typetype/server/services/TypetypeTokenYoutubeSessionClientTest.kt @@ -43,7 +43,7 @@ class TypetypeTokenYoutubeSessionClientTest { "visitorData":"visitor-data", "poToken":"AQ", "streamingPot":"Ag", - "serverAbrStreamingUrl":"https://example.com/sabr", + "serverAbrStreamingUrl":"https://example.com/sabr?cver=2.20260205.04.01", "rawServerAbrStreamingUrl":"https://example.com/raw-sabr", "hlsManifestUrl":"https://example.com/live.m3u8", "videoPlaybackUstreamerConfig":"ustreamer-config", @@ -106,9 +106,10 @@ class TypetypeTokenYoutubeSessionClientTest { assertNotNull(info) assertEquals("video-id", info?.videoId) - assertEquals("https://example.com/raw-sabr", info?.serverAbrStreamingUrl) + assertEquals("https://example.com/sabr?cver=2.20260205.04.01", info?.serverAbrStreamingUrl) + assertEquals("2.20260205.04.01", info?.clientVersion) assertEquals(2, info?.formats?.size) - assertTrue(info?.formats?.any { it.itag == 137 && it.initRangeEnd == 9281L } == true) + assertTrue(info?.formats?.any { it.itag == 137 && it.contentLength == 1_000_000L } == true) assertTrue(info?.formats?.any { it.itag == 140 && it.audioTrackId == "fr-FR.4" } == true) val session = client.fetchPlaybackSession("video-id") diff --git a/src/test/kotlin/dev/typetype/server/services/TypetypeYoutubeSessionPoTokenProviderTest.kt b/src/test/kotlin/dev/typetype/server/services/TypetypeYoutubeSessionPoTokenProviderTest.kt new file mode 100644 index 00000000..01a8bf23 --- /dev/null +++ b/src/test/kotlin/dev/typetype/server/services/TypetypeYoutubeSessionPoTokenProviderTest.kt @@ -0,0 +1,58 @@ +package dev.typetype.server.services + +import org.junit.jupiter.api.Assertions.assertEquals +import org.junit.jupiter.api.Assertions.assertNull +import org.junit.jupiter.api.Test +import org.schabi.newpipe.extractor.localization.ContentCountry +import org.schabi.newpipe.extractor.localization.Localization + +class TypetypeYoutubeSessionPoTokenProviderTest { + @Test + fun `exposes the session token only inside its scope`() { + TypetypeYoutubeSessionPoTokenProvider.withToken(token("visitor", "player-token")) { + assertEquals("visitor", currentToken()?.visitorData) + assertEquals("player-token", currentToken()?.poToken) + } + + assertNull(currentToken()) + } + + @Test + fun `restores the outer token after a nested scope`() { + TypetypeYoutubeSessionPoTokenProvider.withToken(token("outer", "outer-token")) { + TypetypeYoutubeSessionPoTokenProvider.withToken(token("inner", "inner-token")) { + assertEquals("inner", currentToken()?.visitorData) + } + assertEquals("outer", currentToken()?.visitorData) + } + + assertNull(currentToken()) + } + + @Test + fun `clears the token when the scoped call fails`() { + runCatching { + TypetypeYoutubeSessionPoTokenProvider.withToken(token("visitor", "player-token")) { + error("failed") + } + } + + assertNull(currentToken()) + } + + private fun currentToken() = TypetypeYoutubeSessionPoTokenProvider.getSessionPoToken( + "MWEB", + Localization("en", "US"), + ContentCountry("US"), + false, + ) + + private fun token(visitorData: String, playerToken: String) = SabrTokenBundle( + videoId = "video", + visitorBoundPoToken = playerToken, + visitorBoundPoTokenBytes = byteArrayOf(1), + visitorData = visitorData, + videoBoundPoToken = "video-token", + videoBoundPoTokenBytes = byteArrayOf(2), + ) +}