diff --git a/dev/demux.html b/dev/demux.html index cd422de..c6402f5 100644 --- a/dev/demux.html +++ b/dev/demux.html @@ -13,8 +13,11 @@ formats: Mediabunny.ALL_FORMATS, source: new Mediabunny.BlobSource(file), }); - - console.log(await input.getTracks()); + + const videoTrack = await input.getPrimaryVideoTrack(); + const sink = new Mediabunny.EncodedPacketSink(videoTrack); + + console.log(await sink.getKeyPacket(Infinity)); /* diff --git a/src/misc.ts b/src/misc.ts index 3274e7f..bb2d577 100644 --- a/src/misc.ts +++ b/src/misc.ts @@ -248,15 +248,27 @@ export const isAllowSharedBufferSource = (x: unknown) => { export class AsyncMutex { currentPromise = Promise.resolve(); + pending = 0; async acquire() { let resolver: () => void; const nextPromise = new Promise((resolve) => { - resolver = resolve; + let resolved = false; + + resolver = () => { + if (resolved) { + return; + } + + resolve(); + this.pending--; + resolved = true; + }; }); const currentPromiseAlias = this.currentPromise; this.currentPromise = nextPromise; + this.pending++; await currentPromiseAlias; diff --git a/src/mpeg-ts/mpeg-ts-demuxer.ts b/src/mpeg-ts/mpeg-ts-demuxer.ts index 64cdac0..d405146 100644 --- a/src/mpeg-ts/mpeg-ts-demuxer.ts +++ b/src/mpeg-ts/mpeg-ts-demuxer.ts @@ -33,6 +33,7 @@ import { PacketRetrievalOptions } from '../media-sink'; import { DEFAULT_TRACK_DISPOSITION, MetadataTags } from '../metadata'; import { assert, + AsyncMutex, binarySearchLessOrEqual, Bitstream, COLOR_PRIMARIES_MAP_INVERSE, @@ -635,6 +636,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking { referencePesPackets: PesPacketHeader[] = []; endReferencePesPacketAdded = false; readingContexts = new WeakMap(); + mutex = new AsyncMutex(); constructor(public elementaryStream: ElementaryStream) {} @@ -679,7 +681,11 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking { abstract getPacketType(packetData: Uint8Array): PacketType; abstract markNextPacket(context: PacketReadingContext): Promise; - maybeInsertReferencePacket(pesPacketHeader: PesPacketHeader, force: boolean) { + maybeInsertReferencePacket(pesPacketHeader: PesPacketHeader, force: boolean, dropIfMutexLocked: boolean) { + if (dropIfMutexLocked && this.mutex.pending > 0) { + return; // Drop this one to avoid race conditions + } + const index = binarySearchLessOrEqual(this.referencePesPackets, pesPacketHeader.pts, x => x.pts); if (index >= 0) { // Since pts and file position don't necessarily have a monotonic relationship (since pts can go crazy), @@ -697,6 +703,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking { if (index < this.referencePesPackets.length - 1) { const nextEntry = this.referencePesPackets[index + 1]!; if (nextEntry.sectionStartPos < pesPacketHeader.sectionStartPos) { + // Out of order return false; } @@ -721,7 +728,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking { const context = new PacketReadingContext(this, pesPacket, true); await this.markNextPacket(context); - return context.toEncodedPacket(options); + return context.createAndLinkEncodedPacket(context.suppliedPacket, options); } async getNextPacket(packet: EncodedPacket, options: PacketRetrievalOptions): Promise { @@ -733,7 +740,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking { const clone = context.clone(); await this.markNextPacket(clone); - return clone.toEncodedPacket(options); + return clone.createAndLinkEncodedPacket(clone.suppliedPacket, options); } async getNextKeyPacket(packet: EncodedPacket, options: PacketRetrievalOptions): Promise { @@ -777,120 +784,132 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking { const demuxer = this.elementaryStream.demuxer; const reader = demuxer.reader; - if (this.referencePesPackets.length === 0) { - // We've never read a packet, let's read the first one - const firstPacket = await this.getFirstPacket({}); - if (!firstPacket) { - return null; + const release = await this.mutex.acquire(); + let currentPesPacketHeader: PesPacketHeader; + + try { + if (this.referencePesPackets.length === 0) { + const section = this.elementaryStream.firstSection; + assert(section); + + const pesPacketHeader = readPesPacketHeader(section); + assert(pesPacketHeader); + + this.maybeInsertReferencePacket(pesPacketHeader, false, false); + + // @ts-expect-error Faulty inference + assert(this.referencePesPackets.length === 1); } - // @ts-expect-error Faulty inference - assert(this.referencePesPackets.length === 1); - } + let currentIndex = binarySearchLessOrEqual(this.referencePesPackets, searchPts, x => x.pts); + if (currentIndex === -1) { + return null; // We're before the first packet + } - let currentIndex = binarySearchLessOrEqual(this.referencePesPackets, searchPts, x => x.pts); - if (currentIndex === -1) { - return null; // We're before the first packet - } - - // If we're at the end of the reference array, we must make sure we also know about the last packet of the - // track. Without it, we can't perform binary search. This optimization is only possible when we know the file - // size, otherwise the linear refinement will naturally discover the end. - const needsToLookForLastPacket + // If we're at the end of the reference array, we must make sure we also know about the last packet of the + // track. Without it, we can't perform binary search. This optimization is only possible when we know the + // file size, otherwise the linear refinement will naturally discover the end. + const needsToLookForLastPacket = reader.fileSize !== null && currentIndex === this.referencePesPackets.length - 1 && !this.endReferencePesPacketAdded; - if (needsToLookForLastPacket) { - let currentPos = reader.fileSize! - demuxer.packetStride + demuxer.packetOffset; - let packetHeader = await demuxer.readPacketHeader(currentPos); - if (!packetHeader) { - return null; - } - - while (packetHeader.pid !== this.elementaryStream.pid || packetHeader.payloadUnitStartIndicator === 0) { - currentPos -= demuxer.packetStride; - const previousPacketHeader = await demuxer.readPacketHeader(currentPos); - if (!previousPacketHeader) { + if (needsToLookForLastPacket) { + let currentPos = reader.fileSize! - demuxer.packetStride + demuxer.packetOffset; + let packetHeader = await demuxer.readPacketHeader(currentPos); + if (!packetHeader) { return null; } - packetHeader = previousPacketHeader; - } + while (packetHeader.pid !== this.elementaryStream.pid || packetHeader.payloadUnitStartIndicator === 0) { + currentPos -= demuxer.packetStride; + const previousPacketHeader = await demuxer.readPacketHeader(currentPos); + if (!previousPacketHeader) { + return null; + } - const section = await demuxer.readSection(currentPos, false); - assert(section); - - const pesPacketHeader = readPesPacketHeader(section); - if (!pesPacketHeader) { - throw new Error(MISSING_PES_PACKET_ERROR); - } - - this.maybeInsertReferencePacket(pesPacketHeader, true); - this.endReferencePesPacketAdded = true; - } - - // Find the reference point closest to the search timestamp - currentIndex = binarySearchLessOrEqual(this.referencePesPackets, searchPts, x => x.pts); - assert(currentIndex !== -1); - - // Perform binary search based on the reference PES packets, narrowing in to the timestamp we're interested in - while (reader.fileSize !== null) { // Only do the binary search if the file size is known - const currentEntry = this.referencePesPackets[currentIndex]!; - const nextEntry = this.referencePesPackets[currentIndex + 1]; - - if (searchPts - currentEntry.pts < TIMESCALE || !nextEntry) { - // We're at the end or close enough to the entry to the left, stop - break; - } - - // Jump in between the two entries, and then find a fitting packet there - const midpoint = roundToMultiple( - (currentEntry.sectionStartPos + nextEntry.sectionStartPos) / 2, - demuxer.packetStride, - ) + demuxer.packetOffset; - let currentPos = midpoint; - let packetHeader = await demuxer.readPacketHeader(currentPos); - assert(packetHeader); - - while ( - currentPos < nextEntry.sectionStartPos - && (packetHeader.pid !== this.elementaryStream.pid || packetHeader.payloadUnitStartIndicator === 0) - ) { - currentPos += demuxer.packetStride; - const previousPacketHeader = await demuxer.readPacketHeader(currentPos); - if (!previousPacketHeader) { - return null; + packetHeader = previousPacketHeader; } - packetHeader = previousPacketHeader; + const section = await demuxer.readSection(currentPos, false); + assert(section); + + const pesPacketHeader = readPesPacketHeader(section); + if (!pesPacketHeader) { + throw new Error(MISSING_PES_PACKET_ERROR); + } + + this.maybeInsertReferencePacket(pesPacketHeader, true, false); + this.endReferencePesPacketAdded = true; } - if (currentPos >= nextEntry.sectionStartPos) { - // We couldn't find a packet in the middle - break; + // Find the reference point closest to the search timestamp + currentIndex = binarySearchLessOrEqual(this.referencePesPackets, searchPts, x => x.pts); + assert(currentIndex !== -1); + + // Perform binary search based on the reference PES packets, narrowing in to the timestamp we're + // interested in + while (reader.fileSize !== null) { // Only do the binary search if the file size is known + const currentEntry = this.referencePesPackets[currentIndex]!; + const nextEntry = this.referencePesPackets[currentIndex + 1]; + + if (searchPts - currentEntry.pts < TIMESCALE || !nextEntry) { + // We're at the end or close enough to the entry to the left, stop + break; + } + + // Jump in between the two entries, and then find a fitting packet there + const midpoint = roundToMultiple( + (currentEntry.sectionStartPos + nextEntry.sectionStartPos) / 2, + demuxer.packetStride, + ) + demuxer.packetOffset; + let currentPos = midpoint; + let packetHeader = await demuxer.readPacketHeader(currentPos); + assert(packetHeader); + + while ( + currentPos < nextEntry.sectionStartPos + && (packetHeader.pid !== this.elementaryStream.pid || packetHeader.payloadUnitStartIndicator === 0) + ) { + currentPos += demuxer.packetStride; + const previousPacketHeader = await demuxer.readPacketHeader(currentPos); + if (!previousPacketHeader) { + return null; + } + + packetHeader = previousPacketHeader; + } + + if (currentPos >= nextEntry.sectionStartPos) { + // We couldn't find a packet in the middle + break; + } + + const section = await demuxer.readSection(currentPos, false); + assert(section); + + const pesPacketHeader = readPesPacketHeader(section); + if (!pesPacketHeader) { + throw new Error(MISSING_PES_PACKET_ERROR); + } + + const addedPoint = this.maybeInsertReferencePacket(pesPacketHeader, false, false); + if (!addedPoint) { + break; // Should rarely kick + } + + if (pesPacketHeader.pts <= searchPts) { + // The midpoint packet is to the left of our search timestamp, so continue with the right half now + currentIndex++; + } } - const section = await demuxer.readSection(currentPos, false); - assert(section); - - const pesPacketHeader = readPesPacketHeader(section); - if (!pesPacketHeader) { - throw new Error(MISSING_PES_PACKET_ERROR); - } - - const addedPoint = this.maybeInsertReferencePacket(pesPacketHeader, false); - if (!addedPoint) { - break; // Should rarely kick - } - - if (pesPacketHeader.pts <= searchPts) { - // The midpoint packet is to the left of our search timestamp, so continue with the right half now - currentIndex++; - } + currentPesPacketHeader = this.referencePesPackets[currentIndex]!; + assert(currentPesPacketHeader.pts <= searchPts); + } finally { + release(); } - let currentPesPacketHeader = this.referencePesPackets[currentIndex]!; - assert(currentPesPacketHeader.pts <= searchPts); + release(); /** Stores the best PES packet we've found so far (that meets all required criteria). */ let bestPesPacketHeader: PesPacketHeader | null = null; @@ -969,7 +988,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking { if (reader.fileSize === null) { // If the file size is undefined, that means that the binary search step is skipped, meaning no // reference packets are inserted. So, let's instead insert reference packets in the linear search step. - this.maybeInsertReferencePacket(nextPesPacketHeader, false); + this.maybeInsertReferencePacket(nextPesPacketHeader, false, true); } } @@ -1099,32 +1118,37 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking { // Final stage: we found the best PES packet, but that PES packet might contain multiple individual encoded // packets. Or, it might not be the start of an encoded packet, and simply a continuation of a previous one. // So, we have one last search to do. - const searchTimestamp = searchPts / TIMESCALE; // We use this instead of 'timestamp' due to the rounding while (true) { const context = new PacketReadingContext(this, bestPesPacket, false); // Capped context - let bestPacket: EncodedPacket | null = null; + let bestPacket: SuppliedPacket | null = null; + let bestContext: PacketReadingContext | null = null; while (true) { - context.suppliedPacket = null; + // Stupid aliasing trick to make TypeScript not infer stuff wrong + const context2 = context; + context2.suppliedPacket = null; + await this.markNextPacket(context); - const packet = context.toEncodedPacket(options); - if (!packet) { + if (!context.suppliedPacket) { break; } - const eligible = packet.timestamp <= searchTimestamp && (!keyframesOnly || packet.type === 'key'); + const eligible = context.suppliedPacket.pts <= searchPts + && (!keyframesOnly || this.getPacketType(context.suppliedPacket.data) === 'key'); if (!eligible) { continue; } - if (!bestPacket || bestPacket.timestamp < packet.timestamp) { - bestPacket = packet; + if (!bestPacket || bestPacket.pts < context.suppliedPacket.pts) { + bestPacket = context.suppliedPacket; + bestContext = context.clone(); // Kiiiinda ugly } } if (bestPacket) { - return bestPacket; + assert(bestContext); + return bestContext.createAndLinkEncodedPacket(bestPacket, options); } // We didn't find an encoded packet! Let's go to the previous PES packet until we find one. @@ -1272,12 +1296,10 @@ class MpegTsVideoTrackBacking extends MpegTsTrackBacking implements InputVideoTr if (b1 === 0x00 && b2 === 0x00 && b3 === 0x01) { startCodeLength = 4; nalUnitTypeByte = context.readU8(); - // context.skip(-1); } else if (b1 === 0x00 && b2 === 0x01) { // 3-byte start code (0x000001) startCodeLength = 3; nalUnitTypeByte = b3; - // context.skip(-1); // Back up since b3 is the NAL unit type byte } if (startCodeLength === 0) { @@ -1411,6 +1433,13 @@ class MpegTsAudioTrackBacking extends MpegTsTrackBacking implements InputAudioTr } } +type SuppliedPacket = { + pts: number; + intrinsicDuration: number; + data: Uint8Array; + sequenceNumber: number; +}; + /** Stateful context used to extract exact encoded packets from the underlying data stream. */ class PacketReadingContext { backing: MpegTsTrackBacking; @@ -1426,12 +1455,7 @@ class PacketReadingContext { endPos = 0; nextPts = 0; - suppliedPacket: { - pts: number; - intrinsicDuration: number; - data: Uint8Array; - sequenceNumber: number; - } | null = null; + suppliedPacket: SuppliedPacket | null = null; constructor(backing: MpegTsTrackBacking, startingPesPacket: PesPacket, uncapped: boolean) { this.backing = backing; @@ -1442,7 +1466,7 @@ class PacketReadingContext { } clone() { - const clone = new PacketReadingContext(this.backing, this.startingPesPacket, this.uncapped); + const clone = new PacketReadingContext(this.backing, this.startingPesPacket, true); // Close isn't capped clone.currentPos = this.currentPos; clone.pesPackets = [...this.pesPackets]; clone.currentPesPacketIndex = this.currentPesPacketIndex; @@ -1624,7 +1648,7 @@ class PacketReadingContext { return; } - this.backing.maybeInsertReferencePacket(currentPesPacket, false); + this.backing.maybeInsertReferencePacket(currentPesPacket, false, true); const pts = this.nextPts; this.nextPts += intrinsicDuration; @@ -1632,11 +1656,12 @@ class PacketReadingContext { // The sequence number is the starting position of the section the PES packet is in, PLUS the offset within the // PES packet where the packet starts. const sequenceNumber = currentPesPacket.sectionStartPos + (this.currentPos - this.currentPesPacketPos); + const data = this.readBytes(packetLength); this.suppliedPacket = { pts, intrinsicDuration, - data: this.readBytes(packetLength), + data, sequenceNumber, }; @@ -1644,19 +1669,21 @@ class PacketReadingContext { this.currentPesPacketIndex = 0; } - toEncodedPacket(options: PacketRetrievalOptions) { - if (!this.suppliedPacket) { + createAndLinkEncodedPacket(suppliedPacket: SuppliedPacket | null, options: PacketRetrievalOptions) { + if (!suppliedPacket) { return null; } const packet = new EncodedPacket( - options.metadataOnly ? PLACEHOLDER_DATA : this.suppliedPacket.data, - this.backing.getPacketType(this.suppliedPacket.data), - this.suppliedPacket.pts / TIMESCALE, - this.suppliedPacket.intrinsicDuration / TIMESCALE, - this.suppliedPacket.sequenceNumber, - this.suppliedPacket.data.byteLength, + options.metadataOnly ? PLACEHOLDER_DATA : suppliedPacket.data, + this.backing.getPacketType(suppliedPacket.data), + suppliedPacket.pts / TIMESCALE, + suppliedPacket.intrinsicDuration / TIMESCALE, + suppliedPacket.sequenceNumber, + suppliedPacket.data.byteLength, ); + + // Link the context for next packet retrieval this.backing.readingContexts.set(packet, this); return packet; diff --git a/test/node/mpeg-ts-demuxing.test.ts b/test/node/mpeg-ts-demuxing.test.ts index 117df9b..50c7486 100644 --- a/test/node/mpeg-ts-demuxing.test.ts +++ b/test/node/mpeg-ts-demuxing.test.ts @@ -204,8 +204,8 @@ test('MPEG-TS video seeking', async () => { const firstTimestamp = await videoTrack.getFirstTimestamp(); const firstPacket = await sink.getPacket(firstTimestamp); assert(firstPacket); - expect(firstPacket.timestamp).toBe(firstTimestamp); + expect(firstPacket.sequenceNumber).toBe((await sink.getFirstPacket())?.sequenceNumber); const lastPacket = await sink.getPacket(Infinity); assert(lastPacket); @@ -227,6 +227,8 @@ test('MPEG-TS video seeking', async () => { currentPacket = await sink.getNextPacket(currentPacket); } + expect(allPackets).toHaveLength(298); + for (const packet of allPackets) { const seekedPacked = await sink.getPacket(packet.timestamp); assert(seekedPacked); @@ -249,8 +251,8 @@ test('MPEG-TS audio seeking', async () => { const firstTimestamp = await audioTrack.getFirstTimestamp(); const firstPacket = await sink.getPacket(firstTimestamp); assert(firstPacket); - expect(firstPacket.timestamp).toBe(firstTimestamp); + expect(firstPacket.sequenceNumber).toBe((await sink.getFirstPacket())?.sequenceNumber); const lastPacket = await sink.getPacket(Infinity); assert(lastPacket); @@ -272,6 +274,8 @@ test('MPEG-TS audio seeking', async () => { currentPacket = await sink.getNextPacket(currentPacket); } + expect(allPackets).toHaveLength(234); + for (const packet of allPackets) { const seekedPacket = await sink.getPacket(packet.timestamp); assert(seekedPacket); @@ -280,6 +284,38 @@ test('MPEG-TS audio seeking', async () => { } }); +test('MPEG-TS seeking race condition test', async () => { + using input = new Input({ + source: new FilePathSource(path.join(__dirname, '../public/0.ts')), + formats: ALL_FORMATS, + }); + + const videoTrack = await input.getPrimaryVideoTrack(); + assert(videoTrack); + + const sink = new EncodedPacketSink(videoTrack); + + const allPackets: EncodedPacket[] = []; + let currentPacket: EncodedPacket | null = await sink.getFirstPacket(); + + while (currentPacket) { + allPackets.push(currentPacket); + currentPacket = await sink.getNextPacket(currentPacket); + } + + // Perform all seeks concurrently + const seekPromises = allPackets.map(packet => sink.getPacket(packet.timestamp)); + const seekedPackets = await Promise.all(seekPromises); + + for (let i = 0; i < allPackets.length; i++) { + const originalPacket = allPackets[i]!; + const seekedPacket = seekedPackets[i]!; + assert(seekedPacket); + expect(seekedPacket.timestamp).toBe(originalPacket.timestamp); + expect(seekedPacket.sequenceNumber).toBe(originalPacket.sequenceNumber); + } +}); + test('MPEG-TS video key packets', async () => { using input = new Input({ source: new FilePathSource(path.join(__dirname, '../public/193039199_mp4_h264_aac_fhd_7.ts')),