Skip to content

Commit 0963feb

Browse files
committed
修正 RTSP 与宽播测试适配
1 parent 2d52d8e commit 0963feb

3 files changed

Lines changed: 51 additions & 30 deletions

File tree

‎backend/tests/test_rtsp_policy.py‎

Lines changed: 12 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -2,6 +2,7 @@
22

33
import ipaddress
44
import os
5+
import sys
56
import unittest
67
from pathlib import Path
78
from types import SimpleNamespace
@@ -52,11 +53,12 @@ async def test_signed_rtsp_handle_disabled_does_not_create_session(self):
5253
handle = issue_handle(kind="rtsp", url="rtsp://camera.local/live")
5354
with mock.patch.object(main, "config_rtsp_proxy_enabled", return_value=False), \
5455
mock.patch.object(main, "_ensure_rtsp_hls_session", new=mock.AsyncMock()) as ensure:
55-
with self.assertRaises(HTTPException) as ctx:
56-
await media_proxy.media_proxy_rtsp(
57-
handle,
58-
MediaAccessContext(source="anonymous"),
59-
)
56+
with mock.patch.dict(sys.modules, {"main": main}):
57+
with self.assertRaises(HTTPException) as ctx:
58+
await media_proxy.media_proxy_rtsp(
59+
handle,
60+
MediaAccessContext(source="anonymous"),
61+
)
6062
self.assertEqual(ctx.exception.status_code, 503)
6163
ensure.assert_not_awaited()
6264

@@ -78,10 +80,11 @@ async def test_smart_and_handle_paths_share_final_rtsp_boundary(self):
7880
await main.iptv_smart_playlist("camera", None)
7981

8082
handle = issue_handle(kind="rtsp", url="rtsp://camera.local/live")
81-
await media_proxy.media_proxy_rtsp(
82-
handle,
83-
MediaAccessContext(source="anonymous"),
84-
)
83+
with mock.patch.dict(sys.modules, {"main": main}):
84+
await media_proxy.media_proxy_rtsp(
85+
handle,
86+
MediaAccessContext(source="anonymous"),
87+
)
8588

8689
self.assertEqual(safe_host.call_count, 2)
8790
self.assertEqual(ensure.await_count, 2)

‎backend/tests/test_rtsp_startup.py‎

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -47,6 +47,10 @@ async def asyncTearDown(self):
4747
task.cancel()
4848
if pending:
4949
await asyncio.gather(*pending, return_exceptions=True)
50+
await main._stop_all_rtsp_sessions()
51+
self.assertEqual(main._RTSP_HLS_STARTUPS, {})
52+
self.assertEqual(main._RTSP_RESERVED_SESSIONS, set())
53+
self.assertEqual(main.RTSP_HLS_SESSIONS, {})
5054
main._RTSP_HLS_STARTUPS.clear()
5155
main._RTSP_RESERVED_SESSIONS.clear()
5256
main.RTSP_HLS_SESSIONS = self._old_sessions

‎backend/tests/test_wide_playlist.py‎

Lines changed: 35 additions & 21 deletions
Original file line numberDiff line numberDiff line change
@@ -11,6 +11,7 @@
1111
os.environ.setdefault("WAVEFLOW_DB_PATH", ":memory:")
1212

1313
import main
14+
from infrastructure import http_client as media_http
1415
from security.dependencies import MediaAccessContext
1516
from security import proxy_handles, proxy_context
1617

@@ -167,9 +168,11 @@ async def _record_ssrf(url, *a, **kw):
167168
return None
168169

169170
real_ssrf = main.assert_safe_target_url
171+
real_media_ssrf = media_http.assert_safe_target_url
170172
main.assert_safe_target_url = _record_ssrf
173+
media_http.assert_safe_target_url = _record_ssrf
171174
try:
172-
with mock.patch.object(main.http_client, "get", side_effect=fake.get):
175+
with mock.patch.object(main.http_client, "request", side_effect=fake.get):
173176
resp = await main.serve_iptv_wide_playlist_by_source(
174177
upstream_url="https://up.example/live.m3u8",
175178
ctx_id="",
@@ -181,6 +184,7 @@ async def _record_ssrf(url, *a, **kw):
181184
body = resp.body.decode("utf-8")
182185
finally:
183186
main.assert_safe_target_url = real_ssrf
187+
media_http.assert_safe_target_url = real_media_ssrf
184188
return body, ssrf_checked
185189

186190
body, ssrf_checked = asyncio.run(go())
@@ -270,7 +274,7 @@ async def test_same_key_cold_start_is_single_flight_and_single_refresher(self):
270274
refresher_release = asyncio.Event()
271275
calls = 0
272276

273-
async def fake_get(url, **_kwargs):
277+
async def fake_get(_method, url, **_kwargs):
274278
nonlocal calls
275279
calls += 1
276280
if calls == 1:
@@ -282,7 +286,8 @@ async def hold_refresher(*_args):
282286
await refresher_release.wait()
283287

284288
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
285-
mock.patch.object(main.http_client, "get", side_effect=fake_get), \
289+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
290+
mock.patch.object(main.http_client, "request", side_effect=fake_get), \
286291
mock.patch.object(main, "_run_wide_refresher", new=hold_refresher):
287292
first = asyncio.create_task(self._serve("same"))
288293
await entered.wait()
@@ -316,7 +321,7 @@ async def test_different_keys_initialize_in_parallel(self):
316321
maximum_active = 0
317322
refresher_release = asyncio.Event()
318323

319-
async def fake_get(url, **_kwargs):
324+
async def fake_get(_method, url, **_kwargs):
320325
nonlocal active, maximum_active
321326
key = "a" if "/a.m3u8" in url else "b"
322327
calls[key] += 1
@@ -332,7 +337,8 @@ async def hold_refresher(*_args):
332337
await refresher_release.wait()
333338

334339
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
335-
mock.patch.object(main.http_client, "get", side_effect=fake_get), \
340+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
341+
mock.patch.object(main.http_client, "request", side_effect=fake_get), \
336342
mock.patch.object(main, "_run_wide_refresher", new=hold_refresher):
337343
first = asyncio.create_task(self._serve("a"))
338344
second = asyncio.create_task(self._serve("b"))
@@ -352,7 +358,7 @@ async def test_waiter_cancellation_does_not_cancel_initializer(self):
352358
refresher_release = asyncio.Event()
353359
calls = 0
354360

355-
async def fake_get(url, **_kwargs):
361+
async def fake_get(_method, url, **_kwargs):
356362
nonlocal calls
357363
calls += 1
358364
if calls == 1:
@@ -364,7 +370,8 @@ async def hold_refresher(*_args):
364370
await refresher_release.wait()
365371

366372
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
367-
mock.patch.object(main.http_client, "get", side_effect=fake_get), \
373+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
374+
mock.patch.object(main.http_client, "request", side_effect=fake_get), \
368375
mock.patch.object(main, "_run_wide_refresher", new=hold_refresher):
369376
owner = asyncio.create_task(self._serve("cancel-waiter"))
370377
await entered.wait()
@@ -384,12 +391,13 @@ async def hold_refresher(*_args):
384391
async def test_initializer_cancellation_leaves_no_cache_or_refresher(self):
385392
entered = asyncio.Event()
386393

387-
async def blocked_get(_url, **_kwargs):
394+
async def blocked_get(_method, _url, **_kwargs):
388395
entered.set()
389396
await asyncio.Event().wait()
390397

391398
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
392-
mock.patch.object(main.http_client, "get", side_effect=blocked_get):
399+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
400+
mock.patch.object(main.http_client, "request", side_effect=blocked_get):
393401
owner = asyncio.create_task(self._serve("cancel-owner"))
394402
await entered.wait()
395403
owner.cancel()
@@ -402,12 +410,13 @@ async def blocked_get(_url, **_kwargs):
402410
async def test_eviction_cancels_inflight_initializer(self):
403411
entered = asyncio.Event()
404412

405-
async def blocked_get(_url, **_kwargs):
413+
async def blocked_get(_method, _url, **_kwargs):
406414
entered.set()
407415
await asyncio.Event().wait()
408416

409417
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
410-
mock.patch.object(main.http_client, "get", side_effect=blocked_get):
418+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
419+
mock.patch.object(main.http_client, "request", side_effect=blocked_get):
411420
owner = asyncio.create_task(self._serve("evict-owner"))
412421
await entered.wait()
413422
self.assertTrue(main.release_iptv_wide_playlist_by_key(self._cache_key("evict-owner")))
@@ -421,12 +430,13 @@ async def blocked_get(_url, **_kwargs):
421430
async def test_shutdown_cancels_inflight_initializer_before_clearing_state(self):
422431
entered = asyncio.Event()
423432

424-
async def blocked_get(_url, **_kwargs):
433+
async def blocked_get(_method, _url, **_kwargs):
425434
entered.set()
426435
await asyncio.Event().wait()
427436

428437
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
429-
mock.patch.object(main.http_client, "get", side_effect=blocked_get):
438+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
439+
mock.patch.object(main.http_client, "request", side_effect=blocked_get):
430440
owner = asyncio.create_task(self._serve("shutdown-owner"))
431441
await entered.wait()
432442
self.assertTrue(main._wide_initialization_tasks)
@@ -442,7 +452,7 @@ async def test_initializer_failure_releases_key_for_retry(self):
442452
should_fail = True
443453
refresher_release = asyncio.Event()
444454

445-
async def fake_get(url, **_kwargs):
455+
async def fake_get(_method, url, **_kwargs):
446456
nonlocal calls
447457
calls += 1
448458
if should_fail:
@@ -453,7 +463,8 @@ async def hold_refresher(*_args):
453463
await refresher_release.wait()
454464

455465
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
456-
mock.patch.object(main.http_client, "get", side_effect=fake_get), \
466+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
467+
mock.patch.object(main.http_client, "request", side_effect=fake_get), \
457468
mock.patch.object(main, "_run_wide_refresher", new=hold_refresher):
458469
with self.assertRaises(HTTPException) as failure:
459470
await self._serve("retry")
@@ -475,11 +486,12 @@ async def failing_refresher(*_args):
475486
crashed.set()
476487
raise RuntimeError("fixture refresher failure")
477488

478-
async def fake_get(url, **_kwargs):
489+
async def fake_get(_method, url, **_kwargs):
479490
return self._response(url)
480491

481492
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
482-
mock.patch.object(main.http_client, "get", side_effect=fake_get), \
493+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
494+
mock.patch.object(main.http_client, "request", side_effect=fake_get), \
483495
mock.patch.object(main, "_run_wide_refresher", new=failing_refresher):
484496
await self._serve("refresher-failure")
485497
await crashed.wait()
@@ -495,11 +507,12 @@ async def test_eviction_and_shutdown_reclaim_refresher_tasks(self):
495507
async def hold_refresher(*_args):
496508
await refresher_release.wait()
497509

498-
async def fake_get(url, **_kwargs):
510+
async def fake_get(_method, url, **_kwargs):
499511
return self._response(url)
500512

501513
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
502-
mock.patch.object(main.http_client, "get", side_effect=fake_get), \
514+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
515+
mock.patch.object(main.http_client, "request", side_effect=fake_get), \
503516
mock.patch.object(main, "_run_wide_refresher", new=hold_refresher):
504517
await self._serve("evict")
505518
await self._serve("shutdown")
@@ -524,11 +537,12 @@ async def hold_refresher(*_args):
524537
refresher_started.set()
525538
await refresher_release.wait()
526539

527-
async def fake_get(url, **_kwargs):
540+
async def fake_get(_method, url, **_kwargs):
528541
return self._response(url)
529542

530543
with mock.patch.object(main, "assert_safe_target_url", new=self._noop_ssrf), \
531-
mock.patch.object(main.http_client, "get", side_effect=fake_get), \
544+
mock.patch.object(media_http, "assert_safe_target_url", new=self._noop_ssrf), \
545+
mock.patch.object(main.http_client, "request", side_effect=fake_get), \
532546
mock.patch.object(main, "_run_wide_refresher", new=hold_refresher):
533547
await self._serve("recycle")
534548
await refresher_started.wait()

0 commit comments

Comments
 (0)