Files
dy/backend/tests/test_batch_start_api.py
2026-09-01 15:31:05 +08:00

593 lines
22 KiB
Python

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()