This commit is contained in:
Your Name
2026-07-23 17:56:25 +08:00
parent a05dae8412
commit 4970d8f8d3
4262 changed files with 735221 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
+239
View File
@@ -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()
+288
View File
@@ -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()
+154
View File
@@ -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()
+458
View File
@@ -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()
+165
View File
@@ -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()
+676
View File
@@ -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()
+495
View File
@@ -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()