Team Ai
Apppublic

PatSnap/Document-Processing

sourceHugging Faceapache-2.0updated 3mo agoView on Hugging Face
10likes
api_client.py418 linesDownload Raw Back to root
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