fix: harden token pools and CPA export
This commit is contained in:
@@ -189,10 +189,11 @@ def create_standalone_page(proxy: Optional[str] = None, headless: bool = False,
|
||||
pass
|
||||
break
|
||||
|
||||
from .proxyutil import proxy_for_chromium, proxy_log_label, resolve_proxy
|
||||
from .proxyutil import prepare_chromium_proxy, proxy_log_label, resolve_proxy
|
||||
|
||||
resolved = resolve_proxy(proxy)
|
||||
chrome_proxy = proxy_for_chromium(resolved)
|
||||
proxy_bridge = None
|
||||
chrome_proxy, proxy_bridge = prepare_chromium_proxy(resolved, log=logger)
|
||||
if chrome_proxy:
|
||||
options.set_argument("--proxy-server=%s" % chrome_proxy)
|
||||
logger("browser proxy=%s (chromium %s)" % (proxy_log_label(resolved), chrome_proxy))
|
||||
@@ -200,19 +201,50 @@ def create_standalone_page(proxy: Optional[str] = None, headless: bool = False,
|
||||
logger("browser proxy=(none)")
|
||||
|
||||
browser = Chromium(options)
|
||||
if proxy_bridge is not None:
|
||||
try:
|
||||
setattr(browser, "_cpa_proxy_bridge", proxy_bridge)
|
||||
except Exception:
|
||||
pass
|
||||
_register_mint_browser(browser)
|
||||
page = browser.latest_tab
|
||||
logger("standalone chromium started")
|
||||
return browser, page
|
||||
|
||||
|
||||
def close_standalone(browser: Any) -> None:
|
||||
if browser is None:
|
||||
return
|
||||
_unregister_mint_browser(browser)
|
||||
bridge = getattr(browser, "_cpa_proxy_bridge", None)
|
||||
try:
|
||||
browser.quit()
|
||||
except Exception:
|
||||
pass
|
||||
if bridge is not None:
|
||||
try:
|
||||
bridge.stop()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
_mint_tls = threading.local()
|
||||
_mint_registry_lock = threading.Lock()
|
||||
_mint_registry = set()
|
||||
|
||||
|
||||
def _register_mint_browser(browser: Any) -> None:
|
||||
if browser is None:
|
||||
return
|
||||
with _mint_registry_lock:
|
||||
_mint_registry.add(browser)
|
||||
|
||||
|
||||
def _unregister_mint_browser(browser: Any) -> None:
|
||||
if browser is None:
|
||||
return
|
||||
with _mint_registry_lock:
|
||||
_mint_registry.discard(browser)
|
||||
|
||||
|
||||
def _mint_tls_get():
|
||||
@@ -417,8 +449,9 @@ def release_mint_browser(owned: bool, success: bool, log: Optional[LogFn] = None
|
||||
|
||||
def shutdown_mint_browsers() -> None:
|
||||
state = _mint_tls_get()
|
||||
browser = state.get("browser")
|
||||
if browser is not None:
|
||||
with _mint_registry_lock:
|
||||
browsers = list(_mint_registry)
|
||||
for browser in browsers:
|
||||
try:
|
||||
close_standalone(browser)
|
||||
except Exception:
|
||||
@@ -824,7 +857,7 @@ def mint_with_browser(
|
||||
session = request_device_code(proxy=resolved or None)
|
||||
last_error = None
|
||||
break
|
||||
except BaseException as exc:
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
logger("request_device_code attempt %s/3 failed: %s" % (attempt, exc))
|
||||
_sleep(1.5 * attempt)
|
||||
@@ -852,22 +885,28 @@ def mint_with_browser(
|
||||
token_box = {}
|
||||
error_box = {}
|
||||
|
||||
def combined_cancel():
|
||||
return stop_event.is_set() or bool(cancel and cancel())
|
||||
|
||||
def _poll() -> None:
|
||||
try:
|
||||
time.sleep(2)
|
||||
for _ in range(20):
|
||||
if combined_cancel():
|
||||
raise OAuthDeviceError("cancelled")
|
||||
time.sleep(0.1)
|
||||
result = poll_device_token(
|
||||
session.device_code,
|
||||
token_endpoint=session.token_endpoint,
|
||||
interval=max(session.interval, 5),
|
||||
expires_in=min(session.expires_in, int(browser_timeout_sec) + 60),
|
||||
log=logger,
|
||||
cancel=cancel,
|
||||
cancel=combined_cancel,
|
||||
proxy=resolved or None,
|
||||
)
|
||||
token_box["token"] = result
|
||||
stop_event.set()
|
||||
logger("token poll SUCCESS — stop_event set")
|
||||
except BaseException as exc:
|
||||
except Exception as exc:
|
||||
error_box["err"] = exc
|
||||
stop_event.set()
|
||||
|
||||
@@ -897,8 +936,16 @@ def mint_with_browser(
|
||||
)
|
||||
if hard:
|
||||
stop_event.set()
|
||||
thread.join(timeout=5)
|
||||
if thread.is_alive():
|
||||
logger("token poll thread did not stop within 5s after browser failure")
|
||||
raise
|
||||
thread.join(timeout=max(browser_timeout_sec, 60) + 30)
|
||||
if thread.is_alive():
|
||||
stop_event.set()
|
||||
thread.join(timeout=5)
|
||||
if thread.is_alive():
|
||||
raise OAuthDeviceError("token poll thread did not stop after timeout")
|
||||
if "token" in token_box:
|
||||
token_result = token_box["token"]
|
||||
success = True
|
||||
|
||||
+14
-6
@@ -167,7 +167,7 @@ def _post_form(url, form, timeout=30.0, proxy=None, retries=0, retry_sleep=1.5):
|
||||
except urllib.error.HTTPError as exc:
|
||||
body = exc.read().decode("utf-8", errors="replace")
|
||||
status = int(exc.code)
|
||||
except BaseException as exc:
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
if not _is_transient_net_error(exc) or attempt >= int(retries):
|
||||
raise
|
||||
@@ -216,6 +216,14 @@ def request_device_code(client_id=CLIENT_ID, scope=SCOPE, timeout=30.0, proxy=No
|
||||
)
|
||||
|
||||
|
||||
def _sleep_with_cancel(seconds, cancel=None):
|
||||
deadline = time.time() + max(float(seconds), 0.0)
|
||||
while time.time() < deadline:
|
||||
if cancel and cancel():
|
||||
raise OAuthDeviceError("cancelled")
|
||||
time.sleep(min(0.2, max(deadline - time.time(), 0.0)))
|
||||
|
||||
|
||||
def poll_device_token(
|
||||
device_code,
|
||||
token_endpoint,
|
||||
@@ -249,7 +257,7 @@ def poll_device_token(
|
||||
retry_sleep=1.0,
|
||||
)
|
||||
net_streak = 0
|
||||
except BaseException as exc:
|
||||
except Exception as exc:
|
||||
if not _is_transient_net_error(exc):
|
||||
raise
|
||||
net_streak += 1
|
||||
@@ -257,7 +265,7 @@ def poll_device_token(
|
||||
logger("oauth poll network blip (%s/%s): %s — retry in %ss" % (net_streak, max_net_streak, exc, wait_seconds))
|
||||
if net_streak >= max_net_streak:
|
||||
raise OAuthDeviceError("device auth aborted after %s network errors: %s" % (net_streak, exc))
|
||||
time.sleep(wait_seconds)
|
||||
_sleep_with_cancel(wait_seconds, cancel)
|
||||
continue
|
||||
if status == 200 and isinstance(payload, dict) and payload.get("access_token"):
|
||||
access_token = str(payload.get("access_token") or "").strip()
|
||||
@@ -281,7 +289,7 @@ def poll_device_token(
|
||||
if error_code == "slow_down":
|
||||
sleep_seconds = min(sleep_seconds + 5, 30)
|
||||
logger("oauth poll: %s (sleep %ss)" % (error_code, sleep_seconds))
|
||||
time.sleep(sleep_seconds)
|
||||
_sleep_with_cancel(sleep_seconds, cancel)
|
||||
continue
|
||||
if error_code in ("expired_token", "access_denied"):
|
||||
raise OAuthDeviceError("device auth failed: %s: %s" % (error_code, error_description))
|
||||
@@ -293,8 +301,8 @@ def poll_device_token(
|
||||
logger("oauth poll soft HTTP %s: %r — retry in %ss" % (status, payload, wait_seconds))
|
||||
if net_streak >= max_net_streak:
|
||||
raise OAuthDeviceError("device auth aborted after repeated soft HTTP failures status=%s" % status)
|
||||
time.sleep(wait_seconds)
|
||||
_sleep_with_cancel(wait_seconds, cancel)
|
||||
continue
|
||||
logger("oauth poll unexpected HTTP %s: %r" % (status, payload))
|
||||
time.sleep(sleep_seconds)
|
||||
_sleep_with_cancel(sleep_seconds, cancel)
|
||||
raise OAuthDeviceError("device auth timed out waiting for user approval")
|
||||
|
||||
@@ -87,6 +87,8 @@ def build_cpa_xai_auth(
|
||||
expired = ""
|
||||
if exp:
|
||||
expired = datetime.fromtimestamp(exp, tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
|
||||
elif expires_in:
|
||||
expired = datetime.fromtimestamp(time.time() + int(expires_in or 0), tz=timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ")
|
||||
payload = {
|
||||
"type": "xai",
|
||||
"access_token": access,
|
||||
|
||||
+12
-1
@@ -8,13 +8,24 @@ from pathlib import Path
|
||||
from .schema import credential_file_name
|
||||
|
||||
|
||||
def _is_relative_to(path, root):
|
||||
try:
|
||||
path.relative_to(root)
|
||||
return True
|
||||
except ValueError:
|
||||
return False
|
||||
|
||||
|
||||
def write_cpa_xai_auth(auth_dir, payload, filename=None):
|
||||
root = Path(auth_dir).expanduser().resolve()
|
||||
root.mkdir(parents=True, exist_ok=True)
|
||||
target_name = filename or credential_file_name(payload.get("email", ""), payload.get("sub", ""))
|
||||
target_name = Path(str(target_name)).name
|
||||
if not str(target_name).endswith(".json"):
|
||||
target_name = str(target_name) + ".json"
|
||||
destination = root / str(target_name)
|
||||
destination = (root / str(target_name)).resolve()
|
||||
if not _is_relative_to(destination, root):
|
||||
raise ValueError("CPA auth filename must stay inside auth_dir")
|
||||
data = json.dumps(payload, indent=2, ensure_ascii=False) + "\n"
|
||||
file_descriptor, temp_name = tempfile.mkstemp(prefix=".xai-", suffix=".tmp", dir=str(root))
|
||||
try:
|
||||
|
||||
Reference in New Issue
Block a user