Fix Ogg demuxer async race conditions

This commit is contained in:
Vanilagy
2025-05-12 10:40:53 +02:00
parent 9be55666ac
commit 7e1880b916
+50 -23
View File
@@ -3,7 +3,7 @@ import { Demuxer } from '../demuxer';
import { Input } from '../input'; import { Input } from '../input';
import { InputAudioTrack, InputAudioTrackBacking } from '../input-track'; import { InputAudioTrack, InputAudioTrackBacking } from '../input-track';
import { PacketRetrievalOptions } from '../media-sink'; import { PacketRetrievalOptions } from '../media-sink';
import { assert, findLast, roundToPrecision, toDataView, UNDETERMINED_LANGUAGE } from '../misc'; import { assert, AsyncMutex, findLast, roundToPrecision, toDataView, UNDETERMINED_LANGUAGE } from '../misc';
import { EncodedPacket, PLACEHOLDER_DATA } from '../packet'; import { EncodedPacket, PLACEHOLDER_DATA } from '../packet';
import { Reader } from '../reader'; import { Reader } from '../reader';
import { computeOggPageCrc, extractSampleMetadata, OggCodecInfo, parseModesFromVorbisSetupPacket } from './ogg-misc'; import { computeOggPageCrc, extractSampleMetadata, OggCodecInfo, parseModesFromVorbisSetupPacket } from './ogg-misc';
@@ -28,6 +28,11 @@ type Packet = {
export class OggDemuxer extends Demuxer { export class OggDemuxer extends Demuxer {
reader: OggReader; reader: OggReader;
/**
* Lots of reading operations require multiple async reads and thus need to be mutually exclusive to avoid
* conflicts in reader position.
*/
readingMutex = new AsyncMutex();
metadataPromise: Promise<void> | null = null; metadataPromise: Promise<void> | null = null;
fileSize: number | null = null; fileSize: number | null = null;
@@ -491,7 +496,10 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
return encodedPacket; return encodedPacket;
} }
async getFirstPacket(options: PacketRetrievalOptions) { async getFirstPacket(options: PacketRetrievalOptions, exclusive = true) {
const release = exclusive ? await this.demuxer.readingMutex.acquire() : null;
try {
assert(this.bitstream.lastMetadataPacket); assert(this.bitstream.lastMetadataPacket);
const packetPosition = await this.demuxer.findNextPacketStart( const packetPosition = await this.demuxer.findNextPacketStart(
this.demuxer.reader, this.demuxer.reader,
@@ -521,9 +529,15 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
}, },
options, options,
); );
} finally {
release?.();
}
} }
async getNextPacket(prevPacket: EncodedPacket, options: PacketRetrievalOptions) { async getNextPacket(prevPacket: EncodedPacket, options: PacketRetrievalOptions) {
const release = await this.demuxer.readingMutex.acquire();
try {
const prevMetadata = this.encodedPacketToMetadata.get(prevPacket); const prevMetadata = this.encodedPacketToMetadata.get(prevPacket);
if (!prevMetadata) { if (!prevMetadata) {
throw new Error('Packet was not created from this track.'); throw new Error('Packet was not created from this track.');
@@ -550,15 +564,21 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
}, },
options, options,
); );
} finally {
release();
}
} }
async getPacket(timestamp: number, options: PacketRetrievalOptions) { async getPacket(timestamp: number, options: PacketRetrievalOptions) {
const release = await this.demuxer.readingMutex.acquire();
try {
assert(this.demuxer.fileSize !== null); assert(this.demuxer.fileSize !== null);
const timestampInSamples = roundToPrecision(timestamp * this.internalSampleRate, 14); const timestampInSamples = roundToPrecision(timestamp * this.internalSampleRate, 14);
if (timestampInSamples === 0) { if (timestampInSamples === 0) {
// Fast path for timestamp 0 - avoids binary search when playing back from the start // Fast path for timestamp 0 - avoids binary search when playing back from the start
return this.getFirstPacket(options); return this.getFirstPacket(options, false);
} }
if (timestampInSamples < 0) { if (timestampInSamples < 0) {
// There's nothing here // There's nothing here
@@ -581,8 +601,8 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
const lowPages: Page[] = [lowPage]; const lowPages: Page[] = [lowPage];
// First, let's perform a binary serach (bisection search) on the file to find the approximate page where we'll // First, let's perform a binary serach (bisection search) on the file to find the approximate page where
// find the packet. We want to find a page whose end packet position is less than or equal to the // we'll find the packet. We want to find a page whose end packet position is less than or equal to the
// packet position we're searching for. // packet position we're searching for.
// Outer loop: Does the binary serach // Outer loop: Does the binary serach
@@ -616,8 +636,8 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
let pageValid = false; let pageValid = false;
if (page.serialNumber === this.bitstream.serialNumber) { if (page.serialNumber === this.bitstream.serialNumber) {
// Serial numbers are basically random numbers, and the chance of finding a fake page with matching // Serial numbers are basically random numbers, and the chance of finding a fake page with
// serial number is astronomically low, so we can be pretty sure this page is legit. // matching serial number is astronomically low, so we can be pretty sure this page is legit.
pageValid = true; pageValid = true;
} else { } else {
await reader.reader.loadRange(page.headerStartPos, page.headerStartPos + page.totalSize); await reader.reader.loadRange(page.headerStartPos, page.headerStartPos + page.totalSize);
@@ -650,8 +670,8 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
continue; continue;
} }
// The page is valid and belongs to our bitstream; let's check its granule position to see where we need // The page is valid and belongs to our bitstream; let's check its granule position to see where we
// to take the bisection search. // need to take the bisection search.
if (this.granulePositionToTimestampInSamples(page.granulePosition) > timestampInSamples) { if (this.granulePositionToTimestampInSamples(page.granulePosition) > timestampInSamples) {
high = page.headerStartPos; high = page.headerStartPos;
} else { } else {
@@ -663,10 +683,10 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
} }
} }
// Now we have the last page with a packet position <= the packet position we're looking for, but there might // Now we have the last page with a packet position <= the packet position we're looking for, but there
// be multiple pages with the packet position, in which case we actually need to find the first of such pages. // might be multiple pages with the packet position, in which case we actually need to find the first of
// We'll do this in two steps: First, let's find the latest page we know with an earlier packet position, and // such pages. We'll do this in two steps: First, let's find the latest page we know with an earlier packet
// then linear scan ourselves forward until we find the correct page. // position, and then linear scan ourselves forward until we find the correct page.
let lowerPage = startPosition.startPage; let lowerPage = startPosition.startPage;
for (const otherLowPage of lowPages) { for (const otherLowPage of lowPages) {
@@ -733,8 +753,8 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
} }
} }
// This must hold: Since this page has a granule position set, that means there must be a packet that ends // This must hold: Since this page has a granule position set, that means there must be a packet that
// in this page. // ends in this page.
if (currentSegmentIndex === null) { if (currentSegmentIndex === null) {
throw new Error('Invalid page with granule position: no packets end on this page.'); throw new Error('Invalid page with granule position: no packets end on this page.');
} }
@@ -748,8 +768,8 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
const nextPosition = await this.demuxer.findNextPacketStart(reader, pseudopacket); const nextPosition = await this.demuxer.findNextPacketStart(reader, pseudopacket);
if (nextPosition) { if (nextPosition) {
// Let's rewind a single step (packet) - this previous packet ensures that we'll correctly compute the // Let's rewind a single step (packet) - this previous packet ensures that we'll correctly compute
// duration for the packet we're looking for. // the duration for the packet we're looking for.
const endPosition = findPreviousPacketEndPosition(previousPages, currentPage, currentSegmentIndex); const endPosition = findPreviousPacketEndPosition(previousPages, currentPage, currentSegmentIndex);
assert(endPosition); assert(endPosition);
@@ -762,10 +782,12 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
} }
} else { } else {
// There is no next position, which means we're looking for the last packet in the bitstream. The // There is no next position, which means we're looking for the last packet in the bitstream. The
// granule position on the last page tends to be fucky, so let's instead start the search on the page // granule position on the last page tends to be fucky, so let's instead start the search on the
// before that. So let's loop until we find a packet that ends in a previous page. // page before that. So let's loop until we find a packet that ends in a previous page.
while (true) { while (true) {
const endPosition = findPreviousPacketEndPosition(previousPages, currentPage, currentSegmentIndex); const endPosition = findPreviousPacketEndPosition(
previousPages, currentPage, currentSegmentIndex,
);
if (!endPosition) { if (!endPosition) {
break; break;
} }
@@ -792,8 +814,8 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
let lastEncodedPacket: EncodedPacket | null = null; let lastEncodedPacket: EncodedPacket | null = null;
let lastEncodedPacketMetadata: EncodedPacketMetadata | null = null; let lastEncodedPacketMetadata: EncodedPacketMetadata | null = null;
// Alright, now it's time for the final, granular seek: We keep iterating over packets until we've found the one // Alright, now it's time for the final, granular seek: We keep iterating over packets until we've found the
// with the correct timestamp - i.e., the last one with a timestamp <= the timestamp we're looking for. // one with the correct timestamp - i.e., the last one with a timestamp <= the timestamp we're looking for.
while (currentPage !== null) { while (currentPage !== null) {
assert(currentSegmentIndex !== null); assert(currentSegmentIndex !== null);
@@ -826,7 +848,9 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
&& packet.endSegmentIndex === endSegmentIndex && packet.endSegmentIndex === endSegmentIndex
) { ) {
// We know this packet end timestamp can be derived from the page's granule position // We know this packet end timestamp can be derived from the page's granule position
currentTimestampInSamples = this.granulePositionToTimestampInSamples(currentPage.granulePosition); currentTimestampInSamples = this.granulePositionToTimestampInSamples(
currentPage.granulePosition,
);
currentTimestampIsCorrect = true; currentTimestampIsCorrect = true;
// Let's backpatch the packet we just created with the correct timestamp // Let's backpatch the packet we just created with the correct timestamp
@@ -872,6 +896,9 @@ class OggAudioTrackBacking implements InputAudioTrackBacking {
} }
return lastEncodedPacket; return lastEncodedPacket;
} finally {
release();
}
} }
getKeyPacket(timestamp: number, options: PacketRetrievalOptions) { getKeyPacket(timestamp: number, options: PacketRetrievalOptions) {