Add mutex to packet lookup algorithm, fix incorrect next packet retrieval following packet lookup

This commit is contained in:
Vanilagy
2026-01-14 21:53:23 +01:00
parent 4a85aad012
commit 0719a17c85
4 changed files with 211 additions and 133 deletions
+4 -1
View File
@@ -14,7 +14,10 @@
source: new Mediabunny.BlobSource(file), 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));
/* /*
+13 -1
View File
@@ -248,15 +248,27 @@ export const isAllowSharedBufferSource = (x: unknown) => {
export class AsyncMutex { export class AsyncMutex {
currentPromise = Promise.resolve(); currentPromise = Promise.resolve();
pending = 0;
async acquire() { async acquire() {
let resolver: () => void; let resolver: () => void;
const nextPromise = new Promise<void>((resolve) => { const nextPromise = new Promise<void>((resolve) => {
resolver = resolve; let resolved = false;
resolver = () => {
if (resolved) {
return;
}
resolve();
this.pending--;
resolved = true;
};
}); });
const currentPromiseAlias = this.currentPromise; const currentPromiseAlias = this.currentPromise;
this.currentPromise = nextPromise; this.currentPromise = nextPromise;
this.pending++;
await currentPromiseAlias; await currentPromiseAlias;
+70 -43
View File
@@ -33,6 +33,7 @@ import { PacketRetrievalOptions } from '../media-sink';
import { DEFAULT_TRACK_DISPOSITION, MetadataTags } from '../metadata'; import { DEFAULT_TRACK_DISPOSITION, MetadataTags } from '../metadata';
import { import {
assert, assert,
AsyncMutex,
binarySearchLessOrEqual, binarySearchLessOrEqual,
Bitstream, Bitstream,
COLOR_PRIMARIES_MAP_INVERSE, COLOR_PRIMARIES_MAP_INVERSE,
@@ -635,6 +636,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
referencePesPackets: PesPacketHeader[] = []; referencePesPackets: PesPacketHeader[] = [];
endReferencePesPacketAdded = false; endReferencePesPacketAdded = false;
readingContexts = new WeakMap<EncodedPacket, PacketReadingContext>(); readingContexts = new WeakMap<EncodedPacket, PacketReadingContext>();
mutex = new AsyncMutex();
constructor(public elementaryStream: ElementaryStream) {} constructor(public elementaryStream: ElementaryStream) {}
@@ -679,7 +681,11 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
abstract getPacketType(packetData: Uint8Array): PacketType; abstract getPacketType(packetData: Uint8Array): PacketType;
abstract markNextPacket(context: PacketReadingContext): Promise<void>; abstract markNextPacket(context: PacketReadingContext): Promise<void>;
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); const index = binarySearchLessOrEqual(this.referencePesPackets, pesPacketHeader.pts, x => x.pts);
if (index >= 0) { if (index >= 0) {
// Since pts and file position don't necessarily have a monotonic relationship (since pts can go crazy), // 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) { if (index < this.referencePesPackets.length - 1) {
const nextEntry = this.referencePesPackets[index + 1]!; const nextEntry = this.referencePesPackets[index + 1]!;
if (nextEntry.sectionStartPos < pesPacketHeader.sectionStartPos) { if (nextEntry.sectionStartPos < pesPacketHeader.sectionStartPos) {
// Out of order
return false; return false;
} }
@@ -721,7 +728,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
const context = new PacketReadingContext(this, pesPacket, true); const context = new PacketReadingContext(this, pesPacket, true);
await this.markNextPacket(context); await this.markNextPacket(context);
return context.toEncodedPacket(options); return context.createAndLinkEncodedPacket(context.suppliedPacket, options);
} }
async getNextPacket(packet: EncodedPacket, options: PacketRetrievalOptions): Promise<EncodedPacket | null> { async getNextPacket(packet: EncodedPacket, options: PacketRetrievalOptions): Promise<EncodedPacket | null> {
@@ -733,7 +740,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
const clone = context.clone(); const clone = context.clone();
await this.markNextPacket(clone); await this.markNextPacket(clone);
return clone.toEncodedPacket(options); return clone.createAndLinkEncodedPacket(clone.suppliedPacket, options);
} }
async getNextKeyPacket(packet: EncodedPacket, options: PacketRetrievalOptions): Promise<EncodedPacket | null> { async getNextKeyPacket(packet: EncodedPacket, options: PacketRetrievalOptions): Promise<EncodedPacket | null> {
@@ -777,12 +784,18 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
const demuxer = this.elementaryStream.demuxer; const demuxer = this.elementaryStream.demuxer;
const reader = demuxer.reader; const reader = demuxer.reader;
const release = await this.mutex.acquire();
let currentPesPacketHeader: PesPacketHeader;
try {
if (this.referencePesPackets.length === 0) { if (this.referencePesPackets.length === 0) {
// We've never read a packet, let's read the first one const section = this.elementaryStream.firstSection;
const firstPacket = await this.getFirstPacket({}); assert(section);
if (!firstPacket) {
return null; const pesPacketHeader = readPesPacketHeader(section);
} assert(pesPacketHeader);
this.maybeInsertReferencePacket(pesPacketHeader, false, false);
// @ts-expect-error Faulty inference // @ts-expect-error Faulty inference
assert(this.referencePesPackets.length === 1); assert(this.referencePesPackets.length === 1);
@@ -794,8 +807,8 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
} }
// If we're at the end of the reference array, we must make sure we also know about the last packet of the // 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 // track. Without it, we can't perform binary search. This optimization is only possible when we know the
// size, otherwise the linear refinement will naturally discover the end. // file size, otherwise the linear refinement will naturally discover the end.
const needsToLookForLastPacket const needsToLookForLastPacket
= reader.fileSize !== null = reader.fileSize !== null
&& currentIndex === this.referencePesPackets.length - 1 && currentIndex === this.referencePesPackets.length - 1
@@ -825,7 +838,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
throw new Error(MISSING_PES_PACKET_ERROR); throw new Error(MISSING_PES_PACKET_ERROR);
} }
this.maybeInsertReferencePacket(pesPacketHeader, true); this.maybeInsertReferencePacket(pesPacketHeader, true, false);
this.endReferencePesPacketAdded = true; this.endReferencePesPacketAdded = true;
} }
@@ -833,7 +846,8 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
currentIndex = binarySearchLessOrEqual(this.referencePesPackets, searchPts, x => x.pts); currentIndex = binarySearchLessOrEqual(this.referencePesPackets, searchPts, x => x.pts);
assert(currentIndex !== -1); assert(currentIndex !== -1);
// Perform binary search based on the reference PES packets, narrowing in to the timestamp we're interested in // 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 while (reader.fileSize !== null) { // Only do the binary search if the file size is known
const currentEntry = this.referencePesPackets[currentIndex]!; const currentEntry = this.referencePesPackets[currentIndex]!;
const nextEntry = this.referencePesPackets[currentIndex + 1]; const nextEntry = this.referencePesPackets[currentIndex + 1];
@@ -878,7 +892,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
throw new Error(MISSING_PES_PACKET_ERROR); throw new Error(MISSING_PES_PACKET_ERROR);
} }
const addedPoint = this.maybeInsertReferencePacket(pesPacketHeader, false); const addedPoint = this.maybeInsertReferencePacket(pesPacketHeader, false, false);
if (!addedPoint) { if (!addedPoint) {
break; // Should rarely kick break; // Should rarely kick
} }
@@ -889,8 +903,13 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
} }
} }
let currentPesPacketHeader = this.referencePesPackets[currentIndex]!; currentPesPacketHeader = this.referencePesPackets[currentIndex]!;
assert(currentPesPacketHeader.pts <= searchPts); assert(currentPesPacketHeader.pts <= searchPts);
} finally {
release();
}
release();
/** Stores the best PES packet we've found so far (that meets all required criteria). */ /** Stores the best PES packet we've found so far (that meets all required criteria). */
let bestPesPacketHeader: PesPacketHeader | null = null; let bestPesPacketHeader: PesPacketHeader | null = null;
@@ -969,7 +988,7 @@ export abstract class MpegTsTrackBacking implements InputTrackBacking {
if (reader.fileSize === null) { if (reader.fileSize === null) {
// If the file size is undefined, that means that the binary search step is skipped, meaning no // 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. // 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 // 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. // 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. // So, we have one last search to do.
const searchTimestamp = searchPts / TIMESCALE; // We use this instead of 'timestamp' due to the rounding
while (true) { while (true) {
const context = new PacketReadingContext(this, bestPesPacket, false); // Capped context 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) { 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); await this.markNextPacket(context);
const packet = context.toEncodedPacket(options); if (!context.suppliedPacket) {
if (!packet) {
break; 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) { if (!eligible) {
continue; continue;
} }
if (!bestPacket || bestPacket.timestamp < packet.timestamp) { if (!bestPacket || bestPacket.pts < context.suppliedPacket.pts) {
bestPacket = packet; bestPacket = context.suppliedPacket;
bestContext = context.clone(); // Kiiiinda ugly
} }
} }
if (bestPacket) { 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. // 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) { if (b1 === 0x00 && b2 === 0x00 && b3 === 0x01) {
startCodeLength = 4; startCodeLength = 4;
nalUnitTypeByte = context.readU8(); nalUnitTypeByte = context.readU8();
// context.skip(-1);
} else if (b1 === 0x00 && b2 === 0x01) { } else if (b1 === 0x00 && b2 === 0x01) {
// 3-byte start code (0x000001) // 3-byte start code (0x000001)
startCodeLength = 3; startCodeLength = 3;
nalUnitTypeByte = b3; nalUnitTypeByte = b3;
// context.skip(-1); // Back up since b3 is the NAL unit type byte
} }
if (startCodeLength === 0) { 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. */ /** Stateful context used to extract exact encoded packets from the underlying data stream. */
class PacketReadingContext { class PacketReadingContext {
backing: MpegTsTrackBacking; backing: MpegTsTrackBacking;
@@ -1426,12 +1455,7 @@ class PacketReadingContext {
endPos = 0; endPos = 0;
nextPts = 0; nextPts = 0;
suppliedPacket: { suppliedPacket: SuppliedPacket | null = null;
pts: number;
intrinsicDuration: number;
data: Uint8Array;
sequenceNumber: number;
} | null = null;
constructor(backing: MpegTsTrackBacking, startingPesPacket: PesPacket, uncapped: boolean) { constructor(backing: MpegTsTrackBacking, startingPesPacket: PesPacket, uncapped: boolean) {
this.backing = backing; this.backing = backing;
@@ -1442,7 +1466,7 @@ class PacketReadingContext {
} }
clone() { 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.currentPos = this.currentPos;
clone.pesPackets = [...this.pesPackets]; clone.pesPackets = [...this.pesPackets];
clone.currentPesPacketIndex = this.currentPesPacketIndex; clone.currentPesPacketIndex = this.currentPesPacketIndex;
@@ -1624,7 +1648,7 @@ class PacketReadingContext {
return; return;
} }
this.backing.maybeInsertReferencePacket(currentPesPacket, false); this.backing.maybeInsertReferencePacket(currentPesPacket, false, true);
const pts = this.nextPts; const pts = this.nextPts;
this.nextPts += intrinsicDuration; 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 // 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. // PES packet where the packet starts.
const sequenceNumber = currentPesPacket.sectionStartPos + (this.currentPos - this.currentPesPacketPos); const sequenceNumber = currentPesPacket.sectionStartPos + (this.currentPos - this.currentPesPacketPos);
const data = this.readBytes(packetLength);
this.suppliedPacket = { this.suppliedPacket = {
pts, pts,
intrinsicDuration, intrinsicDuration,
data: this.readBytes(packetLength), data,
sequenceNumber, sequenceNumber,
}; };
@@ -1644,19 +1669,21 @@ class PacketReadingContext {
this.currentPesPacketIndex = 0; this.currentPesPacketIndex = 0;
} }
toEncodedPacket(options: PacketRetrievalOptions) { createAndLinkEncodedPacket(suppliedPacket: SuppliedPacket | null, options: PacketRetrievalOptions) {
if (!this.suppliedPacket) { if (!suppliedPacket) {
return null; return null;
} }
const packet = new EncodedPacket( const packet = new EncodedPacket(
options.metadataOnly ? PLACEHOLDER_DATA : this.suppliedPacket.data, options.metadataOnly ? PLACEHOLDER_DATA : suppliedPacket.data,
this.backing.getPacketType(this.suppliedPacket.data), this.backing.getPacketType(suppliedPacket.data),
this.suppliedPacket.pts / TIMESCALE, suppliedPacket.pts / TIMESCALE,
this.suppliedPacket.intrinsicDuration / TIMESCALE, suppliedPacket.intrinsicDuration / TIMESCALE,
this.suppliedPacket.sequenceNumber, suppliedPacket.sequenceNumber,
this.suppliedPacket.data.byteLength, suppliedPacket.data.byteLength,
); );
// Link the context for next packet retrieval
this.backing.readingContexts.set(packet, this); this.backing.readingContexts.set(packet, this);
return packet; return packet;
+38 -2
View File
@@ -204,8 +204,8 @@ test('MPEG-TS video seeking', async () => {
const firstTimestamp = await videoTrack.getFirstTimestamp(); const firstTimestamp = await videoTrack.getFirstTimestamp();
const firstPacket = await sink.getPacket(firstTimestamp); const firstPacket = await sink.getPacket(firstTimestamp);
assert(firstPacket); assert(firstPacket);
expect(firstPacket.timestamp).toBe(firstTimestamp); expect(firstPacket.timestamp).toBe(firstTimestamp);
expect(firstPacket.sequenceNumber).toBe((await sink.getFirstPacket())?.sequenceNumber);
const lastPacket = await sink.getPacket(Infinity); const lastPacket = await sink.getPacket(Infinity);
assert(lastPacket); assert(lastPacket);
@@ -227,6 +227,8 @@ test('MPEG-TS video seeking', async () => {
currentPacket = await sink.getNextPacket(currentPacket); currentPacket = await sink.getNextPacket(currentPacket);
} }
expect(allPackets).toHaveLength(298);
for (const packet of allPackets) { for (const packet of allPackets) {
const seekedPacked = await sink.getPacket(packet.timestamp); const seekedPacked = await sink.getPacket(packet.timestamp);
assert(seekedPacked); assert(seekedPacked);
@@ -249,8 +251,8 @@ test('MPEG-TS audio seeking', async () => {
const firstTimestamp = await audioTrack.getFirstTimestamp(); const firstTimestamp = await audioTrack.getFirstTimestamp();
const firstPacket = await sink.getPacket(firstTimestamp); const firstPacket = await sink.getPacket(firstTimestamp);
assert(firstPacket); assert(firstPacket);
expect(firstPacket.timestamp).toBe(firstTimestamp); expect(firstPacket.timestamp).toBe(firstTimestamp);
expect(firstPacket.sequenceNumber).toBe((await sink.getFirstPacket())?.sequenceNumber);
const lastPacket = await sink.getPacket(Infinity); const lastPacket = await sink.getPacket(Infinity);
assert(lastPacket); assert(lastPacket);
@@ -272,6 +274,8 @@ test('MPEG-TS audio seeking', async () => {
currentPacket = await sink.getNextPacket(currentPacket); currentPacket = await sink.getNextPacket(currentPacket);
} }
expect(allPackets).toHaveLength(234);
for (const packet of allPackets) { for (const packet of allPackets) {
const seekedPacket = await sink.getPacket(packet.timestamp); const seekedPacket = await sink.getPacket(packet.timestamp);
assert(seekedPacket); 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 () => { test('MPEG-TS video key packets', async () => {
using input = new Input({ using input = new Input({
source: new FilePathSource(path.join(__dirname, '../public/193039199_mp4_h264_aac_fhd_7.ts')), source: new FilePathSource(path.join(__dirname, '../public/193039199_mp4_h264_aac_fhd_7.ts')),