"""豆包大模型 API 客户端(共享层). API 和 Worker 两边共用。基于火山引擎方舟平台的 OpenAI 兼容接口。 使用方式: from packages.shared.ai_client import get_doubao_client client = get_doubao_client() if client.is_available: result = client.chat_completion(messages=[...]) """ from __future__ import annotations import logging import time from typing import Any, Optional import httpx from packages.shared.config import get_shared_settings logger = logging.getLogger(__name__) class DoubaoClient: """豆包大模型 API 客户端. 封装 OpenAI 兼容的 Chat Completion 接口,支持自动重试。 未配置 API Key 时 is_available 为 False,调用方应降级处理。 """ def __init__(self) -> None: settings = get_shared_settings() self.api_key: str = settings.doubao_api_key self.model: str = settings.doubao_model self.base_url: str = settings.doubao_base_url.rstrip("/") self.timeout: int = settings.doubao_timeout self.max_retries: int = settings.doubao_max_retries self.vision_model: str = settings.doubao_vision_model def embed_text(self, text: str, timeout: int | None = None) -> list[float] | None: """调用豆包文本 Embedding API,返回浮点向量;失败返回 None。""" if not self.is_available or not text or not text.strip(): return None url = f"{self.base_url}/embeddings" headers = { "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json", } payload: dict[str, Any] = { "model": getattr(self, "embedding_model", None) or "doubao-embedding-large-text-240915", "input": text.strip(), "encoding_format": "float", } req_timeout = timeout or self.timeout last_error: Exception | None = None for attempt in range(self.max_retries + 1): try: resp = httpx.post(url, headers=headers, json=payload, timeout=req_timeout) resp.raise_for_status() data = resp.json() emb_list = data.get("data") or [] if emb_list and isinstance(emb_list, list): vec = emb_list[0].get("embedding") if isinstance(vec, list) and vec: return [float(x) for x in vec] logger.warning("embedding 返回结构异常: %s", str(data)[:200]) return None except Exception as e: last_error = e if attempt < self.max_retries: wait = 0.5 * (2**attempt) logger.warning( "豆包 Embedding 调用失败,%.1fs 后重试 (%d/%d): %s", wait, attempt + 1, self.max_retries + 1, e ) time.sleep(wait) logger.error("豆包 Embedding 调用最终失败: %s", last_error) return None @property def is_available(self) -> bool: """是否可用(配置了 API Key).""" return bool(self.api_key) def chat_completion( self, messages: list[dict[str, str]], temperature: float = 0.7, max_tokens: int = 1024, ) -> Optional[str]: """调用 Chat Completion 接口. Args: messages: 对话消息列表,[{"role": "user"/"system"/"assistant", "content": "..."}] temperature: 采样温度,0-2,默认0.7 max_tokens: 最大生成token数,默认1024 Returns: 模型返回的文本内容,失败返回 None """ if not self.is_available: return None url = f"{self.base_url}/chat/completions" headers = { "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json", } payload: dict[str, Any] = { "model": self.model, "messages": messages, "temperature": temperature, "max_tokens": max_tokens, } last_error: Optional[Exception] = None for attempt in range(self.max_retries + 1): try: response = httpx.post( url, headers=headers, json=payload, timeout=self.timeout, ) response.raise_for_status() data = response.json() content = data["choices"][0]["message"]["content"] return content.strip() except Exception as e: last_error = e if attempt < self.max_retries: wait = 0.5 * (2**attempt) logger.warning( "豆包API调用失败,%.1fs后重试 (第%d/%d次): %s", wait, attempt + 1, self.max_retries + 1, e, ) time.sleep(wait) logger.error("豆包API调用最终失败: %s", last_error) return None def vision_completion( self, messages: list[dict], images: list[str] | None = None, max_tokens: int = 2048, temperature: float = 0.3, timeout: int | None = None, ) -> Optional[str]: """调用豆包视觉理解 API(OpenAI 兼容多模态格式). 将 images 附加到最后一条 user message 的 content 中, 使用 vision_model(默认 doubao-1-5-vision-pro-250915)。 Args: messages: 对话消息列表。最后一条 user message 会被注入图片内容。 images: 图片列表,支持 base64 data URI 或 HTTP(S) URL。 max_tokens: 最大生成 token 数,默认 2048。 temperature: 采样温度,默认 0.3(视觉任务偏低更稳定)。 timeout: 单次请求超时秒数,不传则使用默认 self.timeout。 Returns: 模型返回的文本内容,失败返回 None。 """ if not self.is_available: return None # 构造多模态 content:先追加文本,再追加图片 vision_messages = [] for msg in messages: vision_messages.append(dict(msg)) # 将图片注入最后一条 user message if images and vision_messages: # 找到最后一条 user message for i in range(len(vision_messages) - 1, -1, -1): if vision_messages[i].get("role") == "user": text_content = vision_messages[i].get("content", "") multi_content: list[dict[str, Any]] = [] if text_content: multi_content.append({"type": "text", "text": text_content}) for img in images: if img.startswith("data:") or img.startswith("http://") or img.startswith("https://"): multi_content.append({"type": "image_url", "image_url": {"url": img}}) else: # 当作 base64 编码 multi_content.append( {"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{img}"}} ) vision_messages[i]["content"] = multi_content break url = f"{self.base_url}/chat/completions" headers = { "Authorization": f"Bearer {self.api_key}", "Content-Type": "application/json", } payload: dict[str, Any] = { "model": self.vision_model, "messages": vision_messages, "temperature": temperature, "max_tokens": max_tokens, } req_timeout = timeout or self.timeout last_error: Optional[Exception] = None for attempt in range(self.max_retries + 1): try: response = httpx.post( url, headers=headers, json=payload, timeout=req_timeout, ) response.raise_for_status() data = response.json() content = data["choices"][0]["message"]["content"] return content.strip() except Exception as e: last_error = e if attempt < self.max_retries: wait = 0.5 * (2**attempt) logger.warning( "豆包视觉API调用失败,%.1fs后重试 (第%d/%d次): %s", wait, attempt + 1, self.max_retries + 1, e, ) time.sleep(wait) logger.error("豆包视觉API调用最终失败: %s", last_error) return None # ── 单例 ───────────────────────────────────────────────────────────────────── _client: Optional[DoubaoClient] = None def get_doubao_client() -> DoubaoClient: """获取豆包客户端单例.""" global _client if _client is None: _client = DoubaoClient() return _client