!41 feat: .env example示例
Merge pull request !41 from 郭永昊/feat/async-generate-pipeline
This commit is contained in:
Regular → Executable
Regular → Executable
Regular → Executable
@@ -24,9 +24,10 @@ GEN2D_LLM_TEMPERATURE=0.7
|
|||||||
GEN2D_LLM_MAX_TOKENS=2048
|
GEN2D_LLM_MAX_TOKENS=2048
|
||||||
|
|
||||||
# 文生图模型(OpenAI 兼容 Images API)
|
# 文生图模型(OpenAI 兼容 Images API)
|
||||||
GEN2D_IMAGE_BASE_URL=https://api.suchuang.vip/v1
|
# 模型可选: dall-e-2, dall-e-3, gpt-image-1, gpt-image-2
|
||||||
GEN2D_IMAGE_API_KEY=
|
GEN2D_IMAGE_BASE_URL=https://api.weelinking.com/v1
|
||||||
GEN2D_IMAGE_MODEL=gpt-image-2-token
|
GEN2D_IMAGE_API_KEY=sk-your-api-key
|
||||||
|
GEN2D_IMAGE_MODEL=gpt-image-2
|
||||||
GEN2D_IMAGE_WIDTH=1024
|
GEN2D_IMAGE_WIDTH=1024
|
||||||
GEN2D_IMAGE_HEIGHT=1024
|
GEN2D_IMAGE_HEIGHT=1024
|
||||||
GEN2D_IMAGE_QUALITY=low
|
GEN2D_IMAGE_QUALITY=low
|
||||||
|
|||||||
Regular → Executable
Regular → Executable
Regular → Executable
@@ -25,9 +25,9 @@ llm:
|
|||||||
max_tokens: 2048
|
max_tokens: 2048
|
||||||
|
|
||||||
image_gen:
|
image_gen:
|
||||||
base_url: "https://api.suchuang.vip/v1"
|
base_url: ""
|
||||||
api_key: ""
|
api_key: ""
|
||||||
model: "gpt-image-2-token"
|
model: ""
|
||||||
width: 1024
|
width: 1024
|
||||||
height: 1024
|
height: 1024
|
||||||
quality: "low"
|
quality: "low"
|
||||||
|
|||||||
Regular → Executable
@@ -47,6 +47,7 @@ type TaskResponse struct {
|
|||||||
Prompt string `json:"prompt"`
|
Prompt string `json:"prompt"`
|
||||||
AssetType string `json:"assetType"`
|
AssetType string `json:"assetType"`
|
||||||
Status string `json:"status"`
|
Status string `json:"status"`
|
||||||
|
Stage string `json:"stage,omitempty"`
|
||||||
Progress int `json:"progress"`
|
Progress int `json:"progress"`
|
||||||
RetryCount int `json:"retryCount"`
|
RetryCount int `json:"retryCount"`
|
||||||
Error string `json:"error,omitempty"`
|
Error string `json:"error,omitempty"`
|
||||||
@@ -107,10 +108,15 @@ func Generate(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// runPipelineBg 后台执行生成管线,更新任务状态。
|
// runPipelineBg 后台执行生成管线,更新任务状态。
|
||||||
func runPipelineBg(ctx context.Context, projectID, taskID string, req GenerateRequest) {
|
func runPipelineBg(projectID, taskID string, req GenerateRequest) {
|
||||||
|
// 注入进度上报回调
|
||||||
|
ctx := service.WithProgressReporter(context.Background(), func(stage string, progress int) {
|
||||||
|
updateTaskProgress(taskID, "running", stage, progress)
|
||||||
|
})
|
||||||
|
|
||||||
l := logger.With("task_id", taskID, "project_id", projectID)
|
l := logger.With("task_id", taskID, "project_id", projectID)
|
||||||
|
|
||||||
updateStatus(taskID, "running", 10)
|
updateTaskProgress(taskID, "running", "prompt_builder", 5)
|
||||||
|
|
||||||
in := service.PipelineInput{
|
in := service.PipelineInput{
|
||||||
ProjectID: projectID,
|
ProjectID: projectID,
|
||||||
@@ -138,7 +144,7 @@ func runPipelineBg(ctx context.Context, projectID, taskID string, req GenerateRe
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
updateStatus(taskID, "saving", 80)
|
updateTaskProgress(taskID, "saving", "format_adapter", 90)
|
||||||
|
|
||||||
assets := make([]AssetResponse, len(output.Assets))
|
assets := make([]AssetResponse, len(output.Assets))
|
||||||
for i, a := range output.Assets {
|
for i, a := range output.Assets {
|
||||||
@@ -173,13 +179,14 @@ func runPipelineBg(ctx context.Context, projectID, taskID string, req GenerateRe
|
|||||||
l.Info("task completed", "asset_count", len(assets))
|
l.Info("task completed", "asset_count", len(assets))
|
||||||
}
|
}
|
||||||
|
|
||||||
func updateStatus(taskID, status string, progress int) {
|
func updateTaskProgress(taskID, status, stage string, progress int) {
|
||||||
rec, ok := taskStore.Load(taskID)
|
rec, ok := taskStore.Load(taskID)
|
||||||
if !ok {
|
if !ok {
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
r := rec.(*taskRecord)
|
r := rec.(*taskRecord)
|
||||||
r.task.Status = status
|
r.task.Status = status
|
||||||
|
r.task.Stage = stage
|
||||||
r.task.Progress = progress
|
r.task.Progress = progress
|
||||||
taskStore.Store(taskID, r)
|
taskStore.Store(taskID, r)
|
||||||
}
|
}
|
||||||
|
|||||||
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
+7
@@ -57,8 +57,10 @@ var promptOptimizerNode = compose.InvokableLambda(func(ctx context.Context, in P
|
|||||||
func promptOptimizerPreHandler(ctx context.Context, in PipelineInput, state *PipelineState) (PipelineInput, error) {
|
func promptOptimizerPreHandler(ctx context.Context, in PipelineInput, state *PipelineState) (PipelineInput, error) {
|
||||||
if state.RetryCount == 0 {
|
if state.RetryCount == 0 {
|
||||||
state.Input = in
|
state.Input = in
|
||||||
|
reportProgress(ctx, "prompt_builder", 10)
|
||||||
} else if state.RejectReason != "" {
|
} else if state.RejectReason != "" {
|
||||||
in.RejectReason = state.RejectReason
|
in.RejectReason = state.RejectReason
|
||||||
|
reportProgress(ctx, "prompt_builder", 30+state.RetryCount*10)
|
||||||
}
|
}
|
||||||
return in, nil
|
return in, nil
|
||||||
}
|
}
|
||||||
@@ -66,6 +68,7 @@ func promptOptimizerPreHandler(ctx context.Context, in PipelineInput, state *Pip
|
|||||||
// promptOptimizerPostHandler 将最终提示词写入全局状态。
|
// promptOptimizerPostHandler 将最终提示词写入全局状态。
|
||||||
func promptOptimizerPostHandler(ctx context.Context, out string, state *PipelineState) (string, error) {
|
func promptOptimizerPostHandler(ctx context.Context, out string, state *PipelineState) (string, error) {
|
||||||
state.FinalPrompt = out
|
state.FinalPrompt = out
|
||||||
|
reportProgress(ctx, "asset_generator", 35)
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -82,6 +85,7 @@ var assetGeneratorNode = compose.InvokableLambda(func(ctx context.Context, promp
|
|||||||
// assetGeneratorPostHandler 将原始图片写入全局状态。
|
// assetGeneratorPostHandler 将原始图片写入全局状态。
|
||||||
func assetGeneratorPostHandler(ctx context.Context, out []GeneratedImage, state *PipelineState) ([]GeneratedImage, error) {
|
func assetGeneratorPostHandler(ctx context.Context, out []GeneratedImage, state *PipelineState) ([]GeneratedImage, error) {
|
||||||
state.RawImages = out
|
state.RawImages = out
|
||||||
|
reportProgress(ctx, "quality_supervisor", 60)
|
||||||
return out, nil
|
return out, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -104,11 +108,14 @@ var qualitySupervisorNode = compose.InvokableLambda(func(ctx context.Context, im
|
|||||||
|
|
||||||
if pass {
|
if pass {
|
||||||
state.NextNode = nodeFormatAdapter
|
state.NextNode = nodeFormatAdapter
|
||||||
|
reportProgress(ctx, "format_adapter", 85)
|
||||||
} else if state.RetryCount >= 3 {
|
} else if state.RetryCount >= 3 {
|
||||||
state.NextNode = nodeFormatAdapter
|
state.NextNode = nodeFormatAdapter
|
||||||
|
reportProgress(ctx, "format_adapter", 85)
|
||||||
} else {
|
} else {
|
||||||
state.RetryCount++
|
state.RetryCount++
|
||||||
state.NextNode = nodePromptOptimizer
|
state.NextNode = nodePromptOptimizer
|
||||||
|
reportProgress(ctx, "quality_supervisor", 50+state.RetryCount*10)
|
||||||
}
|
}
|
||||||
|
|
||||||
input = state.Input
|
input = state.Input
|
||||||
|
|||||||
Regular → Executable
+19
@@ -16,6 +16,25 @@ const (
|
|||||||
nodeFormatAdapter = "format_adapter"
|
nodeFormatAdapter = "format_adapter"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
// ProgressReporter 管线进度回调:stage 为当前节点名,progress 为 0-100。
|
||||||
|
type ProgressReporter func(stage string, progress int)
|
||||||
|
|
||||||
|
type progressKeyType struct{}
|
||||||
|
|
||||||
|
var progressCtxKey progressKeyType
|
||||||
|
|
||||||
|
// WithProgressReporter 将进度回调注入 context。
|
||||||
|
func WithProgressReporter(ctx context.Context, r ProgressReporter) context.Context {
|
||||||
|
return context.WithValue(ctx, progressCtxKey, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
// reportProgress 从 context 取出回调上报进度。
|
||||||
|
func reportProgress(ctx context.Context, stage string, progress int) {
|
||||||
|
if r, ok := ctx.Value(progressCtxKey).(ProgressReporter); ok {
|
||||||
|
r(stage, progress)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// NewGenerateGraph 创建生成管线 Graph(PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter)。
|
// NewGenerateGraph 创建生成管线 Graph(PromptOptimizer → AssetGenerator → QualitySupervisor → FormatAdapter)。
|
||||||
//
|
//
|
||||||
// START → PromptOptimizer → AssetGenerator → QualitySupervisor
|
// START → PromptOptimizer → AssetGenerator → QualitySupervisor
|
||||||
|
|||||||
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
+3
-3
@@ -173,9 +173,9 @@ export default function GenerateForm({ onSubmit, submitting }: GenerateFormProps
|
|||||||
type="number"
|
type="number"
|
||||||
value={params.resolution}
|
value={params.resolution}
|
||||||
onChange={e => setParams({ resolution: Number(e.target.value) })}
|
onChange={e => setParams({ resolution: Number(e.target.value) })}
|
||||||
min={32}
|
min={1024}
|
||||||
max={1024}
|
max={1536}
|
||||||
step={32}
|
step={512}
|
||||||
style={{ width: 100 }}
|
style={{ width: 100 }}
|
||||||
/>
|
/>
|
||||||
<span style={{ fontSize: 13, color: 'var(--text-muted)' }}>px</span>
|
<span style={{ fontSize: 13, color: 'var(--text-muted)' }}>px</span>
|
||||||
|
|||||||
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
@@ -19,9 +19,12 @@ export default function GeneratePage() {
|
|||||||
const { style: projectStyle, loadProject } = useProjectStore()
|
const { style: projectStyle, loadProject } = useProjectStore()
|
||||||
const {
|
const {
|
||||||
status,
|
status,
|
||||||
|
stage,
|
||||||
progress,
|
progress,
|
||||||
taskId,
|
taskId,
|
||||||
statusText,
|
statusText,
|
||||||
|
retryCount,
|
||||||
|
rejectReason,
|
||||||
submit,
|
submit,
|
||||||
reset: resetGeneration,
|
reset: resetGeneration,
|
||||||
} = useGenerationStore()
|
} = useGenerationStore()
|
||||||
@@ -83,11 +86,11 @@ export default function GeneratePage() {
|
|||||||
) : (
|
) : (
|
||||||
<div style={{ display: 'flex', flexDirection: 'column', gap: 24 }}>
|
<div style={{ display: 'flex', flexDirection: 'column', gap: 24 }}>
|
||||||
<ProgressBar
|
<ProgressBar
|
||||||
stage={null}
|
stage={stage}
|
||||||
progress={progress}
|
progress={progress}
|
||||||
status={status}
|
status={status}
|
||||||
retryCount={0}
|
retryCount={retryCount}
|
||||||
rejectReason={null}
|
rejectReason={rejectReason}
|
||||||
/>
|
/>
|
||||||
|
|
||||||
{status === 'running' && (
|
{status === 'running' && (
|
||||||
|
|||||||
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
@@ -1,5 +1,5 @@
|
|||||||
import { create } from 'zustand'
|
import { create } from 'zustand'
|
||||||
import type { Asset, GenerateRequest } from '../api/types'
|
import type { Asset, GenerateRequest, PipelineStage } from '../api/types'
|
||||||
import { submitGenerate, getTask, getAssets } from '../api/generate'
|
import { submitGenerate, getTask, getAssets } from '../api/generate'
|
||||||
|
|
||||||
type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed'
|
type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed'
|
||||||
@@ -7,9 +7,12 @@ type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed'
|
|||||||
interface GenerationState {
|
interface GenerationState {
|
||||||
taskId: string | null
|
taskId: string | null
|
||||||
projectId: string | null
|
projectId: string | null
|
||||||
|
stage: PipelineStage | null
|
||||||
progress: number
|
progress: number
|
||||||
status: Status
|
status: Status
|
||||||
statusText: string
|
statusText: string
|
||||||
|
retryCount: number
|
||||||
|
rejectReason: string | null
|
||||||
assets: Asset[]
|
assets: Asset[]
|
||||||
error: string | null
|
error: string | null
|
||||||
submit: (req: GenerateRequest) => Promise<void>
|
submit: (req: GenerateRequest) => Promise<void>
|
||||||
@@ -25,18 +28,21 @@ function stopPolling() {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
export const useGenerationStore = create<GenerationState>((set, get) => ({
|
export const useGenerationStore = create<GenerationState>((set) => ({
|
||||||
taskId: null,
|
taskId: null,
|
||||||
projectId: null,
|
projectId: null,
|
||||||
|
stage: null,
|
||||||
progress: 0,
|
progress: 0,
|
||||||
status: 'idle',
|
status: 'idle',
|
||||||
statusText: '',
|
statusText: '',
|
||||||
|
retryCount: 0,
|
||||||
|
rejectReason: null,
|
||||||
assets: [],
|
assets: [],
|
||||||
error: null,
|
error: null,
|
||||||
|
|
||||||
submit: async (req) => {
|
submit: async (req) => {
|
||||||
stopPolling()
|
stopPolling()
|
||||||
set({ status: 'submitting', error: null, statusText: '提交中...' })
|
set({ status: 'submitting', error: null, statusText: '提交中...', stage: null })
|
||||||
try {
|
try {
|
||||||
const { taskId } = await submitGenerate(req)
|
const { taskId } = await submitGenerate(req)
|
||||||
|
|
||||||
@@ -44,24 +50,26 @@ export const useGenerationStore = create<GenerationState>((set, get) => ({
|
|||||||
taskId,
|
taskId,
|
||||||
projectId: req.projectId,
|
projectId: req.projectId,
|
||||||
status: 'running',
|
status: 'running',
|
||||||
progress: 10,
|
progress: 5,
|
||||||
|
stage: 'prompt_builder',
|
||||||
statusText: '任务已提交,等待生成...',
|
statusText: '任务已提交,等待生成...',
|
||||||
})
|
})
|
||||||
|
|
||||||
// 开始轮询进度
|
|
||||||
pollTimer = setInterval(async () => {
|
pollTimer = setInterval(async () => {
|
||||||
try {
|
try {
|
||||||
const task = await getTask(taskId)
|
const task = await getTask(taskId)
|
||||||
|
|
||||||
set({
|
set((s) => ({
|
||||||
progress: task.progress ?? get().progress,
|
stage: task.stage ?? s.stage,
|
||||||
|
progress: task.progress ?? s.progress,
|
||||||
|
retryCount: task.retryCount ?? s.retryCount,
|
||||||
statusText:
|
statusText:
|
||||||
task.status === 'running'
|
task.status === 'running'
|
||||||
? '生成中...'
|
? stageLabel(task.stage)
|
||||||
: task.status === 'pending'
|
: task.status === 'pending'
|
||||||
? '排队中...'
|
? '排队中...'
|
||||||
: task.status,
|
: task.status,
|
||||||
})
|
}))
|
||||||
|
|
||||||
if (task.status === 'completed') {
|
if (task.status === 'completed') {
|
||||||
stopPolling()
|
stopPolling()
|
||||||
@@ -71,6 +79,7 @@ export const useGenerationStore = create<GenerationState>((set, get) => ({
|
|||||||
set({
|
set({
|
||||||
status: 'completed',
|
status: 'completed',
|
||||||
progress: 100,
|
progress: 100,
|
||||||
|
stage: 'format_adapter',
|
||||||
statusText: '生成完成',
|
statusText: '生成完成',
|
||||||
assets,
|
assets,
|
||||||
})
|
})
|
||||||
@@ -85,7 +94,7 @@ export const useGenerationStore = create<GenerationState>((set, get) => ({
|
|||||||
} catch {
|
} catch {
|
||||||
// 网络错误不中断轮询
|
// 网络错误不中断轮询
|
||||||
}
|
}
|
||||||
}, 2000)
|
}, 1500)
|
||||||
} catch (err) {
|
} catch (err) {
|
||||||
stopPolling()
|
stopPolling()
|
||||||
set({ status: 'failed', error: (err as Error).message, statusText: '提交失败' })
|
set({ status: 'failed', error: (err as Error).message, statusText: '提交失败' })
|
||||||
@@ -97,11 +106,24 @@ export const useGenerationStore = create<GenerationState>((set, get) => ({
|
|||||||
set({
|
set({
|
||||||
taskId: null,
|
taskId: null,
|
||||||
projectId: null,
|
projectId: null,
|
||||||
|
stage: null,
|
||||||
progress: 0,
|
progress: 0,
|
||||||
status: 'idle',
|
status: 'idle',
|
||||||
statusText: '',
|
statusText: '',
|
||||||
|
retryCount: 0,
|
||||||
|
rejectReason: null,
|
||||||
assets: [],
|
assets: [],
|
||||||
error: null,
|
error: null,
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}))
|
}))
|
||||||
|
|
||||||
|
function stageLabel(stage?: string): string {
|
||||||
|
switch (stage) {
|
||||||
|
case 'prompt_builder': return '优化提示词...'
|
||||||
|
case 'asset_generator': return '生成素材中...'
|
||||||
|
case 'quality_supervisor': return '质检中...'
|
||||||
|
case 'format_adapter': return '格式转换中...'
|
||||||
|
default: return '生成中...'
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
Regular → Executable
Regular → Executable
+1
-1
@@ -28,7 +28,7 @@ interface TaskState {
|
|||||||
}
|
}
|
||||||
|
|
||||||
const DEFAULT_PARAMS: TaskParams = {
|
const DEFAULT_PARAMS: TaskParams = {
|
||||||
resolution: 256,
|
resolution: 1024,
|
||||||
format: 'spritesheet',
|
format: 'spritesheet',
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
Regular → Executable
+4
@@ -14,6 +14,10 @@ export default defineConfig({
|
|||||||
target: 'http://localhost:8080',
|
target: 'http://localhost:8080',
|
||||||
changeOrigin: true,
|
changeOrigin: true,
|
||||||
},
|
},
|
||||||
|
'/generation': {
|
||||||
|
target: 'http://localhost:8080',
|
||||||
|
changeOrigin: true,
|
||||||
|
},
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
})
|
})
|
||||||
|
|||||||
Reference in New Issue
Block a user