9a4206e65c
CI/CD Pipeline / Check push changed paths (pull_request) Has been skipped
CI/CD Pipeline / Build Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Check if frontend-only change (pull_request) Successful in 2s
CI/CD Pipeline / Dedup Check - skip PR tests when covered by push pipeline (pull_request) Successful in 3s
CI/CD Pipeline / Build Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Build Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Frontend Lint (pull_request) Has been skipped
CI/CD Pipeline / Frontend Unit Tests (pull_request) Has been skipped
CI/CD Pipeline / PR Build Web Image (pull_request) Has been skipped
CI/CD Pipeline / PR Build API Image (pull_request) Successful in 43s
CI/CD Pipeline / PR Build Worker Image (pull_request) Successful in 1m14s
CI/CD Pipeline / Retag skipped Staging API Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Web Image (pull_request) Has been skipped
CI/CD Pipeline / Retag skipped Staging Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Staging (Watchtower auto-deploy) (pull_request) Has been skipped
CI/CD Pipeline / Staging E2E Tests (pull_request) Has been skipped
CI/CD Pipeline / Staging API Integration Tests (pull_request) Has been skipped
CI/CD Pipeline / ACR Image Cleanup (pull_request) Has been skipped
Preview Deploy / Deploy Preview Environment (pull_request) Successful in 1m21s
PR Automation / Auto Approve on CI Green (pull_request) Successful in 3m6s
CI/CD Pipeline / Integration Tests (pull_request) Successful in 5m41s
CI/CD Pipeline / Validate - Style (pull_request) Successful in 6m22s
CI/CD Pipeline / Validate - Python (mypy + alembic) (pull_request) Successful in 6m36s
AI Code Review / AI Code Review (pull_request) Successful in 6m50s
CI/CD Pipeline / Unit Tests (pull_request) Successful in 13m4s
PR Automation / Auto Merge on CI Green + Approved (pull_request) Successful in 10m51s
CI/CD Pipeline / Validate - Security (pull_request) Successful in 15m30s
CI/CD Pipeline / Build Production API Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Web Image (pull_request) Has been skipped
CI/CD Pipeline / Build Production Worker Image (pull_request) Has been skipped
CI/CD Pipeline / Deploy Production (pull_request) Has been skipped
CI/CD Pipeline / Canary Release to Production (pull_request) Has been skipped
CI/CD Pipeline / Production Browser E2E (pull_request) Has been skipped
CI/CD Pipeline / CI Gate (pull_request) Successful in 1s
ACR Cleanup / ACR Image Cleanup (pull_request_target) Successful in 33s
Preview Cleanup / Cleanup Preview Environment (pull_request) Successful in 1m15s
521 lines
22 KiB
Python
Executable File
521 lines
22 KiB
Python
Executable File
"""豆包大模型 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 os
|
||
import time
|
||
import uuid
|
||
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
|
||
self.vision_lite_model: str = settings.doubao_vision_lite_model
|
||
self.fast_model: str = settings.doubao_fast_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,
|
||
model: str | None = None,
|
||
) -> 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": model or 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,
|
||
model: str | None = None,
|
||
) -> Optional[str]:
|
||
"""调用豆包视觉理解 API(OpenAI 兼容多模态格式).
|
||
|
||
将 images 附加到最后一条 user message 的 content 中,
|
||
使用 vision_model(默认 doubao-1-5-vision-pro-250328)。
|
||
|
||
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": model or 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
|
||
|
||
# ── 视频生成(Seedance 2.5,异步任务)────────────────────────────
|
||
|
||
def video_generation(
|
||
self,
|
||
prompt: str,
|
||
*,
|
||
image_url: str | None = None,
|
||
duration: int = 5,
|
||
ratio: str | None = "9:16",
|
||
resolution: str = "720p",
|
||
generate_audio: bool = True,
|
||
watermark: bool = False,
|
||
output_dir: str | None = None,
|
||
model: str | None = None,
|
||
reference_images: list[str] | None = None,
|
||
reference_audios: list[str] | None = None,
|
||
reference_videos: list[str] | None = None,
|
||
) -> str | None:
|
||
"""调用 Seedance 2.5 生视频(异步任务→轮询→下载),返回本地 MP4 路径;失败返回 None。
|
||
|
||
【v1.6.1 修复】严格按官方 content 数组协议构造请求:
|
||
- 所有参考(图/音/视)必须放进 content 数组并带 role 字段,不能放顶层 reference_audios/reference_videos(非官方字段,会被忽略或导致异常)。
|
||
- 首帧图(first_frame 模式)Seedance 2.5 强制 ratio=adaptive;走 omni_reference(参考生视频)模式时才能指定 9:16/1:1 等具体比例。
|
||
判定:传了参考音频/视频或 ≥1 张多参考图时,走 omni_reference(首张图 role=reference_image);纯首帧无参考时走 first_frame(ratio 强制 adaptive)。
|
||
- 创建任务若因 ratio 报错(HTTP 400),自动回退到 ratio=adaptive 重试一次。
|
||
"""
|
||
if not self.is_available:
|
||
return None
|
||
if not prompt or not prompt.strip():
|
||
return None
|
||
|
||
settings = get_shared_settings()
|
||
poll_interval = getattr(settings, "doubao_video_poll_interval", 10) or 10
|
||
# 收紧总超时:轮询 8min + 下载 2min = 最长 ~10min,防止出现 20min 卡死
|
||
total_timeout = getattr(settings, "doubao_video_timeout", 480) or 480
|
||
default_video_model = getattr(settings, "doubao_video_model", None) or "doubao-seedance-2-5-260628"
|
||
video_model = model or default_video_model
|
||
|
||
ref_audios = [u for u in (reference_audios or [])[:10] if u and isinstance(u, str)]
|
||
ref_videos = [u for u in (reference_videos or [])[:3] if u and isinstance(u, str)]
|
||
ref_imgs = [u for u in (reference_images or [])[:9] if u and isinstance(u, str)]
|
||
|
||
# 判断任务模式:有参考音/视/多图 → omni_reference(支持指定 ratio);纯首帧 → first_frame(ratio=adaptive)
|
||
has_extra_refs = bool(ref_audios or ref_videos or ref_imgs)
|
||
is_first_frame_mode = bool(image_url) and not has_extra_refs
|
||
# 最终 ratio:first_frame 模式强制 adaptive,否则按用户传值(默认 9:16)
|
||
final_ratio = "adaptive" if is_first_frame_mode else (ratio or "9:16")
|
||
|
||
# 构造 content 数组:text + 图 + 音 + 视
|
||
content: list[dict[str, Any]] = [{"type": "text", "text": prompt.strip()}]
|
||
if image_url:
|
||
if has_extra_refs:
|
||
# omni_reference:首张图作为 reference_image,允许指定 ratio
|
||
content.append(
|
||
{
|
||
"type": "image_url",
|
||
"image_url": {"url": image_url},
|
||
"role": "reference_image",
|
||
}
|
||
)
|
||
else:
|
||
# 纯首帧:不带 role,服务端识别为 first_frame(或显式 role=first_frame)
|
||
content.append(
|
||
{
|
||
"type": "image_url",
|
||
"image_url": {"url": image_url},
|
||
"role": "first_frame",
|
||
}
|
||
)
|
||
for u in ref_imgs:
|
||
content.append({"type": "image_url", "image_url": {"url": u}, "role": "reference_image"})
|
||
for u in ref_audios:
|
||
content.append({"type": "audio_url", "audio_url": {"url": u}, "role": "reference_audio"})
|
||
for u in ref_videos:
|
||
content.append({"type": "video_url", "video_url": {"url": u}, "role": "reference_video"})
|
||
|
||
create_payload: dict[str, Any] = {
|
||
"model": video_model,
|
||
"content": content,
|
||
"generate_audio": bool(generate_audio),
|
||
"duration": int(duration),
|
||
"resolution": resolution,
|
||
"watermark": bool(watermark),
|
||
"ratio": final_ratio,
|
||
}
|
||
|
||
headers = {
|
||
"Authorization": f"Bearer {self.api_key}",
|
||
"Content-Type": "application/json",
|
||
}
|
||
create_url = f"{self.base_url}/contents/generations/tasks"
|
||
logger.info(
|
||
"Seedance 创建任务: model=%s dur=%ds ratio=%s mode=%s gen_audio=%s img=%d aud=%d vid=%d",
|
||
video_model,
|
||
duration,
|
||
final_ratio,
|
||
"first_frame" if is_first_frame_mode else "omni_ref",
|
||
generate_audio,
|
||
(1 if image_url else 0) + len(ref_imgs),
|
||
len(ref_audios),
|
||
len(ref_videos),
|
||
)
|
||
|
||
def _do_create(payload: dict) -> tuple[str | None, Exception | None, int, str]:
|
||
"""返回 (task_id, last_err, status_code, body_text)。"""
|
||
last_err: Exception | None = None
|
||
for attempt in range(self.max_retries + 1):
|
||
try:
|
||
resp = httpx.post(create_url, headers=headers, json=payload, timeout=self.timeout)
|
||
sc = int(getattr(resp, "status_code", 0) or 0)
|
||
body = (getattr(resp, "text", "") or "")[:1500]
|
||
if sc >= 400:
|
||
logger.error("Seedance 创建任务 HTTP %d: body=%s", sc, body)
|
||
try:
|
||
resp.raise_for_status()
|
||
except Exception as ee:
|
||
last_err = ee
|
||
if attempt < self.max_retries:
|
||
time.sleep(0.5 * (2**attempt))
|
||
continue
|
||
return None, last_err, sc, body
|
||
data = resp.json()
|
||
tid = data.get("id")
|
||
if tid:
|
||
return tid, None, sc, body
|
||
last_err = RuntimeError(f"create ok but no id: {str(data)[:300]}")
|
||
except Exception as e:
|
||
last_err = e
|
||
if attempt < self.max_retries:
|
||
wait = 0.5 * (2**attempt)
|
||
logger.warning(
|
||
"Seedance 创建任务失败,%.1fs 后重试 (%d/%d): %s",
|
||
wait,
|
||
attempt + 1,
|
||
self.max_retries + 1,
|
||
e,
|
||
)
|
||
time.sleep(wait)
|
||
return None, last_err, 0, ""
|
||
|
||
# 第一次尝试
|
||
task_id, last_err, sc, body = _do_create(create_payload)
|
||
|
||
# ratio 兜底:HTTP 400 且 body 提到 ratio / adaptive → 回退 adaptive 再试一次
|
||
if (
|
||
not task_id
|
||
and sc == 400
|
||
and final_ratio != "adaptive"
|
||
and (
|
||
"ratio" in (body or "").lower()
|
||
or "aspect" in (body or "").lower()
|
||
or "adaptive" in (body or "").lower()
|
||
)
|
||
):
|
||
logger.warning("Seedance 创建因 ratio 失败,回退 ratio=adaptive 重试")
|
||
create_payload["ratio"] = "adaptive"
|
||
task_id, last_err, sc2, body2 = _do_create(create_payload)
|
||
|
||
if not task_id:
|
||
logger.error(
|
||
"Seedance 创建任务最终失败: model=%s base_url=%s err=%s body=%s 【排查】"
|
||
"1) 方舟控制台已开通 doubao-seedance-2-5-260628;2) API Key 有该模型权限;"
|
||
"3) DOUBAO_BASE_URL=https://ark.cn-beijing.volces.com/api/v3;4) 参考素材 URL 公网可访问。",
|
||
video_model,
|
||
self.base_url,
|
||
last_err,
|
||
(body or "")[:500],
|
||
)
|
||
return None
|
||
|
||
logger.info("Seedance 任务已创建: task_id=%s ratio=%s", task_id, create_payload["ratio"])
|
||
|
||
# 2) 轮询状态
|
||
poll_url = f"{create_url}/{task_id}"
|
||
deadline = time.time() + total_timeout
|
||
video_url: str | None = None
|
||
last_status: str = "queued"
|
||
poll_count = 0
|
||
while time.time() < deadline:
|
||
poll_count += 1
|
||
try:
|
||
resp = httpx.get(poll_url, headers=headers, timeout=self.timeout)
|
||
try:
|
||
if int(getattr(resp, "status_code", 200)) >= 400:
|
||
resp.raise_for_status()
|
||
except (TypeError, ValueError):
|
||
pass
|
||
data = resp.json()
|
||
status = data.get("status", "")
|
||
last_status = status
|
||
if status == "succeeded":
|
||
content_obj = data.get("content") or {}
|
||
video_url = content_obj.get("video_url")
|
||
if video_url:
|
||
logger.info("Seedance 任务成功: task_id=%s polls=%d", task_id, poll_count)
|
||
break
|
||
last_err = RuntimeError(f"task succeeded but no video_url: {str(data)[:300]}")
|
||
logger.error("Seedance succeeded 但无 video_url: %s", last_err)
|
||
break
|
||
if status == "failed":
|
||
err = data.get("error") or {}
|
||
last_err = RuntimeError(f"task failed: code={err.get('code','')} msg={err.get('message','')}")
|
||
logger.error("Seedance 任务失败 task_id=%s: %s", task_id, last_err)
|
||
break
|
||
if status in ("expired", "cancelled"):
|
||
last_err = RuntimeError(f"task {status}")
|
||
logger.error("Seedance 任务 %s: task_id=%s", status, task_id)
|
||
break
|
||
# 每 5 次轮询打一次 info 日志,便于观察进度
|
||
if poll_count % 5 == 0:
|
||
logger.info("Seedance 轮询中: task_id=%s status=%s polls=%d", task_id, status, poll_count)
|
||
except httpx.HTTPStatusError as e:
|
||
last_err = e
|
||
logger.warning("Seedance 轮询 HTTP %d: %s", e.response.status_code, (e.response.text or "")[:300])
|
||
except Exception as e:
|
||
last_err = e
|
||
logger.debug("Seedance 轮询异常: %s", e)
|
||
time.sleep(poll_interval)
|
||
|
||
if not video_url:
|
||
logger.error(
|
||
"Seedance 任务未成功: task_id=%s last_status=%s polls=%d err=%s (总等待 %.0fs)",
|
||
task_id,
|
||
last_status,
|
||
poll_count,
|
||
last_err,
|
||
total_timeout,
|
||
)
|
||
return None
|
||
|
||
# 3) 下载到本地(下载超时收紧到 120s)
|
||
try:
|
||
out_dir = output_dir or "/tmp"
|
||
os.makedirs(out_dir, exist_ok=True)
|
||
local_path = f"{out_dir}/seedance_{task_id}_{uuid.uuid4().hex[:8]}.mp4"
|
||
download_timeout = 120.0
|
||
logger.info(
|
||
"Seedance 开始下载: task_id=%s url=%s timeout=%.0fs", task_id, video_url[:120], download_timeout
|
||
)
|
||
with httpx.stream("GET", video_url, timeout=download_timeout) as r:
|
||
r.raise_for_status()
|
||
downloaded = 0
|
||
with open(local_path, "wb") as f:
|
||
for chunk in r.iter_bytes(chunk_size=1024 * 256):
|
||
if chunk:
|
||
f.write(chunk)
|
||
downloaded += len(chunk)
|
||
size = os.path.getsize(local_path)
|
||
logger.info("Seedance 视频下载完成: %s size=%d bytes", local_path, size)
|
||
if size == 0:
|
||
logger.error("Seedance 下载文件大小为 0")
|
||
try:
|
||
os.remove(local_path)
|
||
except Exception:
|
||
pass
|
||
return None
|
||
return local_path
|
||
except Exception as e:
|
||
logger.error("Seedance 视频下载失败: %s", e, exc_info=True)
|
||
return None
|
||
|
||
|
||
# ── 单例 ─────────────────────────────────────────────────────────────────────
|
||
|
||
|
||
_client: Optional[DoubaoClient] = None
|
||
|
||
|
||
def get_doubao_client() -> DoubaoClient:
|
||
"""获取豆包客户端单例."""
|
||
global _client
|
||
if _client is None:
|
||
_client = DoubaoClient()
|
||
return _client
|