client.py 6.9 KB

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