"""火山方舟 (Ark) 通用 HTTP 客户端,共享给 seedream / 未来 seedance 等媒体工具。 litellm 不覆盖豆包的图像/视频生成端点,这里自己用 httpx 直调 OpenAI 兼容路径 /images/generations 与异步任务 /contents/generations/tasks。 """ from __future__ import annotations from dataclasses import dataclass from pathlib import Path from typing import Any, Optional import httpx import yaml from core.paths import ROOT _DOUBAO_YAML = ROOT / "config" / "media" / "doubao.yaml" class ArkError(RuntimeError): """ark API 调用失败的统一异常。""" class ArkTimeoutError(ArkError): """可重试的瞬时失败:请求超时 / 网络抖动(非业务错误)。 HTTP 4xx/5xx 业务错误仍抛普通 ArkError(不该重试,重试也是同样的错)。 caller 可单独 catch 本子类做退避重试;catch ArkError 仍能兜住(isinstance)。 """ @dataclass class ArkConfig: api_key: str base_url: str raw: dict # 完整 yaml 内容(便于 caller 按 image/video 子键再取) api_key_env: str = "" def request_api_key(self) -> str: from core.provider_credentials.runtime import resolve_env_secret return resolve_env_secret(self.api_key_env) if self.api_key_env else self.api_key @classmethod def load(cls, path: Optional[Path] = None) -> Optional["ArkConfig"]: """读媒体 provider yaml(默认 doubao.yaml)+ 解析 env 拿 api_key。 配置键通用名 `api_key_env` / `base_url`(unifyllm.yaml 等新 provider 用); `ark_api_key_env` / `ark_base_url` 是 doubao.yaml 的历史键名,兜底兼容。 api_key env 未设 → 返 None(caller 据此决定是否注册 tool;无 key 用户无感知)。 yaml 不存在 → 返 None。 """ p = path or _DOUBAO_YAML if not p.exists(): return None data = yaml.safe_load(p.read_text(encoding="utf-8")) or {} env = data.get("api_key_env") or data.get("ark_api_key_env") or "ARK_API_KEY" from core.provider_credentials.runtime import resolve_env_secret key = resolve_env_secret(env) if not key: return None base = ( data.get("base_url") or data.get("ark_base_url") or "https://ark.cn-beijing.volces.com/api/v3" ) return cls(api_key=key, base_url=str(base).rstrip("/"), raw=data, api_key_env=env) class ArkClient: """轻量 httpx 封装:统一 base_url + bearer auth + 异常翻译。 成功返 dict(JSON 已解析);非 2xx / 网络异常 / JSON 不可解析都抛 ArkError。 """ def __init__(self, cfg: ArkConfig, timeout_s: float = 60.0) -> None: self.cfg = cfg self.timeout_s = timeout_s self._api_key = cfg.request_api_key() self._client = httpx.Client( base_url=cfg.base_url, headers={ "Authorization": f"Bearer {self._api_key}", }, timeout=timeout_s, ) def post_json(self, path: str, body: dict, *, timeout_s: Optional[float] = None) -> dict: try: resp = self._client.post(path, json=body, timeout=timeout_s or self.timeout_s) except httpx.TimeoutException as e: raise ArkTimeoutError(f"timeout calling POST {path}: {e}") from e except httpx.HTTPError as e: raise ArkTimeoutError(f"network error calling POST {path}: {e}") from e return self._parse(resp, f"POST {path}") def post_multipart( self, path: str, data: dict[str, str], files: dict[str, tuple[str, bytes, str]], *, timeout_s: float | None = None, ) -> dict: """POST multipart/form-data;边界与 Content-Type 交给 httpx 生成。""" try: resp = self._client.post( path, data=data, files=files, timeout=timeout_s or self.timeout_s, ) except httpx.TimeoutException as e: raise ArkTimeoutError(f"timeout calling POST {path}: {e}") from e except httpx.HTTPError as e: raise ArkTimeoutError(f"network error calling POST {path}: {e}") from e return self._parse(resp, f"POST {path}") def get_json(self, path: str, *, timeout_s: Optional[float] = None) -> dict: try: resp = self._client.get(path, timeout=timeout_s or self.timeout_s) except httpx.TimeoutException as e: raise ArkTimeoutError(f"timeout calling GET {path}: {e}") from e except httpx.HTTPError as e: raise ArkTimeoutError(f"network error calling GET {path}: {e}") from e return self._parse(resp, f"GET {path}") def _parse(self, resp: httpx.Response, label: str) -> dict: if resp.status_code >= 400: try: from core.provider_credentials.registry import BY_ENV from core.provider_credentials.service import record_business_failure binding = BY_ENV.get(self.cfg.api_key_env) if binding: record_business_failure( binding[0], status_code=resp.status_code, detail=resp.text[:300] ) except Exception: pass # ark 错误 body 一般是 {"error": {"code": ..., "message": ...}};能解就解 try: err = resp.json().get("error") or {} msg = err.get("message") or resp.text[:300] key = self._api_key msg = str(msg).replace(key, "***") if key else str(msg) code = err.get("code") or resp.status_code raise ArkError(f"{label} → HTTP {resp.status_code} ({code}): {msg}") except ValueError: detail = resp.text[:300] key = self._api_key if key: detail = detail.replace(key, "***") raise ArkError(f"{label} → HTTP {resp.status_code}: {detail}") try: return resp.json() except ValueError as e: raise ArkError(f"{label} → invalid JSON response: {e}") from e def download(self, url: str, dest: Path, *, timeout_s: float = 120.0) -> None: """跨域下载产物(image/video URL 是火山 CDN,不带 ark auth)。""" try: with httpx.stream("GET", url, timeout=timeout_s) as r: if r.status_code >= 400: raise ArkError(f"download {url} → HTTP {r.status_code}") dest.parent.mkdir(parents=True, exist_ok=True) with open(dest, "wb") as f: for chunk in r.iter_bytes(chunk_size=64 * 1024): f.write(chunk) except httpx.HTTPError as e: raise ArkError(f"download {url} failed: {e}") from e def close(self) -> None: self._client.close() def __enter__(self) -> "ArkClient": return self def __exit__(self, *_exc: Any) -> None: self.close()