更新
This commit is contained in:
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,239 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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_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()
|
||||
@@ -0,0 +1,288 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
BACKEND_DIR = Path(__file__).resolve().parents[1]
|
||||
if str(BACKEND_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(BACKEND_DIR))
|
||||
|
||||
from rpa_engine.batch_start import BatchStartQueue
|
||||
from rpa_engine import batch_start as batch_start_module
|
||||
|
||||
|
||||
class BatchStartQueueTests(unittest.IsolatedAsyncioTestCase):
|
||||
def _make_queue(self, handler, *, concurrency: int = 2) -> BatchStartQueue:
|
||||
queue = BatchStartQueue(handler, concurrency=concurrency)
|
||||
self.addAsyncCleanup(queue.stop)
|
||||
return queue
|
||||
|
||||
async def _wait_for_complete(
|
||||
self,
|
||||
queue: BatchStartQueue,
|
||||
batch_id: str,
|
||||
*,
|
||||
timeout: float = 0.5,
|
||||
) -> dict:
|
||||
async def wait() -> dict:
|
||||
while True:
|
||||
snapshot = await queue.get_batch(batch_id, include_items=True)
|
||||
self.assertIsNotNone(snapshot)
|
||||
if snapshot["complete"]:
|
||||
return snapshot
|
||||
await asyncio.sleep(0.001)
|
||||
|
||||
return await asyncio.wait_for(wait(), timeout=timeout)
|
||||
|
||||
async def test_concurrency_limit_is_two(self):
|
||||
active = 0
|
||||
maximum_active = 0
|
||||
started: list[int] = []
|
||||
two_started = asyncio.Event()
|
||||
release = asyncio.Event()
|
||||
|
||||
async def handler(account_id: int) -> dict:
|
||||
nonlocal active, maximum_active
|
||||
active += 1
|
||||
maximum_active = max(maximum_active, active)
|
||||
started.append(account_id)
|
||||
if len(started) == 2:
|
||||
two_started.set()
|
||||
try:
|
||||
await release.wait()
|
||||
return {"message": f"started-{account_id}", "login_mode": "im_direct"}
|
||||
finally:
|
||||
active -= 1
|
||||
|
||||
queue = self._make_queue(handler, concurrency=2)
|
||||
submitted = await queue.submit(list(range(1, 9)))
|
||||
|
||||
await asyncio.wait_for(two_started.wait(), timeout=0.2)
|
||||
await asyncio.sleep(0.02)
|
||||
|
||||
self.assertEqual(len(started), 2)
|
||||
self.assertEqual(maximum_active, 2)
|
||||
self.assertEqual(active, 2)
|
||||
|
||||
release.set()
|
||||
completed = await self._wait_for_complete(queue, submitted["batch_id"])
|
||||
|
||||
self.assertEqual(completed["submitted_count"], 8)
|
||||
self.assertEqual(completed["failed_count"], 0)
|
||||
self.assertEqual(maximum_active, 2)
|
||||
self.assertEqual(active, 0)
|
||||
|
||||
async def test_submit_returns_while_handler_is_blocked(self):
|
||||
handler_started = asyncio.Event()
|
||||
release_handler = asyncio.Event()
|
||||
handler_finished = asyncio.Event()
|
||||
|
||||
async def handler(account_id: int) -> dict:
|
||||
handler_started.set()
|
||||
await release_handler.wait()
|
||||
handler_finished.set()
|
||||
return {"message": f"started-{account_id}"}
|
||||
|
||||
queue = self._make_queue(handler, concurrency=1)
|
||||
|
||||
submitted = await asyncio.wait_for(queue.submit([11]), timeout=0.2)
|
||||
|
||||
self.assertEqual(submitted["accepted_count"], 1)
|
||||
self.assertEqual(submitted["queued_count"], 1)
|
||||
self.assertFalse(submitted["complete"])
|
||||
self.assertFalse(handler_finished.is_set())
|
||||
|
||||
await asyncio.wait_for(handler_started.wait(), timeout=0.2)
|
||||
processing = await queue.get_batch(submitted["batch_id"])
|
||||
self.assertEqual(processing["processing_count"], 1)
|
||||
self.assertFalse(handler_finished.is_set())
|
||||
|
||||
release_handler.set()
|
||||
completed = await self._wait_for_complete(queue, submitted["batch_id"])
|
||||
self.assertTrue(handler_finished.is_set())
|
||||
self.assertEqual(completed["submitted_count"], 1)
|
||||
|
||||
async def test_duplicate_ids_and_overlapping_batches_are_deduplicated(self):
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
calls: list[int] = []
|
||||
|
||||
async def handler(account_id: int) -> dict:
|
||||
calls.append(account_id)
|
||||
if account_id == 7:
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
return {"message": f"started-{account_id}"}
|
||||
|
||||
queue = self._make_queue(handler, concurrency=1)
|
||||
first = await queue.submit([7, 7])
|
||||
self.assertEqual(first["total_count"], 1)
|
||||
self.assertEqual(first["accepted_count"], 1)
|
||||
await asyncio.wait_for(first_started.wait(), timeout=0.2)
|
||||
|
||||
overlapping = await queue.submit([7, 8, 8])
|
||||
|
||||
self.assertEqual(overlapping["total_count"], 2)
|
||||
self.assertEqual(overlapping["accepted_count"], 1)
|
||||
self.assertEqual(overlapping["skipped_count"], 1)
|
||||
by_account = {item["account_id"]: item for item in overlapping["items"]}
|
||||
self.assertEqual(by_account[7]["status"], "already_queued")
|
||||
self.assertEqual(by_account[8]["status"], "queued")
|
||||
|
||||
release_first.set()
|
||||
await self._wait_for_complete(queue, first["batch_id"])
|
||||
completed_overlap = await self._wait_for_complete(
|
||||
queue,
|
||||
overlapping["batch_id"],
|
||||
)
|
||||
|
||||
self.assertEqual(calls, [7, 8])
|
||||
self.assertEqual(completed_overlap["submitted_count"], 1)
|
||||
self.assertEqual(completed_overlap["skipped_count"], 1)
|
||||
|
||||
async def test_failure_does_not_block_following_accounts(self):
|
||||
calls: list[int] = []
|
||||
|
||||
async def handler(account_id: int) -> dict:
|
||||
calls.append(account_id)
|
||||
if account_id == 31:
|
||||
raise RuntimeError("expected startup failure")
|
||||
return {"message": f"started-{account_id}"}
|
||||
|
||||
queue = self._make_queue(handler, concurrency=1)
|
||||
with patch.object(batch_start_module.logger, "exception"):
|
||||
submitted = await queue.submit([31, 32, 33])
|
||||
completed = await self._wait_for_complete(queue, submitted["batch_id"])
|
||||
|
||||
self.assertEqual(calls, [31, 32, 33])
|
||||
self.assertEqual(completed["failed_count"], 1)
|
||||
self.assertEqual(completed["submitted_count"], 2)
|
||||
by_account = {item["account_id"]: item for item in completed["items"]}
|
||||
self.assertEqual(by_account[31]["status"], "failed")
|
||||
self.assertEqual(by_account[31]["message"], "expected startup failure")
|
||||
self.assertEqual(by_account[32]["status"], "submitted")
|
||||
self.assertEqual(by_account[33]["status"], "submitted")
|
||||
|
||||
async def test_failed_account_can_be_submitted_again(self):
|
||||
attempts = 0
|
||||
|
||||
async def handler(account_id: int) -> dict:
|
||||
nonlocal attempts
|
||||
attempts += 1
|
||||
if attempts == 1:
|
||||
raise RuntimeError("fail once")
|
||||
return {"message": f"started-{account_id}"}
|
||||
|
||||
queue = self._make_queue(handler, concurrency=1)
|
||||
with patch.object(batch_start_module.logger, "exception"):
|
||||
first = await queue.submit([40])
|
||||
first_completed = await self._wait_for_complete(queue, first["batch_id"])
|
||||
|
||||
second = await queue.submit([40])
|
||||
second_completed = await self._wait_for_complete(queue, second["batch_id"])
|
||||
|
||||
self.assertEqual(first_completed["failed_count"], 1)
|
||||
self.assertEqual(second["accepted_count"], 1)
|
||||
self.assertEqual(second["skipped_count"], 0)
|
||||
self.assertEqual(second_completed["submitted_count"], 1)
|
||||
self.assertEqual(attempts, 2)
|
||||
|
||||
async def test_cancel_then_immediate_resubmit_keeps_new_job_deduplicated(self):
|
||||
blocker_started = asyncio.Event()
|
||||
release_blocker = asyncio.Event()
|
||||
replacement_started = asyncio.Event()
|
||||
release_replacement = asyncio.Event()
|
||||
calls: list[int] = []
|
||||
|
||||
async def handler(account_id: int) -> dict:
|
||||
calls.append(account_id)
|
||||
if account_id == 60:
|
||||
blocker_started.set()
|
||||
await release_blocker.wait()
|
||||
elif account_id == 61:
|
||||
replacement_started.set()
|
||||
await release_replacement.wait()
|
||||
return {"message": f"started-{account_id}"}
|
||||
|
||||
queue = self._make_queue(handler, concurrency=1)
|
||||
blocker = await queue.submit([60])
|
||||
await asyncio.wait_for(blocker_started.wait(), timeout=0.2)
|
||||
|
||||
cancelled_batch = await queue.submit([61])
|
||||
self.assertEqual(await queue.cancel_account(61), 1)
|
||||
|
||||
replacement = await queue.submit([61])
|
||||
self.assertEqual(replacement["accepted_count"], 1)
|
||||
|
||||
# Releasing the blocker makes the worker consume the stale cancelled
|
||||
# queue entry before it starts the replacement. Cleanup for that old
|
||||
# token must not release the replacement's pending ownership.
|
||||
release_blocker.set()
|
||||
await asyncio.wait_for(replacement_started.wait(), timeout=0.2)
|
||||
|
||||
duplicate = await queue.submit([61])
|
||||
self.assertEqual(duplicate["accepted_count"], 0)
|
||||
self.assertEqual(duplicate["skipped_count"], 1)
|
||||
|
||||
release_replacement.set()
|
||||
await self._wait_for_complete(queue, blocker["batch_id"])
|
||||
cancelled = await self._wait_for_complete(
|
||||
queue,
|
||||
cancelled_batch["batch_id"],
|
||||
)
|
||||
completed = await self._wait_for_complete(queue, replacement["batch_id"])
|
||||
|
||||
self.assertEqual(cancelled["cancelled_count"], 1)
|
||||
self.assertEqual(completed["submitted_count"], 1)
|
||||
self.assertEqual(calls, [60, 61])
|
||||
|
||||
async def test_stop_cancels_active_and_queued_jobs_and_cleans_workers(self):
|
||||
active = 0
|
||||
two_started = asyncio.Event()
|
||||
started: list[int] = []
|
||||
never_release = asyncio.Event()
|
||||
|
||||
async def handler(account_id: int) -> dict:
|
||||
nonlocal active
|
||||
active += 1
|
||||
started.append(account_id)
|
||||
if len(started) == 2:
|
||||
two_started.set()
|
||||
try:
|
||||
await never_release.wait()
|
||||
return {"message": f"started-{account_id}"}
|
||||
finally:
|
||||
active -= 1
|
||||
|
||||
queue = self._make_queue(handler, concurrency=2)
|
||||
submitted = await queue.submit([51, 52, 53, 54])
|
||||
await asyncio.wait_for(two_started.wait(), timeout=0.2)
|
||||
|
||||
await asyncio.wait_for(queue.stop(), timeout=0.2)
|
||||
completed = await queue.get_batch(submitted["batch_id"])
|
||||
|
||||
self.assertTrue(completed["complete"])
|
||||
self.assertEqual(completed["cancelled_count"], 4)
|
||||
self.assertEqual(active, 0)
|
||||
self.assertEqual(queue._pending_jobs, {})
|
||||
self.assertEqual(queue._workers, [])
|
||||
self.assertTrue(queue._queue.empty())
|
||||
await asyncio.wait_for(queue._queue.join(), timeout=0.1)
|
||||
|
||||
leaked = [
|
||||
task
|
||||
for task in asyncio.all_tasks()
|
||||
if task is not asyncio.current_task()
|
||||
and task.get_name().startswith("account-batch-start-")
|
||||
and not task.done()
|
||||
]
|
||||
self.assertEqual(leaked, [])
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,80 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import AsyncMock
|
||||
|
||||
|
||||
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))
|
||||
|
||||
from rpa_engine.douyin_im.http_client import DouyinImHttpClient
|
||||
from rpa_engine.douyin_im.session import DouyinImSession
|
||||
|
||||
|
||||
class ConversationPollBandwidthTests(unittest.IsolatedAsyncioTestCase):
|
||||
def _make_client(self) -> DouyinImHttpClient:
|
||||
return DouyinImHttpClient(
|
||||
DouyinImSession(cookies={"sessionid": "test"}, my_uid=10001),
|
||||
account_id=9,
|
||||
)
|
||||
|
||||
async def test_terminal_token_error_does_not_probe_other_payloads(self):
|
||||
client = self._make_client()
|
||||
client._request = AsyncMock(
|
||||
return_value={
|
||||
"status_code": 500,
|
||||
"error_desc": "empty token",
|
||||
"body": {},
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(await client.get_conversations(), [])
|
||||
client._request.assert_awaited_once()
|
||||
self.assertEqual(client._request.await_args.args[0], "POST")
|
||||
|
||||
async def test_successful_empty_response_stops_after_first_payload(self):
|
||||
client = self._make_client()
|
||||
client._request = AsyncMock(
|
||||
return_value={"status_code": 0, "body": {"conversation_list": []}}
|
||||
)
|
||||
|
||||
self.assertEqual(await client.get_conversations(), [])
|
||||
client._request.assert_awaited_once()
|
||||
|
||||
async def test_parameter_error_can_fall_through_to_compatible_payload(self):
|
||||
client = self._make_client()
|
||||
client._request = AsyncMock(
|
||||
side_effect=[
|
||||
{"status_code": 400, "error_desc": "invalid parameter"},
|
||||
{"status_code": 0, "body": {"conversation_list": []}},
|
||||
]
|
||||
)
|
||||
|
||||
self.assertEqual(await client.get_conversations(), [])
|
||||
self.assertEqual(client._request.await_count, 2)
|
||||
self.assertTrue(
|
||||
all(call.args[0] == "POST" for call in client._request.await_args_list)
|
||||
)
|
||||
|
||||
async def test_get_fallback_only_runs_after_transport_failure(self):
|
||||
client = self._make_client()
|
||||
client._request = AsyncMock(
|
||||
side_effect=[None, {"status_code": 0, "body": {}}]
|
||||
)
|
||||
|
||||
self.assertEqual(await client.get_conversations(), [])
|
||||
self.assertEqual(
|
||||
[call.args[0] for call in client._request.await_args_list],
|
||||
["POST", "GET"],
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,263 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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
|
||||
from rpa_engine import account_profile as account_profile_module
|
||||
|
||||
|
||||
class CookieCredentialLockTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_invalid_cookie_does_not_stop_or_lock_hosting(self):
|
||||
db = SimpleNamespace()
|
||||
user = SimpleNamespace(id=9, role="operator")
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
main,
|
||||
"validate_cookie_json",
|
||||
side_effect=ValueError("invalid json"),
|
||||
),
|
||||
patch.object(main, "get_owned_account", new=AsyncMock()) as get_owned,
|
||||
patch.object(
|
||||
main.batch_start_queue,
|
||||
"cancel_account",
|
||||
new=AsyncMock(),
|
||||
) as cancel,
|
||||
patch.object(main.manager, "stop_worker", new=AsyncMock()) as stop,
|
||||
patch.object(main.manager, "preparation_lock") as preparation_lock,
|
||||
):
|
||||
with self.assertRaises(main.HTTPException) as caught:
|
||||
await main.update_account_cookie(
|
||||
account_id=81,
|
||||
body=main.AccountCookieUpdate(cookie_data="{bad-json"),
|
||||
db=db,
|
||||
user=user,
|
||||
)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 400)
|
||||
get_owned.assert_not_awaited()
|
||||
cancel.assert_not_awaited()
|
||||
stop.assert_not_awaited()
|
||||
preparation_lock.assert_not_called()
|
||||
|
||||
async def test_cookie_update_stays_locked_through_stop_and_commit(self):
|
||||
account = SimpleNamespace(
|
||||
id=82,
|
||||
cookie_data="old-cookie",
|
||||
cookie_path="old-path",
|
||||
cookie_updated_at=None,
|
||||
updated_at=None,
|
||||
)
|
||||
user = SimpleNamespace(id=9, role="operator")
|
||||
events: list[str] = []
|
||||
lock_active = False
|
||||
|
||||
def require_lock(event: str) -> None:
|
||||
self.assertTrue(lock_active, f"{event} ran outside preparation lock")
|
||||
events.append(event)
|
||||
|
||||
@asynccontextmanager
|
||||
async def preparation_lock(account_id: int):
|
||||
nonlocal lock_active
|
||||
self.assertEqual(account_id, 82)
|
||||
lock_active = True
|
||||
events.append("lock-enter")
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
events.append("lock-exit")
|
||||
lock_active = False
|
||||
|
||||
async def get_owned(db, selected_user, account_id, *, write=False):
|
||||
require_lock("authorize")
|
||||
self.assertIs(db, fake_db)
|
||||
self.assertIs(selected_user, user)
|
||||
self.assertEqual(account_id, 82)
|
||||
self.assertTrue(write)
|
||||
return account
|
||||
|
||||
async def cancel_account(account_id: int):
|
||||
require_lock("cancel")
|
||||
self.assertEqual(account_id, 82)
|
||||
return 0
|
||||
|
||||
async def stop_worker(account_id: int):
|
||||
require_lock("stop")
|
||||
self.assertEqual(account_id, 82)
|
||||
return True
|
||||
|
||||
async def execute(_statement):
|
||||
require_lock("clear-profile")
|
||||
return None
|
||||
|
||||
async def commit():
|
||||
require_lock("commit")
|
||||
|
||||
async def refresh(selected_account):
|
||||
require_lock("refresh")
|
||||
self.assertIs(selected_account, account)
|
||||
|
||||
async def apply_profile(db, selected_account, cookie_data):
|
||||
require_lock("sync-profile")
|
||||
self.assertIs(db, fake_db)
|
||||
self.assertIs(selected_account, account)
|
||||
self.assertIn("sessionid", cookie_data)
|
||||
|
||||
fake_db = SimpleNamespace(
|
||||
execute=AsyncMock(side_effect=execute),
|
||||
commit=AsyncMock(side_effect=commit),
|
||||
refresh=AsyncMock(side_effect=refresh),
|
||||
)
|
||||
|
||||
def write_cookie(account_id: int, cookie_data: str) -> str:
|
||||
require_lock("write-cookie")
|
||||
self.assertEqual(account_id, 82)
|
||||
self.assertIn("sessionid", cookie_data)
|
||||
return "new-cookie-path"
|
||||
|
||||
with (
|
||||
patch.object(main, "validate_cookie_json", return_value={"cookies": [{"name": "sessionid", "value": "new"}]}),
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(side_effect=get_owned)) as get_owned_mock,
|
||||
patch.object(main.manager, "preparation_lock", side_effect=preparation_lock),
|
||||
patch.object(main.batch_start_queue, "cancel_account", new=AsyncMock(side_effect=cancel_account)) as cancel,
|
||||
patch.object(main.manager, "stop_worker", new=AsyncMock(side_effect=stop_worker)) as stop,
|
||||
patch.object(main, "write_cookie_file", side_effect=write_cookie),
|
||||
patch.object(account_profile_module, "apply_douyin_profile", new=AsyncMock(side_effect=apply_profile)),
|
||||
patch.object(main, "_build_cookie_response", new=AsyncMock(return_value={"account_id": 82})) as build_response,
|
||||
):
|
||||
response = await main.update_account_cookie(
|
||||
account_id=82,
|
||||
body=main.AccountCookieUpdate(cookie_data="valid-cookie"),
|
||||
db=fake_db,
|
||||
user=user,
|
||||
)
|
||||
|
||||
self.assertEqual(response, {"account_id": 82})
|
||||
self.assertEqual(
|
||||
events,
|
||||
[
|
||||
"lock-enter",
|
||||
"authorize",
|
||||
"cancel",
|
||||
"stop",
|
||||
"write-cookie",
|
||||
"clear-profile",
|
||||
"sync-profile",
|
||||
"commit",
|
||||
"refresh",
|
||||
"lock-exit",
|
||||
],
|
||||
)
|
||||
get_owned_mock.assert_awaited_once()
|
||||
cancel.assert_awaited_once_with(82)
|
||||
stop.assert_awaited_once_with(82)
|
||||
self.assertEqual(account.cookie_path, "new-cookie-path")
|
||||
build_response.assert_awaited_once()
|
||||
|
||||
async def test_cookie_delete_stays_locked_through_stop_and_commit(self):
|
||||
account = SimpleNamespace(
|
||||
id=83,
|
||||
cookie_data="old-cookie",
|
||||
cookie_path="old-path",
|
||||
cookie_updated_at=object(),
|
||||
im_session_data="old-im-session",
|
||||
updated_at=None,
|
||||
)
|
||||
user = SimpleNamespace(id=9, role="operator")
|
||||
events: list[str] = []
|
||||
lock_active = False
|
||||
|
||||
def require_lock(event: str) -> None:
|
||||
self.assertTrue(lock_active, f"{event} ran outside preparation lock")
|
||||
events.append(event)
|
||||
|
||||
@asynccontextmanager
|
||||
async def preparation_lock(account_id: int):
|
||||
nonlocal lock_active
|
||||
self.assertEqual(account_id, 83)
|
||||
lock_active = True
|
||||
events.append("lock-enter")
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
events.append("lock-exit")
|
||||
lock_active = False
|
||||
|
||||
async def get_owned(_db, _user, account_id, *, write=False):
|
||||
require_lock("authorize")
|
||||
self.assertEqual(account_id, 83)
|
||||
self.assertTrue(write)
|
||||
return account
|
||||
|
||||
async def cancel_account(_account_id: int):
|
||||
require_lock("cancel")
|
||||
return 0
|
||||
|
||||
async def stop_worker(_account_id: int):
|
||||
require_lock("stop")
|
||||
return True
|
||||
|
||||
async def execute(_statement):
|
||||
require_lock("clear-profile")
|
||||
|
||||
async def commit():
|
||||
require_lock("commit")
|
||||
|
||||
fake_db = SimpleNamespace(
|
||||
execute=AsyncMock(side_effect=execute),
|
||||
commit=AsyncMock(side_effect=commit),
|
||||
)
|
||||
|
||||
def clear_cookie(account_id: int) -> None:
|
||||
require_lock("clear-cookie-file")
|
||||
self.assertEqual(account_id, 83)
|
||||
|
||||
with (
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(side_effect=get_owned)),
|
||||
patch.object(main.manager, "preparation_lock", side_effect=preparation_lock),
|
||||
patch.object(main.batch_start_queue, "cancel_account", new=AsyncMock(side_effect=cancel_account)),
|
||||
patch.object(main.manager, "stop_worker", new=AsyncMock(side_effect=stop_worker)),
|
||||
patch.object(main, "clear_cookie_file", side_effect=clear_cookie),
|
||||
):
|
||||
response = await main.delete_account_cookie(
|
||||
account_id=83,
|
||||
db=fake_db,
|
||||
user=user,
|
||||
)
|
||||
|
||||
self.assertEqual(response["message"], "Cookie cleared successfully.")
|
||||
self.assertEqual(
|
||||
events,
|
||||
[
|
||||
"lock-enter",
|
||||
"authorize",
|
||||
"cancel",
|
||||
"stop",
|
||||
"clear-cookie-file",
|
||||
"clear-profile",
|
||||
"commit",
|
||||
"lock-exit",
|
||||
],
|
||||
)
|
||||
self.assertIsNone(account.cookie_data)
|
||||
self.assertIsNone(account.cookie_path)
|
||||
self.assertIsNone(account.cookie_updated_at)
|
||||
self.assertIsNone(account.im_session_data)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,202 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from datetime import datetime, timedelta
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, 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 _ProfileResult:
|
||||
def __init__(self, profile):
|
||||
self.profile = profile
|
||||
|
||||
def scalar_one_or_none(self):
|
||||
return self.profile
|
||||
|
||||
|
||||
class DesktopLoginSecUserIdTests(unittest.IsolatedAsyncioTestCase):
|
||||
@staticmethod
|
||||
def _account(*, cookie_data: str | None = "cookie-json"):
|
||||
return SimpleNamespace(
|
||||
id=501,
|
||||
cookie_data=cookie_data,
|
||||
cookie_path=None,
|
||||
cookie_updated_at=datetime(2026, 7, 22, 8, 0, 0),
|
||||
im_session_data=None,
|
||||
)
|
||||
|
||||
async def test_verified_sec_user_id_returns_credential(self):
|
||||
account = self._account()
|
||||
profile = SimpleNamespace(
|
||||
sec_user_id=" MS4wLjAB-valid ",
|
||||
synced_at=account.cookie_updated_at + timedelta(seconds=1),
|
||||
)
|
||||
db = SimpleNamespace(execute=AsyncMock(return_value=_ProfileResult(profile)))
|
||||
response = main.AccountCookieResponse(account_id=501, cookie_data="cookie-json")
|
||||
|
||||
with (
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(return_value=account)),
|
||||
patch.object(main, "_build_cookie_response", new=AsyncMock(return_value=response)) as build,
|
||||
):
|
||||
result = await main.get_desktop_login_credential(
|
||||
account_id=501,
|
||||
db=db,
|
||||
user=SimpleNamespace(id=9, role="operator"),
|
||||
)
|
||||
|
||||
self.assertIs(result, response)
|
||||
db.execute.assert_awaited_once()
|
||||
build.assert_awaited_once_with(account, "cookie-json")
|
||||
|
||||
async def test_missing_or_blank_sec_user_id_is_rejected_without_cookie(self):
|
||||
for missing_value in (None, "", " "):
|
||||
with self.subTest(sec_user_id=missing_value):
|
||||
account = self._account()
|
||||
profile = SimpleNamespace(
|
||||
sec_user_id=missing_value,
|
||||
synced_at=account.cookie_updated_at,
|
||||
)
|
||||
db = SimpleNamespace(
|
||||
execute=AsyncMock(return_value=_ProfileResult(profile))
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(return_value=account)),
|
||||
patch.object(main, "_build_cookie_response", new=AsyncMock()) as build,
|
||||
):
|
||||
with self.assertRaises(main.HTTPException) as caught:
|
||||
await main.get_desktop_login_credential(
|
||||
account_id=501,
|
||||
db=db,
|
||||
user=SimpleNamespace(id=9, role="operator"),
|
||||
)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 409)
|
||||
self.assertIn("缺少 sec_user_id", str(caught.exception.detail))
|
||||
build.assert_not_awaited()
|
||||
|
||||
async def test_missing_profile_or_stale_identity_is_rejected(self):
|
||||
account = self._account()
|
||||
stale_profile = SimpleNamespace(
|
||||
sec_user_id="MS4wLjAB-old",
|
||||
synced_at=account.cookie_updated_at - timedelta(seconds=1),
|
||||
)
|
||||
for profile in (None, stale_profile):
|
||||
with self.subTest(profile=profile):
|
||||
db = SimpleNamespace(
|
||||
execute=AsyncMock(return_value=_ProfileResult(profile))
|
||||
)
|
||||
with (
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(return_value=account)),
|
||||
patch.object(main, "_build_cookie_response", new=AsyncMock()) as build,
|
||||
):
|
||||
with self.assertRaises(main.HTTPException) as caught:
|
||||
await main.get_desktop_login_credential(
|
||||
account_id=501,
|
||||
db=db,
|
||||
user=SimpleNamespace(id=9, role="operator"),
|
||||
)
|
||||
self.assertEqual(caught.exception.status_code, 409)
|
||||
build.assert_not_awaited()
|
||||
|
||||
async def test_profile_check_failure_is_unknown_and_fails_closed(self):
|
||||
account = self._account()
|
||||
db = SimpleNamespace(execute=AsyncMock(side_effect=RuntimeError("db locked")))
|
||||
|
||||
with (
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(return_value=account)),
|
||||
patch.object(main, "_build_cookie_response", new=AsyncMock()) as build,
|
||||
patch.object(main.logger, "exception") as log_exception,
|
||||
):
|
||||
with self.assertRaises(main.HTTPException) as caught:
|
||||
await main.get_desktop_login_credential(
|
||||
account_id=501,
|
||||
db=db,
|
||||
user=SimpleNamespace(id=9, role="operator"),
|
||||
)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 503)
|
||||
self.assertNotIn("缺少 sec_user_id", str(caught.exception.detail))
|
||||
build.assert_not_awaited()
|
||||
log_exception.assert_called_once()
|
||||
|
||||
async def test_no_cookie_keeps_existing_no_login_state_flow(self):
|
||||
account = self._account(cookie_data=None)
|
||||
db = SimpleNamespace(execute=AsyncMock())
|
||||
response = main.AccountCookieResponse(account_id=501, cookie_data=None)
|
||||
|
||||
with (
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(return_value=account)),
|
||||
patch.object(main, "_get_account_cookie_data", return_value=None),
|
||||
patch.object(main, "_build_cookie_response", new=AsyncMock(return_value=response)) as build,
|
||||
):
|
||||
result = await main.get_desktop_login_credential(
|
||||
account_id=501,
|
||||
db=db,
|
||||
user=SimpleNamespace(id=9, role="operator"),
|
||||
)
|
||||
|
||||
self.assertIs(result, response)
|
||||
db.execute.assert_not_awaited()
|
||||
build.assert_awaited_once_with(account, None)
|
||||
|
||||
async def test_normal_cookie_management_endpoint_is_not_guarded(self):
|
||||
account = self._account()
|
||||
db = SimpleNamespace(execute=AsyncMock())
|
||||
response = main.AccountCookieResponse(account_id=501, cookie_data="cookie-json")
|
||||
|
||||
with (
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(return_value=account)),
|
||||
patch.object(main, "_build_cookie_response", new=AsyncMock(return_value=response)) as build,
|
||||
):
|
||||
result = await main.get_account_cookie(
|
||||
account_id=501,
|
||||
purpose="management",
|
||||
db=db,
|
||||
user=SimpleNamespace(id=9, role="operator"),
|
||||
)
|
||||
|
||||
self.assertIs(result, response)
|
||||
db.execute.assert_not_awaited()
|
||||
build.assert_awaited_once_with(account, "cookie-json")
|
||||
|
||||
async def test_legacy_cookie_request_is_guarded_for_existing_desktop_clients(self):
|
||||
account = self._account()
|
||||
profile = SimpleNamespace(
|
||||
sec_user_id=" ",
|
||||
synced_at=account.cookie_updated_at,
|
||||
)
|
||||
db = SimpleNamespace(execute=AsyncMock(return_value=_ProfileResult(profile)))
|
||||
|
||||
with (
|
||||
patch.object(main, "get_owned_account", new=AsyncMock(return_value=account)),
|
||||
patch.object(main, "_build_cookie_response", new=AsyncMock()) as build,
|
||||
):
|
||||
with self.assertRaises(main.HTTPException) as caught:
|
||||
await main.get_account_cookie(
|
||||
account_id=501,
|
||||
purpose=None,
|
||||
db=db,
|
||||
user=SimpleNamespace(id=9, role="operator"),
|
||||
)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 409)
|
||||
self.assertIn("缺少 sec_user_id", str(caught.exception.detail))
|
||||
build.assert_not_awaited()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,154 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import 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 _FakeResponse:
|
||||
def __init__(
|
||||
self,
|
||||
*,
|
||||
final_host: str,
|
||||
headers: dict[str, str] | None = None,
|
||||
chunks: tuple[bytes, ...] = (),
|
||||
) -> None:
|
||||
self.url = SimpleNamespace(host=final_host)
|
||||
self.headers = headers or {}
|
||||
self._chunks = chunks
|
||||
self.raise_for_status = MagicMock()
|
||||
self.iteration_started = False
|
||||
|
||||
async def aiter_bytes(self):
|
||||
self.iteration_started = True
|
||||
for chunk in self._chunks:
|
||||
yield chunk
|
||||
|
||||
|
||||
class _AsyncContext:
|
||||
def __init__(self, value) -> None:
|
||||
self.value = value
|
||||
|
||||
async def __aenter__(self):
|
||||
return self.value
|
||||
|
||||
async def __aexit__(self, exc_type, exc, traceback):
|
||||
return False
|
||||
|
||||
|
||||
class _FakeAsyncClient:
|
||||
def __init__(self, response: _FakeResponse) -> None:
|
||||
self.response = response
|
||||
self.stream_calls: list[tuple[tuple, dict]] = []
|
||||
|
||||
async def __aenter__(self):
|
||||
return self
|
||||
|
||||
async def __aexit__(self, exc_type, exc, traceback):
|
||||
return False
|
||||
|
||||
def stream(self, *args, **kwargs):
|
||||
self.stream_calls.append((args, kwargs))
|
||||
return _AsyncContext(self.response)
|
||||
|
||||
|
||||
class MediaProxyHostGuardTests(unittest.TestCase):
|
||||
def test_allows_each_root_domain_and_its_subdomains(self):
|
||||
for allowed in main._MEDIA_PROXY_HOSTS:
|
||||
with self.subTest(host=allowed):
|
||||
self.assertTrue(main._is_allowed_media_host(allowed))
|
||||
self.assertTrue(main._is_allowed_media_host(f"cdn.images.{allowed}"))
|
||||
|
||||
self.assertTrue(main._is_allowed_media_host("CDN.DOUYINPIC.COM."))
|
||||
|
||||
def test_rejects_empty_suffix_tricks_and_similar_domains(self):
|
||||
rejected = (
|
||||
"",
|
||||
".",
|
||||
"douyin.com.evil",
|
||||
"evildouyin.com",
|
||||
"byteimg.com.evil.example",
|
||||
"evilbyteimg.com",
|
||||
"ibyteimg.comevil",
|
||||
"douyin.co",
|
||||
"douyinpic.co",
|
||||
"douyin-static.com",
|
||||
"snssdk.example",
|
||||
)
|
||||
|
||||
for hostname in rejected:
|
||||
with self.subTest(host=hostname):
|
||||
self.assertFalse(main._is_allowed_media_host(hostname))
|
||||
|
||||
|
||||
class MediaProxyResponseGuardTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def _call_proxy(self, response: _FakeResponse):
|
||||
client = _FakeAsyncClient(response)
|
||||
constructor = MagicMock(return_value=client)
|
||||
with patch.object(main.httpx, "AsyncClient", constructor):
|
||||
result = await main.proxy_media(
|
||||
url="https://cdn.douyinpic.com/media/test.jpg"
|
||||
)
|
||||
return result, client, constructor
|
||||
|
||||
async def test_rejects_redirect_to_non_allowlisted_final_domain(self):
|
||||
response = _FakeResponse(
|
||||
final_host="attacker.example",
|
||||
headers={"content-type": "image/jpeg"},
|
||||
chunks=(b"not-read",),
|
||||
)
|
||||
|
||||
with self.assertRaises(main.HTTPException) as caught:
|
||||
await self._call_proxy(response)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 400)
|
||||
self.assertFalse(response.iteration_started)
|
||||
|
||||
async def test_rejects_oversized_content_length_before_streaming(self):
|
||||
response = _FakeResponse(
|
||||
final_host="cdn.douyinpic.com",
|
||||
headers={"content-length": "9", "content-type": "image/jpeg"},
|
||||
chunks=(b"not-read",),
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(main, "_MEDIA_PROXY_MAX_BYTES", 8),
|
||||
self.assertRaises(main.HTTPException) as caught,
|
||||
):
|
||||
await self._call_proxy(response)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 413)
|
||||
self.assertFalse(response.iteration_started)
|
||||
|
||||
async def test_rejects_stream_when_actual_bytes_exceed_limit(self):
|
||||
response = _FakeResponse(
|
||||
final_host="cdn.douyinpic.com",
|
||||
headers={"content-length": "0", "content-type": "image/jpeg"},
|
||||
chunks=(b"12345", b"6789"),
|
||||
)
|
||||
|
||||
with (
|
||||
patch.object(main, "_MEDIA_PROXY_MAX_BYTES", 8),
|
||||
self.assertRaises(main.HTTPException) as caught,
|
||||
):
|
||||
await self._call_proxy(response)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 413)
|
||||
self.assertTrue(response.iteration_started)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,458 @@
|
||||
import asyncio
|
||||
import sys
|
||||
import unittest
|
||||
from datetime import datetime
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
BACKEND_DIR = Path(__file__).resolve().parents[1]
|
||||
if str(BACKEND_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(BACKEND_DIR))
|
||||
|
||||
from rpa_engine.douyin_im.reply_queue import AccountReplyQueue
|
||||
|
||||
|
||||
class AccountReplyQueueTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def _start_queue(self, account_id: int, **kwargs) -> AccountReplyQueue:
|
||||
queue = AccountReplyQueue(account_id=account_id, **kwargs)
|
||||
await queue.start()
|
||||
self.addAsyncCleanup(queue.stop)
|
||||
return queue
|
||||
|
||||
async def _wait_until_idle(
|
||||
self,
|
||||
queue: AccountReplyQueue,
|
||||
timeout: float = 0.5,
|
||||
) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
deadline = loop.time() + timeout
|
||||
while queue.pending_count and loop.time() < deadline:
|
||||
await asyncio.sleep(0.002)
|
||||
self.assertEqual(queue.pending_count, 0)
|
||||
|
||||
async def test_one_account_runs_three_jobs_at_successive_fifo_slots(self):
|
||||
queue = await self._start_queue(account_id=101)
|
||||
interval = 0.05
|
||||
loop = asyncio.get_running_loop()
|
||||
started_at = loop.time()
|
||||
calls: list[tuple[int, float]] = []
|
||||
finished = asyncio.Event()
|
||||
|
||||
def callback_for(index: int):
|
||||
async def callback() -> None:
|
||||
calls.append((index, loop.time() - started_at))
|
||||
if len(calls) == 3:
|
||||
finished.set()
|
||||
|
||||
return callback
|
||||
|
||||
for index in range(3):
|
||||
await queue.enqueue(interval, callback_for(index), description=str(index))
|
||||
|
||||
await asyncio.wait_for(finished.wait(), timeout=0.75)
|
||||
await self._wait_until_idle(queue)
|
||||
|
||||
self.assertEqual([index for index, _ in calls], [0, 1, 2])
|
||||
elapsed = [timestamp for _, timestamp in calls]
|
||||
for timestamp, expected in zip(elapsed, (interval, interval * 2, interval * 3)):
|
||||
self.assertGreaterEqual(timestamp, expected - 0.015)
|
||||
self.assertLess(timestamp, expected + 0.15)
|
||||
self.assertGreaterEqual(elapsed[1] - elapsed[0], interval - 0.02)
|
||||
self.assertGreaterEqual(elapsed[2] - elapsed[1], interval - 0.02)
|
||||
|
||||
async def test_separate_account_queues_reach_first_slot_without_blocking(self):
|
||||
first_queue = await self._start_queue(account_id=201)
|
||||
second_queue = await self._start_queue(account_id=202)
|
||||
interval = 0.04
|
||||
loop = asyncio.get_running_loop()
|
||||
started_at = loop.time()
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
second_finished = asyncio.Event()
|
||||
timestamps: dict[str, float] = {}
|
||||
|
||||
async def first_callback() -> None:
|
||||
timestamps["first"] = loop.time() - started_at
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
|
||||
async def second_callback() -> None:
|
||||
timestamps["second"] = loop.time() - started_at
|
||||
second_finished.set()
|
||||
|
||||
await asyncio.gather(
|
||||
first_queue.enqueue(interval, first_callback, description="first account"),
|
||||
second_queue.enqueue(interval, second_callback, description="second account"),
|
||||
)
|
||||
|
||||
await asyncio.wait_for(first_started.wait(), timeout=0.4)
|
||||
await asyncio.wait_for(second_finished.wait(), timeout=0.4)
|
||||
self.assertFalse(release_first.is_set())
|
||||
self.assertLess(abs(timestamps["first"] - timestamps["second"]), 0.05)
|
||||
|
||||
release_first.set()
|
||||
await self._wait_until_idle(first_queue)
|
||||
await self._wait_until_idle(second_queue)
|
||||
|
||||
async def test_stop_before_due_discards_pending_job(self):
|
||||
queue = await self._start_queue(account_id=301)
|
||||
callback_called = asyncio.Event()
|
||||
|
||||
async def callback() -> None:
|
||||
callback_called.set()
|
||||
|
||||
await queue.enqueue(0.12, callback, description="must be discarded")
|
||||
await asyncio.sleep(0.02)
|
||||
self.assertEqual(queue.pending_count, 1)
|
||||
|
||||
await queue.stop()
|
||||
|
||||
self.assertEqual(queue.pending_count, 0)
|
||||
await asyncio.sleep(0.13)
|
||||
self.assertFalse(callback_called.is_set())
|
||||
|
||||
async def test_callback_failure_does_not_block_following_job(self):
|
||||
errors: list[tuple[str, str]] = []
|
||||
|
||||
def on_error(description: str, exc: BaseException) -> None:
|
||||
errors.append((description, type(exc).__name__))
|
||||
|
||||
queue = await self._start_queue(account_id=401, on_error=on_error)
|
||||
calls: list[str] = []
|
||||
second_finished = asyncio.Event()
|
||||
|
||||
async def failing_callback() -> None:
|
||||
calls.append("first")
|
||||
raise ValueError("expected test failure")
|
||||
|
||||
async def following_callback() -> None:
|
||||
calls.append("second")
|
||||
second_finished.set()
|
||||
|
||||
with self.assertLogs("douyin_im.reply_queue", level="ERROR"):
|
||||
await queue.enqueue(0.03, failing_callback, description="first")
|
||||
await queue.enqueue(0.03, following_callback, description="second")
|
||||
await asyncio.wait_for(second_finished.wait(), timeout=0.5)
|
||||
|
||||
await self._wait_until_idle(queue)
|
||||
self.assertEqual(calls, ["first", "second"])
|
||||
self.assertEqual(errors, [("first", "ValueError")])
|
||||
|
||||
async def test_snapshot_exposes_waiting_job_details(self):
|
||||
queue = await self._start_queue(account_id=501)
|
||||
|
||||
async def callback() -> None:
|
||||
pass
|
||||
|
||||
await queue.enqueue(
|
||||
0.2,
|
||||
callback,
|
||||
description="回复 测试用户",
|
||||
details={
|
||||
"sender_name": "测试用户",
|
||||
"conversation_id": "conv-501",
|
||||
"incoming_content": "你好",
|
||||
"replies": ["您好"],
|
||||
},
|
||||
)
|
||||
items = await queue.snapshot()
|
||||
|
||||
self.assertEqual(len(items), 1)
|
||||
self.assertEqual(items[0]["position"], 1)
|
||||
self.assertEqual(items[0]["status"], "waiting")
|
||||
self.assertEqual(items[0]["sender_name"], "测试用户")
|
||||
self.assertEqual(items[0]["incoming_content"], "你好")
|
||||
self.assertEqual(items[0]["replies"], ["您好"])
|
||||
self.assertNotIn("callback", items[0])
|
||||
|
||||
async def test_send_now_wakes_first_job_and_moves_later_slots_forward(self):
|
||||
queue = await self._start_queue(account_id=502)
|
||||
interval = 0.2
|
||||
sent = asyncio.Event()
|
||||
|
||||
async def first_callback() -> None:
|
||||
sent.set()
|
||||
|
||||
async def noop() -> None:
|
||||
pass
|
||||
|
||||
await queue.enqueue(interval, first_callback, description="first")
|
||||
await queue.enqueue(interval, noop, description="second")
|
||||
await queue.enqueue(interval, noop, description="third")
|
||||
before = await queue.snapshot()
|
||||
|
||||
result = await queue.send_now(before[0]["job_id"])
|
||||
self.assertEqual(result["status"], "accepted")
|
||||
self.assertEqual(result["shifted_count"], 2)
|
||||
await asyncio.wait_for(sent.wait(), timeout=0.1)
|
||||
await asyncio.sleep(0)
|
||||
|
||||
after = await queue.snapshot()
|
||||
self.assertEqual([item["description"] for item in after], ["second", "third"])
|
||||
for old, new in zip(before[1:], after):
|
||||
old_due = datetime.fromisoformat(old["scheduled_at"]).timestamp()
|
||||
new_due = datetime.fromisoformat(new["scheduled_at"]).timestamp()
|
||||
self.assertAlmostEqual(old_due - new_due, interval, delta=0.05)
|
||||
|
||||
async def test_send_now_middle_runs_first_and_only_shifts_jobs_behind_it(self):
|
||||
queue = await self._start_queue(account_id=503)
|
||||
interval = 0.12
|
||||
calls: list[str] = []
|
||||
all_sent = asyncio.Event()
|
||||
|
||||
def callback_for(name: str):
|
||||
async def callback() -> None:
|
||||
calls.append(name)
|
||||
if len(calls) == 3:
|
||||
all_sent.set()
|
||||
|
||||
return callback
|
||||
|
||||
for name in ("first", "middle", "last"):
|
||||
await queue.enqueue(interval, callback_for(name), description=name)
|
||||
before = await queue.snapshot()
|
||||
result = await queue.send_now(before[1]["job_id"])
|
||||
|
||||
self.assertEqual(result["shifted_count"], 1)
|
||||
await asyncio.wait_for(all_sent.wait(), timeout=0.6)
|
||||
self.assertEqual(calls, ["middle", "first", "last"])
|
||||
|
||||
first_before = datetime.fromisoformat(before[0]["scheduled_at"]).timestamp()
|
||||
# The first normal job keeps its original slot; this is also validated by execution order.
|
||||
self.assertGreater(first_before, datetime.now().timestamp() - 1)
|
||||
|
||||
async def test_send_now_does_not_overlap_active_send(self):
|
||||
queue = await self._start_queue(account_id=504)
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
second_started = asyncio.Event()
|
||||
|
||||
async def first_callback() -> None:
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
|
||||
async def second_callback() -> None:
|
||||
second_started.set()
|
||||
|
||||
await queue.enqueue(0.01, first_callback, description="first")
|
||||
await queue.enqueue(0.2, second_callback, description="second")
|
||||
await asyncio.wait_for(first_started.wait(), timeout=0.2)
|
||||
items = await queue.snapshot()
|
||||
second = next(item for item in items if item["description"] == "second")
|
||||
|
||||
result = await queue.send_now(second["job_id"])
|
||||
self.assertEqual(result["status"], "accepted")
|
||||
await asyncio.sleep(0.03)
|
||||
self.assertFalse(second_started.is_set())
|
||||
|
||||
release_first.set()
|
||||
await asyncio.wait_for(second_started.wait(), timeout=0.2)
|
||||
|
||||
async def test_send_now_tail_releases_slot_for_next_enqueue(self):
|
||||
queue = await self._start_queue(account_id=505)
|
||||
interval = 0.2
|
||||
|
||||
async def noop() -> None:
|
||||
pass
|
||||
|
||||
for name in ("first", "second", "tail"):
|
||||
await queue.enqueue(interval, noop, description=name)
|
||||
before = await queue.snapshot()
|
||||
old_tail_due = datetime.fromisoformat(before[2]["scheduled_at"]).timestamp()
|
||||
|
||||
await queue.send_now(before[2]["job_id"])
|
||||
await queue.enqueue(interval, noop, description="new-tail")
|
||||
after = await queue.snapshot()
|
||||
new_tail = next(item for item in after if item["description"] == "new-tail")
|
||||
new_tail_due = datetime.fromisoformat(new_tail["scheduled_at"]).timestamp()
|
||||
|
||||
self.assertAlmostEqual(new_tail_due, old_tail_due, delta=0.05)
|
||||
|
||||
async def test_repeated_send_now_is_idempotent(self):
|
||||
queue = await self._start_queue(account_id=506)
|
||||
calls = 0
|
||||
finished = asyncio.Event()
|
||||
|
||||
async def callback() -> None:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
finished.set()
|
||||
|
||||
await queue.enqueue(0.2, callback, description="once")
|
||||
item = (await queue.snapshot())[0]
|
||||
first = await queue.send_now(item["job_id"])
|
||||
second = await queue.send_now(item["job_id"])
|
||||
|
||||
self.assertEqual(first["status"], "accepted")
|
||||
self.assertIn(second["status"], {"already_requested", "already_sending"})
|
||||
await asyncio.wait_for(finished.wait(), timeout=0.1)
|
||||
await asyncio.sleep(0.02)
|
||||
self.assertEqual(calls, 1)
|
||||
|
||||
async def test_merge_pending_keeps_one_job_slot_and_callback(self):
|
||||
queue = await self._start_queue(account_id=507)
|
||||
calls = 0
|
||||
finished = asyncio.Event()
|
||||
|
||||
async def callback() -> None:
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
finished.set()
|
||||
|
||||
await queue.enqueue(
|
||||
0.3,
|
||||
callback,
|
||||
description="conversation",
|
||||
details={
|
||||
"incoming_content": "first",
|
||||
"incoming_contents": ["first"],
|
||||
"message_count": 1,
|
||||
"replies": ["one reply"],
|
||||
},
|
||||
merge_key="conv:507",
|
||||
)
|
||||
before = (await queue.snapshot())[0]
|
||||
|
||||
def append_second(details):
|
||||
details["incoming_content"] = "second"
|
||||
details["incoming_contents"].append("second")
|
||||
details["message_count"] = 2
|
||||
return details
|
||||
|
||||
result = await queue.merge_pending("conv:507", append_second)
|
||||
after = await queue.snapshot()
|
||||
|
||||
self.assertEqual(result["status"], "merged")
|
||||
self.assertEqual(result["position"], 1)
|
||||
self.assertEqual(len(after), 1)
|
||||
self.assertEqual(after[0]["job_id"], before["job_id"])
|
||||
self.assertEqual(after[0]["incoming_contents"], ["first", "second"])
|
||||
self.assertEqual(after[0]["replies"], ["one reply"])
|
||||
old_due = datetime.fromisoformat(before["scheduled_at"]).timestamp()
|
||||
new_due = datetime.fromisoformat(after[0]["scheduled_at"]).timestamp()
|
||||
self.assertAlmostEqual(old_due, new_due, delta=0.02)
|
||||
|
||||
# Management snapshots must not expose the queue's mutable nested list.
|
||||
after[0]["incoming_contents"].append("external mutation")
|
||||
self.assertEqual(
|
||||
(await queue.snapshot())[0]["incoming_contents"],
|
||||
["first", "second"],
|
||||
)
|
||||
|
||||
await queue.send_now(before["job_id"])
|
||||
await asyncio.wait_for(finished.wait(), timeout=0.2)
|
||||
self.assertEqual(calls, 1)
|
||||
|
||||
async def test_urgent_job_can_still_merge_before_sending(self):
|
||||
queue = await self._start_queue(account_id=508)
|
||||
blocker_started = asyncio.Event()
|
||||
release_blocker = asyncio.Event()
|
||||
merged_sent = asyncio.Event()
|
||||
merged_calls = 0
|
||||
|
||||
async def blocker() -> None:
|
||||
blocker_started.set()
|
||||
await release_blocker.wait()
|
||||
|
||||
async def merged_callback() -> None:
|
||||
nonlocal merged_calls
|
||||
merged_calls += 1
|
||||
merged_sent.set()
|
||||
|
||||
await queue.enqueue(0.01, blocker, description="blocker")
|
||||
await queue.enqueue(
|
||||
0.3,
|
||||
merged_callback,
|
||||
description="mergeable",
|
||||
details={"incoming_content": "first", "incoming_contents": ["first"]},
|
||||
merge_key="conv:508",
|
||||
)
|
||||
await asyncio.wait_for(blocker_started.wait(), timeout=0.2)
|
||||
mergeable = next(
|
||||
item for item in await queue.snapshot() if item["description"] == "mergeable"
|
||||
)
|
||||
await queue.send_now(mergeable["job_id"])
|
||||
|
||||
def append_second(details):
|
||||
details["incoming_contents"].append("second")
|
||||
details["message_count"] = 2
|
||||
return details
|
||||
|
||||
result = await queue.merge_pending("conv:508", append_second)
|
||||
self.assertEqual(result["status"], "merged")
|
||||
self.assertEqual(result["queue_status"], "ready")
|
||||
release_blocker.set()
|
||||
await asyncio.wait_for(merged_sent.wait(), timeout=0.2)
|
||||
self.assertEqual(merged_calls, 1)
|
||||
|
||||
async def test_active_job_is_sealed_against_merge(self):
|
||||
queue = await self._start_queue(account_id=509)
|
||||
started = asyncio.Event()
|
||||
release = asyncio.Event()
|
||||
|
||||
async def callback() -> None:
|
||||
started.set()
|
||||
await release.wait()
|
||||
|
||||
await queue.enqueue(
|
||||
0.01,
|
||||
callback,
|
||||
details={"incoming_content": "first", "incoming_contents": ["first"]},
|
||||
merge_key="conv:509",
|
||||
)
|
||||
await asyncio.wait_for(started.wait(), timeout=0.2)
|
||||
result = await queue.merge_pending("conv:509", lambda details: details)
|
||||
self.assertEqual(result["status"], "not_found")
|
||||
release.set()
|
||||
|
||||
async def test_merge_keys_bridge_partial_peer_and_conversation_ids(self):
|
||||
queue = await self._start_queue(account_id=510)
|
||||
|
||||
async def callback() -> None:
|
||||
pass
|
||||
|
||||
await queue.enqueue(
|
||||
0.3,
|
||||
callback,
|
||||
details={"incoming_contents": ["first"], "message_count": 1},
|
||||
merge_keys=("peer:123",),
|
||||
)
|
||||
|
||||
def append_second(details):
|
||||
details["incoming_contents"].append("second")
|
||||
details["message_count"] = 2
|
||||
return details
|
||||
|
||||
result = await queue.merge_pending(
|
||||
("conv:510", "peer:123"),
|
||||
append_second,
|
||||
)
|
||||
self.assertEqual(result["status"], "merged")
|
||||
self.assertEqual(queue.pending_count, 1)
|
||||
self.assertEqual(
|
||||
(await queue.snapshot())[0]["incoming_contents"],
|
||||
["first", "second"],
|
||||
)
|
||||
|
||||
async def test_different_conversation_ids_do_not_merge_on_shared_peer(self):
|
||||
queue = await self._start_queue(account_id=511)
|
||||
|
||||
async def callback() -> None:
|
||||
pass
|
||||
|
||||
await queue.enqueue(
|
||||
0.3,
|
||||
callback,
|
||||
details={"incoming_contents": ["first"]},
|
||||
merge_keys=("conv:first", "peer:123"),
|
||||
)
|
||||
result = await queue.merge_pending(
|
||||
("conv:second", "peer:123"),
|
||||
lambda details: details,
|
||||
)
|
||||
self.assertEqual(result["status"], "not_found")
|
||||
self.assertEqual(queue.pending_count, 1)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,165 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, 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
|
||||
|
||||
|
||||
def _queue_item(account_id: int, job_id: str = "job-1") -> dict:
|
||||
return {
|
||||
"job_id": job_id,
|
||||
"account_id": account_id,
|
||||
"position": 1,
|
||||
"status": "waiting",
|
||||
"expedited": False,
|
||||
"description": "回复 测试用户",
|
||||
"sender_name": "测试用户",
|
||||
"sender_id": "peer-1",
|
||||
"sender_avatar": None,
|
||||
"conversation_id": "conv-1",
|
||||
"incoming_content": "你好",
|
||||
"incoming_contents": ["你好", "第二条"],
|
||||
"message_count": 2,
|
||||
"replies": ["您好"],
|
||||
"interval_seconds": 60,
|
||||
"enqueued_at": "2026-07-20T10:00:00+00:00",
|
||||
"scheduled_at": "2026-07-20T10:01:00+00:00",
|
||||
"remaining_seconds": 60,
|
||||
}
|
||||
|
||||
|
||||
class _FakeService:
|
||||
def __init__(self, account_id: int, items=None, action=None):
|
||||
self.account_id = account_id
|
||||
self._running = True
|
||||
self.items = list(items or [])
|
||||
self.action = action or {
|
||||
"status": "accepted",
|
||||
"job_id": "job-1",
|
||||
"shifted_count": 2,
|
||||
}
|
||||
|
||||
async def get_reply_queue_snapshot(self):
|
||||
return list(self.items)
|
||||
|
||||
async def send_queued_reply_now(self, job_id: str):
|
||||
return {**self.action, "job_id": job_id}
|
||||
|
||||
|
||||
class ReplyQueueApiTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncSetUp(self):
|
||||
self.original_workers = main.manager.workers
|
||||
main.manager.workers = {}
|
||||
|
||||
async def asyncTearDown(self):
|
||||
main.manager.workers = self.original_workers
|
||||
|
||||
async def test_summary_filters_accounts_by_owner_scope(self):
|
||||
service_one = _FakeService(1, [_queue_item(1)])
|
||||
service_two = _FakeService(2, [_queue_item(2, "job-2")])
|
||||
main.manager.workers = {
|
||||
1: SimpleNamespace(is_running=True, _im_service=service_one),
|
||||
2: SimpleNamespace(is_running=True, _im_service=service_two),
|
||||
}
|
||||
|
||||
with patch.object(main, "owned_account_ids", AsyncMock(return_value={1})):
|
||||
response = await main.get_reply_queue_summaries(
|
||||
db=object(),
|
||||
user=SimpleNamespace(id=10, role="operator"),
|
||||
)
|
||||
|
||||
self.assertEqual(response.total_pending, 1)
|
||||
self.assertEqual([item.account_id for item in response.items], [1])
|
||||
|
||||
async def test_offline_account_detail_returns_empty_snapshot(self):
|
||||
account = SimpleNamespace(id=1, reply_delay_seconds=0)
|
||||
with (
|
||||
patch.object(main, "get_owned_account", AsyncMock(return_value=account)),
|
||||
patch.object(main, "_global_reply_delay_seconds", return_value=0),
|
||||
):
|
||||
response = await main.get_account_reply_queue(
|
||||
account_id=1,
|
||||
db=object(),
|
||||
user=SimpleNamespace(id=10, role="operator"),
|
||||
)
|
||||
|
||||
self.assertFalse(response.running)
|
||||
self.assertEqual(response.pending_count, 0)
|
||||
self.assertEqual(response.items, [])
|
||||
|
||||
async def test_detail_preserves_merged_incoming_messages(self):
|
||||
account = SimpleNamespace(id=1, reply_delay_seconds=60)
|
||||
service = _FakeService(1, [_queue_item(1)])
|
||||
main.manager.workers = {
|
||||
1: SimpleNamespace(is_running=True, _im_service=service),
|
||||
}
|
||||
with (
|
||||
patch.object(main, "get_owned_account", AsyncMock(return_value=account)),
|
||||
patch.object(main, "_resolve_effective_reply_delay", return_value=60),
|
||||
):
|
||||
response = await main.get_account_reply_queue(
|
||||
account_id=1,
|
||||
db=object(),
|
||||
user=SimpleNamespace(id=10, role="operator"),
|
||||
)
|
||||
|
||||
self.assertEqual(response.pending_count, 1)
|
||||
self.assertEqual(response.items[0].incoming_contents, ["你好", "第二条"])
|
||||
self.assertEqual(response.items[0].message_count, 2)
|
||||
self.assertEqual(response.items[0].replies, ["您好"])
|
||||
|
||||
async def test_send_now_returns_shifted_count(self):
|
||||
service = _FakeService(1)
|
||||
main.manager.workers = {
|
||||
1: SimpleNamespace(is_running=True, _im_service=service),
|
||||
}
|
||||
with (
|
||||
patch.object(main, "get_owned_account", AsyncMock(return_value=SimpleNamespace(id=1))),
|
||||
patch.object(main.system_logger, "record"),
|
||||
):
|
||||
response = await main.send_account_queued_reply_now(
|
||||
account_id=1,
|
||||
job_id="job-1",
|
||||
db=object(),
|
||||
user=SimpleNamespace(id=10, role="operator"),
|
||||
)
|
||||
|
||||
self.assertEqual(response.status, "accepted")
|
||||
self.assertEqual(response.shifted_count, 2)
|
||||
|
||||
async def test_cross_account_or_finished_job_returns_not_found(self):
|
||||
service = _FakeService(1, action={"status": "not_found"})
|
||||
main.manager.workers = {
|
||||
1: SimpleNamespace(is_running=True, _im_service=service),
|
||||
}
|
||||
with patch.object(
|
||||
main,
|
||||
"get_owned_account",
|
||||
AsyncMock(return_value=SimpleNamespace(id=1)),
|
||||
):
|
||||
with self.assertRaises(main.HTTPException) as caught:
|
||||
await main.send_account_queued_reply_now(
|
||||
account_id=1,
|
||||
job_id="other-account-job",
|
||||
db=object(),
|
||||
user=SimpleNamespace(id=10, role="operator"),
|
||||
)
|
||||
|
||||
self.assertEqual(caught.exception.status_code, 404)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,421 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from types import SimpleNamespace
|
||||
from unittest.mock import AsyncMock, Mock
|
||||
|
||||
|
||||
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))
|
||||
|
||||
from auth.system_settings import SystemSettingsData, set_cached_settings
|
||||
from rpa_engine.douyin_im.service import DouyinImService
|
||||
from rpa_engine.playwright_worker import DouyinWorker
|
||||
|
||||
|
||||
class _RecordingQueue:
|
||||
def __init__(self) -> None:
|
||||
self.jobs: list[tuple[float, object, str, dict, frozenset[str]]] = []
|
||||
|
||||
async def enqueue(
|
||||
self,
|
||||
delay_seconds,
|
||||
callback,
|
||||
description="",
|
||||
details=None,
|
||||
merge_key="",
|
||||
merge_keys=None,
|
||||
) -> int:
|
||||
keys = merge_keys if merge_keys is not None else [merge_key]
|
||||
normalized_keys = frozenset(str(key) for key in keys if str(key or "").strip())
|
||||
self.jobs.append(
|
||||
(delay_seconds, callback, description, dict(details or {}), normalized_keys)
|
||||
)
|
||||
return len(self.jobs)
|
||||
|
||||
async def merge_pending(self, merge_key, details_merger):
|
||||
incoming_keys = (
|
||||
frozenset([merge_key])
|
||||
if isinstance(merge_key, str)
|
||||
else frozenset(merge_key)
|
||||
)
|
||||
for index, job in enumerate(self.jobs):
|
||||
existing_conversations = {
|
||||
key for key in job[4] if key.startswith("conv:")
|
||||
}
|
||||
incoming_conversations = {
|
||||
key for key in incoming_keys if key.startswith("conv:")
|
||||
}
|
||||
same_conversation = bool(
|
||||
existing_conversations & incoming_conversations
|
||||
)
|
||||
different_explicit_conversations = bool(
|
||||
existing_conversations
|
||||
and incoming_conversations
|
||||
and not same_conversation
|
||||
)
|
||||
same_peer = bool(
|
||||
{key for key in job[4] if key.startswith("peer:")}
|
||||
& {key for key in incoming_keys if key.startswith("peer:")}
|
||||
)
|
||||
if not same_conversation and (different_explicit_conversations or not same_peer):
|
||||
continue
|
||||
merged_details = details_merger(dict(job[3]))
|
||||
self.jobs[index] = (
|
||||
job[0],
|
||||
job[1],
|
||||
job[2],
|
||||
merged_details,
|
||||
frozenset(job[4] | incoming_keys),
|
||||
)
|
||||
return {
|
||||
"status": "merged",
|
||||
"job_id": f"job-{index + 1}",
|
||||
"position": index + 1,
|
||||
"message_count": merged_details.get("message_count", 1),
|
||||
}
|
||||
return {"status": "not_found"}
|
||||
|
||||
|
||||
def _build_service(delay_seconds: int = 60):
|
||||
match_reply = AsyncMock(return_value=["自动回复"])
|
||||
log_fn = AsyncMock()
|
||||
service = DouyinImService(
|
||||
session=SimpleNamespace(my_uid="999"),
|
||||
match_reply=match_reply,
|
||||
log_fn=log_fn,
|
||||
account_id=1,
|
||||
)
|
||||
service._running = True
|
||||
service._reply_queue = _RecordingQueue()
|
||||
service._resolve_peer_profile = AsyncMock(return_value=("张三", "", "123"))
|
||||
service._resolve_cooldown_seconds = AsyncMock(return_value=0)
|
||||
service._resolve_reply_delay_seconds = AsyncMock(return_value=delay_seconds)
|
||||
return service, match_reply, log_fn
|
||||
|
||||
|
||||
class ReplyQueueIntegrationTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_same_message_from_ws_and_poll_is_queued_once(self):
|
||||
service, match_reply, _ = _build_service()
|
||||
message = {
|
||||
"conversation_id": "conv-1",
|
||||
"sender_uid": "123",
|
||||
"sender_name": "张三",
|
||||
"content": "你好",
|
||||
"server_message_id": "mid-1",
|
||||
}
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await asyncio.gather(
|
||||
service._handle_incoming(dict(message)),
|
||||
service._handle_incoming(dict(message)),
|
||||
)
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 1)
|
||||
self.assertEqual(match_reply.await_count, 1)
|
||||
details = service._reply_queue.jobs[0][3]
|
||||
self.assertEqual(details["sender_name"], "张三")
|
||||
self.assertEqual(details["conversation_id"], "conv-1")
|
||||
self.assertEqual(details["incoming_content"], "你好")
|
||||
self.assertEqual(details["replies"], ["自动回复"])
|
||||
|
||||
async def test_same_conversation_with_different_message_ids_merges(self):
|
||||
service, match_reply, _ = _build_service()
|
||||
base = {
|
||||
"conversation_id": "conv-1",
|
||||
"sender_uid": "123",
|
||||
"sender_name": "张三",
|
||||
"content": "你好",
|
||||
}
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._handle_incoming({**base, "server_message_id": "mid-1"})
|
||||
await service._handle_incoming({**base, "server_message_id": "mid-2"})
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 1)
|
||||
self.assertEqual(match_reply.await_count, 1)
|
||||
details = service._reply_queue.jobs[0][3]
|
||||
self.assertEqual(details["incoming_contents"], ["你好", "你好"])
|
||||
self.assertEqual(details["message_count"], 2)
|
||||
self.assertEqual(details["replies"], ["自动回复"])
|
||||
|
||||
async def test_pending_conversation_merges_before_cooldown_filter(self):
|
||||
service, match_reply, _ = _build_service()
|
||||
service._resolve_cooldown_seconds = AsyncMock(return_value=120)
|
||||
base = {
|
||||
"conversation_id": "conv-1",
|
||||
"sender_uid": "123",
|
||||
"sender_name": "张三",
|
||||
}
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._handle_incoming(
|
||||
{**base, "content": "first", "server_message_id": "mid-1"}
|
||||
)
|
||||
await service._handle_incoming(
|
||||
{
|
||||
**base,
|
||||
"conversation_id": "conv-1",
|
||||
"content": "second",
|
||||
"server_message_id": "mid-2",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 1)
|
||||
self.assertEqual(match_reply.await_count, 1)
|
||||
details = service._reply_queue.jobs[0][3]
|
||||
self.assertEqual(details["incoming_contents"], ["first", "second"])
|
||||
self.assertEqual(details["message_count"], 2)
|
||||
|
||||
async def test_http_previews_from_different_conversations_do_not_collide(self):
|
||||
service, _, _ = _build_service()
|
||||
base = {
|
||||
"sender_uid": "",
|
||||
"sender_name": "同名用户",
|
||||
"content": "相同内容",
|
||||
}
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._handle_incoming({**base, "conversation_id": "conv-a"})
|
||||
await service._handle_incoming({**base, "conversation_id": "conv-b"})
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 2)
|
||||
|
||||
async def test_missing_conversation_id_uses_peer_uid_for_merge(self):
|
||||
service, match_reply, _ = _build_service()
|
||||
base = {
|
||||
"conversation_id": "",
|
||||
"sender_uid": "123",
|
||||
"sender_name": "张三",
|
||||
}
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._handle_incoming(
|
||||
{**base, "content": "first", "server_message_id": "mid-1"}
|
||||
)
|
||||
await service._handle_incoming(
|
||||
{
|
||||
**base,
|
||||
"conversation_id": "conv-1",
|
||||
"content": "second",
|
||||
"server_message_id": "mid-2",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 1)
|
||||
self.assertEqual(match_reply.await_count, 1)
|
||||
self.assertEqual(
|
||||
service._reply_queue.jobs[0][3]["incoming_contents"],
|
||||
["first", "second"],
|
||||
)
|
||||
|
||||
async def test_peer_alias_added_later_merges_with_conversation_only_job(self):
|
||||
service, match_reply, _ = _build_service()
|
||||
|
||||
async def resolve_profile(conv_id, sender_uid, sender, sender_avatar):
|
||||
return "张三", "", str(sender_uid or "")
|
||||
|
||||
service._resolve_peer_profile = AsyncMock(side_effect=resolve_profile)
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._handle_incoming(
|
||||
{
|
||||
"conversation_id": "conv-1",
|
||||
"sender_uid": "",
|
||||
"sender_name": "张三",
|
||||
"content": "first",
|
||||
"server_message_id": "mid-1",
|
||||
}
|
||||
)
|
||||
await service._handle_incoming(
|
||||
{
|
||||
"conversation_id": "conv-1",
|
||||
"sender_uid": "123",
|
||||
"sender_name": "张三",
|
||||
"content": "second",
|
||||
"server_message_id": "mid-2",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 1)
|
||||
self.assertEqual(match_reply.await_count, 1)
|
||||
self.assertEqual(
|
||||
service._reply_queue.jobs[0][3]["incoming_contents"],
|
||||
["first", "second"],
|
||||
)
|
||||
|
||||
async def test_missing_conversation_and_peer_ids_do_not_merge_by_name(self):
|
||||
service, match_reply, _ = _build_service()
|
||||
service._resolve_peer_profile = AsyncMock(return_value=("同名用户", "", ""))
|
||||
base = {
|
||||
"conversation_id": "",
|
||||
"sender_uid": "",
|
||||
"sender_name": "同名用户",
|
||||
}
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._handle_incoming(
|
||||
{**base, "content": "first", "server_message_id": "mid-1"}
|
||||
)
|
||||
await service._handle_incoming(
|
||||
{**base, "content": "second", "server_message_id": "mid-2"}
|
||||
)
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 2)
|
||||
self.assertEqual(match_reply.await_count, 2)
|
||||
|
||||
async def test_slow_first_message_keeps_arrival_order(self):
|
||||
service, _, _ = _build_service()
|
||||
|
||||
async def resolve_profile(conv_id, sender_uid, sender, sender_avatar):
|
||||
if conv_id == "conv-first":
|
||||
await asyncio.sleep(0.03)
|
||||
return "先到", "", "101"
|
||||
return "后到", "", "102"
|
||||
|
||||
service._resolve_peer_profile = AsyncMock(side_effect=resolve_profile)
|
||||
first = {
|
||||
"conversation_id": "conv-first",
|
||||
"sender_uid": "101",
|
||||
"sender_name": "先到",
|
||||
"content": "第一条",
|
||||
"server_message_id": "mid-first",
|
||||
}
|
||||
second = {
|
||||
"conversation_id": "conv-second",
|
||||
"sender_uid": "102",
|
||||
"sender_name": "后到",
|
||||
"content": "第二条",
|
||||
"server_message_id": "mid-second",
|
||||
}
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
first_task = asyncio.create_task(service._handle_incoming(first))
|
||||
await asyncio.sleep(0)
|
||||
second_task = asyncio.create_task(service._handle_incoming(second))
|
||||
await asyncio.gather(first_task, second_task)
|
||||
|
||||
descriptions = [job[2] for job in service._reply_queue.jobs]
|
||||
self.assertEqual(descriptions, ["回复 先到", "回复 后到"])
|
||||
|
||||
async def test_zero_effective_delay_skips_queue_and_sends_immediately(self):
|
||||
service, _, _ = _build_service(delay_seconds=0)
|
||||
service._send_auto_reply = AsyncMock()
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._handle_incoming(
|
||||
{
|
||||
"conversation_id": "conv-1",
|
||||
"sender_uid": "123",
|
||||
"sender_name": "张三",
|
||||
"content": "你好",
|
||||
"server_message_id": "mid-1",
|
||||
}
|
||||
)
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 0)
|
||||
service._send_auto_reply.assert_awaited_once()
|
||||
|
||||
async def test_zero_delay_keeps_separate_immediate_replies(self):
|
||||
service, _, _ = _build_service(delay_seconds=0)
|
||||
service._send_auto_reply = AsyncMock()
|
||||
base = {
|
||||
"conversation_id": "conv-1",
|
||||
"sender_uid": "123",
|
||||
"sender_name": "张三",
|
||||
}
|
||||
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._handle_incoming(
|
||||
{**base, "content": "first", "server_message_id": "mid-1"}
|
||||
)
|
||||
await service._handle_incoming(
|
||||
{**base, "content": "second", "server_message_id": "mid-2"}
|
||||
)
|
||||
|
||||
self.assertEqual(len(service._reply_queue.jobs), 0)
|
||||
self.assertEqual(service._send_auto_reply.await_count, 2)
|
||||
|
||||
async def test_session_stop_during_first_payload_cancels_remaining_payloads(self):
|
||||
service, _, log_fn = _build_service()
|
||||
|
||||
async def fail_and_stop(*args, **kwargs):
|
||||
service.last_error = "INVALID_REQUEST"
|
||||
service._running = False
|
||||
return False, None
|
||||
|
||||
service._send_text = AsyncMock(side_effect=fail_and_stop)
|
||||
with unittest.mock.patch(
|
||||
"rpa_engine.douyin_im.service.system_logger.record", Mock()
|
||||
):
|
||||
await service._send_auto_reply(
|
||||
sender="张三",
|
||||
content="你好",
|
||||
conv_id="conv-1",
|
||||
replies=["第一段", "第二段"],
|
||||
peer_key="123",
|
||||
cooldown=0,
|
||||
log_kwargs={
|
||||
"sender_name": "张三",
|
||||
"sender_id": "123",
|
||||
"sender_avatar": None,
|
||||
"message": "你好",
|
||||
},
|
||||
)
|
||||
|
||||
self.assertEqual(service._send_text.await_count, 1)
|
||||
log_fn.assert_awaited_once()
|
||||
|
||||
|
||||
class ReplyDelayResolutionTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def asyncSetUp(self):
|
||||
self.worker = DouyinWorker(account_id=1)
|
||||
|
||||
async def test_account_value_overrides_system_default(self):
|
||||
set_cached_settings(SystemSettingsData(auto_reply_delay_seconds=60))
|
||||
self.worker.get_reply_delay = AsyncMock(return_value=15)
|
||||
|
||||
self.assertEqual(await self.worker.resolve_reply_delay_seconds(), 15)
|
||||
|
||||
async def test_unset_account_uses_system_default(self):
|
||||
set_cached_settings(SystemSettingsData(auto_reply_delay_seconds=60))
|
||||
self.worker.get_reply_delay = AsyncMock(return_value=None)
|
||||
|
||||
self.assertEqual(await self.worker.resolve_reply_delay_seconds(), 60)
|
||||
|
||||
async def test_both_unset_skip_queue_rule(self):
|
||||
set_cached_settings(SystemSettingsData(auto_reply_delay_seconds=0))
|
||||
self.worker.get_reply_delay = AsyncMock(return_value=None)
|
||||
|
||||
self.assertEqual(await self.worker.resolve_reply_delay_seconds(), 0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,676 @@
|
||||
from __future__ import annotations
|
||||
|
||||
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))
|
||||
|
||||
from rpa_engine.douyin_im.session import DouyinImSession
|
||||
from rpa_engine import account_profile as account_profile_module
|
||||
from rpa_engine.playwright_worker import DouyinWorker
|
||||
|
||||
|
||||
class SecUserIdGuardTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_im_direct_missing_sec_user_id_exits_before_online_state(self):
|
||||
worker = DouyinWorker(account_id=301, login_mode="im_direct")
|
||||
worker._require_sec_user_id = AsyncMock(return_value=False)
|
||||
worker._load_user_agent = AsyncMock(return_value="test-agent")
|
||||
im_session = DouyinImSession(
|
||||
cookies={"sessionid": "test-session"},
|
||||
my_uid=30101,
|
||||
)
|
||||
worker._build_im_session_from_storage = AsyncMock(return_value=im_session)
|
||||
worker._persist_im_session = AsyncMock()
|
||||
worker._run_im_direct_service = AsyncMock()
|
||||
|
||||
with patch(
|
||||
"rpa_engine.playwright_worker.validate_im_session",
|
||||
new=AsyncMock(return_value=(True, "ready")),
|
||||
) as validate:
|
||||
started, reason = await worker._try_cookie_only_im_start(
|
||||
{"cookies": [{"name": "sessionid", "value": "test-session"}]}
|
||||
)
|
||||
|
||||
self.assertFalse(started)
|
||||
self.assertIn("sec_user_id", reason)
|
||||
worker._require_sec_user_id.assert_awaited_once()
|
||||
require_call = worker._require_sec_user_id.await_args
|
||||
self.assertTrue(require_call.kwargs["refresh_if_missing"])
|
||||
self.assertTrue(require_call.kwargs["refresh_if_stale"])
|
||||
worker._build_im_session_from_storage.assert_awaited_once()
|
||||
validate.assert_awaited_once_with(im_session)
|
||||
worker._persist_im_session.assert_not_awaited()
|
||||
worker._run_im_direct_service.assert_not_awaited()
|
||||
|
||||
async def test_running_guard_precedes_disabled_follow_welcome_setting(self):
|
||||
worker = DouyinWorker(account_id=302, login_mode="im_direct")
|
||||
worker.is_running = True
|
||||
worker._im_service = SimpleNamespace(_running=True)
|
||||
worker._require_sec_user_id = AsyncMock(return_value=False)
|
||||
worker.get_db = AsyncMock()
|
||||
|
||||
await worker.follow_welcome_tick()
|
||||
|
||||
worker._require_sec_user_id.assert_awaited_once()
|
||||
require_call = worker._require_sec_user_id.await_args
|
||||
self.assertFalse(require_call.kwargs.get("refresh_if_missing", False))
|
||||
# The identity guard must run before the account's follow-welcome flag
|
||||
# is queried. Otherwise accounts with that feature disabled could stay
|
||||
# hosted indefinitely without a sec_user_id.
|
||||
worker.get_db.assert_not_awaited()
|
||||
|
||||
async def test_blank_sec_user_id_is_missing_and_stops_hosting(self):
|
||||
worker = DouyinWorker(account_id=303, login_mode="im_direct")
|
||||
worker._load_sec_user_id = AsyncMock(return_value=" ")
|
||||
worker._refresh_sec_user_id = AsyncMock(return_value="\t")
|
||||
worker._stop_for_missing_sec_user_id = AsyncMock()
|
||||
|
||||
accepted = await worker._require_sec_user_id(
|
||||
"直连启动前",
|
||||
refresh_if_missing=True,
|
||||
)
|
||||
|
||||
self.assertFalse(accepted)
|
||||
worker._load_sec_user_id.assert_awaited()
|
||||
worker._refresh_sec_user_id.assert_awaited_once_with()
|
||||
worker._stop_for_missing_sec_user_id.assert_awaited_once_with("直连启动前")
|
||||
|
||||
async def test_existing_sec_user_id_does_not_refresh_or_stop(self):
|
||||
worker = DouyinWorker(account_id=304, login_mode="im_direct")
|
||||
worker._load_sec_user_id = AsyncMock(return_value="MS4wLjABAAAA-valid")
|
||||
worker._refresh_sec_user_id = AsyncMock()
|
||||
worker._stop_for_missing_sec_user_id = AsyncMock()
|
||||
|
||||
accepted = await worker._require_sec_user_id(
|
||||
"运行中",
|
||||
refresh_if_missing=True,
|
||||
)
|
||||
|
||||
self.assertTrue(accepted)
|
||||
worker._load_sec_user_id.assert_awaited_once_with()
|
||||
worker._refresh_sec_user_id.assert_not_awaited()
|
||||
worker._stop_for_missing_sec_user_id.assert_not_awaited()
|
||||
|
||||
async def test_identity_read_error_does_not_claim_field_is_missing(self):
|
||||
worker = DouyinWorker(account_id=307, login_mode="im_direct")
|
||||
worker._load_sec_user_id = AsyncMock(
|
||||
side_effect=RuntimeError("database temporarily unavailable")
|
||||
)
|
||||
worker._refresh_sec_user_id = AsyncMock()
|
||||
worker._stop_for_missing_sec_user_id = AsyncMock()
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "database temporarily unavailable"):
|
||||
await worker._require_sec_user_id(
|
||||
"托管运行中",
|
||||
refresh_if_missing=True,
|
||||
)
|
||||
|
||||
worker._refresh_sec_user_id.assert_not_awaited()
|
||||
worker._stop_for_missing_sec_user_id.assert_not_awaited()
|
||||
|
||||
async def test_unknown_profile_refresh_error_does_not_stop_as_missing(self):
|
||||
worker = DouyinWorker(account_id=308, login_mode="im_direct")
|
||||
worker._load_sec_user_id = AsyncMock(return_value="")
|
||||
worker._refresh_sec_user_id = AsyncMock(
|
||||
side_effect=RuntimeError("暂时无法核验 sec_user_id,请稍后重试")
|
||||
)
|
||||
worker._stop_for_missing_sec_user_id = AsyncMock()
|
||||
|
||||
with self.assertRaisesRegex(RuntimeError, "暂时无法核验"):
|
||||
await worker._require_sec_user_id(
|
||||
"启动托管时",
|
||||
refresh_if_missing=True,
|
||||
)
|
||||
|
||||
worker._refresh_sec_user_id.assert_awaited_once_with()
|
||||
worker._stop_for_missing_sec_user_id.assert_not_awaited()
|
||||
|
||||
async def test_stale_cookie_forces_refresh_instead_of_using_cached_identity(self):
|
||||
worker = DouyinWorker(account_id=309, login_mode="im_direct")
|
||||
worker._load_sec_user_id = AsyncMock(return_value="cached-sec-user-id")
|
||||
worker._sec_user_id_is_stale = AsyncMock(return_value=True)
|
||||
worker._refresh_sec_user_id = AsyncMock(return_value="fresh-sec-user-id")
|
||||
worker._stop_for_missing_sec_user_id = AsyncMock()
|
||||
|
||||
resolved = await worker._require_sec_user_id(
|
||||
"启动托管时",
|
||||
refresh_if_missing=True,
|
||||
refresh_if_stale=True,
|
||||
)
|
||||
|
||||
self.assertEqual(resolved, "fresh-sec-user-id")
|
||||
worker._sec_user_id_is_stale.assert_awaited_once_with()
|
||||
worker._refresh_sec_user_id.assert_awaited_once_with()
|
||||
worker._stop_for_missing_sec_user_id.assert_not_awaited()
|
||||
|
||||
async def test_profile_helper_uses_uid_fallback_for_sec_user_id(self):
|
||||
fetch_detail = AsyncMock(
|
||||
return_value={
|
||||
"uid": "123456789",
|
||||
"nickname": "fallback-user",
|
||||
"sec_user_id": "",
|
||||
"fetched": True,
|
||||
"message": "",
|
||||
}
|
||||
)
|
||||
run_in_thread = AsyncMock(return_value=("MS4wLjABAAAA-fallback", True))
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"fetch_douyin_profile_detail",
|
||||
fetch_detail,
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module.asyncio,
|
||||
"to_thread",
|
||||
run_in_thread,
|
||||
),
|
||||
):
|
||||
detail = await account_profile_module.fetch_douyin_profile_detail_with_sec_user_id(
|
||||
"cookie-json",
|
||||
"test-agent",
|
||||
)
|
||||
|
||||
self.assertEqual(detail["sec_user_id"], "MS4wLjABAAAA-fallback")
|
||||
self.assertEqual(detail["sec_user_id_status"], "found")
|
||||
fetch_detail.assert_awaited_once_with("cookie-json", "test-agent")
|
||||
run_in_thread.assert_awaited_once_with(
|
||||
account_profile_module._resolve_sec_user_id_by_uid_sync,
|
||||
"cookie-json",
|
||||
"test-agent",
|
||||
"123456789",
|
||||
)
|
||||
|
||||
async def test_profile_transport_failure_remains_unknown_not_missing(self):
|
||||
fetch_detail = AsyncMock(
|
||||
return_value={
|
||||
"uid": "",
|
||||
"sec_user_id": "",
|
||||
"fetched": False,
|
||||
"message": "profile request timed out",
|
||||
}
|
||||
)
|
||||
run_in_thread = AsyncMock()
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"fetch_douyin_profile_detail",
|
||||
fetch_detail,
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module.asyncio,
|
||||
"to_thread",
|
||||
run_in_thread,
|
||||
),
|
||||
):
|
||||
detail = await account_profile_module.fetch_douyin_profile_detail_with_sec_user_id(
|
||||
"cookie-json",
|
||||
"test-agent",
|
||||
)
|
||||
|
||||
self.assertEqual(detail["sec_user_id_status"], "unknown")
|
||||
run_in_thread.assert_not_awaited()
|
||||
|
||||
async def test_confirmed_empty_uid_fallback_is_missing(self):
|
||||
fetch_detail = AsyncMock(
|
||||
return_value={
|
||||
"uid": "987654321",
|
||||
"sec_user_id": "",
|
||||
"fetched": True,
|
||||
"message": "",
|
||||
}
|
||||
)
|
||||
run_in_thread = AsyncMock(return_value=("", True))
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"fetch_douyin_profile_detail",
|
||||
fetch_detail,
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module.asyncio,
|
||||
"to_thread",
|
||||
run_in_thread,
|
||||
),
|
||||
):
|
||||
detail = await account_profile_module.fetch_douyin_profile_detail_with_sec_user_id(
|
||||
"cookie-json",
|
||||
"test-agent",
|
||||
)
|
||||
|
||||
self.assertEqual(detail["sec_user_id"], "")
|
||||
self.assertEqual(detail["sec_user_id_status"], "missing")
|
||||
run_in_thread.assert_awaited_once()
|
||||
|
||||
async def test_stop_for_missing_sec_user_id_is_offline_and_idempotent(self):
|
||||
worker = DouyinWorker(account_id=310, login_mode="im_direct")
|
||||
worker.is_running = True
|
||||
worker.stopping = False
|
||||
service = SimpleNamespace(_running=True)
|
||||
worker._im_service = service
|
||||
worker.update_account_status = AsyncMock()
|
||||
|
||||
with patch(
|
||||
"rpa_engine.playwright_worker.system_logger.record"
|
||||
) as record_system_log:
|
||||
await worker._stop_for_missing_sec_user_id("托管运行中")
|
||||
await worker._stop_for_missing_sec_user_id("重复检查")
|
||||
|
||||
self.assertTrue(worker.stopping)
|
||||
self.assertFalse(worker.is_running)
|
||||
self.assertFalse(service._running)
|
||||
worker.update_account_status.assert_awaited_once()
|
||||
status_call = worker.update_account_status.await_args
|
||||
self.assertEqual(status_call.args[0], "offline")
|
||||
self.assertIn("sec_user_id", status_call.kwargs["error_msg"])
|
||||
self.assertIn("托管已自动退出", status_call.kwargs["error_msg"])
|
||||
record_system_log.assert_called_once()
|
||||
|
||||
async def test_cookie_uid_without_valid_profile_payload_stays_unknown(self):
|
||||
auth = SimpleNamespace(
|
||||
cookie={},
|
||||
msToken="test-ms-token",
|
||||
get_uid=lambda: "cookie-only-uid",
|
||||
)
|
||||
payloads = [
|
||||
{"status_code": 0, "data": {"status": "ok"}},
|
||||
{"status_code": 0, "message": "success"},
|
||||
{"status_code": 0, "data": {}},
|
||||
]
|
||||
responses = [SimpleNamespace(json=lambda value=value: value) for value in payloads]
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"_build_auth",
|
||||
return_value=(auth, "test-agent"),
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module.requests,
|
||||
"get",
|
||||
side_effect=responses,
|
||||
) as request_get,
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"generate_a_bogus",
|
||||
return_value="a-bogus",
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"generate_webid",
|
||||
return_value="web-id",
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"_requests_proxies",
|
||||
return_value=None,
|
||||
),
|
||||
):
|
||||
raw_detail = account_profile_module.fetch_douyin_profile_detail_sync(
|
||||
"cookie-json",
|
||||
"test-agent",
|
||||
)
|
||||
|
||||
self.assertEqual(request_get.call_count, 3)
|
||||
self.assertEqual(raw_detail["uid"], "cookie-only-uid")
|
||||
self.assertFalse(raw_detail["fetched"])
|
||||
self.assertFalse(raw_detail["profile_response_valid"])
|
||||
|
||||
with patch.object(
|
||||
account_profile_module,
|
||||
"fetch_douyin_profile_detail",
|
||||
new=AsyncMock(return_value=raw_detail),
|
||||
):
|
||||
checked = await account_profile_module.fetch_douyin_profile_detail_with_sec_user_id(
|
||||
"cookie-json",
|
||||
"test-agent",
|
||||
)
|
||||
|
||||
self.assertEqual(checked["sec_user_id_status"], "unknown")
|
||||
|
||||
async def test_invalid_uid_fallback_payloads_stay_unknown(self):
|
||||
auth = SimpleNamespace(cookie={}, msToken="test-ms-token")
|
||||
cases = (
|
||||
(
|
||||
"nonzero-status",
|
||||
{
|
||||
"status_code": 10007,
|
||||
"data": {
|
||||
"user": {
|
||||
"uid": "fallback-uid",
|
||||
"sec_uid": "must-not-be-used",
|
||||
}
|
||||
},
|
||||
},
|
||||
),
|
||||
(
|
||||
"no-user-node",
|
||||
{"status_code": 0, "data": {"status": "ok"}},
|
||||
),
|
||||
)
|
||||
|
||||
for label, payload in cases:
|
||||
with self.subTest(label=label):
|
||||
with (
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"_build_auth",
|
||||
return_value=(auth, "test-agent"),
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"_douyin_get_json",
|
||||
return_value=payload,
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"generate_webid",
|
||||
return_value="web-id",
|
||||
),
|
||||
):
|
||||
resolved, request_completed = (
|
||||
account_profile_module._resolve_sec_user_id_by_uid_sync(
|
||||
"cookie-json",
|
||||
"test-agent",
|
||||
"fallback-uid",
|
||||
)
|
||||
)
|
||||
|
||||
self.assertEqual(resolved, "")
|
||||
self.assertFalse(request_completed)
|
||||
|
||||
fetch_detail = AsyncMock(
|
||||
return_value={
|
||||
"uid": "fallback-uid",
|
||||
"sec_user_id": "",
|
||||
"fetched": True,
|
||||
"message": "",
|
||||
}
|
||||
)
|
||||
run_in_thread = AsyncMock(
|
||||
return_value=(resolved, request_completed)
|
||||
)
|
||||
with (
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"fetch_douyin_profile_detail",
|
||||
fetch_detail,
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module.asyncio,
|
||||
"to_thread",
|
||||
run_in_thread,
|
||||
),
|
||||
):
|
||||
checked = await account_profile_module.fetch_douyin_profile_detail_with_sec_user_id(
|
||||
"cookie-json",
|
||||
"test-agent",
|
||||
)
|
||||
|
||||
self.assertEqual(checked["sec_user_id_status"], "unknown")
|
||||
|
||||
async def test_unknown_profile_sync_preserves_cached_identity_and_videos(self):
|
||||
old_synced_at = object()
|
||||
profile = SimpleNamespace(
|
||||
sec_user_id="cached-sec-user-id",
|
||||
synced_at=old_synced_at,
|
||||
sync_message="previous message",
|
||||
)
|
||||
account = SimpleNamespace(
|
||||
id=411,
|
||||
user_agent="test-agent",
|
||||
douyin_uid="cached-uid",
|
||||
username="cached-user",
|
||||
avatar_url="cached-avatar",
|
||||
)
|
||||
scalar_result = SimpleNamespace(scalar_one_or_none=lambda: profile)
|
||||
db = SimpleNamespace(
|
||||
execute=AsyncMock(return_value=scalar_result),
|
||||
commit=AsyncMock(),
|
||||
add=MagicMock(),
|
||||
)
|
||||
unknown_detail = {
|
||||
"uid": "cookie-uid",
|
||||
"sec_user_id": "",
|
||||
"sec_user_id_status": "unknown",
|
||||
"fetched": False,
|
||||
"message": "profile endpoint temporarily unavailable",
|
||||
}
|
||||
cached_response = {
|
||||
"account_id": 411,
|
||||
"sec_user_id": "cached-sec-user-id",
|
||||
"synced_at": "cached-time",
|
||||
"videos": [{"aweme_id": "cached-video"}],
|
||||
}
|
||||
fetch_videos = AsyncMock()
|
||||
apply_profile = AsyncMock()
|
||||
load_cached = AsyncMock(return_value=dict(cached_response))
|
||||
|
||||
with (
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"fetch_douyin_profile_detail_with_sec_user_id",
|
||||
new=AsyncMock(return_value=unknown_detail),
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"fetch_douyin_user_videos",
|
||||
fetch_videos,
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"apply_douyin_profile",
|
||||
apply_profile,
|
||||
),
|
||||
patch.object(
|
||||
account_profile_module,
|
||||
"load_account_profile_from_db",
|
||||
load_cached,
|
||||
),
|
||||
):
|
||||
result = await account_profile_module.sync_account_profile_to_db(
|
||||
db,
|
||||
account,
|
||||
"cookie-json",
|
||||
)
|
||||
|
||||
self.assertEqual(profile.sec_user_id, "cached-sec-user-id")
|
||||
self.assertIs(profile.synced_at, old_synced_at)
|
||||
fetch_videos.assert_not_awaited()
|
||||
apply_profile.assert_not_awaited()
|
||||
load_cached.assert_awaited_once_with(db, account)
|
||||
self.assertEqual(db.execute.await_count, 1)
|
||||
self.assertNotIn("delete", str(db.execute.await_args.args[0]).lower())
|
||||
self.assertEqual(result["videos"], cached_response["videos"])
|
||||
self.assertEqual(
|
||||
result["message"],
|
||||
"profile endpoint temporarily unavailable",
|
||||
)
|
||||
|
||||
async def test_browser_login_missing_sec_user_id_closes_browser_and_skips_im(self):
|
||||
worker = DouyinWorker(account_id=305, login_mode="browser")
|
||||
worker.is_running = True
|
||||
events: list[str] = []
|
||||
|
||||
async def reject_identity(*_args, **_kwargs):
|
||||
events.append("require-sec-user-id")
|
||||
return False
|
||||
|
||||
worker._load_user_agent = AsyncMock(return_value="test-agent")
|
||||
worker._probe_existing_login = AsyncMock(return_value=True)
|
||||
worker._finalize_login_session = AsyncMock()
|
||||
worker._require_sec_user_id = AsyncMock(side_effect=reject_identity)
|
||||
worker._setup_im_network_listener = AsyncMock()
|
||||
worker._navigate_to_message_center = AsyncMock()
|
||||
worker._harvest_im_credentials = AsyncMock()
|
||||
worker._persist_cookies = AsyncMock()
|
||||
im_session = DouyinImSession(
|
||||
cookies={"sessionid": "test-session"},
|
||||
my_uid=30501,
|
||||
)
|
||||
worker._build_im_session = AsyncMock(return_value=im_session)
|
||||
worker._persist_im_session = AsyncMock()
|
||||
worker._run_im_direct_service = AsyncMock()
|
||||
worker._close_browser_only = AsyncMock()
|
||||
|
||||
page = SimpleNamespace()
|
||||
context = SimpleNamespace(
|
||||
add_init_script=AsyncMock(),
|
||||
new_page=AsyncMock(return_value=page),
|
||||
)
|
||||
browser = SimpleNamespace(new_context=AsyncMock(return_value=context))
|
||||
playwright = SimpleNamespace()
|
||||
playwright_starter = SimpleNamespace(
|
||||
start=AsyncMock(return_value=playwright),
|
||||
)
|
||||
|
||||
@asynccontextmanager
|
||||
async def browser_slot(account_id: int, description: str):
|
||||
self.assertEqual(account_id, 305)
|
||||
self.assertEqual(description, "browser credential harvest")
|
||||
events.append("browser-acquired")
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
events.append("browser-released")
|
||||
|
||||
controller = SimpleNamespace(browser_slot=browser_slot)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"rpa_engine.playwright_worker.async_playwright",
|
||||
return_value=playwright_starter,
|
||||
),
|
||||
patch(
|
||||
"rpa_engine.playwright_worker._launch_chromium",
|
||||
new=AsyncMock(return_value=browser),
|
||||
),
|
||||
patch(
|
||||
"rpa_engine.playwright_worker.get_traffic_controller",
|
||||
return_value=controller,
|
||||
),
|
||||
patch(
|
||||
"rpa_engine.playwright_worker.validate_im_session",
|
||||
new=AsyncMock(return_value=(True, "ready")),
|
||||
),
|
||||
):
|
||||
await worker._run_browser_im_flow(
|
||||
storage_state=None,
|
||||
cookie_info={
|
||||
"cookie_valid": False,
|
||||
"has_sessionid": False,
|
||||
"reason": "no saved login",
|
||||
},
|
||||
)
|
||||
|
||||
worker._finalize_login_session.assert_awaited_once_with()
|
||||
worker._require_sec_user_id.assert_awaited_once()
|
||||
require_call = worker._require_sec_user_id.await_args
|
||||
self.assertTrue(require_call.kwargs["force_refresh"])
|
||||
self.assertEqual(
|
||||
events,
|
||||
["browser-acquired", "browser-released", "require-sec-user-id"],
|
||||
)
|
||||
worker._setup_im_network_listener.assert_awaited_once_with()
|
||||
worker._navigate_to_message_center.assert_awaited_once_with()
|
||||
worker._harvest_im_credentials.assert_awaited_once_with(timeout=25)
|
||||
worker._persist_cookies.assert_awaited_once_with()
|
||||
worker._build_im_session.assert_awaited_once_with()
|
||||
worker._persist_im_session.assert_not_awaited()
|
||||
worker._run_im_direct_service.assert_not_awaited()
|
||||
worker._close_browser_only.assert_awaited_once_with()
|
||||
|
||||
async def test_browser_login_with_sec_user_id_continues_to_im(self):
|
||||
worker = DouyinWorker(account_id=306, login_mode="browser")
|
||||
worker.is_running = True
|
||||
worker._load_user_agent = AsyncMock(return_value="test-agent")
|
||||
worker._probe_existing_login = AsyncMock(return_value=True)
|
||||
worker._finalize_login_session = AsyncMock()
|
||||
worker._require_sec_user_id = AsyncMock(return_value=True)
|
||||
worker._setup_im_network_listener = AsyncMock()
|
||||
worker._navigate_to_message_center = AsyncMock()
|
||||
worker._harvest_im_credentials = AsyncMock()
|
||||
worker._persist_cookies = AsyncMock()
|
||||
im_session = DouyinImSession(
|
||||
cookies={"sessionid": "test-session"},
|
||||
my_uid=30601,
|
||||
)
|
||||
worker._build_im_session = AsyncMock(return_value=im_session)
|
||||
worker._persist_im_session = AsyncMock()
|
||||
worker._run_im_direct_service = AsyncMock()
|
||||
worker._close_browser_only = AsyncMock()
|
||||
|
||||
page = SimpleNamespace()
|
||||
context = SimpleNamespace(
|
||||
add_init_script=AsyncMock(),
|
||||
new_page=AsyncMock(return_value=page),
|
||||
)
|
||||
browser = SimpleNamespace(new_context=AsyncMock(return_value=context))
|
||||
playwright_starter = SimpleNamespace(
|
||||
start=AsyncMock(return_value=SimpleNamespace()),
|
||||
)
|
||||
|
||||
@asynccontextmanager
|
||||
async def browser_slot(_account_id: int, _description: str):
|
||||
yield
|
||||
|
||||
controller = SimpleNamespace(browser_slot=browser_slot)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"rpa_engine.playwright_worker.async_playwright",
|
||||
return_value=playwright_starter,
|
||||
),
|
||||
patch(
|
||||
"rpa_engine.playwright_worker._launch_chromium",
|
||||
new=AsyncMock(return_value=browser),
|
||||
),
|
||||
patch(
|
||||
"rpa_engine.playwright_worker.get_traffic_controller",
|
||||
return_value=controller,
|
||||
),
|
||||
patch(
|
||||
"rpa_engine.playwright_worker.validate_im_session",
|
||||
new=AsyncMock(return_value=(True, "ready")),
|
||||
),
|
||||
):
|
||||
await worker._run_browser_im_flow(
|
||||
storage_state=None,
|
||||
cookie_info={
|
||||
"cookie_valid": False,
|
||||
"has_sessionid": False,
|
||||
"reason": "no saved login",
|
||||
},
|
||||
)
|
||||
|
||||
worker._require_sec_user_id.assert_awaited_once()
|
||||
self.assertTrue(worker._require_sec_user_id.await_args.kwargs["force_refresh"])
|
||||
worker._setup_im_network_listener.assert_awaited_once_with()
|
||||
worker._navigate_to_message_center.assert_awaited_once_with()
|
||||
worker._harvest_im_credentials.assert_awaited_once_with(timeout=25)
|
||||
worker._build_im_session.assert_awaited_once_with()
|
||||
worker._persist_im_session.assert_awaited_once_with(
|
||||
im_session,
|
||||
status="online",
|
||||
clear_error=True,
|
||||
)
|
||||
worker._run_im_direct_service.assert_awaited_once_with(im_session)
|
||||
worker._close_browser_only.assert_awaited()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,284 @@
|
||||
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))
|
||||
|
||||
from rpa_engine.douyin_im import http_client as http_client_module
|
||||
from rpa_engine.douyin_im.http_client import DouyinImHttpClient
|
||||
from rpa_engine.douyin_im.session import DouyinImSession
|
||||
from rpa_engine.playwright_worker import DouyinWorker
|
||||
|
||||
|
||||
class SendTextMessageEntryTests(unittest.IsolatedAsyncioTestCase):
|
||||
def _make_client(self, account_id: int = 41) -> DouyinImHttpClient:
|
||||
session = DouyinImSession(
|
||||
cookies={"sessionid": "test-session"},
|
||||
my_uid=10001,
|
||||
conv_meta={"existing": {"ticket": "old-ticket"}},
|
||||
)
|
||||
return DouyinImHttpClient(session, account_id=account_id)
|
||||
|
||||
async def test_public_entry_submits_the_whole_send_to_outbound_queue(self):
|
||||
client = self._make_client(account_id=73)
|
||||
submit = AsyncMock(return_value=True)
|
||||
|
||||
with patch(
|
||||
"rpa_engine.douyin_im.traffic_control.submit_outbound",
|
||||
submit,
|
||||
):
|
||||
sent = await client.send_text_message(
|
||||
"0:1:10001:20002",
|
||||
"hello",
|
||||
conversation_short_id="short-1",
|
||||
)
|
||||
|
||||
self.assertTrue(sent)
|
||||
submit.assert_awaited_once()
|
||||
account_id, operation = submit.await_args.args[:2]
|
||||
self.assertEqual(account_id, 73)
|
||||
self.assertTrue(callable(operation))
|
||||
self.assertIn("20002", submit.await_args.kwargs["description"])
|
||||
|
||||
async def test_bypass_path_does_not_submit_again(self):
|
||||
client = self._make_client()
|
||||
client._resolve_authoritative_uid = MagicMock(return_value=0)
|
||||
submit = AsyncMock(return_value=True)
|
||||
|
||||
with patch(
|
||||
"rpa_engine.douyin_im.traffic_control.submit_outbound",
|
||||
submit,
|
||||
):
|
||||
sent = await client.send_text_message(
|
||||
"0:1:10001:20002",
|
||||
"hello",
|
||||
_bypass_global_queue=True,
|
||||
)
|
||||
|
||||
self.assertFalse(sent)
|
||||
submit.assert_not_awaited()
|
||||
client._resolve_authoritative_uid.assert_called_once()
|
||||
|
||||
async def test_queued_client_result_state_is_copied_to_calling_client(self):
|
||||
client = self._make_client(account_id=88)
|
||||
queued_meta = {
|
||||
"0:1:10001:20002": {
|
||||
"ticket": "fresh-ticket",
|
||||
"conversation_short_id": "fresh-short-id",
|
||||
}
|
||||
}
|
||||
queued_client = SimpleNamespace(
|
||||
send_text_message=AsyncMock(return_value=False),
|
||||
last_send_meta=queued_meta,
|
||||
last_error="credential expired",
|
||||
last_send_needs_refresh=True,
|
||||
last_request_debug="response status=401",
|
||||
)
|
||||
queued_context = MagicMock()
|
||||
queued_context.__aenter__ = AsyncMock(return_value=queued_client)
|
||||
queued_context.__aexit__ = AsyncMock(return_value=None)
|
||||
queued_factory = MagicMock(return_value=queued_context)
|
||||
|
||||
async def execute_submission(account_id, operation, description=""):
|
||||
self.assertEqual(account_id, 88)
|
||||
self.assertIn("20002", description)
|
||||
return await operation()
|
||||
|
||||
submit = AsyncMock(side_effect=execute_submission)
|
||||
|
||||
with (
|
||||
patch(
|
||||
"rpa_engine.douyin_im.traffic_control.submit_outbound",
|
||||
submit,
|
||||
),
|
||||
patch.object(
|
||||
http_client_module,
|
||||
"DouyinImHttpClient",
|
||||
queued_factory,
|
||||
),
|
||||
):
|
||||
sent = await client.send_text_message(
|
||||
"0:1:10001:20002",
|
||||
"queued hello",
|
||||
conversation_short_id="short-before-send",
|
||||
)
|
||||
|
||||
self.assertFalse(sent)
|
||||
submit.assert_awaited_once()
|
||||
queued_factory.assert_called_once_with(client.session, account_id=88)
|
||||
queued_client.send_text_message.assert_awaited_once_with(
|
||||
"0:1:10001:20002",
|
||||
"queued hello",
|
||||
conversation_short_id="short-before-send",
|
||||
_bypass_global_queue=True,
|
||||
)
|
||||
self.assertEqual(client.last_send_meta, queued_meta)
|
||||
self.assertIsNot(client.last_send_meta, queued_meta)
|
||||
self.assertEqual(client.last_error, "credential expired")
|
||||
self.assertTrue(client.last_send_needs_refresh)
|
||||
self.assertEqual(client.last_request_debug, "response status=401")
|
||||
|
||||
|
||||
class WorkerLifecycleTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_start_saves_task_and_stop_waits_until_it_is_done(self):
|
||||
worker = DouyinWorker(account_id=919, login_mode="im_direct")
|
||||
loop_started = asyncio.Event()
|
||||
loop_finished = asyncio.Event()
|
||||
never_finishes_without_cancellation = asyncio.Event()
|
||||
|
||||
async def controlled_run_loop() -> None:
|
||||
loop_started.set()
|
||||
try:
|
||||
await never_finishes_without_cancellation.wait()
|
||||
finally:
|
||||
loop_finished.set()
|
||||
|
||||
worker._run_loop = controlled_run_loop
|
||||
worker.update_account_status = AsyncMock()
|
||||
worker.cleanup = AsyncMock()
|
||||
|
||||
await worker.start()
|
||||
task = worker._task
|
||||
|
||||
self.assertIsInstance(task, asyncio.Task)
|
||||
self.assertEqual(task.get_name(), "douyin-worker-919")
|
||||
self.assertFalse(task.done())
|
||||
await asyncio.wait_for(loop_started.wait(), timeout=0.2)
|
||||
|
||||
# Starting an already-running worker must retain the same owned task.
|
||||
await worker.start()
|
||||
self.assertIs(worker._task, task)
|
||||
|
||||
await asyncio.wait_for(worker.stop(), timeout=0.5)
|
||||
|
||||
self.assertTrue(task.done())
|
||||
self.assertTrue(task.cancelled())
|
||||
self.assertTrue(loop_finished.is_set())
|
||||
self.assertIsNone(worker._task)
|
||||
self.assertTrue(worker.stopping)
|
||||
self.assertFalse(worker.is_running)
|
||||
worker.update_account_status.assert_awaited_once_with("offline")
|
||||
worker.cleanup.assert_awaited_once_with()
|
||||
leaked = [
|
||||
running
|
||||
for running in asyncio.all_tasks()
|
||||
if running is not asyncio.current_task()
|
||||
and running.get_name() == "douyin-worker-919"
|
||||
and not running.done()
|
||||
]
|
||||
self.assertEqual(leaked, [])
|
||||
|
||||
async def test_browser_flow_closes_resources_before_releasing_slot_on_error(self):
|
||||
worker = DouyinWorker(account_id=920, login_mode="browser")
|
||||
events: list[str] = []
|
||||
|
||||
# A truthy browser resource makes the flow's finally block responsible
|
||||
# for cleanup even though credential preparation fails midway.
|
||||
worker.page = object()
|
||||
|
||||
async def failing_prepare(storage_state, cookie_info):
|
||||
events.append("prepare")
|
||||
raise RuntimeError("credential harvest failed")
|
||||
|
||||
async def close_browser_only():
|
||||
events.append("close")
|
||||
worker.page = None
|
||||
|
||||
@asynccontextmanager
|
||||
async def browser_slot(account_id, description):
|
||||
self.assertEqual(account_id, 920)
|
||||
self.assertEqual(description, "browser credential harvest")
|
||||
events.append("acquire")
|
||||
try:
|
||||
yield
|
||||
finally:
|
||||
events.append("release")
|
||||
|
||||
controller = SimpleNamespace(browser_slot=browser_slot)
|
||||
worker._prepare_browser_im_flow = AsyncMock(side_effect=failing_prepare)
|
||||
worker._close_browser_only = AsyncMock(side_effect=close_browser_only)
|
||||
|
||||
with patch(
|
||||
"rpa_engine.playwright_worker.get_traffic_controller",
|
||||
return_value=controller,
|
||||
):
|
||||
with self.assertRaisesRegex(RuntimeError, "credential harvest failed"):
|
||||
await worker._run_browser_im_flow(
|
||||
storage_state={"cookies": []},
|
||||
cookie_info={"has_sessionid": True},
|
||||
)
|
||||
|
||||
self.assertEqual(events, ["acquire", "prepare", "close", "release"])
|
||||
worker._prepare_browser_im_flow.assert_awaited_once_with(
|
||||
{"cookies": []},
|
||||
{"has_sessionid": True},
|
||||
)
|
||||
worker._close_browser_only.assert_awaited_once_with()
|
||||
|
||||
async def test_close_browser_only_closes_every_resource_and_clears_references(self):
|
||||
worker = DouyinWorker(account_id=921, login_mode="browser")
|
||||
events: list[str] = []
|
||||
|
||||
async def record(name: str, *, fail: bool = False) -> None:
|
||||
events.append(name)
|
||||
if fail:
|
||||
raise RuntimeError(f"{name} close failed")
|
||||
|
||||
async def close_page() -> None:
|
||||
await record("page")
|
||||
|
||||
async def close_context() -> None:
|
||||
await record("context", fail=True)
|
||||
|
||||
async def close_browser() -> None:
|
||||
await record("browser")
|
||||
|
||||
async def stop_playwright() -> None:
|
||||
await record("playwright")
|
||||
|
||||
page = SimpleNamespace(
|
||||
close=AsyncMock(side_effect=close_page),
|
||||
)
|
||||
context = SimpleNamespace(
|
||||
# One failed close must not keep the remaining processes alive.
|
||||
close=AsyncMock(side_effect=close_context),
|
||||
)
|
||||
browser = SimpleNamespace(
|
||||
close=AsyncMock(side_effect=close_browser),
|
||||
)
|
||||
playwright = SimpleNamespace(
|
||||
stop=AsyncMock(side_effect=stop_playwright),
|
||||
)
|
||||
worker.page = page
|
||||
worker.context = context
|
||||
worker.browser = browser
|
||||
worker.playwright = playwright
|
||||
|
||||
await worker._close_browser_only()
|
||||
|
||||
self.assertEqual(events, ["page", "context", "browser", "playwright"])
|
||||
page.close.assert_awaited_once_with()
|
||||
context.close.assert_awaited_once_with()
|
||||
browser.close.assert_awaited_once_with()
|
||||
playwright.stop.assert_awaited_once_with()
|
||||
self.assertIsNone(worker.page)
|
||||
self.assertIsNone(worker.context)
|
||||
self.assertIsNone(worker.browser)
|
||||
self.assertIsNone(worker.playwright)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
@@ -0,0 +1,495 @@
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
import unittest
|
||||
from pathlib import Path
|
||||
from unittest.mock import patch
|
||||
|
||||
|
||||
BACKEND_DIR = Path(__file__).resolve().parents[1]
|
||||
if str(BACKEND_DIR) not in sys.path:
|
||||
sys.path.insert(0, str(BACKEND_DIR))
|
||||
|
||||
from rpa_engine.douyin_im.traffic_control import (
|
||||
GlobalSendQueue,
|
||||
TrafficController,
|
||||
get_traffic_controller,
|
||||
)
|
||||
|
||||
|
||||
class GlobalSendQueueTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def _make_queue(self, interval_seconds: float = 0.0) -> GlobalSendQueue:
|
||||
queue = GlobalSendQueue(interval_seconds=interval_seconds)
|
||||
self.addAsyncCleanup(queue.stop)
|
||||
return queue
|
||||
|
||||
async def _wait_for_queued_count(
|
||||
self,
|
||||
queue: GlobalSendQueue,
|
||||
expected: int,
|
||||
timeout: float = 0.5,
|
||||
) -> None:
|
||||
loop = asyncio.get_running_loop()
|
||||
deadline = loop.time() + timeout
|
||||
while loop.time() < deadline:
|
||||
if (await queue.snapshot())["queued_count"] == expected:
|
||||
return
|
||||
await asyncio.sleep(0.001)
|
||||
self.assertEqual((await queue.snapshot())["queued_count"], expected)
|
||||
|
||||
async def test_global_send_concurrency_is_one_across_accounts(self):
|
||||
queue = await self._make_queue()
|
||||
active = 0
|
||||
maximum_active = 0
|
||||
|
||||
async def operation(index: int) -> int:
|
||||
nonlocal active, maximum_active
|
||||
active += 1
|
||||
maximum_active = max(maximum_active, active)
|
||||
await asyncio.sleep(0.015)
|
||||
active -= 1
|
||||
return index
|
||||
|
||||
tasks = [
|
||||
asyncio.create_task(
|
||||
queue.submit(
|
||||
account_id=(index % 3) + 1,
|
||||
operation=lambda index=index: operation(index),
|
||||
description=f"send-{index}",
|
||||
)
|
||||
)
|
||||
for index in range(9)
|
||||
]
|
||||
|
||||
results = await asyncio.wait_for(asyncio.gather(*tasks), timeout=1.0)
|
||||
|
||||
self.assertEqual(sorted(results), list(range(9)))
|
||||
self.assertEqual(maximum_active, 1)
|
||||
self.assertEqual(active, 0)
|
||||
|
||||
async def test_same_account_keeps_fifo_order(self):
|
||||
queue = await self._make_queue()
|
||||
calls: list[int] = []
|
||||
|
||||
async def operation(index: int) -> int:
|
||||
calls.append(index)
|
||||
await asyncio.sleep(0)
|
||||
return index
|
||||
|
||||
tasks = [
|
||||
asyncio.create_task(
|
||||
queue.submit(
|
||||
account_id=101,
|
||||
operation=lambda index=index: operation(index),
|
||||
description=str(index),
|
||||
)
|
||||
)
|
||||
for index in range(6)
|
||||
]
|
||||
|
||||
results = await asyncio.wait_for(asyncio.gather(*tasks), timeout=0.5)
|
||||
|
||||
self.assertEqual(calls, list(range(6)))
|
||||
self.assertEqual(results, list(range(6)))
|
||||
|
||||
async def test_busy_accounts_are_dispatched_round_robin(self):
|
||||
queue = await self._make_queue()
|
||||
calls: list[str] = []
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
|
||||
async def first_operation() -> str:
|
||||
calls.append("a1")
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
return "a1"
|
||||
|
||||
async def operation(name: str) -> str:
|
||||
calls.append(name)
|
||||
await asyncio.sleep(0)
|
||||
return name
|
||||
|
||||
first = asyncio.create_task(queue.submit(1, first_operation, "a1"))
|
||||
await asyncio.wait_for(first_started.wait(), timeout=0.2)
|
||||
|
||||
# Account 1 is deliberately noisy. Accounts 2 and 3 must each get a
|
||||
# turn before account 1 is allowed to dispatch its next waiting job.
|
||||
remaining_specs = [
|
||||
(1, "a2"),
|
||||
(1, "a3"),
|
||||
(2, "b1"),
|
||||
(2, "b2"),
|
||||
(3, "c1"),
|
||||
]
|
||||
remaining = [
|
||||
asyncio.create_task(
|
||||
queue.submit(
|
||||
account_id,
|
||||
lambda name=name: operation(name),
|
||||
description=name,
|
||||
)
|
||||
)
|
||||
for account_id, name in remaining_specs
|
||||
]
|
||||
await self._wait_for_queued_count(queue, expected=len(remaining))
|
||||
|
||||
release_first.set()
|
||||
results = await asyncio.wait_for(
|
||||
asyncio.gather(first, *remaining),
|
||||
timeout=0.5,
|
||||
)
|
||||
|
||||
self.assertEqual(calls, ["a1", "b1", "c1", "a2", "b2", "a3"])
|
||||
self.assertCountEqual(results, ["a1", "a2", "a3", "b1", "b2", "c1"])
|
||||
|
||||
async def test_send_starts_are_separated_by_configured_interval(self):
|
||||
interval = 0.08
|
||||
queue = await self._make_queue(interval_seconds=interval)
|
||||
loop = asyncio.get_running_loop()
|
||||
started_at: list[float] = []
|
||||
|
||||
async def operation() -> None:
|
||||
started_at.append(loop.time())
|
||||
|
||||
tasks = [
|
||||
asyncio.create_task(queue.submit(index + 1, operation, str(index)))
|
||||
for index in range(3)
|
||||
]
|
||||
await asyncio.wait_for(asyncio.gather(*tasks), timeout=0.7)
|
||||
|
||||
self.assertEqual(len(started_at), 3)
|
||||
gaps = [later - earlier for earlier, later in zip(started_at, started_at[1:])]
|
||||
for gap in gaps:
|
||||
# The Windows selector clock may wake one ~15.6 ms tick early.
|
||||
self.assertGreaterEqual(gap, interval - 0.025)
|
||||
|
||||
async def test_operation_failure_does_not_block_later_job(self):
|
||||
queue = await self._make_queue()
|
||||
calls: list[str] = []
|
||||
|
||||
async def failing_operation() -> None:
|
||||
calls.append("failed")
|
||||
raise ValueError("expected failure")
|
||||
|
||||
async def following_operation() -> str:
|
||||
calls.append("continued")
|
||||
return "ok"
|
||||
|
||||
failed = asyncio.create_task(queue.submit(1, failing_operation, "failed"))
|
||||
continued = asyncio.create_task(queue.submit(2, following_operation, "continued"))
|
||||
|
||||
results = await asyncio.wait_for(
|
||||
asyncio.gather(failed, continued, return_exceptions=True),
|
||||
timeout=0.2,
|
||||
)
|
||||
|
||||
self.assertIsInstance(results[0], ValueError)
|
||||
self.assertEqual(str(results[0]), "expected failure")
|
||||
self.assertEqual(results[1], "ok")
|
||||
self.assertEqual(calls, ["failed", "continued"])
|
||||
|
||||
async def test_cancelling_waiting_submit_prevents_operation(self):
|
||||
queue = await self._make_queue()
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
cancelled_operation_called = asyncio.Event()
|
||||
|
||||
async def blocking_operation() -> None:
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
|
||||
async def must_not_run() -> None:
|
||||
cancelled_operation_called.set()
|
||||
|
||||
blocker = asyncio.create_task(queue.submit(1, blocking_operation, "blocker"))
|
||||
await asyncio.wait_for(first_started.wait(), timeout=0.2)
|
||||
waiting = asyncio.create_task(queue.submit(2, must_not_run, "cancelled"))
|
||||
await self._wait_for_queued_count(queue, expected=1)
|
||||
|
||||
waiting.cancel()
|
||||
with self.assertRaises(asyncio.CancelledError):
|
||||
await waiting
|
||||
release_first.set()
|
||||
await asyncio.wait_for(blocker, timeout=0.2)
|
||||
await asyncio.sleep(0.02)
|
||||
|
||||
self.assertFalse(cancelled_operation_called.is_set())
|
||||
self.assertEqual((await queue.snapshot())["pending_count"], 0)
|
||||
|
||||
async def test_cancelling_active_submit_waits_for_real_result_without_overlap(self):
|
||||
queue = await self._make_queue()
|
||||
active = 0
|
||||
maximum_active = 0
|
||||
first_started = asyncio.Event()
|
||||
release_first = asyncio.Event()
|
||||
second_started = asyncio.Event()
|
||||
|
||||
async def first_operation() -> str:
|
||||
nonlocal active, maximum_active
|
||||
active += 1
|
||||
maximum_active = max(maximum_active, active)
|
||||
first_started.set()
|
||||
await release_first.wait()
|
||||
active -= 1
|
||||
return "delivered"
|
||||
|
||||
async def second_operation() -> str:
|
||||
nonlocal active, maximum_active
|
||||
active += 1
|
||||
maximum_active = max(maximum_active, active)
|
||||
second_started.set()
|
||||
await asyncio.sleep(0)
|
||||
active -= 1
|
||||
return "next"
|
||||
|
||||
first = asyncio.create_task(queue.submit(1, first_operation, "active"))
|
||||
await asyncio.wait_for(first_started.wait(), timeout=0.2)
|
||||
second = asyncio.create_task(queue.submit(2, second_operation, "next"))
|
||||
await self._wait_for_queued_count(queue, expected=1)
|
||||
|
||||
first.cancel()
|
||||
await asyncio.sleep(0.02)
|
||||
|
||||
# Once the network operation has begun, cancelling its caller must not
|
||||
# release the single lane or report a false cancellation to the caller.
|
||||
self.assertFalse(first.done())
|
||||
self.assertFalse(second_started.is_set())
|
||||
self.assertEqual(active, 1)
|
||||
|
||||
release_first.set()
|
||||
self.assertEqual(await asyncio.wait_for(first, timeout=0.2), "delivered")
|
||||
self.assertEqual(await asyncio.wait_for(second, timeout=0.2), "next")
|
||||
self.assertTrue(second_started.is_set())
|
||||
self.assertEqual(maximum_active, 1)
|
||||
self.assertEqual(active, 0)
|
||||
|
||||
async def test_cancel_account_removes_waiting_jobs_without_running_them(self):
|
||||
queue = await self._make_queue()
|
||||
blocker_started = asyncio.Event()
|
||||
release_blocker = asyncio.Event()
|
||||
removed_calls: list[str] = []
|
||||
|
||||
async def blocking_operation() -> None:
|
||||
blocker_started.set()
|
||||
await release_blocker.wait()
|
||||
|
||||
async def removed_operation(name: str) -> None:
|
||||
removed_calls.append(name)
|
||||
|
||||
blocker = asyncio.create_task(queue.submit(1, blocking_operation, "blocker"))
|
||||
await asyncio.wait_for(blocker_started.wait(), timeout=0.2)
|
||||
removed = [
|
||||
asyncio.create_task(
|
||||
queue.submit(
|
||||
77,
|
||||
lambda name=name: removed_operation(name),
|
||||
name,
|
||||
)
|
||||
)
|
||||
for name in ("waiting-1", "waiting-2")
|
||||
]
|
||||
await self._wait_for_queued_count(queue, expected=2)
|
||||
|
||||
self.assertEqual(await queue.cancel_account(77), 2)
|
||||
results = await asyncio.wait_for(
|
||||
asyncio.gather(*removed, return_exceptions=True),
|
||||
timeout=0.1,
|
||||
)
|
||||
snapshot = await queue.snapshot()
|
||||
|
||||
self.assertTrue(all(isinstance(result, asyncio.CancelledError) for result in results))
|
||||
self.assertEqual(removed_calls, [])
|
||||
self.assertNotIn(77, snapshot["per_account"])
|
||||
self.assertEqual(snapshot["queued_count"], 0)
|
||||
self.assertEqual(snapshot["active_account_id"], 1)
|
||||
|
||||
release_blocker.set()
|
||||
await asyncio.wait_for(blocker, timeout=0.2)
|
||||
await asyncio.sleep(0)
|
||||
self.assertEqual((await queue.snapshot())["pending_count"], 0)
|
||||
|
||||
async def test_cancel_account_releases_active_caller_but_drains_lane(self):
|
||||
queue = await self._make_queue()
|
||||
active = 0
|
||||
maximum_active = 0
|
||||
active_started = asyncio.Event()
|
||||
release_active = asyncio.Event()
|
||||
active_drained = asyncio.Event()
|
||||
next_started = asyncio.Event()
|
||||
calls: list[str] = []
|
||||
|
||||
async def active_operation() -> str:
|
||||
nonlocal active, maximum_active
|
||||
active += 1
|
||||
maximum_active = max(maximum_active, active)
|
||||
calls.append("active-start")
|
||||
active_started.set()
|
||||
await release_active.wait()
|
||||
calls.append("active-drained")
|
||||
active -= 1
|
||||
active_drained.set()
|
||||
return "discarded-result"
|
||||
|
||||
async def next_operation() -> str:
|
||||
nonlocal active, maximum_active
|
||||
active += 1
|
||||
maximum_active = max(maximum_active, active)
|
||||
calls.append("next")
|
||||
next_started.set()
|
||||
await asyncio.sleep(0)
|
||||
active -= 1
|
||||
return "next-result"
|
||||
|
||||
active_submit = asyncio.create_task(
|
||||
queue.submit(88, active_operation, "active-account")
|
||||
)
|
||||
await asyncio.wait_for(active_started.wait(), timeout=0.2)
|
||||
next_submit = asyncio.create_task(queue.submit(89, next_operation, "other-account"))
|
||||
await self._wait_for_queued_count(queue, expected=1)
|
||||
|
||||
self.assertEqual(await queue.cancel_account(88), 1)
|
||||
cancelled_result = await asyncio.wait_for(
|
||||
asyncio.gather(active_submit, return_exceptions=True),
|
||||
timeout=0.1,
|
||||
)
|
||||
|
||||
self.assertIsInstance(cancelled_result[0], asyncio.CancelledError)
|
||||
self.assertTrue(active_submit.cancelled())
|
||||
self.assertFalse(active_drained.is_set())
|
||||
self.assertFalse(next_started.is_set())
|
||||
self.assertEqual(active, 1)
|
||||
|
||||
# The cancelled account's network operation still owns the lane until
|
||||
# it actually drains; only then can a different account begin.
|
||||
release_active.set()
|
||||
self.assertEqual(
|
||||
await asyncio.wait_for(next_submit, timeout=0.2),
|
||||
"next-result",
|
||||
)
|
||||
self.assertTrue(active_drained.is_set())
|
||||
self.assertEqual(calls, ["active-start", "active-drained", "next"])
|
||||
self.assertEqual(maximum_active, 1)
|
||||
self.assertEqual(active, 0)
|
||||
|
||||
# A fresh submission proves the dispatcher remains usable after the
|
||||
# cancelled active job and its queued successor have both completed.
|
||||
async def recovered_operation() -> str:
|
||||
calls.append("recovered")
|
||||
return "recovered-result"
|
||||
|
||||
self.assertEqual(
|
||||
await asyncio.wait_for(
|
||||
queue.submit(90, recovered_operation, "recovered"),
|
||||
timeout=0.2,
|
||||
),
|
||||
"recovered-result",
|
||||
)
|
||||
self.assertEqual(calls[-1], "recovered")
|
||||
self.assertEqual((await queue.snapshot())["pending_count"], 0)
|
||||
|
||||
|
||||
class BackgroundTrafficLimitTests(unittest.IsolatedAsyncioTestCase):
|
||||
async def test_background_slot_respects_configured_concurrency_limit(self):
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{"KEFU_BACKGROUND_NETWORK_CONCURRENCY": "2"},
|
||||
):
|
||||
controller = TrafficController()
|
||||
self.addAsyncCleanup(controller.stop)
|
||||
|
||||
active = 0
|
||||
maximum_active = 0
|
||||
entered = 0
|
||||
all_finished = asyncio.Event()
|
||||
|
||||
async def request(account_id: int) -> None:
|
||||
nonlocal active, maximum_active, entered
|
||||
async with controller.background_slot(account_id, "poll"):
|
||||
active += 1
|
||||
entered += 1
|
||||
maximum_active = max(maximum_active, active)
|
||||
await asyncio.sleep(0.02)
|
||||
active -= 1
|
||||
if entered == 8 and active == 0:
|
||||
all_finished.set()
|
||||
|
||||
tasks = [asyncio.create_task(request(index + 1)) for index in range(8)]
|
||||
await asyncio.wait_for(asyncio.gather(*tasks), timeout=0.5)
|
||||
|
||||
self.assertEqual(entered, 8)
|
||||
self.assertEqual(maximum_active, 2)
|
||||
self.assertEqual(active, 0)
|
||||
self.assertEqual(controller.background_active, 0)
|
||||
self.assertEqual(controller.background_waiting, 0)
|
||||
|
||||
async def test_nested_background_slot_in_same_task_is_reentrant(self):
|
||||
with patch.dict(
|
||||
os.environ,
|
||||
{"KEFU_BACKGROUND_NETWORK_CONCURRENCY": "1"},
|
||||
):
|
||||
controller = TrafficController()
|
||||
self.addAsyncCleanup(controller.stop)
|
||||
|
||||
tasks_seen: list[asyncio.Task] = []
|
||||
|
||||
async def nested_request() -> None:
|
||||
tasks_seen.append(asyncio.current_task())
|
||||
async with controller.background_slot(1, "outer"):
|
||||
self.assertEqual(controller.background_active, 1)
|
||||
self.assertEqual(controller.background_waiting, 0)
|
||||
self.assertEqual(controller._background._value, 0)
|
||||
|
||||
tasks_seen.append(asyncio.current_task())
|
||||
async with controller.background_slot(1, "inner"):
|
||||
tasks_seen.append(asyncio.current_task())
|
||||
self.assertEqual(controller.background_active, 1)
|
||||
self.assertEqual(controller.background_waiting, 0)
|
||||
self.assertEqual(controller._background._value, 0)
|
||||
|
||||
self.assertEqual(controller.background_active, 1)
|
||||
self.assertEqual(controller._background._value, 0)
|
||||
|
||||
self.assertEqual(controller.background_active, 0)
|
||||
self.assertEqual(controller.background_waiting, 0)
|
||||
self.assertEqual(controller._background._value, 1)
|
||||
|
||||
await asyncio.wait_for(nested_request(), timeout=0.2)
|
||||
|
||||
self.assertEqual(len(tasks_seen), 3)
|
||||
self.assertTrue(all(task is tasks_seen[0] for task in tasks_seen))
|
||||
|
||||
|
||||
class TrafficControllerLoopIsolationTests(unittest.TestCase):
|
||||
def test_get_traffic_controller_does_not_reuse_asyncio_primitives(self):
|
||||
async def capture_controller_state():
|
||||
controller = get_traffic_controller()
|
||||
return {
|
||||
"loop": asyncio.get_running_loop(),
|
||||
"controller": controller,
|
||||
"send_lock": controller.send_queue._state_lock,
|
||||
"send_wake": controller.send_queue._wake,
|
||||
"background": controller._background,
|
||||
"browser": controller._browser,
|
||||
}
|
||||
|
||||
# IsolatedAsyncioTestCase uses a fresh asyncio.Runner per test. Keep
|
||||
# both closed loop objects alive so the WeakKeyDictionary is exercised
|
||||
# with two distinct keys rather than relying on garbage collection.
|
||||
with asyncio.Runner() as first_runner:
|
||||
first = first_runner.run(capture_controller_state())
|
||||
with asyncio.Runner() as second_runner:
|
||||
second = second_runner.run(capture_controller_state())
|
||||
|
||||
for key in (
|
||||
"loop",
|
||||
"controller",
|
||||
"send_lock",
|
||||
"send_wake",
|
||||
"background",
|
||||
"browser",
|
||||
):
|
||||
self.assertIsNot(first[key], second[key], key)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user