client.py 6.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178
  1. """HTTP 客户端:JSON 请求、线上站点之间出声的 fallback、token 显示辅助。"""
  2. import base64
  3. import json
  4. import os
  5. import sys
  6. import time
  7. import urllib.error
  8. import urllib.parse
  9. import urllib.request
  10. from datetime import datetime, timezone
  11. from creds import bucket_name_for, get_bucket, load_creds, resolve_api_url, save_creds
  12. from errors import ApiError, WpError
  13. from sites import ONLINE_URLS, SITES, site_label
  14. TOKEN_REFRESH_MARGIN = 3600
  15. DEFAULT_TIMEOUT = 30
  16. WRITE_TIMEOUT = 120
  17. DEFAULT_BATCH = 50
  18. def mask(token):
  19. if not token:
  20. return "(无)"
  21. if len(token) <= 16:
  22. return token[:4] + "…"
  23. return token[:8] + "…" + token[-4:]
  24. def jwt_payload(token):
  25. """不验签地读出 JWT payload,仅用于显示有效期。"""
  26. try:
  27. part = token.split(".")[1]
  28. part += "=" * (-len(part) % 4)
  29. return json.loads(base64.urlsafe_b64decode(part.encode("ascii")))
  30. except Exception:
  31. return {}
  32. def token_expiry(token):
  33. exp = jwt_payload(token).get("exp")
  34. return int(exp) if isinstance(exp, (int, float)) else None
  35. def fmt_ts(ts):
  36. if not ts:
  37. return "未知"
  38. return datetime.fromtimestamp(ts, tz=timezone.utc).astimezone().strftime("%Y-%m-%d %H:%M")
  39. def note(msg):
  40. print(msg, file=sys.stderr)
  41. def http_json(api_url, method, path, token=None, body=None, query=None, timeout=DEFAULT_TIMEOUT):
  42. """发一个 JSON 请求,返回解析后的响应体(dict)。
  43. 网络层失败抛 urllib 的异常(由 Client 决定是否 fallback);
  44. HTTP 层失败抛 ApiError,带上服务端 message。
  45. """
  46. # 路径里可能有巴利词(parivāsa),urllib 只接受 ASCII,必须先百分号编码
  47. url = api_url + "/" + urllib.parse.quote(path.lstrip("/"), safe="/")
  48. if query:
  49. url += "?" + urllib.parse.urlencode({k: v for k, v in query.items() if v is not None})
  50. data = None
  51. headers = {"Accept": "application/json", "User-Agent": "wikipali-write-skill"}
  52. if body is not None:
  53. data = json.dumps(body, ensure_ascii=False).encode("utf-8")
  54. headers["Content-Type"] = "application/json"
  55. if token:
  56. headers["Authorization"] = "Bearer " + token
  57. req = urllib.request.Request(url, data=data, headers=headers, method=method)
  58. try:
  59. with urllib.request.urlopen(req, timeout=timeout) as resp:
  60. raw = resp.read().decode("utf-8", "replace")
  61. status = resp.status
  62. except urllib.error.HTTPError as exc:
  63. raw = exc.read().decode("utf-8", "replace")
  64. status = exc.code
  65. payload = safe_json(raw)
  66. message = payload.get("message") if isinstance(payload, dict) else None
  67. raise ApiError(status, message or f"HTTP {status}", url=url, body=payload or raw)
  68. payload = safe_json(raw)
  69. if not isinstance(payload, dict):
  70. raise ApiError(status, f"响应不是 JSON:{raw[:200]}", url=url, body=raw)
  71. if not payload.get("ok", False):
  72. raise ApiError(status, payload.get("message") or "请求失败", url=url, body=payload)
  73. return payload.get("data")
  74. def safe_json(raw):
  75. try:
  76. return json.loads(raw)
  77. except ValueError:
  78. return None
  79. class Client:
  80. """按站点收发请求,并在线上地址之间做出声的 fallback。"""
  81. def __init__(self, api_url, source, creds, allow_fallback=True):
  82. self.api_url = api_url
  83. self.source = source
  84. self.creds = creds
  85. self.bucket_name = bucket_name_for(api_url)
  86. self.bucket = get_bucket(creds, self.bucket_name, api_url)
  87. self.allow_fallback = allow_fallback and api_url in ONLINE_URLS
  88. # -- 凭据 ---------------------------------------------------------------
  89. @property
  90. def user_token(self):
  91. token = (self.bucket.get("user") or {}).get("token")
  92. if not token:
  93. raise WpError(
  94. "尚未登录。请自己执行(Claude Code 里用 ! 前缀):\n"
  95. " ! python3 " + os.path.join(os.path.dirname(os.path.abspath(__file__)), "wp_login.py")
  96. )
  97. return token
  98. @property
  99. def model(self):
  100. model = self.bucket.get("model") or {}
  101. if not model.get("token"):
  102. raise WpError("尚未取得模型身份 token。请先跑:python3 wp.py ensure-model --name <模型名>")
  103. return model
  104. def save(self):
  105. save_creds(self.creds)
  106. # -- 请求 ---------------------------------------------------------------
  107. def fallback_order(self):
  108. """同版本的另一域名 → 另一版本的同域名 → 其余。绝不含 local。"""
  109. cur = next((s for s in SITES if s["url"] == self.api_url), None)
  110. if not cur:
  111. return []
  112. others = [s for s in SITES if s["key"] != "local" and s["url"] != self.api_url]
  113. others.sort(
  114. key=lambda s: (
  115. 0 if s["version"] == cur["version"] else 1,
  116. 0 if s["domain"] == cur["domain"] else 1,
  117. )
  118. )
  119. return [s["url"] for s in others]
  120. def call(self, method, path, token=None, body=None, query=None, timeout=DEFAULT_TIMEOUT):
  121. urls = [self.api_url] + (self.fallback_order() if self.allow_fallback else [])
  122. last = None
  123. for idx, url in enumerate(urls):
  124. try:
  125. data = http_json(url, method, path, token=token, body=body, query=query, timeout=timeout)
  126. except (urllib.error.URLError, TimeoutError, OSError) as exc:
  127. # 仅网络层不可达才换站点;HTTP 错误是服务端的明确答复,不该被掩盖
  128. last = exc
  129. reason = getattr(exc, "reason", exc)
  130. if idx + 1 < len(urls):
  131. note(f"⚠ {url} 连接失败({reason}),改用 {urls[idx + 1]}")
  132. continue
  133. if url != self.api_url:
  134. # fallback 成功后本次会话都用它,但不写回凭据文件
  135. note(f"⚠ 本次请求实际发往 {url}({site_label(url)})")
  136. self.api_url = url
  137. return data
  138. raise WpError(f"所有可用站点都连不上,最后一次错误:{last}")
  139. def api_note(self):
  140. src = {"cli": "--api", "env": "环境变量", "creds": "凭据文件", "default": "内置默认"}[self.source]
  141. return f"{self.api_url}({site_label(self.api_url)},来源:{src})"
  142. def make_client(args, allow_fallback=True):
  143. creds = load_creds()
  144. api_url, source = resolve_api_url(getattr(args, "api", None), creds)
  145. return Client(api_url, source, creds, allow_fallback=allow_fallback)