fix: finish shared registration flow hardening
This commit is contained in:
@@ -1,15 +1,30 @@
|
||||
import sys
|
||||
import types
|
||||
import unittest
|
||||
from unittest.mock import patch
|
||||
|
||||
# Keep this unit test independent from optional browser/network dependencies.
|
||||
drission = types.ModuleType("DrissionPage")
|
||||
drission.Chromium = type("Chromium", (), {})
|
||||
drission.ChromiumOptions = type("ChromiumOptions", (), {})
|
||||
drission_errors = types.ModuleType("DrissionPage.errors")
|
||||
drission_errors.PageDisconnectedError = type("PageDisconnectedError", (Exception,), {})
|
||||
curl_cffi = types.ModuleType("curl_cffi")
|
||||
curl_cffi.requests = types.SimpleNamespace()
|
||||
sys.modules.setdefault("DrissionPage", drission)
|
||||
sys.modules.setdefault("DrissionPage.errors", drission_errors)
|
||||
sys.modules.setdefault("curl_cffi", curl_cffi)
|
||||
|
||||
import grok_register_ttk as app
|
||||
|
||||
|
||||
class DummyResponse:
|
||||
def __init__(self, payload=None, status_code=200, reason=""):
|
||||
def __init__(self, payload=None, status_code=200, reason="", headers=None, text=""):
|
||||
self._payload = payload or {}
|
||||
self.status_code = status_code
|
||||
self.reason = reason
|
||||
self.text = ""
|
||||
self.headers = headers or {}
|
||||
self.text = text
|
||||
|
||||
def raise_for_status(self):
|
||||
if self.status_code >= 400:
|
||||
@@ -26,12 +41,17 @@ class Grok2ApiRemotePoolTests(unittest.TestCase):
|
||||
def tearDown(self):
|
||||
app.config = self.original_config
|
||||
|
||||
def test_remote_pool_falls_back_to_admin_api_prefix_when_root_tokens_add_is_404(self):
|
||||
def _configure(self, **overrides):
|
||||
app.config.update({
|
||||
"grok2api_remote_base": "https://grok.example.com",
|
||||
"grok2api_remote_app_key": "app-secret",
|
||||
"grok2api_pool_name": "ssoBasic",
|
||||
"grok2api_allow_legacy_full_save": False,
|
||||
**overrides,
|
||||
})
|
||||
|
||||
def test_remote_pool_falls_back_to_admin_api_prefix_when_root_tokens_add_is_404(self):
|
||||
self._configure()
|
||||
calls = []
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
@@ -56,11 +76,10 @@ class Grok2ApiRemotePoolTests(unittest.TestCase):
|
||||
})
|
||||
|
||||
def test_remote_pool_does_not_duplicate_admin_api_prefix_when_base_already_points_to_admin_api(self):
|
||||
app.config.update({
|
||||
"grok2api_remote_base": "https://grok.example.com/admin/api",
|
||||
"grok2api_remote_app_key": "app-secret",
|
||||
"grok2api_pool_name": "ssoSuper",
|
||||
})
|
||||
self._configure(
|
||||
grok2api_remote_base="https://grok.example.com/admin/api",
|
||||
grok2api_pool_name="ssoSuper",
|
||||
)
|
||||
calls = []
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
@@ -76,12 +95,8 @@ class Grok2ApiRemotePoolTests(unittest.TestCase):
|
||||
])
|
||||
self.assertEqual(calls[0][1]["json"]["pool"], "super")
|
||||
|
||||
def test_remote_pool_full_save_fallback_tries_admin_api_tokens_path(self):
|
||||
app.config.update({
|
||||
"grok2api_remote_base": "https://grok.example.com",
|
||||
"grok2api_remote_app_key": "app-secret",
|
||||
"grok2api_pool_name": "ssoBasic",
|
||||
})
|
||||
def test_remote_pool_full_save_fallback_requires_opt_in_and_uses_etag(self):
|
||||
self._configure(grok2api_allow_legacy_full_save=True)
|
||||
get_calls = []
|
||||
post_calls = []
|
||||
|
||||
@@ -96,7 +111,10 @@ class Grok2ApiRemotePoolTests(unittest.TestCase):
|
||||
def fake_get(url, **kwargs):
|
||||
get_calls.append((url, kwargs))
|
||||
if url == "https://grok.example.com/admin/api/tokens":
|
||||
return DummyResponse({"tokens": {"ssoBasic": []}})
|
||||
return DummyResponse(
|
||||
{"tokens": {"ssoBasic": []}},
|
||||
headers={"ETag": '"version-7"'},
|
||||
)
|
||||
return DummyResponse(status_code=404)
|
||||
|
||||
with patch.object(app, "http_post", side_effect=fake_post), \
|
||||
@@ -109,10 +127,45 @@ class Grok2ApiRemotePoolTests(unittest.TestCase):
|
||||
"https://grok.example.com/admin/api/tokens",
|
||||
])
|
||||
self.assertEqual(post_calls[-1][0], "https://grok.example.com/admin/api/tokens")
|
||||
self.assertEqual(post_calls[-1][1]["headers"]["If-Match"], '"version-7"')
|
||||
self.assertEqual(post_calls[-1][1]["json"], {
|
||||
"ssoBasic": [{"token": "fallback123", "tags": ["auto-register"], "note": "a@example.com"}],
|
||||
})
|
||||
|
||||
def test_remote_pool_legacy_fallback_is_disabled_by_default(self):
|
||||
self._configure()
|
||||
with patch.object(app, "http_post", return_value=DummyResponse(status_code=404)), \
|
||||
patch.object(app, "http_get") as get_mock:
|
||||
with self.assertRaises(app.RemoteTokenCompatibilityError):
|
||||
app.add_token_to_grok2api_remote_pool("abc")
|
||||
get_mock.assert_not_called()
|
||||
|
||||
def test_remote_pool_500_does_not_fallback(self):
|
||||
self._configure(grok2api_allow_legacy_full_save=True)
|
||||
with patch.object(app, "http_post", return_value=DummyResponse(status_code=500, text="boom")), \
|
||||
patch.object(app, "http_get") as get_mock:
|
||||
with self.assertRaises(app.RemoteTokenRequestError):
|
||||
app.add_token_to_grok2api_remote_pool("abc")
|
||||
get_mock.assert_not_called()
|
||||
|
||||
def test_remote_pool_legacy_fallback_rejects_missing_etag(self):
|
||||
self._configure(grok2api_allow_legacy_full_save=True)
|
||||
|
||||
def fake_post(url, **kwargs):
|
||||
if url.endswith("/tokens/add"):
|
||||
return DummyResponse(status_code=404)
|
||||
return DummyResponse({"status": "success"})
|
||||
|
||||
def fake_get(url, **kwargs):
|
||||
if url.endswith("/tokens"):
|
||||
return DummyResponse({"tokens": {"ssoBasic": []}})
|
||||
return DummyResponse(status_code=404)
|
||||
|
||||
with patch.object(app, "http_post", side_effect=fake_post), \
|
||||
patch.object(app, "http_get", side_effect=fake_get):
|
||||
with self.assertRaises(app.RemoteTokenCompatibilityError):
|
||||
app.add_token_to_grok2api_remote_pool("abc")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
@@ -0,0 +1,118 @@
|
||||
import unittest
|
||||
|
||||
from registration_flow import (
|
||||
RegistrationCallbacks,
|
||||
RegistrationOperations,
|
||||
run_batch,
|
||||
)
|
||||
|
||||
|
||||
class Cancelled(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class Retryable(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class FakeOps:
|
||||
def __init__(self, save_ok=True, observer_events=None):
|
||||
self.events = []
|
||||
self.save_ok = save_ok
|
||||
self.observer_events = observer_events if observer_events is not None else []
|
||||
self.account_no = 0
|
||||
|
||||
def operations(self):
|
||||
return RegistrationOperations(
|
||||
start_browser=lambda: self.events.append("start"),
|
||||
restart_browser=lambda: self.events.append("restart"),
|
||||
browser_missing=lambda: False,
|
||||
open_signup_page=lambda: self.events.append("open"),
|
||||
fill_email_and_submit=self._email,
|
||||
save_mail_credential=lambda email, token: True,
|
||||
fill_code_and_submit=lambda email, token: "123456",
|
||||
fill_profile_and_submit=lambda: {"given_name": "A", "family_name": "B", "password": "pw"},
|
||||
wait_for_sso_cookie=lambda: "sso-token",
|
||||
enable_nsfw=lambda sso: (True, "ok"),
|
||||
persist_account_line=self._persist,
|
||||
queue_unsaved_result=lambda payload, error: True,
|
||||
add_tokens=lambda sso, email: {
|
||||
"local": {"enabled": False, "ok": None, "error": None},
|
||||
"remote": {"enabled": False, "ok": None, "error": None},
|
||||
},
|
||||
export_cpa=lambda email, password, sso: {"ok": False, "skipped": True},
|
||||
cleanup=lambda reason: self.events.append(("cleanup", reason)),
|
||||
sleep=lambda seconds: self.events.append(("sleep", seconds)),
|
||||
cancelled_exception=Cancelled,
|
||||
retry_exception=Retryable,
|
||||
)
|
||||
|
||||
def _email(self):
|
||||
self.account_no += 1
|
||||
return f"user{self.account_no}@example.com", "mail-token"
|
||||
|
||||
def _persist(self, email, password, sso):
|
||||
if not self.save_ok:
|
||||
raise OSError("disk full")
|
||||
self.events.append(("persist", email))
|
||||
|
||||
|
||||
class RegistrationFlowTests(unittest.TestCase):
|
||||
def callbacks(self, logs=None):
|
||||
logs = logs if logs is not None else []
|
||||
return RegistrationCallbacks(log=logs.append, cancelled=lambda: False)
|
||||
|
||||
def test_start_failure_still_runs_cleanup(self):
|
||||
fake = FakeOps()
|
||||
ops = fake.operations()
|
||||
ops.start_browser = lambda: (_ for _ in ()).throw(RuntimeError("start failed"))
|
||||
with self.assertRaises(RuntimeError):
|
||||
run_batch(1, self.callbacks(), lambda *args: None, ops)
|
||||
self.assertEqual(fake.events, [("cleanup", "任务结束")])
|
||||
|
||||
def test_last_account_does_not_restart_browser(self):
|
||||
fake = FakeOps()
|
||||
batch = run_batch(1, self.callbacks(), lambda *args: None, fake.operations())
|
||||
self.assertEqual(batch.success_count, 1)
|
||||
self.assertNotIn("restart", fake.events)
|
||||
self.assertEqual(fake.events[-1], ("cleanup", "任务结束"))
|
||||
|
||||
def test_cleanup_interval_does_not_repeat_after_unsaved_result(self):
|
||||
fake = FakeOps(save_ok=True)
|
||||
ops = fake.operations()
|
||||
original_persist = ops.persist_account_line
|
||||
calls = {"count": 0}
|
||||
|
||||
def persist(email, password, sso):
|
||||
calls["count"] += 1
|
||||
if calls["count"] == 2:
|
||||
raise OSError("disk full")
|
||||
original_persist(email, password, sso)
|
||||
|
||||
ops.persist_account_line = persist
|
||||
batch = run_batch(2, self.callbacks(), lambda *args: None, ops, cleanup_interval=1)
|
||||
interval_cleanups = [
|
||||
event for event in fake.events
|
||||
if isinstance(event, tuple)
|
||||
and len(event) > 1
|
||||
and isinstance(event[1], str)
|
||||
and "已成功" in event[1]
|
||||
]
|
||||
self.assertEqual(len(interval_cleanups), 1)
|
||||
self.assertEqual(batch.success_count, 1)
|
||||
self.assertEqual(batch.registered_unsaved_count, 1)
|
||||
|
||||
def test_observer_failure_is_logged_and_batch_continues(self):
|
||||
fake = FakeOps()
|
||||
logs = []
|
||||
|
||||
def broken_observer(*args):
|
||||
raise RuntimeError("ui broke")
|
||||
|
||||
batch = run_batch(1, self.callbacks(logs), broken_observer, fake.operations())
|
||||
self.assertEqual(batch.success_count, 1)
|
||||
self.assertTrue(any("observer 执行失败" in line for line in logs))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
Reference in New Issue
Block a user