baodan/api/insurance/poster/image_generator.py

291 lines
12 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

"""海报图片生成器 — 支持多供应商生图模型OpenAI / 豆包 / 智谱等)。"""
import os
import base64
import io
import logging
logger = logging.getLogger(__name__)
# 用户尺寸 → API 尺寸映射
SIZE_MAP = {
"1080x1920": "1024x1792",
"900x500": "1792x1024",
"1080x1080": "1024x1024",
"800x1200": "1024x1792",
}
def _load_image_config() -> dict:
"""从数据库读取海报图片模型配置。"""
try:
from insurance.db.compat import db
from insurance.models.system_setting import SystemSetting
keys = ["poster_image_provider", "poster_image_model",
"poster_image_api_key", "poster_image_base_url"]
settings = {s.key: s.value for s in SystemSetting.query.filter(SystemSetting.key.in_(keys)).all()}
return settings
except Exception:
return {}
class PosterImageGenerator:
"""海报图片生成器,支持 OpenAI 兼容接口的多供应商生图模型。"""
def __init__(self):
self._client = None
self._config = None
def _get_config(self) -> dict:
if self._config is None:
self._config = _load_image_config()
return self._config
def _get_client(self):
if self._client is None:
import openai
cfg = self._get_config()
api_key = cfg.get("poster_image_api_key") or os.getenv("OPENAI_API_KEY", "")
base_url = cfg.get("poster_image_base_url") or os.getenv("OPENAI_BASE_URL", "https://api.openai.com/v1")
# 安全校验拒绝私网地址SEC-P1-02
from insurance.utils.security import is_safe_base_url
is_safe, err_msg = is_safe_base_url(base_url)
if not is_safe:
raise ValueError(f"不安全的 Base URL: {err_msg}")
self._client = openai.OpenAI(api_key=api_key, base_url=base_url)
return self._client
def _get_model(self) -> str:
cfg = self._get_config()
return cfg.get("poster_image_model") or "gpt-image-1"
def build_prompt(self, template: dict = None, product: dict = None,
company: dict = None, copy: dict = None, size: str = "1024x1792",
poster_content: dict = None) -> str:
"""组装完整的生图 promptAI 只生成无字视觉图)。"""
parts = [
"你是一位专业的保险营销海报设计师。请根据以下信息生成一张保险营销背景视觉图。",
"注意只生成背景场景图不包含任何文字、数字、Logo。",
"",
]
# 风格要求
if template:
parts.append("【风格要求】")
parts.append(template.get("styleDescription", ""))
if template.get("colorScheme"):
parts.append(f"配色方案:{template['colorScheme']}")
parts.append("")
# 产品信息
if product:
parts.append("【产品信息】")
parts.append(f"产品名称:{product.get('displayName', '')}")
parts.append(f"所属公司:{company.get('displayName', '') if company else ''}")
if product.get("manualParsedRules"):
rules = product["manualParsedRules"]
if isinstance(rules, str):
import json
try:
rules = json.loads(rules)
except Exception:
rules = {}
features = rules.get("features", [])
if features:
parts.append("产品亮点:" + "".join(f.get("title", "") for f in features[:3]))
parts.append("")
# 营销文案(仅用于风格参考,不生成文字)
if copy:
style_hint = copy.get("headline", "")
if style_hint:
parts.append("【风格参考】")
parts.append(f"主题方向:{style_hint}")
parts.append("(注意:只生成背景视觉图,不要在图中绘制任何文字)")
parts.append("")
# 小册子产品卖点(用于视觉风格方向)
if poster_content and poster_content.get("features"):
features = poster_content["features"]
parts.append("【产品卖点方向】")
parts.append("以下卖点决定视觉风格方向(不生成文字):")
for f in features[:5]:
title = f.get("title", "")
if title:
parts.append(f"- {title}")
parts.append("")
# 输出要求AI 只生成无字视觉图,文字由 HTML 渲染)
parts.append("【输出要求】")
parts.append(f"- 尺寸:{size} 像素")
parts.append("- 纯视觉图片不包含任何文字、数字、Logo 或图表")
parts.append("- 适合保险行业专业风格的背景或场景图")
parts.append("- 色调与配色方案一致")
return "\n".join(parts)
def _extract_image_bytes(self, response) -> bytes:
"""从图片 API 响应中提取图片 bytes支持 b64_json 和 URL 两种格式。"""
data_item = response.data[0] if response.data else None
if not data_item:
raise ValueError("图片 API 返回空数据")
# 优先使用 b64_json
if hasattr(data_item, 'b64_json') and data_item.b64_json:
return base64.b64decode(data_item.b64_json)
# 回退到 URL 下载
if hasattr(data_item, 'url') and data_item.url:
logger.info(f"图片 API 返回 URL正在下载: {data_item.url[:80]}...")
from insurance.utils.security import safe_httpx_client
with safe_httpx_client() as client:
resp = client.get(data_item.url, timeout=60, follow_redirects=True)
resp.raise_for_status()
return resp.content
raise ValueError("图片 API 响应既无 b64_json 也无 url")
def generate(self, prompt: str, size: str = "1024x1792", reference_image: str = None) -> tuple:
"""调用配置的生图模型生成海报图片。
返回:
(image_bytes, provider_info) — PNG 图片的 bytes 和生成信息
"""
api_size = SIZE_MAP.get(size, "1024x1792")
client = self._get_client()
model = self._get_model()
provider_info = {"provider": self._get_config().get("poster_image_provider", "openai"), "model": model}
kwargs = {
"model": model,
"prompt": prompt,
"n": 1,
"size": api_size,
}
# 后台模板可配置本地参考图或公网 URL失败时降级为纯 prompt 生图。
if reference_image:
image_file = None
close_file = False
try:
if reference_image.startswith("data:"):
# 前端上传的 base64 data URL
header, b64data = reference_image.split(",", 1)
image_bytes = base64.b64decode(b64data)
if len(image_bytes) > 10 * 1024 * 1024:
raise ValueError("参考图超过 10MB")
image_file = io.BytesIO(image_bytes)
image_file.name = "reference.png"
elif reference_image.startswith(("http://", "https://")):
from insurance.utils.security import is_safe_base_url, safe_httpx_client
is_safe, error_message = is_safe_base_url(reference_image)
if not is_safe:
raise ValueError(error_message)
with safe_httpx_client() as client:
image_response = client.get(
reference_image, timeout=30, follow_redirects=False
)
image_response.raise_for_status()
if len(image_response.content) > 10 * 1024 * 1024:
raise ValueError("参考图超过 10MB")
image_file = io.BytesIO(image_response.content)
image_file.name = "reference.png"
elif os.path.exists(reference_image):
from insurance.config import get_storage_root
ref_abs = os.path.abspath(reference_image)
storage_root = os.path.abspath(get_storage_root())
if os.path.commonpath([ref_abs, storage_root]) != storage_root:
raise ValueError("参考图不在允许的存储目录")
image_file = open(ref_abs, "rb")
close_file = True
if image_file is not None:
response = client.images.edit(
model=model,
image=image_file,
prompt=prompt,
n=1,
size=api_size,
)
return self._extract_image_bytes(response), provider_info
except Exception as e:
logger.warning(f"参考图 edit 失败,降级为纯 prompt 生成: {e}")
finally:
if close_file and image_file is not None:
image_file.close()
response = client.images.generate(**kwargs)
return self._extract_image_bytes(response), provider_info
def generate_fallback(copy_content: dict, size: str = "1024x1792") -> bytes:
"""降级方案:使用 Pillow 生成基础排版图。"""
from PIL import Image, ImageDraw, ImageFont
w, h = 1024, 1792
if "x" in size:
try:
parts = size.split("x")
w, h = int(parts[0]), int(parts[1])
except Exception:
pass
# 缩放到 API 支持的尺寸
w, h = min(w, 1792), min(h, 1792)
img = Image.new("RGB", (w, h), color=(26, 26, 46))
draw = ImageDraw.Draw(img)
# 简单排版
headline = copy_content.get("headline", "")
body = copy_content.get("body", "")
cta = copy_content.get("call_to_action", "")
# 尝试加载支持中文的字体
chinese_fonts = [
"/usr/share/fonts/truetype/noto/NotoSansCJK-Regular.ttc",
"/usr/share/fonts/opentype/noto/NotoSansCJK-Regular.ttc",
"/usr/share/fonts/noto-cjk/NotoSansCJK-Regular.ttc",
"/usr/share/fonts/truetype/wqy/wqy-zenhei.ttc",
"/usr/share/fonts/wqy-zenhei/wqy-zenhei.ttc",
"C:/Windows/Fonts/msyh.ttc", # Windows 微软雅黑
"C:/Windows/Fonts/simsun.ttc", # Windows 宋体
]
font_large = font_medium = font_small = None
for font_path in chinese_fonts:
if os.path.exists(font_path):
try:
font_large = ImageFont.truetype(font_path, 48)
font_medium = ImageFont.truetype(font_path, 28)
font_small = ImageFont.truetype(font_path, 22)
break
except Exception:
continue
if font_large is None:
font_large = ImageFont.load_default()
font_medium = ImageFont.load_default()
font_small = ImageFont.load_default()
# 绘制文字
y = h // 4
# anchor 参数在 Pillow < 8.0 的默认字体上不支持,安全降级
use_anchor = hasattr(font_large, 'getbbox') # TrueType 字体支持 anchor
anchor_kw = {"anchor": "mm"} if use_anchor else {}
if headline:
draw.text((w // 2, y), headline, fill="white", font=font_large, **anchor_kw)
y += 80
if body:
# 简单换行
for line in body.split("\n")[:4]:
draw.text((w // 2, y), line, fill=(200, 200, 200), font=font_medium, **anchor_kw)
y += 40
if cta:
draw.text((w // 2, h - 200), cta, fill=(102, 126, 234), font=font_small, **anchor_kw)
import io
buf = io.BytesIO()
img.save(buf, format="PNG")
return buf.getvalue()