593 lines
22 KiB
Python
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()
|