Document-Processing / api_client.py
Jiawei Dong
api fixed
d77d8e5
Raw
History Blame Contribute Delete
13.3 kB
"""HTTP clients for HIRO Translation and Smart Doc (public connect gateway)."""
from __future__ import annotations
import json
import os
import re
import time
from collections.abc import Iterator
from pathlib import Path
from typing import Any
import requests
# Public gateways (matches connect.zhihuiya.com curl examples).
DEFAULT_BASE = "https://connect.zhihuiya.com/hiro_translation"
SMARTDOC_URL = "https://connect.zhihuiya.com/rd-llm/v1/documents/doc_parsing"
API_KEY_ENV = "HIRO_API_KEY"
API_KEY_ENV_FALLBACK = "RD_LLM_API_KEY"
MAX_SMARTDOC_BYTES = 10 * 1024 * 1024
# Async fast: submit + poll (aligned with hiro-translation-api).
DEFAULT_POLL_INTERVAL_S = 1.5
DEFAULT_POLL_TIMEOUT_S = 1800
def resolve_api_key(api_key: str | None = None) -> str:
"""Prefer explicit key, then HIRO_API_KEY, then RD_LLM_API_KEY."""
if api_key is not None and str(api_key).strip():
return str(api_key).strip()
return (
os.environ.get(API_KEY_ENV, "").strip()
or os.environ.get(API_KEY_ENV_FALLBACK, "").strip()
)
def lang_to_codes(lang: str) -> tuple[str, str]:
src, _, tgt = lang.partition("2")
if not src or not tgt:
raise ValueError(f"invalid lang {lang!r}; expected format like zh2en")
return src, tgt
def translate_payload(content: str, lang: str) -> dict[str, Any]:
"""Request body for POST /translate and POST /translate/async (no mode field)."""
source, target = lang_to_codes(lang)
return {
"content": content,
"sourceLanguageCode": source,
"targetLanguageCode": target,
}
def _api_url(base_url: str, path: str) -> str:
return f"{base_url.rstrip('/')}{path}"
def _request_headers(
api_key: str | None = None,
*,
json_body: bool = True,
) -> dict[str, str]:
headers: dict[str, str] = {}
if json_body:
headers["Content-Type"] = "application/json"
key = resolve_api_key(api_key)
if key:
headers["Authorization"] = f"Bearer {key}"
return headers
def _translated_text(data: dict[str, Any]) -> str:
value = data.get("textTranslated", data.get("text_translated", ""))
return value if isinstance(value, str) else ""
def _original_text(data: dict[str, Any], fallback: str = "") -> str:
value = data.get("textOriginal", data.get("text_original", fallback))
return value if isinstance(value, str) else fallback
def _http_error_message(resp: requests.Response) -> str:
try:
body = resp.json()
if isinstance(body, dict):
return str(body.get("error") or body.get("message") or body)
except (json.JSONDecodeError, ValueError):
pass
return resp.text[:500]
def health_ok(
base_url: str = DEFAULT_BASE,
*,
api_key: str | None = None,
) -> tuple[bool, str]:
try:
r = requests.get(
_api_url(base_url, "/health"),
headers=_request_headers(api_key),
timeout=12,
)
r.raise_for_status()
data = r.json()
if data.get("status") != "OK":
return False, f"Unexpected response: {r.text[:200]}"
upstream = data.get("upstream", "UNKNOWN")
if upstream == "OK":
return True, "Healthy · upstream OK"
return False, f"Gateway up, upstream unavailable ({upstream})"
except requests.RequestException as exc:
return False, str(exc)
def _gateway_error_message(data: Any, fallback: str = "") -> str | None:
"""Extract connect-gateway style errors (error_msg / error_code)."""
if not isinstance(data, dict):
return None
if data.get("error_msg") or data.get("error_code") is not None:
msg = data.get("error_msg") or data.get("error") or "gateway error"
code = data.get("error_code")
if code is not None:
return f"[{code}] {msg}"
return str(msg)
if data.get("error"):
return str(data["error"])
if data.get("status") is False:
return fallback or str(data)
return None
def submit_async_translate(
text: str,
lang: str,
*,
base_url: str = DEFAULT_BASE,
api_key: str | None = None,
timeout: int = 60,
) -> dict[str, Any]:
"""POST /translate/async → {taskId, state}."""
r = requests.post(
_api_url(base_url, "/translate/async"),
json=translate_payload(text, lang),
headers=_request_headers(api_key),
timeout=timeout,
)
if r.status_code >= 400:
raise RuntimeError(_http_error_message(r))
try:
data = r.json()
except (json.JSONDecodeError, ValueError) as exc:
raise RuntimeError(f"invalid JSON from async submit: {r.text[:500]}") from exc
gateway_err = _gateway_error_message(data)
if gateway_err:
raise RuntimeError(gateway_err)
task_id = data.get("taskId")
if not task_id:
raise RuntimeError(f"missing taskId in submit response: {data}")
return {
"task_id": str(task_id),
"state": data.get("state", "pending"),
"billing_amount": r.headers.get("X-Openapi-Amount"),
"raw": data,
}
def get_async_translate_result(
task_id: str,
*,
base_url: str = DEFAULT_BASE,
api_key: str | None = None,
timeout: int = 60,
) -> dict[str, Any]:
"""GET /translate/async/{taskId}."""
r = requests.get(
_api_url(base_url, f"/translate/async/{task_id}"),
headers=_request_headers(api_key, json_body=False),
timeout=timeout,
)
if r.status_code == 404:
raise RuntimeError(f"task not found: {_http_error_message(r)}")
if r.status_code >= 400:
raise RuntimeError(_http_error_message(r))
data = r.json()
return {
"task_id": data.get("taskId", task_id),
"state": data.get("state"),
"text_translated": _translated_text(data),
"translated_character_count": data.get("translatedCharacterCount"),
"error": data.get("error"),
"billing_amount": r.headers.get("X-Openapi-Amount"),
"raw": data,
}
def iter_translate_fast(
text: str,
lang: str,
*,
base_url: str = DEFAULT_BASE,
api_key: str | None = None,
poll_interval_s: float = DEFAULT_POLL_INTERVAL_S,
poll_timeout_s: float = DEFAULT_POLL_TIMEOUT_S,
) -> Iterator[dict[str, Any]]:
"""Submit async fast translation and poll until succeeded/failed.
Yields status events, then a final event with ``done=True``.
"""
submitted = submit_async_translate(text, lang, base_url=base_url, api_key=api_key)
task_id = submitted["task_id"]
yield {
"done": False,
"task_id": task_id,
"state": submitted.get("state", "pending"),
"phase": "submitted",
}
deadline = time.perf_counter() + poll_timeout_s
while True:
if time.perf_counter() > deadline:
raise RuntimeError(
f"async translation timed out after {poll_timeout_s:.0f}s "
f"(taskId={task_id})"
)
time.sleep(poll_interval_s)
result = get_async_translate_result(
task_id, base_url=base_url, api_key=api_key
)
state = result.get("state")
if state in ("succeeded", "failed"):
if state == "failed":
err = result.get("error") or "async translation failed"
raise RuntimeError(str(err))
translated = result.get("text_translated") or ""
if not str(translated).strip():
raise RuntimeError("empty translation from async result")
yield {
"done": True,
"task_id": task_id,
"state": "succeeded",
"phase": "done",
"text_original": text,
"text_translated": translated,
"translated_character_count": result.get(
"translated_character_count"
),
"billing_amount": result.get("billing_amount"),
"raw": result.get("raw"),
}
return
yield {
"done": False,
"task_id": task_id,
"state": state or "pending",
"phase": "polling",
}
def translate_fast(
text: str,
lang: str,
*,
base_url: str = DEFAULT_BASE,
api_key: str | None = None,
poll_interval_s: float = DEFAULT_POLL_INTERVAL_S,
poll_timeout_s: float = DEFAULT_POLL_TIMEOUT_S,
timeout: int | None = None,
) -> dict[str, Any]:
"""Async fast translate: submit + poll until complete (blocking)."""
if timeout is not None:
poll_timeout_s = float(timeout)
final: dict[str, Any] | None = None
for event in iter_translate_fast(
text,
lang,
base_url=base_url,
api_key=api_key,
poll_interval_s=poll_interval_s,
poll_timeout_s=poll_timeout_s,
):
if event.get("done"):
final = event
if final is None:
raise RuntimeError("async translation ended without a result")
return {
"state": "success",
"task_id": final.get("task_id"),
"text_original": final.get("text_original", text),
"text_translated": final.get("text_translated", ""),
"translated_character_count": final.get("translated_character_count"),
"billing_amount": final.get("billing_amount"),
"raw": final.get("raw"),
}
def _iter_sse_json(resp: requests.Response) -> Iterator[dict[str, Any]]:
for raw in resp.iter_lines(decode_unicode=True):
if not raw:
continue
line = raw.strip()
if not line.startswith("data:"):
continue
payload = line[5:].lstrip()
if not payload:
continue
chunk = json.loads(payload)
if isinstance(chunk, dict) and chunk.get("error"):
raise RuntimeError(str(chunk["error"]))
if isinstance(chunk, dict):
yield {
"state": chunk.get("state", "success"),
"text_original": _original_text(chunk),
"text_translated": _translated_text(chunk),
"translated_character_count": chunk.get("translatedCharacterCount"),
"progress": chunk.get("progress"),
"raw": chunk,
}
def stream_translate(
text: str,
lang: str,
*,
base_url: str = DEFAULT_BASE,
api_key: str | None = None,
timeout: int = 1800,
) -> Iterator[dict[str, Any]]:
"""POST /translate (stream-only; do not send mode)."""
with requests.post(
_api_url(base_url, "/translate"),
json=translate_payload(text, lang),
headers=_request_headers(api_key),
stream=True,
timeout=timeout,
) as resp:
if resp.status_code >= 400:
raise RuntimeError(_http_error_message(resp))
yield from _iter_sse_json(resp)
def normalize_smartdoc_markdown(markdown: str) -> str:
"""Fix gateway-escaped newlines without breaking LaTeX commands like \\neq."""
text = markdown
text = re.sub(r"\\r\\n(?![A-Za-z])", "\n", text)
text = re.sub(r"\\n(?![A-Za-z])", "\n", text)
text = re.sub(r"\\r(?![A-Za-z])", "\n", text)
return text
def parse_document(
file_path: str,
*,
api_key: str | None = None,
output_format: str = "both",
timeout: int = 300,
max_bytes: int = MAX_SMARTDOC_BYTES,
) -> dict[str, Any]:
"""POST multipart to Smart Doc doc_parsing endpoint."""
path = Path(file_path)
if not path.is_file():
raise RuntimeError(f"file not found: {file_path}")
size = path.stat().st_size
if size > max_bytes:
raise RuntimeError(
f"The file must not exceed {max_bytes // (1024 * 1024)} MB"
)
key = resolve_api_key(api_key)
if not key:
raise RuntimeError(
f"API key missing. Set {API_KEY_ENV} (or {API_KEY_ENV_FALLBACK}) "
"under Space Settings → Variables / Secrets"
)
with path.open("rb") as fh:
files = {"file": (path.name, fh)}
data = {"output_format": output_format}
r = requests.post(
SMARTDOC_URL,
headers=_request_headers(key, json_body=False),
files=files,
data=data,
timeout=timeout,
)
try:
body = r.json()
except json.JSONDecodeError as exc:
raise RuntimeError(f"HTTP {r.status_code}: {r.text[:500]}") from exc
if r.status_code >= 400 or body.get("status") != "success":
raise RuntimeError(
body.get("error_msg")
or body.get("message")
or body.get("error")
or f"HTTP {r.status_code}"
)
payload = body.get("data") if isinstance(body.get("data"), dict) else body
markdown = payload.get("markdown") if isinstance(payload, dict) else ""
if not isinstance(markdown, str):
markdown = ""
results = payload.get("results") if isinstance(payload, dict) else []
if not isinstance(results, list):
results = []
total_pages = payload.get("total_pages") if isinstance(payload, dict) else 0
return {
"markdown": normalize_smartdoc_markdown(markdown),
"results": results,
"total_pages": total_pages or 0,
"raw": body,
}