from __future__ import annotations import asyncio 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_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()