from __future__ import annotations import asyncio import json import os import sys import unittest from contextlib import asynccontextmanager from pathlib import Path from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock, patch BACKEND_DIR = Path(__file__).resolve().parents[1] os.environ["KEFU_DB_TYPE"] = "sqlite" os.environ["KEFU_DATABASE_URL"] = "" os.environ["KEFU_DB_PATH"] = str(BACKEND_DIR / "kefu.db") if str(BACKEND_DIR) not in sys.path: sys.path.insert(0, str(BACKEND_DIR)) import main class _RowsResult: def __init__(self, rows): self._rows = list(rows) def all(self): return list(self._rows) def _fake_db(rows): return SimpleNamespace(execute=AsyncMock(return_value=_RowsResult(rows))) class BatchStartApiTests(unittest.IsolatedAsyncioTestCase): async def test_shutdown_continues_when_batch_queue_cleanup_times_out(self): original_workers = main.manager.workers original_flush_task = main._system_log_flush_task main.manager.workers = {} main._system_log_flush_task = None stop_started = asyncio.Event() async def blocked_batch_stop(): stop_started.set() await asyncio.Event().wait() try: with ( patch.dict( os.environ, {"KEFU_BATCH_STOP_TIMEOUT_SECONDS": "1"}, ), patch.object( main.batch_start_queue, "stop", AsyncMock(side_effect=blocked_batch_stop), ), patch( "rpa_engine.douyin_im.traffic_control.shutdown_traffic_controller", AsyncMock(), ) as stop_traffic, ): await asyncio.wait_for(main.shutdown(), timeout=2.0) self.assertTrue(stop_started.is_set()) stop_traffic.assert_awaited_once_with() finally: main.manager.workers = original_workers main._system_log_flush_task = original_flush_task async def test_shutdown_stops_many_accounts_with_bounded_parallelism(self): active = 0 maximum_active = 0 stopped: list[int] = [] original_workers = main.manager.workers original_flush_task = main._system_log_flush_task main.manager.workers = { account_id: SimpleNamespace(is_running=True) for account_id in range(601, 613) } main._system_log_flush_task = None async def stop_worker(account_id: int): nonlocal active, maximum_active active += 1 maximum_active = max(maximum_active, active) try: await asyncio.sleep(0.005) stopped.append(account_id) main.manager.workers.pop(account_id, None) return True finally: active -= 1 try: with ( patch.dict( os.environ, { "KEFU_SHUTDOWN_CONCURRENCY": "3", "KEFU_SHUTDOWN_TIMEOUT_SECONDS": "5", }, ), patch.object(main.batch_start_queue, "stop", AsyncMock()), patch.object(main.manager, "stop_worker", AsyncMock(side_effect=stop_worker)), patch( "rpa_engine.douyin_im.traffic_control.shutdown_traffic_controller", AsyncMock(), ), ): await main.shutdown() self.assertEqual(len(stopped), 12) self.assertEqual(maximum_active, 3) finally: main.manager.workers = original_workers main._system_log_flush_task = original_flush_task async def test_worker_manager_waits_for_full_ready_and_reuses_validation(self): worker = SimpleNamespace( is_running=True, start=AsyncMock(), wait_until_ready=AsyncMock(), ) manager = main.WorkerManager() with patch.object(main, "DouyinWorker", return_value=worker) as worker_factory: started = await manager.start_worker( 501, login_mode="im_direct", wait_until_ready=True, credential_prevalidated=True, ) self.assertTrue(started) worker_factory.assert_called_once_with( 501, login_mode="im_direct", credential_prevalidated=True, ) worker.start.assert_awaited_once_with() worker.wait_until_ready.assert_awaited_once_with() self.assertIs(manager.workers[501], worker) async def test_cancelled_ready_wait_stops_and_removes_detached_worker(self): wait_started = asyncio.Event() waiting = asyncio.Event() async def wait_forever(): wait_started.set() await waiting.wait() worker = SimpleNamespace( is_running=True, start=AsyncMock(), wait_until_ready=AsyncMock(side_effect=wait_forever), stop=AsyncMock(), ) manager = main.WorkerManager() with patch.object(main, "DouyinWorker", return_value=worker): task = asyncio.create_task( manager.start_worker( 502, login_mode="im_direct", wait_until_ready=True, credential_prevalidated=True, ) ) await asyncio.wait_for(wait_started.wait(), timeout=0.2) task.cancel() with self.assertRaises(asyncio.CancelledError): await task worker.stop.assert_awaited_once_with() self.assertNotIn(502, manager.workers) async def test_batch_start_waits_for_ready_and_skips_duplicate_validation(self): account = SimpleNamespace( id=503, status="offline", qr_code_base64=None, error_message=None, im_session_data="saved-session", ) db = SimpleNamespace(commit=AsyncMock()) assessment = { "login_mode": "im_direct", "should_reset": False, "can_skip_browser": True, "message": "ready", "cookie_valid": True, "im_ready": True, } with ( patch.object(main.manager, "is_running", return_value=False), patch.object(main.manager, "start_worker", AsyncMock(return_value=True)) as start, patch.object(main, "_get_account_cookie_data", return_value="{}"), patch.object(main, "assess_account_credential", AsyncMock(return_value=assessment)), ): result = await main._start_account_rpa_impl( account, db, wait_for_ready=True, ) start.assert_awaited_once_with( 503, login_mode="im_direct", wait_until_ready=True, credential_prevalidated=True, ) self.assertEqual(result["status"], "running") async def test_batch_start_holds_no_db_connection_while_it_waits(self): """A queued start must not pin one of the few pooled connections. Credential validation and readiness waiting take seconds per account. Holding a session open across them exhausted the pool during a bulk start, so every unrelated request waited out ``pool_timeout``. """ account = SimpleNamespace( id=505, status="starting", qr_code_base64=None, error_message=None, im_session_data="saved-session", ) events: list[str] = [] assessment = { "login_mode": "im_direct", "should_reset": False, "can_skip_browser": True, "message": "ready", "cookie_valid": True, "im_ready": True, } async def commit(): events.append("release") async def assess(*_args, **_kwargs): events.append("assess") return assessment async def start_worker(*_args, **_kwargs): events.append("start-worker") return True db = SimpleNamespace(commit=AsyncMock(side_effect=commit)) with ( patch.object(main.manager, "is_running", return_value=False), patch.object(main.manager, "start_worker", AsyncMock(side_effect=start_worker)), patch.object(main, "_get_account_cookie_data", return_value="{}"), patch.object(main, "assess_account_credential", AsyncMock(side_effect=assess)), ): await main._start_account_rpa_impl(account, db, wait_for_ready=True) # "starting" was already persisted, so the only commits here exist to # return the connection: one before validation, one before the wait. self.assertEqual(events, ["release", "assess", "release", "start-worker"]) async def test_changed_egress_preserves_valid_credentials(self): ready_assessment = { "login_mode": "im_direct", "should_reset": False, "can_skip_browser": True, "message": "ready", "cookie_valid": True, "im_ready": True, } scenarios = ( ({"cookies": []}, "47.96.154.74"), ({"cookies": [], "credential_egress_public_ip": "116.62.23.103"}, "47.96.154.74"), ({"cookies": [], "credential_egress_public_ip": "47.96.154.74"}, ""), ) modes = (("im_direct", False), (None, False), (None, True)) for storage, selected_ip in scenarios: for requested_mode, wait_for_ready in modes: with self.subTest(storage=storage, mode=requested_mode, batch=wait_for_ready): cookie_data = json.dumps(storage) account = SimpleNamespace( id=506, status="offline", qr_code_base64=None, error_message="old channel warning", cookie_data=cookie_data, im_session_data="saved-session", egress_public_ip=selected_ip, ) db = SimpleNamespace(commit=AsyncMock()) with ( patch.object(main.manager, "is_running", return_value=False), patch.object(main.manager, "start_worker", AsyncMock(return_value=True)) as start, patch.object(main, "_get_account_cookie_data", return_value=cookie_data), patch.object(main, "_reset_account_credentials", AsyncMock()) as reset, patch.object(main, "assess_account_credential", AsyncMock(return_value=ready_assessment)) as assess, ): result = await main._start_account_rpa_impl( account, db, requested_mode, wait_for_ready=wait_for_ready ) reset.assert_not_awaited() assess.assert_awaited_once_with( cookie_data, "saved-session", startup_priority=True, egress_public_ip=selected_ip, ) start.assert_awaited_once_with( 506, login_mode="im_direct", wait_until_ready=wait_for_ready, credential_prevalidated=True, ) self.assertEqual(account.cookie_data, cookie_data) self.assertEqual(account.im_session_data, "saved-session") self.assertIsNone(account.error_message) self.assertTrue(result["skip_qr"]) self.assertTrue(result["skip_browser"]) async def test_changed_egress_still_rejects_invalid_im_credentials(self): account = SimpleNamespace( id=506, status="offline", qr_code_base64=None, error_message=None, im_session_data="saved-session", egress_public_ip="47.96.154.74", ) db = SimpleNamespace(commit=AsyncMock()) invalid_assessment = { "login_mode": "browser", "should_reset": False, "can_skip_browser": False, "message": "缺少 IM 签名密钥(web_protect/keys),请用浏览器登录补全", "cookie_valid": True, "im_ready": False, } with ( patch.object(main.manager, "is_running", return_value=False), patch.object(main.manager, "start_worker", AsyncMock()) as start, patch.object(main, "_get_account_cookie_data", return_value='{"cookies": []}'), patch.object(main, "_reset_account_credentials", AsyncMock()) as reset, patch.object(main, "assess_account_credential", AsyncMock(return_value=invalid_assessment)), ): with self.assertRaises(main.HTTPException) as error: await main._start_account_rpa_impl(account, db, "im_direct") self.assertEqual(error.exception.status_code, 400) self.assertEqual(error.exception.detail, invalid_assessment["message"]) reset.assert_not_awaited() start.assert_not_awaited() async def test_batch_start_does_not_launch_interactive_browser_login(self): account = SimpleNamespace( id=504, status="offline", qr_code_base64=None, error_message=None, im_session_data=None, ) db = SimpleNamespace(commit=AsyncMock()) assessment = { "login_mode": "browser", "should_reset": False, "can_skip_browser": False, "message": "login required", "cookie_valid": False, "im_ready": False, } with ( patch.object(main.manager, "is_running", return_value=False), patch.object(main.manager, "start_worker", AsyncMock()) as start, patch.object(main, "_get_account_cookie_data", return_value=None), patch.object(main, "assess_account_credential", AsyncMock(return_value=assessment)), ): with self.assertRaisesRegex(RuntimeError, "批量启动"): await main._start_account_rpa_impl( account, db, wait_for_ready=True, ) start.assert_not_awaited() async def test_start_all_uses_lightweight_select_and_submits_once(self): db = _fake_db([(1, False), (2, True), (3, False), (4, False)]) user = SimpleNamespace(id=9, role="admin") submitted_response = { "batch_id": "all-batch", "accepted_count": 2, "complete": False, } submit = AsyncMock(return_value=submitted_response) with ( patch.object(main.batch_start_queue, "submit", submit), patch.object( main.manager, "is_running", side_effect=lambda account_id: account_id == 3, ) as is_running, patch.object(main, "_build_account_response") as build_response, ): response = await main.submit_account_start_batch( body=main.BatchStartRequest(all_accounts=True), db=db, user=user, ) self.assertEqual(response, submitted_response) db.execute.assert_awaited_once() submit.assert_awaited_once_with( [1, 4], owner_id=9, metadata={ "requested_count": 4, "accessible_count": 4, "skipped_running_count": 1, "skipped_disabled_count": 1, }, ) self.assertEqual(is_running.call_count, 3) build_response.assert_not_called() statement = db.execute.await_args.args[0] selected_names = [entry.get("name") for entry in statement.column_descriptions] self.assertEqual(selected_names, ["id", "quota_disabled"]) self.assertNotIn("cookie_data", str(statement).lower()) self.assertNotIn("im_session_data", str(statement).lower()) async def test_selected_ids_are_deduplicated_scoped_and_filtered(self): db = _fake_db([(5, False), (3, False)]) user = SimpleNamespace(id=42, role="operator") submit = AsyncMock( return_value={ "batch_id": "selected-batch", "accepted_count": 2, "complete": False, } ) with ( patch.object(main.batch_start_queue, "submit", submit), patch.object(main.manager, "is_running", return_value=False), ): await main.submit_account_start_batch( body=main.BatchStartRequest( account_ids=[5, 5, 3, -1, 999], all_accounts=False, ), db=db, user=user, ) submit.assert_awaited_once_with( [5, 3], owner_id=42, metadata={ "requested_count": 3, "accessible_count": 2, "skipped_running_count": 0, "skipped_disabled_count": 0, }, ) statement = db.execute.await_args.args[0] compiled_params = statement.compile().params parameter_values = list(compiled_params.values()) self.assertIn(42, parameter_values) self.assertIn([5, 3, 999], parameter_values) sql = str(statement).lower() self.assertIn("owner_id", sql) self.assertIn("accounts.id in", sql) async def test_metadata_counts_inaccessible_running_and_disabled_accounts(self): db = _fake_db([(101, False), (102, True), (103, False)]) user = SimpleNamespace(id=7, role="operator") submit = AsyncMock(return_value={"batch_id": "metadata-batch"}) with ( patch.object(main.batch_start_queue, "submit", submit), patch.object( main.manager, "is_running", side_effect=lambda account_id: account_id == 103, ), ): await main.submit_account_start_batch( body=main.BatchStartRequest(account_ids=[101, 102, 103, 104]), db=db, user=user, ) submit.assert_awaited_once_with( [101], owner_id=7, metadata={ "requested_count": 4, "accessible_count": 3, "skipped_running_count": 1, "skipped_disabled_count": 1, }, ) async def test_batch_lookup_is_isolated_by_owner(self): async def get_batch(batch_id, owner_id, include_items=False): if batch_id == "owned-batch" and owner_id == 11: return { "batch_id": batch_id, "accepted_count": 2, "complete": True, } return None lookup = AsyncMock(side_effect=get_batch) with patch.object(main.batch_start_queue, "get_batch", lookup): owned = await main.get_account_start_batch( batch_id="owned-batch", user=SimpleNamespace(id=11, role="operator"), ) with self.assertRaises(main.HTTPException) as caught: await main.get_account_start_batch( batch_id="owned-batch", user=SimpleNamespace(id=12, role="operator"), ) self.assertEqual(owned["batch_id"], "owned-batch") self.assertEqual(caught.exception.status_code, 404) self.assertEqual( lookup.await_args_list[0].kwargs, {"owner_id": 11, "include_items": False}, ) self.assertEqual( lookup.await_args_list[1].kwargs, {"owner_id": 12, "include_items": False}, ) async def test_single_account_start_cancels_queued_batch_job_first(self): account = SimpleNamespace(id=77) db = SimpleNamespace() user = SimpleNamespace(id=5, role="operator") events: list[str] = [] async def cancel_account(account_id: int): events.append("cancel") return 1 @asynccontextmanager async def preparation_lock(account_id: int): self.assertEqual(account_id, 77) events.append("lock-enter") try: yield finally: events.append("lock-exit") async def start_impl(selected_account, selected_db, requested_login_mode): events.append("start") self.assertIs(selected_account, account) self.assertIs(selected_db, db) self.assertEqual(requested_login_mode, "im_direct") return {"status": "starting", "login_mode": "im_direct"} cancel = AsyncMock(side_effect=cancel_account) get_owned = AsyncMock(return_value=account) start = AsyncMock(side_effect=start_impl) with ( patch.object(main, "get_owned_account", get_owned), patch.object(main.batch_start_queue, "cancel_account", cancel), patch.object(main.manager, "preparation_lock", side_effect=preparation_lock), patch.object(main, "_start_account_rpa_impl", start), ): response = await main.start_account_rpa( account_id=77, body=main.StartAccountRequest(login_mode="im_direct"), db=db, user=user, ) self.assertEqual(response["status"], "starting") get_owned.assert_awaited_once_with(db, user, 77, write=True) cancel.assert_awaited_once_with(77) start.assert_awaited_once_with(account, db, "im_direct") self.assertEqual(events, ["cancel", "lock-enter", "start", "lock-exit"]) if __name__ == "__main__": unittest.main()