AI 生图流水线系统 — 前后台详细设计¶
文档定位:详设背景。现行架构以 product/backend/design.md + architecture/messaging.md 为准;本文提供端到端流水线细节(Word 解析 → 多模型并行 → SSE → OSS),product 未单独成文。
从 Word 文档提取提示词 + 样图 → 多模型并行调度 → SSE 实时进度 → 结果持久化
1. 整体架构¶
flowchart TB
subgraph Frontend["🖥 前端 (React + Vite :5173)"]
UPLOAD["Word 上传组件"]
PROMPT_VIEW["提示词 + 样图预览"]
DASHBOARD["生成进度面板 (SSE)"]
GALLERY["结果画廊"]
end
subgraph Backend["⚙ 后端 (FastAPI :8000/:8001)"]
API["REST API Layer"]
PARSER["Word 解析器"]
DISPATCHER["调度引擎 (Async)"]
SSE_HUB["SSE 推送中心"]
end
subgraph Messaging["📨 消息层 (Redis Streams)"]
TASK_Q["task:image_gen 队列"]
RESULT_Q["result:image_gen 队列"]
end
subgraph Workers["🔧 Worker 进程"]
W1["Worker 1 (模型 A)"]
W2["Worker 2 (模型 B)"]
W3["Worker 3 (模型 C)"]
end
subgraph Storage["💾 存储"]
OSS["OSS 微服务 (:8002)"]
DB["PostgreSQL"]
end
UPLOAD -->|"POST /api/doc/upload"| API
API --> PARSER
PARSER -->|"提示词 JSON + 图片"| DB
PARSER -->|"样图"| OSS
API -->|"返回 session_id"| UPLOAD
UPLOAD -->|"POST /api/gen/start"| DISPATCHER
DISPATCHER -->|"XADD 子任务"| TASK_Q
DISPATCHER -->|"订阅结果"| RESULT_Q
TASK_Q -->|"XREADGROUP"| W1
TASK_Q -->|"XREADGROUP"| W2
TASK_Q -->|"XREADGROUP"| W3
W1 & W2 & W3 -->|"XADD 结果"| RESULT_Q
RESULT_Q -->|"消费"| SSE_HUB
SSE_HUB -->|"SSE 推送"| DASHBOARD
W1 & W2 & W3 -->|"上传生成图"| OSS
DISPATCHER --> DB
核心数据流:
Word 文档 → 解析(提示词+样图) → 用户确认 → 调度分发 → 多Worker并行生图 → SSE实时推送 → 前端展示 → 保存
2. Word 文档解析¶
2.1 提示词格式约定¶
Word 文档中提示词按以下格式组织(中文序号):
第一张:一只在月球上跳舞的宇航猫,赛博朋克风格,霓虹灯管,8K超清
第二张:水墨山水画,孤舟蓑笠翁,独钓寒江雪,留白构图
第三张:未来城市天际线,日落时分,飞行汽车穿梭,广角镜头
...
2.2 解析流程¶
flowchart LR
DOCX["Word 文档\n(.docx)"] --> UNZIP["解压 ZIP"]
UNZIP --> XML["解析 document.xml"]
UNZIP --> MEDIA["提取 word/media/* 图片"]
XML --> REGEX["正则匹配\n第[一二三四五六七八九十百千0-9]+张"]
REGEX --> PROMPTS["提示词列表"]
MEDIA --> IMAGES["样图列表"]
PROMPTS --> MATCH["按出现顺序\n关联提示词与图片"]
IMAGES --> MATCH
MATCH --> RESULT["PromptItem[]"]
2.3 核心代码设计¶
文件位置: src/myapp/core/services/doc_parser.py
import re
import asyncio
from dataclasses import dataclass, field
from io import BytesIO
from pathlib import Path
from zipfile import ZipFile
from typing import Optional
from docx import Document
from docx.opc.constants import RELATIONSHIP_TYPE as RT
@dataclass
class PromptItem:
"""单个提示词条目"""
index: int # 序号 (0-based)
label: str # "第一张" / "第二张" ...
prompt_text: str # 提示词正文
reference_image_path: Optional[str] = None # 关联样图 (OSS key)
reference_image_data: Optional[bytes] = None # 原始图片数据
@dataclass
class ParseResult:
"""解析结果"""
session_id: str
prompts: list[PromptItem]
total: int
class WordPromptParser:
"""从 Word 文档中提取提示词和图片"""
# 匹配 "第X张"、"第XX张" — 支持中文数字 + 阿拉伯数字
PROMPT_PATTERN = re.compile(
r'(第[一二三四五六七八九十百千\d]+张)[::]\s*(.+?)(?=第[一二三四五六七八九十百千\d]+张[::]|$)',
re.DOTALL
)
# 中文数字 → 阿拉伯数字映射
CN_NUM_MAP = {
'一': 1, '二': 2, '三': 3, '四': 4, '五': 5,
'六': 6, '七': 7, '八': 8, '九': 9, '十': 10,
}
@staticmethod
def _cn_to_int(cn: str) -> int:
"""中文数字转整数,如 '十二' → 12"""
# 简化实现;生产环境建议用 cn2an 库
result = 0
for char in cn:
if char in WordPromptParser.CN_NUM_MAP:
result = result * 10 + WordPromptParser.CN_NUM_MAP[char]
elif char.isdigit():
result = result * 10 + int(char)
return result
def parse(self, file_bytes: bytes) -> ParseResult:
"""解析 Word 文档字节流"""
doc = Document(BytesIO(file_bytes))
# 1. 提取全文
full_text = "\n".join(p.text for p in doc.paragraphs)
# 2. 正则匹配提示词
matches = self.PROMPT_PATTERN.findall(full_text)
prompts: list[PromptItem] = []
for idx, (label_raw, text) in enumerate(matches):
# 提取序号数字
num_part = label_raw[1:-1] # 去掉 "第" 和 "张"
num = self._cn_to_int(num_part) if num_part else idx + 1
prompts.append(PromptItem(
index=num - 1, # 转 0-based
label=label_raw,
prompt_text=text.strip(),
))
# 3. 提取内嵌图片
images = self._extract_images(doc)
# 4. 按顺序关联(图片按文档中出现顺序分配)
# 规则:如果图片数量 ≤ 提示词数量,一对一按序分配
for i, img_data in enumerate(images):
if i < len(prompts):
prompts[i].reference_image_data = img_data
return ParseResult(
session_id="", # 由上层填充
prompts=sorted(prompts, key=lambda p: p.index),
total=len(prompts),
)
def _extract_images(self, doc: Document) -> list[bytes]:
"""从 docx 中提取内嵌图片"""
images: list[bytes] = []
for rel in doc.part.rels.values():
if "image" in rel.reltype:
images.append(rel.target_part.blob)
return images
2.4 样图上传到 OSS¶
# src/myapp/core/services/prompt_service.py
class PromptService:
def __init__(self, storage_client: StorageClient, db: AsyncSession):
self.storage = storage_client
self.db = db
async def process_upload(self, file_bytes: bytes) -> ParseResult:
parser = WordPromptParser()
result = parser.parse(file_bytes)
# 为 session 生成唯一 ID
result.session_id = uuid4().hex[:12]
# 上传样图到 OSS
for item in result.prompts:
if item.reference_image_data:
key = f"prompts/{result.session_id}/{item.index:03d}_ref.png"
item.reference_image_path = await self.storage.upload(
key=key,
data=item.reference_image_data,
content_type="image/png",
)
# 清空原始数据,避免内存占用
item.reference_image_data = None
# 持久化到 DB
await self._save_to_db(result)
return result
3. 多模型调度引擎¶
3.1 设计目标¶
| 需求 | 方案 |
|---|---|
| 多个后台模型并行生图 | 每个模型独立 Worker,通过 Redis Streams 消费 |
| 请求分配 | 轮询 + 权重 + 模型亲和性(不同提示词适合不同模型) |
| 失败重试 | 指数退避,最多 3 次,超过进入死信队列 |
| 前端实时进度 | SSE 推送每个子任务的状态变更 |
3.2 模型注册表¶
# src/myapp/core/services/model_registry.py
from enum import Enum
from dataclasses import dataclass
class ModelProvider(str, Enum):
OPENAI_DALLE = "openai_dalle"
STABLE_DIFFUSION = "stable_diffusion"
MIDJOURNEY_API = "midjourney_api"
FLUX = "flux"
KOLORS = "kolors"
@dataclass
class ModelConfig:
provider: ModelProvider
model_name: str # "dall-e-3", "sdxl", "flux.1-pro"
endpoint: str # API 地址
api_key_env: str # 环境变量名
max_concurrent: int = 3 # 最大并发数
weight: float = 1.0 # 调度权重(越大越优先)
supported_styles: list[str] = field(default_factory=list) # 擅长风格
rate_limit_per_min: int = 10
timeout_seconds: int = 120
# 配置示例
MODEL_POOL: list[ModelConfig] = [
ModelConfig(
provider=ModelProvider.OPENAI_DALLE,
model_name="dall-e-3",
endpoint="https://api.openai.com/v1/images/generations",
api_key_env="OPENAI_API_KEY",
max_concurrent=2,
weight=1.2,
supported_styles=["写实", "摄影", "3D渲染"],
rate_limit_per_min=5,
),
ModelConfig(
provider=ModelProvider.STABLE_DIFFUSION,
model_name="sdxl-1.0",
endpoint="https://api.stability.ai/v1/generation/...",
api_key_env="STABILITY_API_KEY",
max_concurrent=5,
weight=1.0,
supported_styles=["插画", "概念艺术", "水墨"],
),
ModelConfig(
provider=ModelProvider.FLUX,
model_name="flux.1-pro",
endpoint="https://api.bfl.ml/v1/flux-pro",
api_key_env="FLUX_API_KEY",
max_concurrent=3,
weight=1.5,
supported_styles=["写实", "赛博朋克", "风景"],
),
]
3.3 调度策略¶
flowchart TB
START["收到 N 个 PromptItem"] --> MATCH["计算 Prompt-Model 亲和度"]
MATCH --> ALLOC["为每个 Prompt 选 Top-3 模型\n按权重 + 风格匹配排序"]
ALLOC --> QUEUE["生成 M 个子任务 (N×3)\n写入 Redis Stream"]
QUEUE --> CONSUME["Worker Consumer Group\n抢占消费"]
CONSUME --> DEDUP{"同 Prompt\n已有成功结果?"}
DEDUP -->|是| DROP["丢弃重复任务"]
DEDUP -->|否| EXEC["调用模型 API"]
EXEC -->|成功| PUBLISH["发布成功事件"]
EXEC -->|失败| RETRY{"重试次数 < 3?"}
RETRY -->|是| BACKOFF["指数退避\n2^n 秒后重新入队"]
BACKOFF --> QUEUE
RETRY -->|否| DLQ["进入死信队列\n标记为最终失败"]
核心代码:
# src/myapp/core/services/dispatcher.py
import asyncio
import json
from datetime import datetime
from myapp.core.services.model_registry import MODEL_POOL, ModelConfig
class GenerationDispatcher:
"""生图调度引擎"""
RETRY_MAX = 3
RETRY_BASE_DELAY = 2 # 秒
def __init__(self, redis_client, sse_hub, db_session):
self.redis = redis_client
self.sse = sse_hub
self.db = db_session
self._model_pool = MODEL_POOL
async def dispatch(self, session_id: str, prompts: list[PromptItem]) -> dict:
"""
核心调度方法:
1. 为每个提示词选择候选模型
2. 生成子任务写入 Redis Stream
3. 返回任务总览
"""
task_id = f"gen_{session_id}_{datetime.now():%Y%m%d%H%M%S}"
subtasks: list[dict] = []
for prompt in prompts:
# 选择候选模型(Top-3 按亲和度排序)
candidates = self._select_models(prompt)
for model in candidates:
subtask = {
"task_id": task_id,
"subtask_id": f"{task_id}_p{prompt.index}_m{model.provider.value}",
"session_id": session_id,
"prompt_index": prompt.index,
"prompt_label": prompt.label,
"prompt_text": prompt.prompt_text,
"reference_image_path": prompt.reference_image_path,
"model_provider": model.provider.value,
"model_name": model.model_name,
"endpoint": model.endpoint,
"retry_count": 0,
"created_at": datetime.now().isoformat(),
}
subtasks.append(subtask)
# 写入 Redis Stream
pipeline = self.redis.pipeline()
for st in subtasks:
pipeline.xadd(
"task:image_gen",
{"data": json.dumps(st, ensure_ascii=False)},
maxlen=10000,
)
pipeline.execute()
# 保存任务元信息到 DB
await self._save_task_meta(task_id, session_id, len(prompts), len(subtasks))
# 通知前端:任务已创建
await self.sse.publish(session_id, {
"event": "task_created",
"task_id": task_id,
"total_prompts": len(prompts),
"total_subtasks": len(subtasks),
})
return {
"task_id": task_id,
"total_prompts": len(prompts),
"total_subtasks": len(subtasks),
}
def _select_models(self, prompt: PromptItem, top_k: int = 3) -> list[ModelConfig]:
"""为提示词选择最优模型"""
scored = []
for model in self._model_pool:
score = model.weight
# 风格匹配加分
for style in model.supported_styles:
if style in prompt.prompt_text:
score += 0.5
scored.append((score, model))
scored.sort(key=lambda x: x[0], reverse=True)
return [m for _, m in scored[:top_k]]
3.4 Worker 实现(含重试)¶
# src/messaging/workers/image_gen_worker.py
import asyncio
import json
import time
from typing import Optional
import httpx
from redis.asyncio import Redis
class ImageGenWorker:
"""生图 Worker — 每个模型实例一个"""
def __init__(
self,
redis: Redis,
consumer_group: str,
consumer_name: str,
provider: str,
api_client: httpx.AsyncClient,
storage_client,
sse_hub,
):
self.redis = redis
self.group = consumer_group
self.name = consumer_name
self.provider = provider
self.http = api_client
self.storage = storage_client
self.sse = sse_hub
self.max_retries = 3
self.base_delay = 2
async def run(self):
"""主循环:持续消费任务"""
# 确保 consumer group 存在
try:
await self.redis.xgroup_create(
"task:image_gen", self.group, id="0", mkstream=True
)
except Exception:
pass # group 已存在
while True:
# 读取新消息(阻塞 5s)
messages = await self.redis.xreadgroup(
groupname=self.group,
consumername=self.name,
streams={"task:image_gen": ">"},
count=1,
block=5000,
)
if not messages:
continue
for stream, entries in messages:
for msg_id, fields in entries:
raw = fields.get(b"data", b"{}").decode()
subtask = json.loads(raw)
await self._process(msg_id, subtask)
async def _process(self, msg_id: bytes, subtask: dict):
"""处理单个子任务"""
subtask_id = subtask["subtask_id"]
session_id = subtask["session_id"]
# 检查是否已有成功结果(去重)
if await self._is_duplicate(subtask):
await self.redis.xack("task:image_gen", self.group, msg_id)
return
# 通知:开始生成
await self.sse.publish(session_id, {
"event": "generating",
"subtask_id": subtask_id,
"prompt_index": subtask["prompt_index"],
"prompt_label": subtask["prompt_label"],
"model": subtask["model_name"],
})
try:
# 调用 AI 模型 API
image_bytes: Optional[bytes] = None
if subtask["model_provider"] == "openai_dalle":
image_bytes = await self._call_dalle(subtask)
elif subtask["model_provider"] == "stable_diffusion":
image_bytes = await self._call_sd(subtask)
elif subtask["model_provider"] == "flux":
image_bytes = await self._call_flux(subtask)
else:
image_bytes = await self._call_generic(subtask)
if image_bytes:
# 上传到 OSS
oss_key = f"results/{session_id}/{subtask_id}.png"
image_url = await self.storage.upload(
key=oss_key, data=image_bytes, content_type="image/png"
)
# 通知:生成成功
await self.sse.publish(session_id, {
"event": "completed",
"subtask_id": subtask_id,
"prompt_index": subtask["prompt_index"],
"prompt_label": subtask["prompt_label"],
"model": subtask["model_name"],
"image_url": image_url,
})
# ACK
await self.redis.xack("task:image_gen", self.group, msg_id)
except Exception as e:
retry_count = subtask.get("retry_count", 0)
if retry_count < self.max_retries:
# 指数退避重试
delay = self.base_delay ** (retry_count + 1)
await self.sse.publish(session_id, {
"event": "retrying",
"subtask_id": subtask_id,
"prompt_label": subtask["prompt_label"],
"model": subtask["model_name"],
"retry": retry_count + 1,
"delay_seconds": delay,
"error": str(e),
})
await asyncio.sleep(delay)
# 重新入队
subtask["retry_count"] = retry_count + 1
await self.redis.xadd(
"task:image_gen",
{"data": json.dumps(subtask, ensure_ascii=False)},
)
await self.redis.xack("task:image_gen", self.group, msg_id)
else:
# 最终失败 → 死信
await self.sse.publish(session_id, {
"event": "failed",
"subtask_id": subtask_id,
"prompt_label": subtask["prompt_label"],
"model": subtask["model_name"],
"error": str(e),
})
await self.redis.xadd(
"task:image_gen_dlq",
{"data": json.dumps(subtask, ensure_ascii=False)},
)
await self.redis.xack("task:image_gen", self.group, msg_id)
3.5 各模型 API 适配¶
async def _call_dalle(self, subtask: dict) -> bytes:
"""OpenAI DALL-E 3"""
resp = await self.http.post(
subtask["endpoint"],
json={
"model": "dall-e-3",
"prompt": subtask["prompt_text"],
"n": 1,
"size": "1024x1024",
"quality": "hd",
},
timeout=120,
)
resp.raise_for_status()
data = resp.json()
image_url = data["data"][0]["url"]
# 下载生成的图片
img_resp = await self.http.get(image_url)
return img_resp.content
async def _call_sd(self, subtask: dict) -> bytes:
"""Stable Diffusion (Stability AI API)"""
form_data = {
"prompt": subtask["prompt_text"],
"output_format": "png",
}
# 如果有参考图,走 img2img
if subtask.get("reference_image_path"):
# 从 OSS 下载参考图
ref_bytes = await self.storage.download(subtask["reference_image_path"])
files = {"image": ("ref.png", BytesIO(ref_bytes), "image/png")}
resp = await self.http.post(
subtask["endpoint"],
data=form_data,
files=files,
timeout=120,
)
else:
resp = await self.http.post(
subtask["endpoint"],
data=form_data,
timeout=120,
)
resp.raise_for_status()
return resp.content
async def _call_flux(self, subtask: dict) -> bytes:
"""Flux Pro (BFL API)"""
resp = await self.http.post(
subtask["endpoint"],
json={
"prompt": subtask["prompt_text"],
"width": 1024,
"height": 1024,
"steps": 28,
},
timeout=180,
)
resp.raise_for_status()
# Flux 可能返回轮询 URL
data = resp.json()
if "result" in data:
# 轮询获取结果
poll_url = data["result"]
for _ in range(30): # 最多等 5min
await asyncio.sleep(10)
poll_resp = await self.http.get(poll_url)
poll_data = poll_resp.json()
if poll_data.get("status") == "Ready":
img_url = poll_data["result"]["sample"]
img_resp = await self.http.get(img_url)
return img_resp.content
# 同步返回
return resp.content
4. SSE 实时进度推送¶
4.1 事件类型定义¶
# src/messaging/streams/sse_events.py
class SSEEventType(str, Enum):
TASK_CREATED = "task_created" # 调度完成,子任务入队
GENERATING = "generating" # 某个子任务开始生成
PROGRESS = "progress" # 进度百分比(可选实现)
RETRYING = "retrying" # 失败重试中
COMPLETED = "completed" # 某个子任务成功
FAILED = "failed" # 某个子任务最终失败
ALL_DONE = "all_done" # 全部子任务结束
4.2 SSE 推送中心¶
利用项目已有的 src/messaging/streams/ 基础设施:
# src/messaging/streams/sse_hub.py
import asyncio
import json
from collections import defaultdict
class SSEHub:
"""
SSE 连接管理器 — 每个 session_id 可以有多个前端连接
"""
def __init__(self):
# session_id → list[asyncio.Queue]
self._connections: dict[str, list[asyncio.Queue]] = defaultdict(list)
async def subscribe(self, session_id: str) -> asyncio.Queue:
"""前端订阅某个 session 的 SSE 流"""
queue: asyncio.Queue = asyncio.Queue(maxsize=256)
self._connections[session_id].append(queue)
return queue
def unsubscribe(self, session_id: str, queue: asyncio.Queue):
"""前端断开连接"""
if session_id in self._connections:
self._connections[session_id].remove(queue)
async def publish(self, session_id: str, event: dict):
"""向所有订阅者推送事件"""
queues = self._connections.get(session_id, [])
dead: list[asyncio.Queue] = []
payload = f"data: {json.dumps(event, ensure_ascii=False)}\n\n"
for q in queues:
try:
q.put_nowait(payload)
except asyncio.QueueFull:
dead.append(q)
for q in dead:
self._connections[session_id].remove(q)
4.3 SSE API 端点¶
# src/myapp/api/routes/generation.py
from fastapi import APIRouter, Request
from fastapi.responses import StreamingResponse
from myapp.composition.dependencies import get_sse_hub
router = APIRouter(prefix="/api/gen", tags=["generation"])
@router.get("/stream/{session_id}")
async def stream_progress(session_id: str):
"""前端 SSE 订阅端点"""
hub = get_sse_hub()
queue = await hub.subscribe(session_id)
async def event_generator():
try:
while True:
# 检查客户端是否断开
data = await asyncio.wait_for(queue.get(), timeout=30)
yield data
except asyncio.TimeoutError:
yield "data: {\"event\":\"heartbeat\"}\n\n"
finally:
hub.unsubscribe(session_id, queue)
return StreamingResponse(
event_generator(),
media_type="text/event-stream",
headers={
"Cache-Control": "no-cache",
"Connection": "keep-alive",
"X-Accel-Buffering": "no", # nginx 不缓冲
},
)
@router.post("/start/{session_id}")
async def start_generation(session_id: str):
"""触发批量生图"""
result = await generation_dispatcher.dispatch_by_session(session_id)
return {"code": 0, "data": result}
@router.get("/status/{session_id}")
async def get_status(session_id: str):
"""轮询方式获取当前进度(SSE 的降级方案)"""
status = await generation_service.get_progress(session_id)
return {"code": 0, "data": status}
5. 前端设计¶
5.1 页面结构¶
flowchart TB
subgraph Page["AI 生图流水线页面"]
direction TB
STEP1["🔼 Step 1: 上传 Word 文档"]
STEP2["📋 Step 2: 提示词预览 & 确认"]
STEP3["⚡ Step 3: 生成进度面板 (SSE)"]
STEP4["🖼 Step 4: 结果画廊 & 下载"]
end
STEP1 --> STEP2 --> STEP3 --> STEP4
5.2 组件树¶
features/image-gen/
├── routes/
│ └── ImageGenPage.tsx # 主页面(含步骤状态机)
├── components/
│ ├── WordUploader.tsx # 文件拖拽上传
│ ├── PromptPreview.tsx # 提示词卡片列表 + 样图预览
│ ├── ModelSelector.tsx # 选择启用哪些模型
│ ├── GenerationDashboard.tsx # 实时进度面板(核心组件)
│ ├── ProgressCard.tsx # 单个提示词的进度卡片
│ ├── ResultGallery.tsx # 生成结果画廊
│ └── ResultImage.tsx # 单张结果图(点击放大/下载)
├── hooks/
│ ├── useSSEStream.ts # SSE 订阅 hook
│ ├── useGenerationTask.ts # 任务状态管理
│ └── useWordUpload.ts # 文件上传 hook
├── stores/
│ └── generationStore.ts # Zustand 全局状态
├── types/
│ └── generation.ts # 类型定义
└── api/
└── generationApi.ts # API 调用封装
5.3 核心 Hook:SSE 实时订阅¶
// src/agent-system/src/features/image-gen/hooks/useSSEStream.ts
import { useEffect, useRef, useCallback } from 'react';
interface SSEEvent {
event: string;
subtask_id: string;
prompt_index: number;
prompt_label: string;
model: string;
image_url?: string;
error?: string;
retry?: number;
delay_seconds?: number;
}
type EventHandler = (event: SSEEvent) => void;
export function useSSEStream(sessionId: string | null, onEvent: EventHandler) {
const eventSourceRef = useRef<EventSource | null>(null);
const reconnectTimerRef = useRef<number>();
const connect = useCallback(() => {
if (!sessionId) return;
// 关闭旧连接
eventSourceRef.current?.close();
const es = new EventSource(`/api/gen/stream/${sessionId}`);
eventSourceRef.current = es;
es.onmessage = (msg) => {
try {
const event: SSEEvent = JSON.parse(msg.data);
onEvent(event);
} catch {}
};
es.onerror = () => {
es.close();
// 自动重连(指数退避)
reconnectTimerRef.current = window.setTimeout(connect, 3000);
};
}, [sessionId, onEvent]);
useEffect(() => {
connect();
return () => {
eventSourceRef.current?.close();
clearTimeout(reconnectTimerRef.current);
};
}, [connect]);
const disconnect = useCallback(() => {
eventSourceRef.current?.close();
clearTimeout(reconnectTimerRef.current);
}, []);
return { disconnect };
}
5.4 核心 Hook:任务状态管理¶
// src/agent-system/src/features/image-gen/hooks/useGenerationTask.ts
import { useCallback, useRef, useState } from 'react';
import { useSSEStream } from './useSSEStream';
export type SubTaskStatus = 'queued' | 'generating' | 'completed' | 'failed' | 'retrying';
export interface SubTaskState {
subtaskId: string;
promptIndex: number;
promptLabel: string;
model: string;
status: SubTaskStatus;
imageUrl?: string;
error?: string;
retryCount: number;
}
export function useGenerationTask(sessionId: string | null) {
const [subtasks, setSubtasks] = useState<Map<string, SubTaskState>>(new Map());
const [allDone, setAllDone] = useState(false);
// 统计
const stats = {
total: subtasks.size,
completed: [...subtasks.values()].filter(s => s.status === 'completed').length,
failed: [...subtasks.values()].filter(s => s.status === 'failed').length,
generating: [...subtasks.values()].filter(s => s.status === 'generating').length,
};
// 按提示词分组(每个提示词对应多个模型的结果)
const groupedByPrompt = [...subtasks.values()].reduce((acc, st) => {
if (!acc[st.promptLabel]) acc[st.promptLabel] = [];
acc[st.promptLabel].push(st);
return acc;
}, {} as Record<string, SubTaskState[]>);
const handleEvent = useCallback((event: SSEEvent) => {
setSubtasks(prev => {
const next = new Map(prev);
const id = event.subtask_id;
switch (event.event) {
case 'task_created':
// 初始化所有子任务为 queued
// (需要从另一个 API 获取完整列表,这里简化)
break;
case 'generating':
next.set(id, {
subtaskId: id,
promptIndex: event.prompt_index,
promptLabel: event.prompt_label,
model: event.model,
status: 'generating',
retryCount: 0,
});
break;
case 'retrying':
next.set(id, {
...(next.get(id) || {} as SubTaskState),
status: 'retrying',
retryCount: event.retry || 0,
error: event.error,
});
break;
case 'completed':
next.set(id, {
...(next.get(id) || {} as SubTaskState),
status: 'completed',
imageUrl: event.image_url,
});
break;
case 'failed':
next.set(id, {
...(next.get(id) || {} as SubTaskState),
status: 'failed',
error: event.error,
});
break;
case 'all_done':
setAllDone(true);
break;
}
return next;
});
}, []);
const { disconnect } = useSSEStream(sessionId, handleEvent);
return { subtasks, stats, groupedByPrompt, allDone, disconnect };
}
5.5 核心组件:生成进度面板¶
// src/agent-system/src/features/image-gen/components/GenerationDashboard.tsx
import React from 'react';
import { ProgressCard } from './ProgressCard';
import { useGenerationTask } from '../hooks/useGenerationTask';
import styles from './GenerationDashboard.module.css';
interface Props {
sessionId: string;
}
export const GenerationDashboard: React.FC<Props> = ({ sessionId }) => {
const { subtasks, stats, groupedByPrompt, allDone } = useGenerationTask(sessionId);
return (
<div className={styles.dashboard}>
{/* 总体进度条 */}
<div className={styles.overallProgress}>
<div className={styles.stats}>
<span>总进度: {stats.completed + stats.failed <span>总进度: {stats.completed + stats.failed}/{stats.total}</span>
<span className={styles.green}>✅ {stats.completed}</span>
<span className={styles.red}>❌ {stats.failed}</span>
<span className={styles.blue}>⚡ {stats.generating}</span>
</div>
<div className={styles.progressBar}>
<div
className={styles.progressFill}
style={{
width: `${stats.total > 0 ? ((stats.completed + stats.failed) / stats.total) * 100 : 0}%`
}}
/>
</div>
{allDone && <div className={styles.doneBadge}>🎉 全部完成!</div>}
</div>
{/* 按提示词分组展示 */}
<div className={styles.promptGroups}>
{Object.entries(groupedByPrompt).map(([label, items]) => (
<ProgressCard key={label} label={label} subtasks={items} />
))}
</div>
</div>
);
};
5.6 单提示词进度卡片¶
// ProgressCard.tsx
export const ProgressCard: React.FC<{ label: string; subtasks: SubTaskState[] }> = ({
label,
subtasks,
}) => {
const bestResult = subtasks.find(s => s.status === 'completed');
return (
<div className={styles.card}>
<div className={styles.cardHeader}>
<h4>{label}</h4>
{bestResult && (
<img
src={bestResult.imageUrl}
alt={label}
className={styles.thumbnail}
/>
)}
</div>
<div className={styles.modelList}>
{subtasks.map(st => (
<div key={st.subtaskId} className={styles.modelRow}>
<span className={styles.modelName}>{st.model}</span>
<StatusBadge status={st.status} />
{st.status === 'retrying' && (
<span className={styles.retryInfo}>
重试 {st.retryCount}/3
</span>
)}
{st.status === 'failed' && (
<span className={styles.error} title={st.error}>
{st.error?.slice(0, 40)}...
</span>
)}
</div>
))}
</div>
</div>
);
};
const StatusBadge: React.FC<{ status: SubTaskStatus }> = ({ status }) => {
const map = {
queued: { text: '排队中', className: 'gray', icon: '⏳' },
generating: { text: '生成中', className: 'blue', icon: '⚡' },
completed: { text: '完成', className: 'green', icon: '✅' },
failed: { text: '失败', className: 'red', icon: '❌' },
retrying: { text: '重试中', className: 'orange', icon: '🔄' },
};
const item = map[status];
return <span className={`badge ${item.className}`}>{item.icon} {item.text}</span>;
};
5.7 结果画廊¶
// ResultGallery.tsx
export const ResultGallery: React.FC<{ sessionId: string }> = ({ sessionId }) => {
const { groupedByPrompt } = useGenerationTask(sessionId);
const [selectedImage, setSelectedImage] = useState<string | null>(null);
const [viewMode, setViewMode] = useState<'grid' | 'compare'>('grid');
return (
<div className={styles.gallery}>
{/* 工具栏 */}
<div className={styles.toolbar}>
<button onClick={() => setViewMode('grid')}>网格视图</button>
<button onClick={() => setViewMode('compare')}>模型对比</button>
<button onClick={() => downloadAll(sessionId)}>📥 全部下载</button>
</div>
{viewMode === 'grid' ? (
<div className={styles.grid}>
{Object.entries(groupedByPrompt).map(([label, items]) => {
const completed = items.filter(i => i.status === 'completed');
return completed.map(item => (
<div
key={item.subtaskId}
className={styles.gridItem}
onClick={() => setSelectedImage(item.imageUrl!)}
>
<img src={item.imageUrl} alt={`${label} - ${item.model}`} />
<div className={styles.caption}>
<span>{label}</span>
<span className={styles.modelTag}>{item.model}</span>
</div>
</div>
));
})}
</div>
) : (
<CompareView groupedByPrompt={groupedByPrompt} />
)}
{/* 图片灯箱 */}
{selectedImage && (
<Lightbox imageUrl={selectedImage} onClose={() => setSelectedImage(null)} />
)}
</div>
);
};
6. API 设计总览¶
| 端点 | 方法 | 用途 |
|---|---|---|
/api/doc/upload |
POST | 上传 Word,返回 {session_id, prompts[], images[]} |
/api/gen/start/{session_id} |
POST | 触发批量生图调度 |
/api/gen/stream/{session_id} |
GET | SSE 端点,实时推送进度 |
/api/gen/status/{session_id} |
GET | 轮询状态(SSE 降级) |
/api/gen/results/{session_id} |
GET | 获取全部结果(含 OSS URL) |
/api/gen/results/{session_id}/download |
GET | 打包下载 ZIP |
/api/gen/retry/{session_id}/{subtask_id} |
POST | 手动重试失败子任务 |
/api/models |
GET | 可用模型列表 |
请求/响应示例¶
POST /api/doc/upload
// Request: multipart/form-data (file: docx)
// Response:
{
"code": 0,
"data": {
"session_id": "a1b2c3d4e5f6",
"prompts": [
{
"index": 0,
"label": "第一张",
"prompt_text": "一只在月球上跳舞的宇航猫,赛博朋克风格...",
"reference_image_url": "https://upload.harness.local/prompts/a1b2c3/000_ref.png"
},
{
"index": 1,
"label": "第二张",
"prompt_text": "水墨山水画,孤舟蓑笠翁...",
"reference_image_url": null
}
],
"total": 3
}
}
POST /api/gen/start/{session_id}
// Request (可选指定模型):
{
"model_filter": ["openai_dalle", "flux"], // 仅用指定模型
"priority": "speed" // "speed" | "quality" | "balanced"
}
// Response:
{
"code": 0,
"data": {
"task_id": "gen_a1b2c3d4e5f6_20260703120000",
"total_prompts": 3,
"total_subtasks": 6
}
}
7. 数据库设计¶
-- 生图会话
CREATE TABLE gen_sessions (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
session_id VARCHAR(12) UNIQUE NOT NULL, -- 短ID供前端使用
original_filename VARCHAR(255),
prompt_count INTEGER NOT NULL,
status VARCHAR(20) DEFAULT 'pending', -- pending | running | done | partial_fail | failed
created_at TIMESTAMPTZ DEFAULT now(),
completed_at TIMESTAMPTZ
);
-- 提示词条目
CREATE TABLE gen_prompts (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
session_id VARCHAR(12) REFERENCES gen_sessions(session_id),
prompt_index INTEGER NOT NULL,
label VARCHAR(50) NOT NULL, -- "第一张"
prompt_text TEXT NOT NULL,
reference_image_path VARCHAR(500), -- OSS key
created_at TIMESTAMPTZ DEFAULT now()
);
-- 子任务(每个 Prompt×Model 组合)
CREATE TABLE gen_subtasks (
id UUID PRIMARY KEY DEFAULT gen_random_uuid(),
subtask_id VARCHAR(100) UNIQUE NOT NULL,
task_id VARCHAR(100) NOT NULL,
session_id VARCHAR(12) REFERENCES gen_sessions(session_id),
prompt_index INTEGER NOT NULL,
model_provider VARCHAR(50) NOT NULL,
model_name VARCHAR(100) NOT NULL,
status VARCHAR(20) DEFAULT 'queued', -- queued | generating | completed | failed
retry_count INTEGER DEFAULT 0,
result_image_path VARCHAR(500), -- OSS key
error_message TEXT,
started_at TIMESTAMPTZ,
completed_at TIMESTAMPTZ,
created_at TIMESTAMPTZ DEFAULT now()
);
CREATE INDEX idx_subtasks_session ON gen_subtasks(session_id);
CREATE INDEX idx_subtasks_status ON gen_subtasks(status);
8. 失败重试机制总结¶
sequenceDiagram
participant W as Worker
participant API as AI Model API
participant R as Redis Stream
participant DLQ as Dead Letter Queue
W->>R: XREADGROUP 取任务
W->>API: POST 生图请求
alt 成功
API-->>W: 200 + 图片
W->>R: XACK
W->>SSE: completed 事件
else 第 1 次失败
API-->>W: 5xx / timeout
W->>SSE: retrying (第1次)
W->>W: sleep(2s)
W->>R: XADD 重新入队
W->>R: XACK 旧消息
else 第 2 次失败
API-->>W: 5xx / timeout
W->>SSE: retrying (第2次)
W->>W: sleep(4s)
W->>R: XADD 重新入队
W->>R: XACK 旧消息
else 第 3 次失败
API-->>W: 5xx / timeout
W->>SSE: retrying (第3次)
W->>W: sleep(8s)
W->>R: XADD 重新入队
W->>R: XACK 旧消息
else 第 4 次失败 (超过上限)
API-->>W: 5xx / timeout
W->>SSE: failed 事件
W->>DLQ: XADD 死信队列
W->>R: XACK 旧消息
end
关键参数:
- 最大重试次数:3 次
- 退避策略:指数退避 2^retry_count 秒(2s → 4s → 8s)
- 去重:同一 Prompt 只要有一个模型成功,同 Prompt 其他模型的排队任务自动丢弃
- 死信:3 次重试仍失败 → 写入 task:image_gen_dlq,前端展示失败,支持手动重试
- 熔断:某模型连续失败 5 个任务 → 临时禁用该模型 60 秒,通知前端
9. 文件结构与集成¶
基于项目现有架构,新增/修改的文件:
src/myapp/
├── api/routes/
│ └── generation.py # [NEW] 生图 API 路由
├── core/services/
│ ├── doc_parser.py # [NEW] Word 文档解析
│ ├── prompt_service.py # [NEW] 提示词管理
│ ├── model_registry.py # [NEW] 模型注册表
│ ├── dispatcher.py # [NEW] 调度引擎
│ └── generation_service.py # [NEW] 生图流程编排
├── db/
│ ├── models/
│ │ └── generation.py # [NEW] SQLAlchemy 模型
│ └── repositories/
│ └── generation_repo.py # [NEW] 数据访问层
├── schemas/
│ └── generation.py # [NEW] Pydantic 请求/响应模型
└── composition/
└── dependencies.py # [MODIFY] 新增依赖注入
src/messaging/
├── streams/
│ └── sse_hub.py # [NEW] SSE 推送中心
└── workers/
└── image_gen_worker.py # [NEW] 生图 Worker
src/agent-system/src/features/
└── image-gen/ # [NEW] 前端功能模块
├── routes/ImageGenPage.tsx
├── components/
│ ├── WordUploader.tsx
│ ├── PromptPreview.tsx
│ ├── ModelSelector.tsx
│ ├── GenerationDashboard.tsx
│ ├── ProgressCard.tsx
│ └── ResultGallery.tsx
├── hooks/
│ ├── useSSEStream.ts
│ ├── useGenerationTask.ts
│ └── useWordUpload.ts
├── stores/generationStore.ts
└── api/generationApi.ts
deploy/k3s/base/
└── image-gen-worker-deployment.yaml # [NEW] Worker K8s Deployment
10. 快速启动流程¶
# 1. 启动依赖
make up # postgres + redis + myapp + storage
# 2. 启动生图 Worker(每个模型一个进程)
poetry run python -m messaging.workers.image_gen_worker --provider openai_dalle &
poetry run python -m messaging.workers.image_gen_worker --provider flux &
poetry run python -m messaging.workers.image_gen_worker --provider stable_diffusion &
# 3. 启动前端
make agent-dev # :5173
# 4. 访问 http://agent.localhost:5173/image-gen
附录 A:提示词解析正则的边界情况¶
| 输入 | 预期 |
|---|---|
第一张:xxx |
匹配,「第一张」→ index 0 |
第1张:xxx |
匹配,「第1张」→ index 0 |
第十二张:xxx |
匹配,「十二」→ index 11 |
第X张 (英文) |
不匹配,但可在正则中扩展 |
| 空行分隔的纯文本(无「第X张」标记) | 每段作为一个提示词,自动编号 |
| Word 中无图片 | reference_image_path 为 null |
附录 B:模型亲和度匹配逻辑¶
STYLE_KEYWORDS = {
"写实": ["写实", "照片级", "摄影", "真实", "超写实", "realistic"],
"插画": ["插画", "扁平", "卡通", "漫画", "illustration"],
"水墨": ["水墨", "国画", "山水", "工笔", "ink wash"],
"赛博朋克": ["赛博朋克", "cyberpunk", "霓虹", "neon"],
"3D": ["3D", "三维", "渲染", "C4D", "blender"],
"概念艺术": ["概念艺术", "concept art", "场景设计"],
}
def match_style(prompt_text: str, model_styles: list[str]) -> float:
score = 0.0
for style in model_styles:
keywords = STYLE_KEYWORDS.get(style, [style])
for kw in keywords:
if kw.lower() in prompt_text.lower():
score += 1.0
break
return score
11. 图片格式与分辨率规范¶
11.1 规则¶
除 AI 生成的最终图片外,所有图片统一转换为 WebP / 72 DPI。
| 图片类型 | 来源 | 格式 | 分辨率 | 说明 |
|---|---|---|---|---|
| 样图(参考图) | Word 文档内嵌图片 | WebP | 72 DPI | 上传时即时转换 |
| 前端缩略图 | 样图 / 结果图 | WebP | 72 DPI | 缩略图接口按需生成 |
| 前端预览图 | 样图 | WebP | 72 DPI | 同上 |
| AI 生成结果图 | Worker 调用模型 API | 原格式不变 | 原始分辨率 | PNG/原样保存,不做任何转换 |
AI 生成的结果图保持原始质量和分辨率,因为这是最终交付物。其余图片转为 WebP/72dpi 是为了减少存储成本和加速前端加载。
11.2 处理管线¶
flowchart LR
DOCX["Word 文档\n内嵌图片 (PNG/JPG)"] --> PARSE["WordParser\n_extract_images()"]
PARSE --> RAW["原始 bytes"]
RAW --> CONVERT["ImageProcessor\nconvert_to_webp()"]
CONVERT --> WEBP["WebP bytes\n(quality=80, dpi=72)"]
WEBP --> OSS["OSS 上传\nprompts/{session_id}/xxx.webp"]
GEN["AI 生成图\n(PNG/原格式)"] --> GEN_OSS["OSS 上传\nresults/{session_id}/xxx.png"]
GEN_OSS -->|"❌ 不转换"| STORE["按原格式存储"]
11.3 图片处理器¶
新增工具模块,所有非生成图片的统一转换入口:
# src/myapp/utils/image_processor.py
import io
from PIL import Image
class ImageProcessor:
"""
图片格式统一处理器。
规则:除 AI 生成结果外,全部转为 WebP / 72 DPI。
"""
WEBP_QUALITY = 80 # WebP 压缩质量 (0-100)
TARGET_DPI = 72 # 目标分辨率
MAX_DIMENSION = 2048 # 最大边长(超过等比缩放)
@classmethod
def convert_to_webp(cls, raw_bytes: bytes, source_format: str = None) -> bytes:
"""
将任意图片转为 WebP / 72 DPI。
Args:
raw_bytes: 原始图片字节
source_format: 可选,源格式提示(如 'png', 'jpg')
Returns:
WebP 格式字节
"""
img = Image.open(io.BytesIO(raw_bytes))
# 转 RGB(WebP 不支持 CMYK / RGBA 调色板模式)
if img.mode in ("RGBA", "LA", "P"):
img = img.convert("RGBA")
elif img.mode != "RGB":
img = img.convert("RGB")
# 限制最大尺寸(等比缩放)
w, h = img.size
if max(w, h) > cls.MAX_DIMENSION:
ratio = cls.MAX_DIMENSION / max(w, h)
img = img.resize((int(w * ratio), int(h * ratio)), Image.LANCZOS)
# 输出 WebP
output = io.BytesIO()
img.save(
output,
format="WEBP",
quality=cls.WEBP_QUALITY,
dpi=(cls.TARGET_DPI, cls.TARGET_DPI),
method=6, # 最慢但压缩率最高
)
return output.getvalue()
@classmethod
def generate_thumbnail(cls, raw_bytes: bytes, size: int = 320) -> bytes:
"""
生成缩略图(WebP / 72 DPI),用于前端列表预览。
"""
img = Image.open(io.BytesIO(raw_bytes))
if img.mode in ("RGBA", "LA", "P"):
img = img.convert("RGBA")
elif img.mode != "RGB":
img = img.convert("RGB")
# 居中裁剪正方形
w, h = img.size
s = min(w, h)
left = (w - s) // 2
top = (h - s) // 2
img = img.crop((left, top, left + s, top + s))
img = img.resize((size, size), Image.LANCZOS)
output = io.BytesIO()
img.save(output, format="WEBP", quality=70, dpi=(72, 72))
return output.getvalue()
11.4 更新 WordParser:提取时即时转换¶
# src/myapp/core/services/doc_parser.py(修改 _extract_images)
from myapp.utils.image_processor import ImageProcessor
class WordPromptParser:
# ... 其他代码不变 ...
def _extract_images(self, doc: Document) -> list[bytes]:
"""从 docx 中提取内嵌图片,即时转为 WebP"""
images: list[bytes] = []
for rel in doc.part.rels.values():
if "image" in rel.reltype:
raw = rel.target_part.blob
# ⬇ 关键:即时转换
webp = ImageProcessor.convert_to_webp(raw)
images.append(webp)
return images
11.5 更新 PromptService:上传路径改为 .webp¶
# src/myapp/core/services/prompt_service.py(修改 process_upload)
async def process_upload(self, file_bytes: bytes) -> ParseResult:
parser = WordPromptParser()
result = parser.parse(file_bytes)
result.session_id = uuid4().hex[:12]
for item in result.prompts:
if item.reference_image_data:
# WebP 格式,扩展名 .webp
key = f"prompts/{result.session_id}/{item.index:03d}_ref.webp"
item.reference_image_path = await self.storage.upload(
key=key,
data=item.reference_image_data, # 已是 WebP bytes
content_type="image/webp",
)
# 同时生成缩略图
thumb_data = ImageProcessor.generate_thumbnail(item.reference_image_data)
thumb_key = f"prompts/{result.session_id}/{item.index:03d}_thumb.webp"
await self.storage.upload(
key=thumb_key,
data=thumb_data,
content_type="image/webp",
)
item.reference_image_data = None
await self._save_to_db(result)
return result
11.6 Worker 生成图:明确标记「不做转换」¶
# src/messaging/workers/image_gen_worker.py(_process 方法中上传生成图)
# 生成图保持原始格式,不做 WebP/72dpi 转换
oss_key = f"results/{session_id}/{subtask_id}.png"
image_url = await self.storage.upload(
key=oss_key,
data=image_bytes, # AI 返回的原始 PNG bytes,原样存储
content_type="image/png", # 原格式
)
# 前端缩略图额外生成一份(这个可以做 WebP)
thumb_data = ImageProcessor.generate_thumbnail(image_bytes)
thumb_key = f"results/{session_id}/{subtask_id}_thumb.webp"
thumb_url = await self.storage.upload(
key=thumb_key,
data=thumb_data,
content_type="image/webp",
)
11.7 存储路径一览¶
OSS 存储结构
├── prompts/{session_id}/ ← 样图区(全部 WebP / 72 DPI)
│ ├── 000_ref.webp # 第一张提示词的参考图
│ ├── 000_thumb.webp # 缩略图
│ ├── 001_ref.webp
│ └── 001_thumb.webp
│
├── results/{session_id}/ ← 结果区
│ ├── gen_xxx_p0_m_dalle.png # AI 生成图 → png 原样
│ ├── gen_xxx_p0_m_dalle_thumb.webp # 缩略图 → webp
│ ├── gen_xxx_p1_m_flux.png
│ └── gen_xxx_p1_m_flux_thumb.webp
│
└── temp/ ← 临时文件(24h TTL)
11.8 前端加载策略¶
前端请求图时,优先用缩略图(列表/卡片),点开才加载原图:
┌─ ProgressCard ─────────────────────┐
│ [img: 000_thumb.webp] ← 列表用 │
│ 第一张:xxx │
│ ✅ dalle ⚡ flux ⏳ sd │
└────────────────────────────────────┘
│ 点击
▼
┌─ Lightbox ────────────────────────┐
│ [img: 原图 .png] ← 灯箱用 │
│ 模型: dalle | 下载 ⬇ │
└───────────────────────────────────┘
前端 API 响应中同时返回 image_url 和 thumb_url,组件自行选择加载哪个。
附录 C:Pillow 依赖¶
# pyproject.toml
[tool.poetry.dependencies]
Pillow = "^10.3.0" # WebP 编解码 + DPI 设置
WebP 编码需要系统级 libwebp,已在 Python 基础镜像中预装(python:3.12-slim 需额外 apt install libwebp-dev)。