1166 lines
49 KiB
Python
1166 lines
49 KiB
Python
"""Offline regression checks: python -m unittest discover -s src -v."""
|
|
|
|
import asyncio
|
|
import json
|
|
import sqlite3
|
|
import sys
|
|
import tempfile
|
|
import unittest
|
|
import zipfile
|
|
from pathlib import Path
|
|
from unittest.mock import AsyncMock, patch
|
|
|
|
from account_browser import (
|
|
AppLock,
|
|
BrowserManager,
|
|
browser_executable,
|
|
fingerprint_seed,
|
|
validate_config,
|
|
)
|
|
from account_engine import Engine
|
|
from account_session import Session, SessionError, validate_endpoint
|
|
from account_store import (
|
|
DEFAULT_RULE,
|
|
Store,
|
|
notice_kind,
|
|
notice_preview,
|
|
validate_rule,
|
|
)
|
|
from build_windows import extract_browser
|
|
|
|
|
|
def rule(**overrides):
|
|
return {
|
|
**DEFAULT_RULE,
|
|
"enabled": True,
|
|
"follow": True,
|
|
"dm": True,
|
|
"require_follow": False,
|
|
"text": "offline-only-test",
|
|
**overrides,
|
|
}
|
|
|
|
|
|
def notice(uid, nid="111", target="999"):
|
|
return {"user_id": uid, "nid_str": nid, "follow": {"from_user": [{"uid": target}]}}
|
|
|
|
|
|
def work_notice(uid, nid, target, work, kind="digg"):
|
|
return {
|
|
"user_id": uid,
|
|
"nid_str": nid,
|
|
kind: {
|
|
"from_user": [{"uid": target, "nickname": "完整昵称"}],
|
|
"aweme": {"aweme_id": work, "desc": "完整作品描述"},
|
|
},
|
|
}
|
|
|
|
|
|
class StoreTest(unittest.TestCase):
|
|
def setUp(self):
|
|
self.temp = tempfile.TemporaryDirectory()
|
|
self.store = Store(self.temp.name)
|
|
self.main = self.store.add("main", "大号 A", "101")
|
|
self.other = self.store.add("main", "大号 B", "102")
|
|
self.worker = self.store.add("worker", "小号", "201", self.main)
|
|
self.store.set_rule(self.main, rule())
|
|
|
|
def tearDown(self):
|
|
self.store.close()
|
|
self.temp.cleanup()
|
|
|
|
def test_notification_kinds_split_favorite_and_share(self):
|
|
for raw_kind, expected in (
|
|
("digg", "digg"),
|
|
("comment", "comment"),
|
|
("favorite", "favorite"),
|
|
("collect", "favorite"),
|
|
("share", "share"),
|
|
):
|
|
raw = work_notice("101", raw_kind, "902", "700", raw_kind)
|
|
self.assertEqual(notice_kind(raw), expected)
|
|
self.assertEqual(notice_preview(raw)["kind"], expected)
|
|
validate_rule(rule(kinds=["favorite", "share"]))
|
|
|
|
def test_defaults_validation_unique_identity(self):
|
|
with self.assertRaises(sqlite3.IntegrityError):
|
|
self.store.add("worker", "重复", "101")
|
|
with self.assertRaises(sqlite3.IntegrityError):
|
|
self.store.add("worker", "非法归属", "202", self.worker)
|
|
with self.assertRaises(ValueError):
|
|
validate_rule(rule(text=""))
|
|
with self.assertRaises(ValueError):
|
|
validate_rule(rule(interval=True))
|
|
with self.assertRaises(ValueError):
|
|
validate_rule(rule(cooldown=True))
|
|
normalized = validate_rule(
|
|
{k: v for k, v in rule(require_follow=True).items() if k != "cooldown"}
|
|
)
|
|
self.assertFalse(normalized["require_follow"])
|
|
self.assertEqual(normalized["cooldown"], 14400)
|
|
self.assertFalse(json.loads(self.store.account(self.other)["rule"])["enabled"])
|
|
with self.assertRaises(ValueError):
|
|
self.store.add("worker", "坏UID", "01")
|
|
|
|
def test_name_only_binding_duplicate_and_identity_lock(self):
|
|
a = self.store.add("main", "无需 UID")
|
|
b = self.store.add("main", "另一个")
|
|
self.assertIsNone(self.store.account(a)["uid"])
|
|
self.assertIsNone(self.store.account(b)["uid"])
|
|
profile = {
|
|
"uid": "301",
|
|
"nickname": "自动昵称",
|
|
"avatar": "https://example.test/avatar",
|
|
"follower_count": 123,
|
|
"following_count": 12,
|
|
"aweme_count": 4,
|
|
"total_favorited": 1000,
|
|
"cookie": "must-not-store",
|
|
"phone": "must-not-store",
|
|
}
|
|
self.store.bind_profile(a, profile)
|
|
saved = self.store.account(a)
|
|
self.assertEqual(saved["uid"], "301")
|
|
self.assertEqual(saved["nickname"], "自动昵称")
|
|
self.assertEqual(json.loads(saved["profile"])["follower_count"], 123)
|
|
stored_profile = json.loads(saved["profile"])
|
|
self.assertEqual(stored_profile["cookie"], "[凭据已隐藏]")
|
|
self.assertEqual(stored_profile["phone"], "must-not-store")
|
|
with self.assertRaises(ValueError):
|
|
self.store.bind_profile(b, profile)
|
|
self.assertIsNone(self.store.account(b)["uid"])
|
|
with self.assertRaises(ValueError):
|
|
self.store.bind_profile(a, {**profile, "uid": "302"})
|
|
self.assertEqual(self.store.account(a)["uid"], "301")
|
|
|
|
def test_unbound_worker_not_dispatched(self):
|
|
self.store.move(self.worker, None)
|
|
unbound = self.store.add("worker", "等待登录", owner=self.main)
|
|
self.store.ingest(self.main, notice("101"))
|
|
self.assertEqual(self.store.tasks(), [])
|
|
self.assertIsNone(self.store.claim(unbound))
|
|
self.store.bind_profile(unbound, {"uid": "301"})
|
|
self.store.dispatch(self.main)
|
|
self.assertEqual(len(self.store.tasks()), 2)
|
|
|
|
def test_stopped_worker_is_persistent_and_excluded_from_dispatch(self):
|
|
self.store.set_worker_stopped(self.worker, True)
|
|
self.store.ingest(self.main, notice("101"))
|
|
self.assertEqual(self.store.tasks(), [])
|
|
self.assertTrue(self.store.account(self.worker)["task_stopped"])
|
|
self.store.close()
|
|
self.store = Store(self.temp.name)
|
|
self.assertTrue(self.store.account(self.worker)["task_stopped"])
|
|
self.store.set_worker_stopped(self.worker, False)
|
|
self.assertEqual(len(self.store.tasks()), 2)
|
|
|
|
def test_upgrade_preserves_identity_tasks_and_creates_backup(self):
|
|
self.store.ingest(self.main, notice("101"))
|
|
expected = self.store.tasks()
|
|
db = self.store.db
|
|
db.execute("PRAGMA foreign_keys=OFF")
|
|
db.executescript("""
|
|
CREATE TABLE legacy_accounts(id TEXT PRIMARY KEY, role TEXT NOT NULL,
|
|
name TEXT NOT NULL, uid TEXT NOT NULL UNIQUE, owner TEXT REFERENCES accounts(id),
|
|
avatar TEXT NOT NULL DEFAULT '', nickname TEXT NOT NULL DEFAULT '',
|
|
rule TEXT NOT NULL, rr INTEGER NOT NULL DEFAULT 0);
|
|
INSERT INTO legacy_accounts SELECT id,role,name,uid,owner,avatar,nickname,rule,rr FROM accounts;
|
|
DROP TABLE accounts;
|
|
ALTER TABLE legacy_accounts RENAME TO accounts;
|
|
""")
|
|
self.store.close()
|
|
self.store = Store(self.temp.name)
|
|
self.assertEqual(self.store.account(self.main)["uid"], "101")
|
|
self.assertEqual(self.store.account(self.worker)["owner"], self.main)
|
|
self.assertEqual(self.store.tasks(), expected)
|
|
self.assertEqual(
|
|
self.store.db.execute("PRAGMA foreign_key_check").fetchall(), []
|
|
)
|
|
self.assertEqual(
|
|
len(list(Path(self.temp.name).glob("before-onboarding-*.sqlite3"))), 1
|
|
)
|
|
pending = self.store.add("main", "升级后新账号")
|
|
self.assertIsNone(self.store.account(pending)["uid"])
|
|
with self.assertRaises(sqlite3.IntegrityError):
|
|
self.store.add("worker", "不能归属小号", owner=self.worker)
|
|
|
|
def test_dedup_round_robin_and_parallel_batch(self):
|
|
second = self.store.add("worker", "第二小号", "202", self.main)
|
|
for nid, target in (
|
|
("111", "901"),
|
|
("111", "901"),
|
|
("112", "902"),
|
|
("113", "903"),
|
|
):
|
|
self.store.ingest(self.main, notice("101", nid, target))
|
|
tasks = sorted(self.store.tasks(), key=lambda t: t["id"])
|
|
self.assertEqual(len(tasks), 6)
|
|
self.assertEqual(
|
|
[tasks[i]["worker"] for i in (0, 2, 4)], [self.worker, second, self.worker]
|
|
)
|
|
batch = self.store.claim_batch(self.worker)
|
|
self.assertEqual({task["action"] for task in batch}, {"follow", "dm"})
|
|
self.assertEqual(
|
|
{
|
|
task["status"]
|
|
for task in self.store.tasks()
|
|
if task["id"] in {t["id"] for t in batch}
|
|
},
|
|
{"running"},
|
|
)
|
|
self.assertIsNone(self.store.claim(self.worker))
|
|
self.store.finish(batch[0]["id"], "failed", {})
|
|
self.store.finish(batch[1]["id"], "unknown", {})
|
|
next_batch = self.store.claim_batch(self.worker)
|
|
self.assertEqual({task["target"] for task in next_batch}, {"903"})
|
|
|
|
def test_selected_work_filter_covers_work_interactions_but_not_follow(self):
|
|
self.store.set_rule(
|
|
self.main,
|
|
rule(
|
|
dm=False,
|
|
cooldown=0,
|
|
kinds=["digg", "follow", "comment", "general_notice"],
|
|
work_mode="selected",
|
|
work_ids=["700"],
|
|
),
|
|
)
|
|
self.assertTrue(
|
|
self.store.ingest(self.main, work_notice("101", "1011", "901", "700"))
|
|
)
|
|
self.assertTrue(
|
|
self.store.ingest(self.main, work_notice("101", "1012", "902", "701"))
|
|
)
|
|
self.assertTrue(
|
|
self.store.ingest(
|
|
self.main, work_notice("101", "1013", "903", "700", "favorite")
|
|
)
|
|
)
|
|
self.assertTrue(self.store.ingest(self.main, notice("101", "1014", "904")))
|
|
self.assertEqual(len(self.store.tasks()), 3)
|
|
states = dict(
|
|
self.store.db.execute(
|
|
"SELECT nid,state FROM events WHERE source=?", (self.main,)
|
|
)
|
|
)
|
|
self.assertEqual(states["1012"], "ignored")
|
|
self.assertEqual(
|
|
{states[nid] for nid in ("1011", "1013", "1014")}, {"dispatched"}
|
|
)
|
|
|
|
def test_persistent_caches_merge_incrementally_filter_and_hide_credentials(self):
|
|
self.store.cache_works(
|
|
self.main,
|
|
[
|
|
{
|
|
"aweme_id": "700",
|
|
"create_time": 100,
|
|
"desc": "旧描述",
|
|
"statistics": {"digg_count": 1},
|
|
"cover": "https://example.invalid/cover",
|
|
"business": {"Authorization": "secret", "desc": "旧描述"},
|
|
}
|
|
],
|
|
)
|
|
self.store.cache_works(
|
|
self.main,
|
|
[
|
|
{
|
|
"aweme_id": "700",
|
|
"create_time": 100,
|
|
"desc": "新描述",
|
|
"statistics": {"digg_count": 2},
|
|
"cover": "https://example.invalid/cover",
|
|
"business": {"Authorization": "secret", "desc": "新描述"},
|
|
}
|
|
],
|
|
)
|
|
selected = work_notice("101", "3001", "901", "700")
|
|
other = work_notice("101", "3002", "902", "701")
|
|
followed = notice("101", "3003", "903")
|
|
selected["Authorization"] = "secret"
|
|
self.store.cache_notices(self.main, [selected, other, followed], "history")
|
|
self.store.cache_notices(self.main, [selected], "live")
|
|
self.store.set_work_filter(self.main, "selected", ["700"], 0)
|
|
|
|
self.store.close()
|
|
self.store = Store(self.temp.name)
|
|
|
|
works = self.store.cached_works(self.main)
|
|
self.assertEqual(len(works), 1)
|
|
self.assertEqual(works[0]["desc"], "新描述")
|
|
self.assertEqual(works[0]["statistics"]["digg_count"], 2)
|
|
work_business = json.loads(
|
|
self.store.db.execute(
|
|
"SELECT business FROM cached_works WHERE source=? AND aweme_id=?",
|
|
(self.main, "700"),
|
|
).fetchone()["business"]
|
|
)
|
|
self.assertEqual(work_business["Authorization"], "[凭据已隐藏]")
|
|
cached = {item["nid"]: item for item in self.store.cached_notices(self.main)}
|
|
self.assertEqual(set(cached), {"3001", "3003"})
|
|
self.assertEqual(cached["3001"]["origin"], "live")
|
|
business = self.store.cached_notice(self.main, "3001")
|
|
assert business is not None
|
|
self.assertEqual(business["Authorization"], "[凭据已隐藏]")
|
|
self.assertEqual(
|
|
json.loads(self.store.account(self.main)["rule"])["works_refresh_interval"],
|
|
0,
|
|
)
|
|
|
|
def test_live_tasks_preempt_history_and_promote_same_notification(self):
|
|
self.store.set_rule(self.main, rule(dm=False, cooldown=0))
|
|
self.store.ingest(self.main, notice("101", "2001", "901"), origin="history")
|
|
self.store.ingest(self.main, notice("101", "2002", "902"), origin="history")
|
|
self.store.ingest(self.main, notice("101", "2003", "903"), origin="live")
|
|
first = self.store.claim_batch(self.worker)
|
|
self.assertEqual({task["target"] for task in first}, {"903"})
|
|
for task in first:
|
|
self.store.finish(task["id"], "succeeded", {})
|
|
self.store.ingest(self.main, notice("101", "2004", "901"), origin="live")
|
|
old = self.store.db.execute(
|
|
"SELECT status FROM tasks WHERE event=(SELECT id FROM events WHERE nid='2001')"
|
|
).fetchall()
|
|
self.assertEqual({row["status"] for row in old}, {"cancelled"})
|
|
promoted = notice("101", "2005", "904")
|
|
self.store.ingest(self.main, promoted, origin="history")
|
|
self.store.ingest(self.main, promoted, origin="live")
|
|
event = self.store.db.execute(
|
|
"SELECT id,origin FROM events WHERE source=? AND nid='2005'", (self.main,)
|
|
).fetchone()
|
|
self.assertEqual(event["origin"], "live")
|
|
self.assertEqual(
|
|
{
|
|
row["priority"]
|
|
for row in self.store.db.execute(
|
|
"SELECT priority FROM tasks WHERE event=?", (event["id"],)
|
|
)
|
|
},
|
|
{100},
|
|
)
|
|
|
|
def test_group_cooldown_persists_is_shared_and_other_groups_are_independent(self):
|
|
second = self.store.add("worker", "第二小号", "202", self.main)
|
|
other_worker = self.store.add("worker", "其他组小号", "203", self.other)
|
|
self.store.set_rule(self.other, rule())
|
|
with patch("account_store.time.time", return_value=100):
|
|
self.store.ingest(self.main, notice("101", "111", "999"))
|
|
self.assertEqual(
|
|
self.store.db.execute("SELECT count(*) FROM cooldowns").fetchone()[0], 0
|
|
)
|
|
# A second notification before execution cannot duplicate the queued bundle.
|
|
self.store.ingest(self.main, notice("101", "112", "999"))
|
|
self.assertEqual(len(self.store.tasks()), 2)
|
|
batch = self.store.claim_batch(self.worker)
|
|
self.assertEqual(len(batch), 2)
|
|
self.assertEqual(
|
|
self.store.db.execute(
|
|
"SELECT last_at FROM cooldowns WHERE source=? AND target='999'",
|
|
(self.main,),
|
|
).fetchone()[0],
|
|
100,
|
|
)
|
|
for task in batch:
|
|
self.store.finish(task["id"], "succeeded", {})
|
|
# The same UID in a different main-account group is independent.
|
|
self.store.ingest(self.other, notice("102", "211", "999"))
|
|
self.assertEqual(
|
|
len([t for t in self.store.tasks() if t["worker"] == other_worker]), 2
|
|
)
|
|
self.store.close()
|
|
self.store = Store(self.temp.name)
|
|
with patch("account_store.time.time", return_value=200):
|
|
self.store.ingest(self.main, notice("101", "113", "999"))
|
|
self.assertEqual(
|
|
len([t for t in self.store.tasks() if t["source"] == self.main]), 2
|
|
)
|
|
with patch("account_store.time.time", return_value=14501):
|
|
self.store.ingest(self.main, notice("101", "114", "999"))
|
|
new_tasks = [t for t in self.store.tasks() if t["source"] == self.main]
|
|
self.assertEqual(len(new_tasks), 4)
|
|
self.assertEqual({t["worker"] for t in new_tasks[:2]}, {second})
|
|
|
|
def test_zero_cooldown_allows_later_event_after_first_bundle_finishes(self):
|
|
self.store.set_rule(self.main, rule(cooldown=0))
|
|
self.store.ingest(self.main, notice("101", "111", "999"))
|
|
batch = self.store.claim_batch(self.worker)
|
|
for task in batch:
|
|
self.store.finish(task["id"], "succeeded", {})
|
|
self.store.ingest(self.main, notice("101", "112", "999"))
|
|
self.assertEqual(len(self.store.tasks()), 4)
|
|
self.assertEqual(
|
|
self.store.db.execute("SELECT count(*) FROM cooldowns").fetchone()[0], 0
|
|
)
|
|
|
|
def test_move_cancels_only_pending_and_is_atomic(self):
|
|
self.store.ingest(self.main, notice("101"))
|
|
task = self.store.claim(self.worker)
|
|
assert task is not None
|
|
with self.assertRaises(ValueError):
|
|
self.store.move(self.worker, self.other)
|
|
self.assertEqual(self.store.account(self.worker)["owner"], self.main)
|
|
self.store.finish(task["id"], "unknown", {})
|
|
self.store.move(self.worker, self.other)
|
|
self.assertEqual(
|
|
{t["status"] for t in self.store.tasks()}, {"cancelled", "unknown"}
|
|
)
|
|
self.assertIsNone(self.store.claim(self.worker))
|
|
self.store.move(self.worker, None)
|
|
self.assertIsNone(self.store.account(self.worker)["owner"])
|
|
|
|
def test_restart_running_batch_becomes_unknown_never_resends(self):
|
|
self.store.ingest(self.main, notice("101"))
|
|
batch = self.store.claim_batch(self.worker)
|
|
self.assertEqual(len(batch), 2)
|
|
self.store.close()
|
|
self.store = Store(self.temp.name)
|
|
self.store.recover()
|
|
self.assertEqual({t["status"] for t in self.store.tasks()}, {"unknown"})
|
|
self.assertEqual(self.store.claim_batch(self.worker), [])
|
|
self.assertEqual(self.store.claim_batch(self.worker), [])
|
|
|
|
def test_parallel_results_are_independent_and_never_resend(self):
|
|
self.store.ingest(self.main, notice("101"))
|
|
batch = self.store.claim_batch(self.worker)
|
|
self.assertEqual(len(batch), 2)
|
|
self.store.finish(batch[0]["id"], "failed", {})
|
|
self.store.finish(batch[1]["id"], "unknown", {})
|
|
self.assertEqual(self.store.claim_batch(self.worker), [])
|
|
self.assertEqual(
|
|
{t["status"] for t in self.store.tasks()}, {"failed", "unknown"}
|
|
)
|
|
|
|
def test_disabled_and_cross_group_identity(self):
|
|
with self.assertRaises(ValueError):
|
|
self.store.ingest(self.main, notice("102"))
|
|
self.store.set_rule(self.main, rule(enabled=False))
|
|
self.store.ingest(self.main, notice("101"))
|
|
self.assertEqual(self.store.tasks(), [])
|
|
self.store.set_rule(self.main, rule())
|
|
self.store.dispatch(self.main)
|
|
self.assertEqual(self.store.tasks(), [])
|
|
self.store.ingest(self.main, notice("101", "112"))
|
|
self.store.set_rule(self.main, rule(enabled=False))
|
|
self.assertIsNone(self.store.claim(self.worker))
|
|
|
|
def test_missing_detail_backoff_persists_and_does_not_starve_new_ids(self):
|
|
self.store.record_push(self.main, ["111"])
|
|
with patch("account_store.time.time", return_value=100):
|
|
self.assertEqual(self.store.due_details(self.main), ["111"])
|
|
self.store.defer_details(self.main, ["111"])
|
|
self.store.close()
|
|
self.store = Store(self.temp.name)
|
|
with patch("account_store.time.time", return_value=110):
|
|
self.assertEqual(self.store.due_details(self.main), [])
|
|
self.store.record_push(self.main, ["111", "112"])
|
|
self.assertEqual(self.store.due_details(self.main), ["112"])
|
|
self.store.ingest(self.main, notice("101", "112"))
|
|
with patch("account_store.time.time", return_value=130):
|
|
self.assertEqual(self.store.due_details(self.main), ["111"])
|
|
self.store.defer_details(self.main, ["111"])
|
|
row = self.store.db.execute(
|
|
"SELECT retry_count,next_retry_at FROM inbox WHERE nid='111'"
|
|
).fetchone()
|
|
self.assertEqual(tuple(row), (2, 190))
|
|
self.assertEqual(self.store.pending_details_count(self.main), 1)
|
|
|
|
def test_missing_detail_retry_is_capped_and_never_deletes_id(self):
|
|
self.store.record_push(self.main, ["111"])
|
|
with patch("account_store.time.time", return_value=100):
|
|
for _ in range(20):
|
|
self.store.defer_details(self.main, ["111"])
|
|
row = self.store.db.execute(
|
|
"SELECT state,retry_count,next_retry_at FROM inbox WHERE nid='111'"
|
|
).fetchone()
|
|
self.assertEqual(tuple(row), ("pending", 20, 400))
|
|
|
|
def test_rule_migration_adds_default_cooldown_and_disables_old_dependency(self):
|
|
legacy = rule(require_follow=True)
|
|
legacy.pop("cooldown")
|
|
self.store.db.execute(
|
|
"UPDATE accounts SET rule=? WHERE id=?", (json.dumps(legacy), self.main)
|
|
)
|
|
self.store.close()
|
|
self.store = Store(self.temp.name)
|
|
migrated = json.loads(self.store.account(self.main)["rule"])
|
|
self.assertEqual(migrated["cooldown"], 14400)
|
|
self.assertFalse(migrated["require_follow"])
|
|
self.assertIsNotNone(
|
|
self.store.db.execute(
|
|
"SELECT 1 FROM sqlite_master WHERE type='table' AND name='cooldowns'"
|
|
).fetchone()
|
|
)
|
|
|
|
def test_legacy_inbox_migration_preserves_pending_ids(self):
|
|
self.store.record_push(self.main, ["111"])
|
|
self.store.db.executescript("""
|
|
CREATE TABLE legacy_inbox(source TEXT REFERENCES accounts(id), nid TEXT,
|
|
state TEXT NOT NULL DEFAULT 'pending', PRIMARY KEY(source,nid));
|
|
INSERT INTO legacy_inbox SELECT source,nid,state FROM inbox;
|
|
DROP TABLE inbox;
|
|
ALTER TABLE legacy_inbox RENAME TO inbox;
|
|
""")
|
|
self.store.close()
|
|
self.store = Store(self.temp.name)
|
|
row = self.store.db.execute(
|
|
"SELECT state,retry_count,next_retry_at FROM inbox WHERE nid='111'"
|
|
).fetchone()
|
|
self.assertEqual(tuple(row), ("pending", 0, 0))
|
|
self.assertEqual(self.store.due_details(self.main), ["111"])
|
|
|
|
def test_waiting_without_workers_and_inbox_durability(self):
|
|
self.store.move(self.worker, None)
|
|
self.store.record_push(self.main, ["111", "111"])
|
|
self.store.ingest(self.main, notice("101"))
|
|
self.assertEqual(self.store.tasks(), [])
|
|
self.store.move(self.worker, self.main)
|
|
self.store.dispatch(self.main)
|
|
self.assertEqual(len(self.store.tasks()), 2)
|
|
self.assertEqual(
|
|
self.store.db.execute("SELECT state FROM inbox").fetchone()[0], "done"
|
|
)
|
|
|
|
def test_other_connection_cannot_duplicate_claim(self):
|
|
self.store.ingest(self.main, notice("101"))
|
|
other = Store(self.temp.name)
|
|
try:
|
|
first = self.store.claim(self.worker)
|
|
self.assertIsNone(other.claim(self.worker))
|
|
self.assertIsNotNone(first)
|
|
finally:
|
|
other.close()
|
|
|
|
|
|
class BrowserSafetyTest(unittest.TestCase):
|
|
def test_endpoint_and_arguments_fail_closed(self):
|
|
self.assertEqual(
|
|
validate_endpoint("http://127.0.0.1:9222"), "http://127.0.0.1:9222"
|
|
)
|
|
for endpoint in (
|
|
"http://10.1.1.1:9222",
|
|
"http://localhost",
|
|
"http://user@localhost:12",
|
|
"https://localhost:123",
|
|
"http://localhost:1/path",
|
|
):
|
|
with self.assertRaises(ValueError):
|
|
validate_endpoint(endpoint)
|
|
for arg in (
|
|
"--user-data-dir=/tmp",
|
|
"--remote-debugging-port=9222",
|
|
"--no-sandbox",
|
|
"--disable-web-security",
|
|
"--proxy-server=evil",
|
|
"--load-extension=x",
|
|
"--headless",
|
|
"--enable-automation",
|
|
"--fingerprint=123",
|
|
"--fingerprint-platform=linux",
|
|
):
|
|
with self.assertRaises(ValueError):
|
|
validate_config({"extra_args": [arg]})
|
|
validate_config({"extra_args": ["--lang=zh-CN", "--window-size=1280,900"]})
|
|
with self.assertRaises(ValueError):
|
|
validate_config({"executable_path": "/definitely/not/a/browser"})
|
|
|
|
def test_fingerprint_seed_stable_and_no_system_fallback(self):
|
|
self.assertEqual(fingerprint_seed("a" * 32), fingerprint_seed("a" * 32))
|
|
self.assertNotEqual(fingerprint_seed("a" * 32), fingerprint_seed("b" * 32))
|
|
self.assertTrue(0 < fingerprint_seed("a" * 32) <= 0x7FFFFFFF)
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
root = Path(directory)
|
|
system = root / "system/Google/Chrome/Application/chrome.exe"
|
|
system.parent.mkdir(parents=True)
|
|
system.touch()
|
|
with (
|
|
patch.dict(
|
|
"os.environ", {"PROGRAMFILES": str(root / "system")}, clear=True
|
|
),
|
|
patch.object(sys, "frozen", True, create=True),
|
|
patch.object(sys, "executable", str(root / "app.exe")),
|
|
):
|
|
with self.assertRaises(ValueError):
|
|
browser_executable({})
|
|
bundled = root / "fingerprint-browser/chrome.exe"
|
|
bundled.parent.mkdir()
|
|
bundled.touch()
|
|
self.assertEqual(browser_executable({}), bundled.resolve())
|
|
|
|
def test_fingerprint_archive_flattening_and_path_rejection(self):
|
|
with tempfile.TemporaryDirectory() as directory:
|
|
root = Path(directory)
|
|
archive = root / "browser.zip"
|
|
with zipfile.ZipFile(archive, "w") as zipped:
|
|
zipped.writestr("release/chrome.exe", b"offline-test")
|
|
zipped.writestr("release/resources/data", b"data")
|
|
extract_browser(archive, root / "out", "release")
|
|
self.assertEqual((root / "out/chrome.exe").read_bytes(), b"offline-test")
|
|
for name in (
|
|
"release/../escape",
|
|
"other/file",
|
|
"release/file:stream",
|
|
"release/bad\\name",
|
|
):
|
|
with zipfile.ZipFile(archive, "w") as zipped:
|
|
entry = zipfile.ZipInfo(name)
|
|
entry.filename = (
|
|
name # Preserve malformed wire separators on Windows too.
|
|
)
|
|
zipped.writestr(entry, b"x")
|
|
with self.assertRaises(RuntimeError):
|
|
extract_browser(archive, root / "out", "release")
|
|
self.assertFalse((root / "escape").exists())
|
|
|
|
def test_directory_and_instance_lock(self):
|
|
with tempfile.TemporaryDirectory() as root:
|
|
lock = AppLock(Path(root))
|
|
try:
|
|
with self.assertRaises(RuntimeError):
|
|
AppLock(Path(root))
|
|
finally:
|
|
lock.close()
|
|
manager = BrowserManager(root)
|
|
with self.assertRaises(ValueError):
|
|
manager.account_dir("../other")
|
|
ident = "a" * 32
|
|
path = manager.account_dir(ident)
|
|
path.mkdir()
|
|
(path / "local-only").write_text("test")
|
|
with (
|
|
patch.object(manager, "live", return_value=[object()]),
|
|
self.assertRaises(ValueError),
|
|
):
|
|
manager.erase(ident)
|
|
self.assertTrue(path.exists())
|
|
with patch.object(manager, "live", return_value=[]):
|
|
manager.erase(ident)
|
|
self.assertFalse(path.exists())
|
|
|
|
|
|
class EngineTest(unittest.IsolatedAsyncioTestCase):
|
|
async def asyncSetUp(self):
|
|
self.temp = tempfile.TemporaryDirectory()
|
|
self.engine = Engine(self.temp.name, None)
|
|
self.main = self.engine.store.add("main", "主", "101")
|
|
self.worker = self.engine.store.add("worker", "子", "201", self.main)
|
|
self.engine.store.set_rule(self.main, rule())
|
|
|
|
async def asyncTearDown(self):
|
|
await self.engine.shutdown()
|
|
self.temp.cleanup()
|
|
|
|
async def test_identity_mismatch_blocks_session(self):
|
|
session = Session(None, "http://127.0.0.1:9222", "201")
|
|
session.profile = AsyncMock(return_value={"uid": "301"})
|
|
with self.assertRaises(SessionError):
|
|
await session.identity()
|
|
|
|
async def test_legacy_browser_settings_not_silently_reused(self):
|
|
self.engine.store.set_setting("browser", {"executable_path": "old-chrome.exe"})
|
|
self.engine.store.set_setting(
|
|
"browser:" + self.worker, {"executable_path": "old-account-chrome.exe"}
|
|
)
|
|
self.assertEqual(self.engine.config(self.worker), {})
|
|
self.engine.store.set_setting(
|
|
"fingerprint_browser", {"extra_args": ["--lang=zh-CN"]}
|
|
)
|
|
self.assertEqual(
|
|
self.engine.config(self.worker), {"extra_args": ["--lang=zh-CN"]}
|
|
)
|
|
self.assertEqual(
|
|
self.engine.store.setting("browser"), {"executable_path": "old-chrome.exe"}
|
|
)
|
|
|
|
async def test_patchright_world_selection_explicit(self):
|
|
session = Session(None, "http://127.0.0.1:9222", "101")
|
|
page = AsyncMock()
|
|
page.is_closed = lambda: False
|
|
session.page = page
|
|
await session.evaluate("1")
|
|
page.evaluate.assert_awaited_with("1", None, isolated_context=True)
|
|
await session.evaluate("window.sdk", main_world=True)
|
|
page.evaluate.assert_awaited_with("window.sdk", None, isolated_context=False)
|
|
|
|
async def test_unbound_session_can_only_read_profile(self):
|
|
session = Session(None, "http://127.0.0.1:9222", None)
|
|
session.evaluate = AsyncMock(
|
|
return_value={
|
|
"status": 200,
|
|
"body": json.dumps(
|
|
{
|
|
"status_code": 0,
|
|
"user": {
|
|
"uid": "301",
|
|
"nickname": "读取昵称",
|
|
"follower_count": 42,
|
|
},
|
|
}
|
|
),
|
|
}
|
|
)
|
|
self.assertEqual((await session.profile())["follower_count"], 42)
|
|
session.evaluate.reset_mock()
|
|
for action in (
|
|
session.identity(),
|
|
session.follow("999"),
|
|
session.im("send", "999", "test", True),
|
|
session.install(),
|
|
):
|
|
with self.assertRaises(SessionError):
|
|
await action
|
|
session.evaluate.assert_not_called()
|
|
|
|
async def test_name_only_enrollment_waits_then_backfills_and_starts(self):
|
|
ident = self.engine.store.add("main", "只填名称")
|
|
session = AsyncMock()
|
|
session.endpoint = "http://127.0.0.1:9555"
|
|
session.connect.side_effect = [
|
|
SessionError("等待登录"),
|
|
{"uid": "301", "nickname": "回填", "follower_count": 42},
|
|
]
|
|
self.engine.sessions[ident] = session
|
|
self.engine.manager.start = AsyncMock(return_value=session.endpoint)
|
|
with (
|
|
patch.object(self.engine.manager, "live", return_value=[object()]),
|
|
patch("account_engine.asyncio.sleep", new=AsyncMock()),
|
|
):
|
|
await self.engine.enroll(ident)
|
|
saved = self.engine.store.account(ident)
|
|
self.assertEqual(saved["uid"], "301")
|
|
self.assertEqual(json.loads(saved["profile"])["follower_count"], 42)
|
|
self.assertEqual(session.uid, "301")
|
|
self.assertIn(ident, self.engine.desired)
|
|
self.assertEqual(self.engine.store.tasks(), [])
|
|
session.im.assert_not_called()
|
|
session.follow.assert_not_called()
|
|
|
|
async def test_unbound_group_cannot_start_and_add_starts_login_only(self):
|
|
with patch.object(self.engine, "begin_login") as login:
|
|
ident = await self.engine.command(
|
|
"add", {"role": "main", "name": "界面名称"}
|
|
)
|
|
login.assert_called_once_with(ident)
|
|
with self.assertRaises(ValueError):
|
|
await self.engine.start([ident])
|
|
self.assertEqual(self.engine.desired, set())
|
|
|
|
async def test_auto_start_skips_unbound_and_stopped_workers(self):
|
|
unbound = self.engine.store.add("worker", "未登录", owner=self.main)
|
|
stopped = self.engine.store.add("worker", "手动停止", "202", self.main)
|
|
self.engine.store.set_worker_stopped(stopped, True)
|
|
await self.engine.start([self.main], automatic=True)
|
|
self.assertIn(self.main, self.engine.desired)
|
|
self.assertIn(self.worker, self.engine.desired)
|
|
self.assertNotIn(unbound, self.engine.desired)
|
|
self.assertNotIn(stopped, self.engine.desired)
|
|
|
|
async def test_worker_start_stop_commands_persist_and_control_runtime(self):
|
|
self.engine.desired.add(self.main)
|
|
await self.engine.command("worker_stop", {"id": self.worker})
|
|
self.assertTrue(self.engine.store.account(self.worker)["task_stopped"])
|
|
self.assertNotIn(self.worker, self.engine.desired)
|
|
with patch.object(self.engine, "_activate") as activate:
|
|
await self.engine.command("worker_start", {"id": self.worker})
|
|
self.assertFalse(self.engine.store.account(self.worker)["task_stopped"])
|
|
activate.assert_called_once_with(self.worker)
|
|
|
|
async def test_im_result_classification_keeps_business_data_and_hides_credentials(
|
|
self,
|
|
):
|
|
session = AsyncMock()
|
|
task = {"params": json.dumps(rule()), "action": "dm", "target": "999"}
|
|
session.im.return_value = {
|
|
"success": True,
|
|
"message": {
|
|
"client_id": "c1",
|
|
"server_id": "s1",
|
|
"content": "完整私信正文",
|
|
},
|
|
"token": "AUTH-CREDENTIAL",
|
|
}
|
|
status, result = await self.engine.perform(session, task)
|
|
self.assertEqual(status, "succeeded")
|
|
self.assertIn("完整私信正文", json.dumps(result, ensure_ascii=False))
|
|
self.assertNotIn("AUTH-CREDENTIAL", json.dumps(result))
|
|
session.im.return_value = {"error": "SDK_REQUEST_FAILED"}
|
|
self.assertEqual((await self.engine.perform(session, task))[0], "unknown")
|
|
session.im.return_value = {"error": "LOGIN_REQUIRED"}
|
|
self.assertEqual((await self.engine.perform(session, task))[0], "failed")
|
|
|
|
async def test_pause_main_prevents_worker_claim(self):
|
|
self.engine.store.ingest(self.main, notice("101"))
|
|
self.engine.desired.add(self.worker)
|
|
self.engine.open = AsyncMock(side_effect=AssertionError("must not open/send"))
|
|
task = asyncio.create_task(self.engine.worker_loop(self.worker))
|
|
await asyncio.sleep(0.05)
|
|
self.engine.desired.clear()
|
|
await task
|
|
self.engine.open.assert_not_called()
|
|
self.assertTrue(
|
|
all(t["status"] == "pending" for t in self.engine.store.tasks())
|
|
)
|
|
|
|
async def test_detail_null_empty_and_partial_are_valid(self):
|
|
session = Session(None, "http://127.0.0.1:9222", "101")
|
|
session.identity = AsyncMock()
|
|
for rows in (None, [], [notice("101", "112")]):
|
|
session.evaluate = AsyncMock(
|
|
return_value={
|
|
"status": 200,
|
|
"body": json.dumps({"status_code": 0, "notice_list_v2": rows}),
|
|
}
|
|
)
|
|
self.assertEqual(await session.details(["111", "112"]), rows or [])
|
|
|
|
async def test_detail_real_errors_and_foreign_identity_still_fail_closed(self):
|
|
session = Session(None, "http://127.0.0.1:9222", "101")
|
|
session.identity = AsyncMock()
|
|
for body in (
|
|
{"status_code": 13, "notice_list_v2": None},
|
|
{"status_code": 0},
|
|
{"status_code": 0, "notice_list_v2": {}},
|
|
{"status_code": 0, "notice_list_v2": [notice("102", "111")]},
|
|
{"status_code": 0, "notice_list_v2": [notice("101", "999")]},
|
|
):
|
|
session.evaluate = AsyncMock(
|
|
return_value={"status": 200, "body": json.dumps(body)}
|
|
)
|
|
with self.assertRaises(SessionError):
|
|
await session.details(["111"])
|
|
|
|
async def test_partial_details_ingest_available_and_defer_only_missing(self):
|
|
self.engine.store.record_push(self.main, ["111", "112"])
|
|
session = AsyncMock()
|
|
session.details.return_value = [notice("101", "112")]
|
|
self.assertEqual(await self.engine.collect_details(self.main, session), 1)
|
|
states = {
|
|
r["nid"]: (r["state"], r["retry_count"])
|
|
for r in self.engine.store.db.execute("SELECT * FROM inbox")
|
|
}
|
|
self.assertEqual(states, {"111": ("pending", 1), "112": ("done", 0)})
|
|
self.assertEqual(len(self.engine.store.tasks()), 2)
|
|
session.details.reset_mock()
|
|
await self.engine.collect_details(self.main, session)
|
|
session.details.assert_not_called()
|
|
|
|
async def test_unavailable_detail_keeps_listener_ready_for_new_pushes(self):
|
|
self.engine.store.record_push(self.main, ["111"])
|
|
self.engine.desired.add(self.main)
|
|
session = AsyncMock()
|
|
session.details.side_effect = [[], [notice("101", "112")]]
|
|
calls = 0
|
|
|
|
async def wait():
|
|
nonlocal calls
|
|
calls += 1
|
|
self.assertIn(self.main, self.engine.ready)
|
|
if calls == 1:
|
|
return [
|
|
{
|
|
"kind": "push",
|
|
"service": 20313,
|
|
"payload": json.dumps(
|
|
{
|
|
"notices": [
|
|
{"notice_id_str": "112", "effect_groups": [960]}
|
|
]
|
|
}
|
|
),
|
|
}
|
|
]
|
|
self.engine.desired.discard(self.main)
|
|
return []
|
|
|
|
session.wait.side_effect = wait
|
|
self.engine.open = AsyncMock(return_value=session)
|
|
await asyncio.wait_for(self.engine.main_loop(self.main), 2)
|
|
self.engine.open.assert_awaited_once()
|
|
session.install.assert_awaited_once()
|
|
session.uninstall.assert_awaited_once()
|
|
self.assertEqual(session.details.await_count, 2)
|
|
self.assertEqual(self.engine.store.pending_details_count(self.main), 1)
|
|
self.assertEqual(len(self.engine.store.tasks()), 2)
|
|
self.assertIn("不阻塞新通知", self.engine.states[self.main])
|
|
|
|
async def test_history_and_works_paginate_read_only_with_exact_ids(self):
|
|
session = Session(None, "http://127.0.0.1:9222", "101")
|
|
first = {
|
|
"status": 200,
|
|
"body": json.dumps(
|
|
{
|
|
"status_code": 0,
|
|
"notice_list_v2": [notice("101", "9223372036854775701")],
|
|
"has_more": 1,
|
|
"min_time": 100,
|
|
"max_time": 200,
|
|
}
|
|
),
|
|
}
|
|
second = {
|
|
"status": 200,
|
|
"body": json.dumps(
|
|
{
|
|
"status_code": 0,
|
|
"notice_list_v2": [notice("101", "9223372036854775702")],
|
|
"has_more": 0,
|
|
"min_time": 90,
|
|
"max_time": 190,
|
|
}
|
|
),
|
|
}
|
|
session.json = AsyncMock(side_effect=[first, second])
|
|
history = await session.history_notices()
|
|
self.assertEqual(
|
|
[row["nid_str"] for row in history],
|
|
["9223372036854775701", "9223372036854775702"],
|
|
)
|
|
self.assertTrue(
|
|
all(
|
|
"is_mark_read" in call.args[0] and "700" in call.args[0]
|
|
for call in session.json.await_args_list
|
|
)
|
|
)
|
|
session.identity = AsyncMock(return_value={"uid": "101", "sec_uid": "SEC-UID"})
|
|
session.json = AsyncMock(
|
|
side_effect=[
|
|
{
|
|
"status": 200,
|
|
"body": json.dumps(
|
|
{
|
|
"status_code": 0,
|
|
"aweme_list": [
|
|
{
|
|
"aweme_id": "9223372036854775703",
|
|
"desc": "完整作品",
|
|
"create_time": 123,
|
|
"author": {"uid": "101"},
|
|
"statistics": {"digg_count": 9},
|
|
"video": {
|
|
"cover": {
|
|
"url_list": [
|
|
"https://example.invalid/full.jpg"
|
|
]
|
|
}
|
|
},
|
|
}
|
|
],
|
|
"has_more": 1,
|
|
"max_cursor": 321,
|
|
}
|
|
),
|
|
},
|
|
{
|
|
"status": 200,
|
|
"body": json.dumps(
|
|
{
|
|
"status_code": 0,
|
|
"aweme_list": [],
|
|
"has_more": 0,
|
|
"max_cursor": 0,
|
|
}
|
|
),
|
|
},
|
|
]
|
|
)
|
|
works = await session.works(True)
|
|
self.assertEqual(works[0]["aweme_id"], "9223372036854775703")
|
|
self.assertEqual(works[0]["desc"], "完整作品")
|
|
self.assertEqual(works[0]["cover"], "https://example.invalid/full.jpg")
|
|
self.assertEqual(session.json.await_count, 2)
|
|
|
|
async def test_history_and_works_reject_cursor_loops_and_foreign_identity(self):
|
|
session = Session(None, "http://127.0.0.1:9222", "101")
|
|
repeating = {
|
|
"status": 200,
|
|
"body": json.dumps(
|
|
{
|
|
"status_code": 0,
|
|
"notice_list_v2": [],
|
|
"has_more": 1,
|
|
"min_time": 0,
|
|
"max_time": 0,
|
|
}
|
|
),
|
|
}
|
|
session.json = AsyncMock(return_value=repeating)
|
|
with self.assertRaisesRegex(SessionError, "游标未推进"):
|
|
await session.history_notices()
|
|
session.identity = AsyncMock(return_value={"uid": "101", "sec_uid": "SEC-UID"})
|
|
session.json = AsyncMock(
|
|
return_value={
|
|
"status": 200,
|
|
"body": json.dumps(
|
|
{
|
|
"status_code": 0,
|
|
"aweme_list": [{"aweme_id": "700", "author": {"uid": "other"}}],
|
|
"has_more": 0,
|
|
}
|
|
),
|
|
}
|
|
)
|
|
with self.assertRaisesRegex(SessionError, "身份不符"):
|
|
await session.works(False)
|
|
|
|
async def test_detail_64bit_ids_preserved(self):
|
|
session = Session(None, "http://127.0.0.1:9222", "101")
|
|
session.identity = AsyncMock(return_value={"uid": "101"})
|
|
nid = "9223372036854775701"
|
|
session.evaluate = AsyncMock(
|
|
return_value={
|
|
"status": 200,
|
|
"body": json.dumps(
|
|
{
|
|
"status_code": 0,
|
|
"notice_list_v2": [{"nid": int(nid), "user_id": "101"}],
|
|
}
|
|
),
|
|
}
|
|
)
|
|
result = await session.details([nid])
|
|
self.assertEqual(str(result[0]["nid"]), nid)
|
|
|
|
async def test_history_preview_confirmation_and_work_filter_commands(self):
|
|
session = AsyncMock()
|
|
history = [work_notice("101", "3001", "901", "700")]
|
|
works = [
|
|
{
|
|
"aweme_id": "700",
|
|
"desc": "完整作品描述",
|
|
"create_time": 123,
|
|
"statistics": {"digg_count": 5},
|
|
"cover": "https://example.invalid/cover.jpg",
|
|
"business": {"aweme_id": "700", "desc": "完整作品描述"},
|
|
}
|
|
]
|
|
session.history_notices.return_value = history
|
|
session.works.return_value = works
|
|
self.engine.store.set_rule(
|
|
self.main,
|
|
rule(kinds=["digg", "follow", "comment", "general_notice"]),
|
|
)
|
|
self.engine.open = AsyncMock(return_value=session)
|
|
preview = await self.engine.command("history_fetch", {"id": self.main})
|
|
self.assertEqual(preview["items"][0]["nid"], "3001")
|
|
self.assertEqual(preview["items"][0]["work_id"], "700")
|
|
self.assertEqual(preview["items"][0]["actor_names"], ["完整昵称"])
|
|
self.assertEqual(self.engine.store.tasks(), [])
|
|
with self.assertRaises(ValueError):
|
|
await self.engine.command(
|
|
"history_enqueue", {"id": self.main, "ids": ["not-loaded"]}
|
|
)
|
|
result = await self.engine.command(
|
|
"history_enqueue", {"id": self.main, "ids": ["3001"]}
|
|
)
|
|
self.assertEqual(result, {"selected": 1, "events": 1, "tasks": 2})
|
|
fetched = await self.engine.command("works_open", {"id": self.main})
|
|
self.assertEqual(fetched["items"][0]["aweme_id"], "700")
|
|
self.assertEqual(fetched["refresh_interval"], 3600)
|
|
saved = await self.engine.command(
|
|
"work_filter",
|
|
{
|
|
"id": self.main,
|
|
"mode": "selected",
|
|
"ids": ["700"],
|
|
"refresh_interval": 120,
|
|
},
|
|
)
|
|
self.assertEqual(
|
|
saved, {"mode": "selected", "count": 1, "refresh_interval": 120}
|
|
)
|
|
rule_value = json.loads(self.engine.store.account(self.main)["rule"])
|
|
self.assertEqual(rule_value["work_mode"], "selected")
|
|
self.assertEqual(rule_value["work_ids"], ["700"])
|
|
self.assertEqual(rule_value["works_refresh_interval"], 120)
|
|
session.works.assert_awaited_once_with(True, None)
|
|
session.history_notices.assert_awaited_once_with(set())
|
|
|
|
async def test_move_during_identity_check_does_not_cross_groups(self):
|
|
other = self.engine.store.add("main", "另一组", "102")
|
|
self.engine.store.set_rule(other, rule())
|
|
self.engine.store.ingest(self.main, notice("101"))
|
|
self.engine.desired.update((self.main, self.worker))
|
|
self.engine.ready.add(self.main)
|
|
main = AsyncMock()
|
|
worker = AsyncMock()
|
|
worker.browser.is_connected = lambda: True
|
|
worker.page.is_closed = lambda: False
|
|
self.engine.sessions[self.main] = main
|
|
self.engine.open = AsyncMock(return_value=worker)
|
|
switched = asyncio.Event()
|
|
|
|
async def switch_during_check():
|
|
self.engine.store.move(self.worker, other)
|
|
self.engine.store.ingest(other, notice("102", "112"))
|
|
switched.set()
|
|
return {"uid": "201"}
|
|
|
|
worker.identity.side_effect = switch_during_check
|
|
runner = asyncio.create_task(self.engine.worker_loop(self.worker))
|
|
self.engine.runners[self.worker] = runner
|
|
await asyncio.wait_for(switched.wait(), 2)
|
|
await asyncio.sleep(0.05)
|
|
await self.engine.stop()
|
|
worker.follow.assert_not_called()
|
|
worker.im.assert_not_called()
|
|
self.assertTrue(
|
|
all(
|
|
t["status"] == "pending"
|
|
for t in self.engine.store.tasks()
|
|
if t["source"] == other
|
|
)
|
|
)
|
|
|
|
async def test_stop_waits_inflight_and_persists(self):
|
|
self.engine.store.ingest(self.main, notice("101"))
|
|
self.engine.desired.update((self.main, self.worker))
|
|
self.engine.ready.add(self.main)
|
|
main = AsyncMock()
|
|
worker = AsyncMock()
|
|
worker.browser.is_connected = lambda: True
|
|
worker.page.is_closed = lambda: False
|
|
self.engine.sessions[self.main] = main
|
|
self.engine.open = AsyncMock(return_value=worker)
|
|
both_entered = asyncio.Event()
|
|
release = asyncio.Event()
|
|
entered = 0
|
|
|
|
async def action(kind, *args):
|
|
nonlocal entered
|
|
entered += 1
|
|
if entered == 2:
|
|
both_entered.set()
|
|
await release.wait()
|
|
return (
|
|
{"status": "succeeded"}
|
|
if kind == "follow"
|
|
else {
|
|
"success": True,
|
|
"message": {"client_id": "client", "server_id": "server"},
|
|
}
|
|
)
|
|
|
|
async def follow_action(*args, **kwargs):
|
|
return await action("follow", *args)
|
|
|
|
async def dm_action(*args, **kwargs):
|
|
return await action("dm", *args)
|
|
|
|
worker.follow.side_effect = follow_action
|
|
worker.im.side_effect = dm_action
|
|
runner = asyncio.create_task(self.engine.worker_loop(self.worker))
|
|
self.engine.runners[self.worker] = runner
|
|
await asyncio.wait_for(both_entered.wait(), 2)
|
|
stopping = asyncio.create_task(self.engine.stop())
|
|
await asyncio.sleep(0.02)
|
|
self.assertFalse(stopping.done())
|
|
release.set()
|
|
await asyncio.wait_for(stopping, 2)
|
|
tasks = self.engine.store.tasks()
|
|
self.assertEqual({t["status"] for t in tasks}, {"succeeded"})
|
|
self.assertEqual(worker.follow.await_count, 1)
|
|
self.assertEqual(worker.im.await_count, 1)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|