refactor: modularize registration runtime safely

This commit is contained in:
github-actions[bot]
2026-07-14 18:10:13 +00:00
parent df525d0962
commit b845be012a
23 changed files with 3945 additions and 8518 deletions
+36
View File
@@ -0,0 +1,36 @@
import unittest
from cpa_xai import browser_session
class Bridge:
def __init__(self):
self.stops = 0
def stop(self):
self.stops += 1
class Browser:
def __init__(self, bridge=None):
self._cpa_proxy_bridge = bridge
self.quits = 0
def quit(self):
self.quits += 1
class BrowserSessionTests(unittest.TestCase):
def test_close_standalone_closes_browser_and_bridge(self):
bridge = Bridge()
browser = Browser(bridge)
browser_session._register_mint_browser(browser)
browser_session.close_standalone(browser)
self.assertEqual(browser.quits, 1)
self.assertEqual(bridge.stops, 1)
def test_normalize_cookies_rejects_invalid_items(self):
value = browser_session.normalize_cookies([None, {"name": "a", "value": "b"}, "bad"])
self.assertEqual(len(value), 1)
self.assertEqual(value[0]["name"], "a")
if __name__ == "__main__":
unittest.main()
+32
View File
@@ -0,0 +1,32 @@
import tempfile
import unittest
from pathlib import Path
from unittest.mock import patch
from cpa_xai.mint import mint_and_export
from cpa_xai.schema import build_cpa_xai_auth, jwt_payload
from cpa_xai.writer import write_cpa_xai_auth
class CpaCoreTests(unittest.TestCase):
def test_schema_rejects_missing_tokens(self):
with self.assertRaises(ValueError):
build_cpa_xai_auth("a@example.com", "", "refresh")
with self.assertRaises(ValueError):
jwt_payload("not-a-jwt")
def test_writer_failure_does_not_leave_temp_file(self):
with tempfile.TemporaryDirectory() as directory:
with patch("cpa_xai.writer.os.replace", side_effect=OSError("disk")):
with self.assertRaises(OSError):
write_cpa_xai_auth(directory, {"email": "a@example.com"}, "a.json")
self.assertEqual([p.name for p in Path(directory).iterdir()], [])
def test_mint_rejects_missing_identity_without_browser(self):
result = mint_and_export("", "", tempfile.gettempdir())
self.assertFalse(result["ok"])
self.assertIn("missing", result["error"])
if __name__ == "__main__":
unittest.main()
+32
View File
@@ -0,0 +1,32 @@
import unittest
from unittest.mock import patch
import app_config
import grok_register_ttk as app
import mail_service
import registration_browser
from cpa_xai import browser_confirm
from cpa_xai import browser_session
class ModuleCompatibilityTests(unittest.TestCase):
def test_config_object_is_shared(self):
self.assertIs(app.config, app_config.config)
def test_original_public_functions_remain_available(self):
names = ['_normalize_sso_token', '_pick_list_payload', 'add_token_to_grok2api_local_pool', 'add_token_to_grok2api_pools', 'add_token_to_grok2api_remote_pool', 'build_profile', 'cleanup_runtime_memory', 'click_email_signup_button', 'cloudflare_apply_auth_params', 'cloudflare_build_headers', 'cloudflare_create_account', 'cloudflare_create_temp_address', 'cloudflare_get_domains', 'cloudflare_get_message_detail', 'cloudflare_get_messages', 'cloudflare_get_oai_code', 'cloudflare_get_token', 'cloudflare_is_admin_create_path', 'cloudflare_next_default_domain', 'cloudmail_get_email_and_token', 'cloudmail_get_messages', 'cloudmail_get_oai_code', 'cloudmail_next_domain', 'create_account', 'create_browser_options', 'duckmail_get_oai_code', 'enable_nsfw_for_token', 'encode_grpc_nsfw_settings', 'extract_verification_code', 'fill_code_and_submit', 'fill_email_and_submit', 'fill_profile_and_submit', 'generate_random_birthdate', 'generate_username', 'getTurnstileToken', 'get_cloudflare_api_base', 'get_cloudflare_api_key', 'get_cloudflare_auth_mode', 'get_cloudflare_path', 'get_cloudmail_api_base', 'get_cloudmail_path', 'get_cloudmail_public_token', 'get_domains', 'get_duckmail_api_key', 'get_email_and_token', 'get_email_provider', 'get_grok2api_remote_api_bases', 'get_message_detail', 'get_messages', 'get_oai_code', 'get_token', 'get_user_agent', 'get_yyds_api_key', 'get_yyds_jwt', 'has_profile_form', 'http_get', 'http_post', 'is_cloudflare_block_response', 'open_signup_page', 'pick_domain', 'refresh_active_page', 'resolve_grok2api_local_token_file', 'response_preview', 'restart_browser', 'set_birth_date', 'set_tos_accepted', 'start_browser', 'stop_browser', 'stop_browser_proxy_bridge', 'update_nsfw_settings', 'wait_for_sso_cookie', 'yyds_create_account', 'yyds_generate_username', 'yyds_get_domains', 'yyds_get_email_and_token', 'yyds_get_message_detail', 'yyds_get_messages', 'yyds_get_oai_code', 'yyds_get_token', 'yyds_pick_domain']
for name in names:
self.assertTrue(callable(getattr(app, name)), name)
def test_mail_wrapper_delegates(self):
with patch.object(mail_service, "get_email_provider", return_value="duckmail") as mocked:
self.assertEqual(app.get_email_provider(), "duckmail")
mocked.assert_called_once_with()
def test_browser_confirm_reexports_session_api(self):
self.assertIs(browser_confirm.close_standalone, browser_session.close_standalone)
self.assertIs(browser_confirm.normalize_cookies, browser_session.normalize_cookies)
if __name__ == "__main__":
unittest.main()
+69
View File
@@ -0,0 +1,69 @@
import json
import unittest
from unittest.mock import patch
from cpa_xai import oauth_device as oauth
class Response:
def __init__(self, body, status=200):
self.body = body.encode("utf-8")
self.status = status
def __enter__(self):
return self
def __exit__(self, *args):
return False
def read(self):
return self.body
class Opener:
def __init__(self, actions):
self.actions = list(actions)
self.calls = 0
def open(self, request, timeout=None):
self.calls += 1
action = self.actions.pop(0)
if isinstance(action, BaseException):
raise action
return action
class OAuthDeviceTests(unittest.TestCase):
def test_discovery_success(self):
payload = {"device_authorization_endpoint": "https://auth.x.ai/device", "token_endpoint": "https://auth.x.ai/token"}
opener = Opener([Response(json.dumps(payload))])
with patch.object(oauth, "_build_opener", return_value=opener):
self.assertEqual(oauth.discover(retries=0)["token_endpoint"], payload["token_endpoint"])
def test_discovery_cancelled_before_request(self):
with self.assertRaisesRegex(oauth.OAuthDeviceError, "cancelled"):
oauth.discover(cancel=lambda: True)
def test_discovery_retries_transient_error(self):
payload = {"device_authorization_endpoint": "https://auth.x.ai/device", "token_endpoint": "https://auth.x.ai/token"}
opener = Opener([TimeoutError("slow"), Response(json.dumps(payload))])
with patch.object(oauth, "_build_opener", return_value=opener), patch.object(oauth, "_sleep_with_cancel"):
oauth.discover(retries=1)
self.assertEqual(opener.calls, 2)
def test_post_form_returns_non_json_body(self):
opener = Opener([Response("not-json", status=502)])
with patch.object(oauth, "_build_opener", return_value=opener):
status, payload = oauth._post_form("https://auth.x.ai/token", {}, retries=0)
self.assertEqual((status, payload), (502, "not-json"))
def test_slow_down_increases_wait(self):
responses = [
(400, {"error": "slow_down"}),
(200, {"access_token": "a", "refresh_token": "r"}),
]
waits = []
with patch.object(oauth, "_post_form", side_effect=responses), patch.object(oauth, "_sleep_with_cancel", side_effect=lambda seconds, cancel=None: waits.append(seconds)):
result = oauth.poll_device_token("d", "https://auth.x.ai/token", interval=1, expires_in=60)
self.assertEqual(result.refresh_token, "r")
self.assertEqual(waits, [6])
if __name__ == "__main__":
unittest.main()
+37
View File
@@ -155,6 +155,43 @@ class RegistrationFlowTests(unittest.TestCase):
self.assertEqual(batch.fail_count, 0)
self.assertEqual(batch.postprocess_warning_count, 1)
def test_cleanup_failure_does_not_change_success_statistics(self):
fake = FakeOps()
ops = fake.operations()
base_cleanup = ops.cleanup
def cleanup(reason):
if "已成功" in reason:
raise RuntimeError("cleanup failed")
base_cleanup(reason)
ops.cleanup = cleanup
batch = run_batch(2, self.callbacks(), lambda *args: None, ops, cleanup_interval=1)
self.assertEqual((batch.success_count, batch.fail_count, batch.processed_count), (2, 0, 2))
def test_cancel_during_next_account_wait_is_normal_cancellation(self):
fake = FakeOps()
ops = fake.operations()
ops.sleep = lambda seconds: (_ for _ in ()).throw(Cancelled())
batch = run_batch(2, self.callbacks(), lambda *args: None, ops)
self.assertTrue(batch.cancelled)
self.assertEqual(batch.processed_count, 1)
def test_final_cleanup_does_not_mask_original_error(self):
fake = FakeOps()
ops = fake.operations()
ops.start_browser = lambda: (_ for _ in ()).throw(ValueError("original"))
ops.cleanup = lambda reason: (_ for _ in ()).throw(RuntimeError("cleanup"))
with self.assertRaisesRegex(ValueError, "original"):
run_batch(1, self.callbacks(), lambda *args: None, ops)
def test_optional_postprocessing_exceptions_become_warning(self):
fake = FakeOps()
ops = fake.operations()
ops.add_tokens = lambda sso, email: (_ for _ in ()).throw(RuntimeError("pool"))
ops.export_cpa = lambda email, password, sso: (_ for _ in ()).throw(RuntimeError("cpa"))
batch = run_batch(1, self.callbacks(), lambda *args: None, ops)
self.assertEqual(batch.success_count, 1)
self.assertEqual(batch.postprocess_warning_count, 1)
if __name__ == "__main__":
unittest.main()