Make PUBLIC_MODE per-visit isolated sessions like reclip

This commit is contained in:
TonyBlu 2026-04-30 16:11:06 +08:00
parent b6da1e0e7a
commit 0c3860cf9b
3 changed files with 59 additions and 27 deletions

View file

@ -446,50 +446,59 @@ def _migrate_legacy_request(post: dict) -> dict:
return post return post
public_client_downloads: dict[str, set[str]] = {} public_visit_downloads: dict[str, set[str]] = {}
public_download_owner: dict[str, str] = {} public_download_owner: dict[str, str] = {}
public_download_owner_by_id: dict[str, str] = {} public_download_owner_by_id: dict[str, str] = {}
public_visit_sids: dict[str, set[str]] = {}
public_sid_visit: dict[str, str] = {}
def _register_public_download(sid: str, url: str): def _register_public_download(visit_id: str, url: str):
if not sid or not url: if not visit_id or not url:
return return
public_client_downloads.setdefault(sid, set()).add(url) public_visit_downloads.setdefault(visit_id, set()).add(url)
public_download_owner[url] = sid public_download_owner[url] = visit_id
def _forget_public_download(url: str | None = None, *, dl_id: str | None = None): def _forget_public_download(url: str | None = None, *, dl_id: str | None = None):
sid = None visit_id = None
if dl_id: if dl_id:
sid = public_download_owner_by_id.pop(dl_id, None) visit_id = public_download_owner_by_id.pop(dl_id, None)
if sid is None and url: if visit_id is None and url:
sid = public_download_owner.pop(url, None) visit_id = public_download_owner.pop(url, None)
if not sid: if not visit_id:
return return
urls = public_client_downloads.get(sid) urls = public_visit_downloads.get(visit_id)
if not urls: if not urls:
return return
urls.discard(url) if url:
urls.discard(url)
if not urls: if not urls:
public_client_downloads.pop(sid, None) public_visit_downloads.pop(visit_id, None)
async def _emit_download_event(event: str, payload, *, url: str | None = None, dl_id: str | None = None): async def _emit_download_event(event: str, payload, *, url: str | None = None, dl_id: str | None = None):
if not config.PUBLIC_MODE: if not config.PUBLIC_MODE:
await sio.emit(event, payload) await sio.emit(event, payload)
return return
if not url: visit_id = (public_download_owner_by_id.get(dl_id) if dl_id else None) or (public_download_owner.get(url) if url else None)
if not visit_id:
return return
sid = (public_download_owner_by_id.get(dl_id) if dl_id else None) or (public_download_owner.get(url) if url else None) for sid in list(public_visit_sids.get(visit_id, set())):
if sid:
await sio.emit(event, payload, to=sid) await sio.emit(event, payload, to=sid)
def _public_visit_from_environ(environ: dict) -> str:
qs = parse_qs(environ.get('QUERY_STRING', ''))
return (qs.get('visit_id', [''])[0] or '').strip()
class Notifier(DownloadQueueNotifier): class Notifier(DownloadQueueNotifier):
async def added(self, dl): async def added(self, dl):
log.info(f"Notifier: Download added - {dl.title}") log.info(f"Notifier: Download added - {dl.title}")
sid = public_download_owner.get(dl.url) visit_id = public_download_owner.get(dl.url)
if sid: if visit_id:
public_download_owner_by_id[dl.id] = sid public_download_owner_by_id[dl.id] = visit_id
await _emit_download_event('added', serializer.encode(dl), url=dl.url, dl_id=dl.id) await _emit_download_event('added', serializer.encode(dl), url=dl.url, dl_id=dl.id)
async def updated(self, dl): async def updated(self, dl):
@ -770,10 +779,10 @@ async def add(request):
o['auto_start'], o['auto_start'],
) )
if config.PUBLIC_MODE: if config.PUBLIC_MODE:
sid = request.headers.get('X-Client-Id', '').strip() visit_id = request.headers.get('X-Visit-Id', '').strip()
if not sid: if not visit_id:
raise web.HTTPBadRequest(reason='missing X-Client-Id in public mode') raise web.HTTPBadRequest(reason='missing X-Visit-Id in public mode')
_register_public_download(sid, o['url']) _register_public_download(visit_id, o['url'])
status = await dqueue.add( status = await dqueue.add(
o['url'], o['url'],
o['download_type'], o['download_type'],
@ -1016,6 +1025,11 @@ async def history(request):
@sio.event @sio.event
async def connect(sid, environ): async def connect(sid, environ):
log.info(f"Client connected: {sid}") log.info(f"Client connected: {sid}")
if config.PUBLIC_MODE:
visit_id = _public_visit_from_environ(environ)
if visit_id:
public_sid_visit[sid] = visit_id
public_visit_sids.setdefault(visit_id, set()).add(sid)
if config.PUBLIC_MODE: if config.PUBLIC_MODE:
await sio.emit('all', serializer.encode(([], [])), to=sid) await sio.emit('all', serializer.encode(([], [])), to=sid)
else: else:
@ -1032,10 +1046,21 @@ async def connect(sid, environ):
async def disconnect(sid): async def disconnect(sid):
if not config.PUBLIC_MODE: if not config.PUBLIC_MODE:
return return
owned = list(public_client_downloads.pop(sid, set())) visit_id = public_sid_visit.pop(sid, '')
if not visit_id:
return
sids = public_visit_sids.get(visit_id, set())
sids.discard(sid)
if sids:
return
public_visit_sids.pop(visit_id, None)
owned = list(public_visit_downloads.pop(visit_id, set()))
for url in owned: for url in owned:
public_download_owner.pop(url, None) public_download_owner.pop(url, None)
public_download_owner_by_id.pop(url, None) # remove any lingering id ownership for this visit
for dl_id, owner in list(public_download_owner_by_id.items()):
if owner == visit_id:
public_download_owner_by_id.pop(dl_id, None)
if owned: if owned:
await dqueue.cancel(owned) await dqueue.cancel(owned)

View file

@ -152,7 +152,7 @@ export class DownloadsService {
const ce = payload.clipEnd?.trim(); const ce = payload.clipEnd?.trim();
if (cs) body['clip_start'] = cs; if (cs) body['clip_start'] = cs;
if (ce) body['clip_end'] = ce; if (ce) body['clip_end'] = ce;
const headers = this.socket.ioSocket?.id ? { 'X-Client-Id': this.socket.ioSocket.id } : undefined; const headers = { 'X-Visit-Id': this.socket.getVisitId() };
return this.http.post<Status>('add', body, { headers }).pipe( return this.http.post<Status>('add', body, { headers }).pipe(
catchError(this.handleHTTPError) catchError(this.handleHTTPError)
); );

View file

@ -6,12 +6,19 @@ import { Socket } from 'ngx-socket-io';
{ providedIn: 'root' } { providedIn: 'root' }
) )
export class MeTubeSocket extends Socket { export class MeTubeSocket extends Socket {
private readonly visitId: string;
getVisitId(): string {
return this.visitId;
}
constructor() { constructor() {
const appRef = inject(ApplicationRef); const appRef = inject(ApplicationRef);
const path = const path =
document.location.pathname.replace(/share-target/, '') + 'socket.io'; document.location.pathname.replace(/share-target/, '') + 'socket.io';
super({ url: '', options: { path } }, appRef); const visitId = globalThis.crypto?.randomUUID?.() ?? Math.random().toString(36).slice(2);
super({ url: '', options: { path, query: { visit_id: visitId } } }, appRef);
this.visitId = visitId;
} }
} }