fix: complete remaining registration hardening
This commit is contained in:
+57
-37
@@ -72,34 +72,56 @@ def _validate_endpoint(raw_url, field_name):
|
||||
return value
|
||||
|
||||
|
||||
def discover(proxy=None, timeout=30.0):
|
||||
opener = _build_opener(proxy)
|
||||
def discover(proxy=None, timeout=30.0, cancel=None, retries=2):
|
||||
request = urllib.request.Request(
|
||||
DISCOVERY_URL,
|
||||
method="GET",
|
||||
headers={"Accept": "application/json", "User-Agent": "grok-register-cpa/1.0"},
|
||||
)
|
||||
try:
|
||||
with opener.open(request, timeout=timeout) as response:
|
||||
body = response.read().decode("utf-8", errors="replace")
|
||||
status = int(getattr(response, "status", 200) or 200)
|
||||
except urllib.error.HTTPError as exc:
|
||||
body = exc.read().decode("utf-8", errors="replace")
|
||||
raise OAuthDeviceError("xAI discovery failed HTTP %s: %s" % (exc.code, body))
|
||||
except Exception as exc:
|
||||
raise OAuthDeviceError("xAI discovery request failed: %s" % exc)
|
||||
if status != 200:
|
||||
raise OAuthDeviceError("xAI discovery failed HTTP %s: %s" % (status, body))
|
||||
try:
|
||||
payload = json.loads(body)
|
||||
except Exception as exc:
|
||||
raise OAuthDeviceError("xAI discovery parse failed: %s" % exc)
|
||||
return {
|
||||
"device_authorization_endpoint": _validate_endpoint(
|
||||
payload.get("device_authorization_endpoint"), "device_authorization_endpoint"
|
||||
),
|
||||
"token_endpoint": _validate_endpoint(payload.get("token_endpoint"), "token_endpoint"),
|
||||
}
|
||||
last_error = None
|
||||
for attempt in range(max(int(retries), 0) + 1):
|
||||
_check_cancel(cancel)
|
||||
opener = _build_opener(proxy)
|
||||
try:
|
||||
with opener.open(request, timeout=float(timeout)) as response:
|
||||
body = response.read().decode("utf-8", errors="replace")
|
||||
status = int(getattr(response, "status", 200) or 200)
|
||||
_check_cancel(cancel)
|
||||
except urllib.error.HTTPError as exc:
|
||||
body = exc.read().decode("utf-8", errors="replace")
|
||||
raise OAuthDeviceError("xAI discovery failed HTTP %s: %s" % (exc.code, body))
|
||||
except Exception as exc:
|
||||
last_error = exc
|
||||
if not _is_transient_net_error(exc) or attempt >= int(retries):
|
||||
raise OAuthDeviceError("xAI discovery request failed: %s" % exc)
|
||||
_sleep_with_cancel(1.0 * (attempt + 1), cancel)
|
||||
continue
|
||||
if status != 200:
|
||||
raise OAuthDeviceError("xAI discovery failed HTTP %s: %s" % (status, body))
|
||||
try:
|
||||
payload = json.loads(body)
|
||||
except Exception as exc:
|
||||
raise OAuthDeviceError("xAI discovery parse failed: %s" % exc)
|
||||
return {
|
||||
"device_authorization_endpoint": _validate_endpoint(
|
||||
payload.get("device_authorization_endpoint"), "device_authorization_endpoint"
|
||||
),
|
||||
"token_endpoint": _validate_endpoint(payload.get("token_endpoint"), "token_endpoint"),
|
||||
}
|
||||
raise OAuthDeviceError("xAI discovery failed: %s" % last_error)
|
||||
|
||||
|
||||
def _check_cancel(cancel):
|
||||
if cancel and cancel():
|
||||
raise OAuthDeviceError("cancelled")
|
||||
|
||||
|
||||
def _sleep_with_cancel(seconds, cancel=None):
|
||||
deadline = time.time() + max(float(seconds), 0.0)
|
||||
while time.time() < deadline:
|
||||
_check_cancel(cancel)
|
||||
time.sleep(min(0.2, max(deadline - time.time(), 0.0)))
|
||||
_check_cancel(cancel)
|
||||
|
||||
|
||||
def _is_transient_net_error(exc):
|
||||
@@ -145,7 +167,7 @@ def _is_transient_net_error(exc):
|
||||
return False
|
||||
|
||||
|
||||
def _post_form(url, form, timeout=30.0, proxy=None, retries=0, retry_sleep=1.5):
|
||||
def _post_form(url, form, timeout=30.0, proxy=None, retries=0, retry_sleep=1.5, cancel=None):
|
||||
data = urllib.parse.urlencode(form).encode("utf-8")
|
||||
request = urllib.request.Request(
|
||||
url,
|
||||
@@ -159,11 +181,13 @@ def _post_form(url, form, timeout=30.0, proxy=None, retries=0, retry_sleep=1.5):
|
||||
)
|
||||
last_error = None
|
||||
for attempt in range(max(int(retries), 0) + 1):
|
||||
_check_cancel(cancel)
|
||||
opener = _build_opener(proxy)
|
||||
try:
|
||||
with opener.open(request, timeout=timeout) as response:
|
||||
body = response.read().decode("utf-8", errors="replace")
|
||||
status = int(getattr(response, "status", 200) or 200)
|
||||
_check_cancel(cancel)
|
||||
except urllib.error.HTTPError as exc:
|
||||
body = exc.read().decode("utf-8", errors="replace")
|
||||
status = int(exc.code)
|
||||
@@ -171,7 +195,7 @@ def _post_form(url, form, timeout=30.0, proxy=None, retries=0, retry_sleep=1.5):
|
||||
last_error = exc
|
||||
if not _is_transient_net_error(exc) or attempt >= int(retries):
|
||||
raise
|
||||
time.sleep(float(retry_sleep) * (attempt + 1))
|
||||
_sleep_with_cancel(float(retry_sleep) * (attempt + 1), cancel)
|
||||
continue
|
||||
try:
|
||||
return status, json.loads(body)
|
||||
@@ -182,8 +206,9 @@ def _post_form(url, form, timeout=30.0, proxy=None, retries=0, retry_sleep=1.5):
|
||||
raise OAuthDeviceError("form request failed without response")
|
||||
|
||||
|
||||
def request_device_code(client_id=CLIENT_ID, scope=SCOPE, timeout=30.0, proxy=None):
|
||||
discovery = discover(proxy=proxy, timeout=timeout)
|
||||
def request_device_code(client_id=CLIENT_ID, scope=SCOPE, timeout=15.0, proxy=None, cancel=None, retries=2):
|
||||
discovery = discover(proxy=proxy, timeout=timeout, cancel=cancel, retries=retries)
|
||||
_check_cancel(cancel)
|
||||
device_endpoint = discovery["device_authorization_endpoint"]
|
||||
token_endpoint = discovery["token_endpoint"]
|
||||
status, payload = _post_form(
|
||||
@@ -191,9 +216,11 @@ def request_device_code(client_id=CLIENT_ID, scope=SCOPE, timeout=30.0, proxy=No
|
||||
{"client_id": client_id, "scope": scope},
|
||||
timeout=timeout,
|
||||
proxy=proxy,
|
||||
retries=2,
|
||||
retries=retries,
|
||||
retry_sleep=1.0,
|
||||
cancel=cancel,
|
||||
)
|
||||
_check_cancel(cancel)
|
||||
if status != 200 or not isinstance(payload, dict):
|
||||
raise OAuthDeviceError("device code request failed HTTP %s: %r" % (status, payload))
|
||||
device_code = str(payload.get("device_code") or "").strip()
|
||||
@@ -216,14 +243,6 @@ 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,
|
||||
@@ -251,10 +270,11 @@ def poll_device_token(
|
||||
"device_code": str(device_code).strip(),
|
||||
"client_id": client_id,
|
||||
},
|
||||
timeout=min(float(timeout), 5.0),
|
||||
timeout=float(timeout),
|
||||
proxy=proxy,
|
||||
retries=0,
|
||||
retry_sleep=1.0,
|
||||
cancel=cancel,
|
||||
)
|
||||
net_streak = 0
|
||||
except Exception as exc:
|
||||
|
||||
Reference in New Issue
Block a user