PatSnap/Document-Processing
10
1"""HTTP clients for HIRO Translation and Smart Doc (public connect gateway)."""2 3from __future__ import annotations4 5import json6import os7import re8import time9from collections.abc import Iterator10from pathlib import Path11from typing import Any12 13import requests14 15# Public gateways (matches connect.zhihuiya.com curl examples).16DEFAULT_BASE = "https://connect.zhihuiya.com/hiro_translation"17SMARTDOC_URL = "https://connect.zhihuiya.com/rd-llm/v1/documents/doc_parsing"18 19API_KEY_ENV = "HIRO_API_KEY"20API_KEY_ENV_FALLBACK = "RD_LLM_API_KEY"21 22MAX_SMARTDOC_BYTES = 10 * 1024 * 102423 24# Async fast: submit + poll (aligned with hiro-translation-api).25DEFAULT_POLL_INTERVAL_S = 1.526DEFAULT_POLL_TIMEOUT_S = 180027 28 29def resolve_api_key(api_key: str | None = None) -> str:30 """Prefer explicit key, then HIRO_API_KEY, then RD_LLM_API_KEY."""31 if api_key is not None and str(api_key).strip():32 return str(api_key).strip()33 return (34 os.environ.get(API_KEY_ENV, "").strip()35 or os.environ.get(API_KEY_ENV_FALLBACK, "").strip()36 )37 38 39def lang_to_codes(lang: str) -> tuple[str, str]:40 src, _, tgt = lang.partition("2")41 if not src or not tgt:42 raise ValueError(f"invalid lang {lang!r}; expected format like zh2en")43 return src, tgt44 45 46def translate_payload(content: str, lang: str) -> dict[str, Any]:47 """Request body for POST /translate and POST /translate/async (no mode field)."""48 source, target = lang_to_codes(lang)49 return {50 "content": content,51 "sourceLanguageCode": source,52 "targetLanguageCode": target,53 }54 55 56def _api_url(base_url: str, path: str) -> str:57 return f"{base_url.rstrip('/')}{path}"58 59 60def _request_headers(61 api_key: str | None = None,62 *,63 json_body: bool = True,64) -> dict[str, str]:65 headers: dict[str, str] = {}66 if json_body:67 headers["Content-Type"] = "application/json"68 key = resolve_api_key(api_key)69 if key:70 headers["Authorization"] = f"Bearer {key}"71 return headers72 73 74def _translated_text(data: dict[str, Any]) -> str:75 value = data.get("textTranslated", data.get("text_translated", ""))76 return value if isinstance(value, str) else ""77 78 79def _original_text(data: dict[str, Any], fallback: str = "") -> str:80 value = data.get("textOriginal", data.get("text_original", fallback))81 return value if isinstance(value, str) else fallback82 83 84def _http_error_message(resp: requests.Response) -> str:85 try:86 body = resp.json()87 if isinstance(body, dict):88 return str(body.get("error") or body.get("message") or body)89 except (json.JSONDecodeError, ValueError):90 pass91 return resp.text[:500]92 93 94def health_ok(95 base_url: str = DEFAULT_BASE,96 *,97 api_key: str | None = None,98) -> tuple[bool, str]:99 try:100 r = requests.get(101 _api_url(base_url, "/health"),102 headers=_request_headers(api_key),103 timeout=12,104 )105 r.raise_for_status()106 data = r.json()107 if data.get("status") != "OK":108 return False, f"Unexpected response: {r.text[:200]}"109 upstream = data.get("upstream", "UNKNOWN")110 if upstream == "OK":111 return True, "Healthy · upstream OK"112 return False, f"Gateway up, upstream unavailable ({upstream})"113 except requests.RequestException as exc:114 return False, str(exc)115 116 117def _gateway_error_message(data: Any, fallback: str = "") -> str | None:118 """Extract connect-gateway style errors (error_msg / error_code)."""119 if not isinstance(data, dict):120 return None121 if data.get("error_msg") or data.get("error_code") is not None:122 msg = data.get("error_msg") or data.get("error") or "gateway error"123 code = data.get("error_code")124 if code is not None:125 return f"[{code}] {msg}"126 return str(msg)127 if data.get("error"):128 return str(data["error"])129 if data.get("status") is False:130 return fallback or str(data)131 return None132 133 134def submit_async_translate(135 text: str,136 lang: str,137 *,138 base_url: str = DEFAULT_BASE,139 api_key: str | None = None,140 timeout: int = 60,141) -> dict[str, Any]:142 """POST /translate/async → {taskId, state}."""143 r = requests.post(144 _api_url(base_url, "/translate/async"),145 json=translate_payload(text, lang),146 headers=_request_headers(api_key),147 timeout=timeout,148 )149 if r.status_code >= 400:150 raise RuntimeError(_http_error_message(r))151 try:152 data = r.json()153 except (json.JSONDecodeError, ValueError) as exc:154 raise RuntimeError(f"invalid JSON from async submit: {r.text[:500]}") from exc155 156 gateway_err = _gateway_error_message(data)157 if gateway_err:158 raise RuntimeError(gateway_err)159 160 task_id = data.get("taskId")161 if not task_id:162 raise RuntimeError(f"missing taskId in submit response: {data}")163 return {164 "task_id": str(task_id),165 "state": data.get("state", "pending"),166 "billing_amount": r.headers.get("X-Openapi-Amount"),167 "raw": data,168 }169 170 171def get_async_translate_result(172 task_id: str,173 *,174 base_url: str = DEFAULT_BASE,175 api_key: str | None = None,176 timeout: int = 60,177) -> dict[str, Any]:178 """GET /translate/async/{taskId}."""179 r = requests.get(180 _api_url(base_url, f"/translate/async/{task_id}"),181 headers=_request_headers(api_key, json_body=False),182 timeout=timeout,183 )184 if r.status_code == 404:185 raise RuntimeError(f"task not found: {_http_error_message(r)}")186 if r.status_code >= 400:187 raise RuntimeError(_http_error_message(r))188 data = r.json()189 return {190 "task_id": data.get("taskId", task_id),191 "state": data.get("state"),192 "text_translated": _translated_text(data),193 "translated_character_count": data.get("translatedCharacterCount"),194 "error": data.get("error"),195 "billing_amount": r.headers.get("X-Openapi-Amount"),196 "raw": data,197 }198 199 200def iter_translate_fast(201 text: str,202 lang: str,203 *,204 base_url: str = DEFAULT_BASE,205 api_key: str | None = None,206 poll_interval_s: float = DEFAULT_POLL_INTERVAL_S,207 poll_timeout_s: float = DEFAULT_POLL_TIMEOUT_S,208) -> Iterator[dict[str, Any]]:209 """Submit async fast translation and poll until succeeded/failed.210 211 Yields status events, then a final event with ``done=True``.212 """213 submitted = submit_async_translate(text, lang, base_url=base_url, api_key=api_key)214 task_id = submitted["task_id"]215 yield {216 "done": False,217 "task_id": task_id,218 "state": submitted.get("state", "pending"),219 "phase": "submitted",220 }221 222 deadline = time.perf_counter() + poll_timeout_s223 while True:224 if time.perf_counter() > deadline:225 raise RuntimeError(226 f"async translation timed out after {poll_timeout_s:.0f}s "227 f"(taskId={task_id})"228 )229 time.sleep(poll_interval_s)230 result = get_async_translate_result(231 task_id, base_url=base_url, api_key=api_key232 )233 state = result.get("state")234 if state in ("succeeded", "failed"):235 if state == "failed":236 err = result.get("error") or "async translation failed"237 raise RuntimeError(str(err))238 translated = result.get("text_translated") or ""239 if not str(translated).strip():240 raise RuntimeError("empty translation from async result")241 yield {242 "done": True,243 "task_id": task_id,244 "state": "succeeded",245 "phase": "done",246 "text_original": text,247 "text_translated": translated,248 "translated_character_count": result.get(249 "translated_character_count"250 ),251 "billing_amount": result.get("billing_amount"),252 "raw": result.get("raw"),253 }254 return255 yield {256 "done": False,257 "task_id": task_id,258 "state": state or "pending",259 "phase": "polling",260 }261 262 263def translate_fast(264 text: str,265 lang: str,266 *,267 base_url: str = DEFAULT_BASE,268 api_key: str | None = None,269 poll_interval_s: float = DEFAULT_POLL_INTERVAL_S,270 poll_timeout_s: float = DEFAULT_POLL_TIMEOUT_S,271 timeout: int | None = None,272) -> dict[str, Any]:273 """Async fast translate: submit + poll until complete (blocking)."""274 if timeout is not None:275 poll_timeout_s = float(timeout)276 final: dict[str, Any] | None = None277 for event in iter_translate_fast(278 text,279 lang,280 base_url=base_url,281 api_key=api_key,282 poll_interval_s=poll_interval_s,283 poll_timeout_s=poll_timeout_s,284 ):285 if event.get("done"):286 final = event287 if final is None:288 raise RuntimeError("async translation ended without a result")289 return {290 "state": "success",291 "task_id": final.get("task_id"),292 "text_original": final.get("text_original", text),293 "text_translated": final.get("text_translated", ""),294 "translated_character_count": final.get("translated_character_count"),295 "billing_amount": final.get("billing_amount"),296 "raw": final.get("raw"),297 }298 299 300def _iter_sse_json(resp: requests.Response) -> Iterator[dict[str, Any]]:301 for raw in resp.iter_lines(decode_unicode=True):302 if not raw:303 continue304 line = raw.strip()305 if not line.startswith("data:"):306 continue307 payload = line[5:].lstrip()308 if not payload:309 continue310 chunk = json.loads(payload)311 if isinstance(chunk, dict) and chunk.get("error"):312 raise RuntimeError(str(chunk["error"]))313 if isinstance(chunk, dict):314 yield {315 "state": chunk.get("state", "success"),316 "text_original": _original_text(chunk),317 "text_translated": _translated_text(chunk),318 "translated_character_count": chunk.get("translatedCharacterCount"),319 "progress": chunk.get("progress"),320 "raw": chunk,321 }322 323 324def stream_translate(325 text: str,326 lang: str,327 *,328 base_url: str = DEFAULT_BASE,329 api_key: str | None = None,330 timeout: int = 1800,331) -> Iterator[dict[str, Any]]:332 """POST /translate (stream-only; do not send mode)."""333 with requests.post(334 _api_url(base_url, "/translate"),335 json=translate_payload(text, lang),336 headers=_request_headers(api_key),337 stream=True,338 timeout=timeout,339 ) as resp:340 if resp.status_code >= 400:341 raise RuntimeError(_http_error_message(resp))342 yield from _iter_sse_json(resp)343 344 345def normalize_smartdoc_markdown(markdown: str) -> str:346 """Fix gateway-escaped newlines without breaking LaTeX commands like \\neq."""347 text = markdown348 text = re.sub(r"\\r\\n(?![A-Za-z])", "\n", text)349 text = re.sub(r"\\n(?![A-Za-z])", "\n", text)350 text = re.sub(r"\\r(?![A-Za-z])", "\n", text)351 return text352 353 354def parse_document(355 file_path: str,356 *,357 api_key: str | None = None,358 output_format: str = "both",359 timeout: int = 300,360 max_bytes: int = MAX_SMARTDOC_BYTES,361) -> dict[str, Any]:362 """POST multipart to Smart Doc doc_parsing endpoint."""363 path = Path(file_path)364 if not path.is_file():365 raise RuntimeError(f"file not found: {file_path}")366 size = path.stat().st_size367 if size > max_bytes:368 raise RuntimeError(369 f"The file must not exceed {max_bytes // (1024 * 1024)} MB"370 )371 372 key = resolve_api_key(api_key)373 if not key:374 raise RuntimeError(375 f"API key missing. Set {API_KEY_ENV} (or {API_KEY_ENV_FALLBACK}) "376 "under Space Settings → Variables / Secrets"377 )378 379 with path.open("rb") as fh:380 files = {"file": (path.name, fh)}381 data = {"output_format": output_format}382 r = requests.post(383 SMARTDOC_URL,384 headers=_request_headers(key, json_body=False),385 files=files,386 data=data,387 timeout=timeout,388 )389 390 try:391 body = r.json()392 except json.JSONDecodeError as exc:393 raise RuntimeError(f"HTTP {r.status_code}: {r.text[:500]}") from exc394 395 if r.status_code >= 400 or body.get("status") != "success":396 raise RuntimeError(397 body.get("error_msg")398 or body.get("message")399 or body.get("error")400 or f"HTTP {r.status_code}"401 )402 403 payload = body.get("data") if isinstance(body.get("data"), dict) else body404 markdown = payload.get("markdown") if isinstance(payload, dict) else ""405 if not isinstance(markdown, str):406 markdown = ""407 results = payload.get("results") if isinstance(payload, dict) else []408 if not isinstance(results, list):409 results = []410 total_pages = payload.get("total_pages") if isinstance(payload, dict) else 0411 412 return {413 "markdown": normalize_smartdoc_markdown(markdown),414 "results": results,415 "total_pages": total_pages or 0,416 "raw": body,417 }418 