feat(generate): 生成管线改为异步,前端接入轮询进度
后端: - POST /api/v1/generate 改为异步模式,立即返回 taskId - 管线在后台 goroutine 执行,状态通过 GET /tasks/:taskId 轮询 - 任务状态流转: pending → running → saving → completed/failed 前端: - generation store: submit 后每 2s 轮询 getTask,完成时自动 getAssets - GeneratePage: 实时显示轮询状态文本 + 进度条 - ResultPage: 挂载时轮询任务,未完成显示骨架屏+进度,完成自动加载素材 - Types: Task.status 增加 submitted / saving 状态
This commit is contained in:
@@ -1,7 +1,9 @@
|
|||||||
package handler
|
package handler
|
||||||
|
|
||||||
import (
|
import (
|
||||||
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"log"
|
||||||
"net/http"
|
"net/http"
|
||||||
"os"
|
"os"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
@@ -32,8 +34,6 @@ type GenerateRequest struct {
|
|||||||
// GenerateResponse 素材生成响应体。
|
// GenerateResponse 素材生成响应体。
|
||||||
type GenerateResponse struct {
|
type GenerateResponse struct {
|
||||||
TaskID string `json:"taskId"`
|
TaskID string `json:"taskId"`
|
||||||
Assets []AssetResponse `json:"assets"`
|
|
||||||
Metadata service.AssetMetadata `json:"metadata"`
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// AssetResponse 单个素材响应。
|
// AssetResponse 单个素材响应。
|
||||||
@@ -72,7 +72,8 @@ var (
|
|||||||
taskStore = sync.Map{} // taskID → *taskRecord
|
taskStore = sync.Map{} // taskID → *taskRecord
|
||||||
)
|
)
|
||||||
|
|
||||||
// Generate 素材生成接口。
|
// Generate 素材生成接口(异步)。
|
||||||
|
// 立即返回 taskId,后台执行管线,前端通过 GET /tasks/:taskId 轮询进度。
|
||||||
func Generate(c *gin.Context) {
|
func Generate(c *gin.Context) {
|
||||||
var req GenerateRequest
|
var req GenerateRequest
|
||||||
if err := c.ShouldBindJSON(&req); err != nil {
|
if err := c.ShouldBindJSON(&req); err != nil {
|
||||||
@@ -85,6 +86,31 @@ func Generate(c *gin.Context) {
|
|||||||
projectID = "default"
|
projectID = "default"
|
||||||
}
|
}
|
||||||
taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli())
|
taskID := fmt.Sprintf("task-%d", time.Now().UnixMilli())
|
||||||
|
createdAt := time.Now().Format(time.RFC3339)
|
||||||
|
|
||||||
|
// 存入 pending 状态
|
||||||
|
taskStore.Store(taskID, &taskRecord{
|
||||||
|
task: TaskResponse{
|
||||||
|
ID: taskID,
|
||||||
|
ProjectID: projectID,
|
||||||
|
Prompt: req.Prompt,
|
||||||
|
AssetType: req.AssetType,
|
||||||
|
Status: "pending",
|
||||||
|
Progress: 0,
|
||||||
|
CreatedAt: createdAt,
|
||||||
|
},
|
||||||
|
})
|
||||||
|
|
||||||
|
// 返回 taskId
|
||||||
|
c.JSON(http.StatusOK, model.OK(GenerateResponse{TaskID: taskID}))
|
||||||
|
|
||||||
|
// 后台执行管线
|
||||||
|
go runPipelineBg(projectID, taskID, req)
|
||||||
|
}
|
||||||
|
|
||||||
|
// runPipelineBg 后台执行生成管线,更新任务状态。
|
||||||
|
func runPipelineBg(projectID, taskID string, req GenerateRequest) {
|
||||||
|
updateStatus(taskID, "running", 10)
|
||||||
|
|
||||||
in := service.PipelineInput{
|
in := service.PipelineInput{
|
||||||
ProjectID: projectID,
|
ProjectID: projectID,
|
||||||
@@ -105,16 +131,20 @@ func Generate(c *gin.Context) {
|
|||||||
},
|
},
|
||||||
}
|
}
|
||||||
|
|
||||||
output, err := service.RunPipeline(c.Request.Context(), in)
|
ctx := context.Background()
|
||||||
|
output, err := service.RunPipeline(ctx, in)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "素材生成失败: "+err.Error()))
|
log.Printf("[generate] task %s failed: %v", taskID, err)
|
||||||
|
updateFailed(taskID, err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
|
updateStatus(taskID, "saving", 80)
|
||||||
|
|
||||||
// 保存图片到 ../generation/{projectId}/{taskId}/
|
// 保存图片到 ../generation/{projectId}/{taskId}/
|
||||||
genDir := filepath.Join("..", "generation", projectID, taskID)
|
genDir := filepath.Join("..", "generation", projectID, taskID)
|
||||||
if err := os.MkdirAll(genDir, 0755); err != nil {
|
if err := os.MkdirAll(genDir, 0755); err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "创建输出目录失败: "+err.Error()))
|
updateFailed(taskID, "创建输出目录失败: "+err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -123,7 +153,7 @@ func Generate(c *gin.Context) {
|
|||||||
filename := fmt.Sprintf("%d.%s", i, a.Format)
|
filename := fmt.Sprintf("%d.%s", i, a.Format)
|
||||||
filePath := filepath.Join(genDir, filename)
|
filePath := filepath.Join(genDir, filename)
|
||||||
if err := os.WriteFile(filePath, a.Data, 0644); err != nil {
|
if err := os.WriteFile(filePath, a.Data, 0644); err != nil {
|
||||||
c.JSON(http.StatusInternalServerError, model.Fail(http.StatusInternalServerError, "保存图片失败: "+err.Error()))
|
updateFailed(taskID, "保存图片失败: "+err.Error())
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
assets[i] = AssetResponse{
|
assets[i] = AssetResponse{
|
||||||
@@ -132,7 +162,7 @@ func Generate(c *gin.Context) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// 存储任务记录到内存
|
// 更新为完成状态
|
||||||
taskStore.Store(taskID, &taskRecord{
|
taskStore.Store(taskID, &taskRecord{
|
||||||
task: TaskResponse{
|
task: TaskResponse{
|
||||||
ID: taskID,
|
ID: taskID,
|
||||||
@@ -147,11 +177,29 @@ func Generate(c *gin.Context) {
|
|||||||
metadata: output.Metadata,
|
metadata: output.Metadata,
|
||||||
})
|
})
|
||||||
|
|
||||||
c.JSON(http.StatusOK, model.OK(GenerateResponse{
|
log.Printf("[generate] task %s completed, %d assets", taskID, len(assets))
|
||||||
TaskID: taskID,
|
}
|
||||||
Assets: assets,
|
|
||||||
Metadata: output.Metadata,
|
func updateStatus(taskID, status string, progress int) {
|
||||||
}))
|
rec, ok := taskStore.Load(taskID)
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r := rec.(*taskRecord)
|
||||||
|
r.task.Status = status
|
||||||
|
r.task.Progress = progress
|
||||||
|
taskStore.Store(taskID, r)
|
||||||
|
}
|
||||||
|
|
||||||
|
func updateFailed(taskID, errMsg string) {
|
||||||
|
rec, ok := taskStore.Load(taskID)
|
||||||
|
if !ok {
|
||||||
|
return
|
||||||
|
}
|
||||||
|
r := rec.(*taskRecord)
|
||||||
|
r.task.Status = "failed"
|
||||||
|
r.task.Error = errMsg
|
||||||
|
taskStore.Store(taskID, r)
|
||||||
}
|
}
|
||||||
|
|
||||||
// GetTask 查询任务信息。
|
// GetTask 查询任务信息。
|
||||||
@@ -175,6 +223,10 @@ func GetAssets(c *gin.Context) {
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
r := rec.(*taskRecord)
|
r := rec.(*taskRecord)
|
||||||
|
if r.task.Status != "completed" {
|
||||||
|
c.JSON(http.StatusBadRequest, model.Fail(http.StatusBadRequest, "任务尚未完成,当前状态: "+r.task.Status))
|
||||||
|
return
|
||||||
|
}
|
||||||
c.JSON(http.StatusOK, model.OK(AssetsResponse{
|
c.JSON(http.StatusOK, model.OK(AssetsResponse{
|
||||||
Assets: r.assets,
|
Assets: r.assets,
|
||||||
Metadata: r.metadata,
|
Metadata: r.metadata,
|
||||||
|
|||||||
@@ -19,13 +19,6 @@ export async function getTask(taskId: string): Promise<Task> {
|
|||||||
|
|
||||||
export async function getAssets(taskId: string): Promise<Asset[]> {
|
export async function getAssets(taskId: string): Promise<Asset[]> {
|
||||||
const resp = await get<AssetsResponse>(`/api/v1/tasks/${taskId}/assets`)
|
const resp = await get<AssetsResponse>(`/api/v1/tasks/${taskId}/assets`)
|
||||||
return toAssetList(resp)
|
|
||||||
}
|
|
||||||
|
|
||||||
/** 将响应转为 Asset[] 供前端组件使用 */
|
|
||||||
function toAssetList(
|
|
||||||
resp: AssetsResponse | GenerateResponse,
|
|
||||||
): Asset[] {
|
|
||||||
return resp.assets.map((a, i) => ({
|
return resp.assets.map((a, i) => ({
|
||||||
id: `asset-${i}`,
|
id: `asset-${i}`,
|
||||||
url: a.url,
|
url: a.url,
|
||||||
|
|||||||
@@ -67,7 +67,7 @@ export interface Task {
|
|||||||
projectId: string
|
projectId: string
|
||||||
prompt: string
|
prompt: string
|
||||||
assetType: string
|
assetType: string
|
||||||
status: 'pending' | 'running' | 'completed' | 'failed'
|
status: 'pending' | 'submitted' | 'running' | 'completed' | 'failed'
|
||||||
stage?: PipelineStage
|
stage?: PipelineStage
|
||||||
progress?: number
|
progress?: number
|
||||||
retryCount?: number
|
retryCount?: number
|
||||||
@@ -106,19 +106,9 @@ export interface GenerateRequest {
|
|||||||
format?: 'spritesheet' | 'individual'
|
format?: 'spritesheet' | 'individual'
|
||||||
}
|
}
|
||||||
|
|
||||||
// 生成响应 — 对应 POST /api/v1/generate 返回
|
// 生成响应 — 对应 POST /api/v1/generate 返回(异步,仅含 taskId)
|
||||||
export interface GenerateResponse {
|
export interface GenerateResponse {
|
||||||
taskId: string
|
taskId: string
|
||||||
assets: {
|
|
||||||
url: string
|
|
||||||
format: string
|
|
||||||
}[]
|
|
||||||
metadata: {
|
|
||||||
frameWidth: number
|
|
||||||
frameHeight: number
|
|
||||||
frameCount: number
|
|
||||||
directions: number
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
|
|
||||||
// 素材列表响应 — 对应 GET /api/v1/tasks/:taskId/assets
|
// 素材列表响应 — 对应 GET /api/v1/tasks/:taskId/assets
|
||||||
|
|||||||
@@ -21,6 +21,7 @@ export default function GeneratePage() {
|
|||||||
status,
|
status,
|
||||||
progress,
|
progress,
|
||||||
taskId,
|
taskId,
|
||||||
|
statusText,
|
||||||
submit,
|
submit,
|
||||||
reset: resetGeneration,
|
reset: resetGeneration,
|
||||||
} = useGenerationStore()
|
} = useGenerationStore()
|
||||||
@@ -91,7 +92,7 @@ export default function GeneratePage() {
|
|||||||
|
|
||||||
{status === 'running' && (
|
{status === 'running' && (
|
||||||
<p style={{ textAlign: 'center', color: 'var(--text-secondary)' }}>
|
<p style={{ textAlign: 'center', color: 'var(--text-secondary)' }}>
|
||||||
管线执行中,请稍候...
|
{statusText || '管线执行中,请稍候...'}
|
||||||
</p>
|
</p>
|
||||||
)}
|
)}
|
||||||
|
|
||||||
|
|||||||
@@ -1,4 +1,4 @@
|
|||||||
import { useEffect, useState } from 'react'
|
import { useEffect, useState, useRef } from 'react'
|
||||||
import { Link, useParams } from 'react-router-dom'
|
import { Link, useParams } from 'react-router-dom'
|
||||||
import { getTask, getAssets } from '../api/generate'
|
import { getTask, getAssets } from '../api/generate'
|
||||||
import type { Asset, Task } from '../api/types'
|
import type { Asset, Task } from '../api/types'
|
||||||
@@ -10,19 +10,47 @@ export default function ResultPage() {
|
|||||||
const [task, setTask] = useState<Task | null>(null)
|
const [task, setTask] = useState<Task | null>(null)
|
||||||
const [assets, setAssets] = useState<Asset[]>([])
|
const [assets, setAssets] = useState<Asset[]>([])
|
||||||
const [loading, setLoading] = useState(true)
|
const [loading, setLoading] = useState(true)
|
||||||
|
const [polling, setPolling] = useState(false)
|
||||||
|
const pollRef = useRef<ReturnType<typeof setInterval>>()
|
||||||
|
|
||||||
useEffect(() => {
|
useEffect(() => {
|
||||||
if (!taskId) return
|
if (!taskId) return
|
||||||
setLoading(true)
|
setLoading(true)
|
||||||
Promise.all([getTask(taskId), getAssets(taskId)])
|
|
||||||
.then(([t, a]) => {
|
const fetchTask = async () => {
|
||||||
|
try {
|
||||||
|
const t = await getTask(taskId)
|
||||||
setTask(t)
|
setTask(t)
|
||||||
|
|
||||||
|
if (t.status === 'completed') {
|
||||||
|
if (pollRef.current) clearInterval(pollRef.current)
|
||||||
|
setPolling(false)
|
||||||
|
const a = await getAssets(taskId)
|
||||||
setAssets(a)
|
setAssets(a)
|
||||||
})
|
setLoading(false)
|
||||||
.finally(() => setLoading(false))
|
} else if (t.status === 'failed') {
|
||||||
|
if (pollRef.current) clearInterval(pollRef.current)
|
||||||
|
setPolling(false)
|
||||||
|
setLoading(false)
|
||||||
|
} else if (!pollRef.current) {
|
||||||
|
// 开始轮询
|
||||||
|
setPolling(true)
|
||||||
|
pollRef.current = setInterval(fetchTask, 2000)
|
||||||
|
}
|
||||||
|
} catch {
|
||||||
|
// 出错也停止加载态
|
||||||
|
setLoading(false)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
fetchTask()
|
||||||
|
|
||||||
|
return () => {
|
||||||
|
if (pollRef.current) clearInterval(pollRef.current)
|
||||||
|
}
|
||||||
}, [taskId])
|
}, [taskId])
|
||||||
|
|
||||||
if (loading) {
|
if (loading || polling) {
|
||||||
return (
|
return (
|
||||||
<div className="container page-enter" style={{ paddingTop: 40 }}>
|
<div className="container page-enter" style={{ paddingTop: 40 }}>
|
||||||
<div className="card" style={{ marginBottom: 24 }}>
|
<div className="card" style={{ marginBottom: 24 }}>
|
||||||
@@ -30,6 +58,13 @@ export default function ResultPage() {
|
|||||||
<div style={{ marginTop: 16 }}>
|
<div style={{ marginTop: 16 }}>
|
||||||
<Skeleton variant="text" lines={4} />
|
<Skeleton variant="text" lines={4} />
|
||||||
</div>
|
</div>
|
||||||
|
{task && (
|
||||||
|
<p style={{ textAlign: 'center', color: 'var(--text-secondary)', marginTop: 16 }}>
|
||||||
|
{task.status === 'pending' && '任务排队中...'}
|
||||||
|
{task.status === 'running' && `生成中... ${task.progress ?? 0}%`}
|
||||||
|
{task.status === 'submitted' && '已提交,等待处理...'}
|
||||||
|
</p>
|
||||||
|
)}
|
||||||
</div>
|
</div>
|
||||||
<div className="card">
|
<div className="card">
|
||||||
<Skeleton variant="card" />
|
<Skeleton variant="card" />
|
||||||
@@ -53,27 +88,35 @@ export default function ResultPage() {
|
|||||||
<div className="container page-enter" style={{ paddingTop: 24, paddingBottom: 40 }}>
|
<div className="container page-enter" style={{ paddingTop: 24, paddingBottom: 40 }}>
|
||||||
<h1 style={{ fontSize: 24, marginBottom: 32 }}>生成结果</h1>
|
<h1 style={{ fontSize: 24, marginBottom: 32 }}>生成结果</h1>
|
||||||
|
|
||||||
{/* 任务信息 */}
|
|
||||||
<section className="card" style={{ marginBottom: 24 }}>
|
<section className="card" style={{ marginBottom: 24 }}>
|
||||||
<h2 style={{ fontSize: 16, marginBottom: 16 }}>任务信息</h2>
|
<div style={{ display: 'flex', justifyContent: 'space-between', alignItems: 'center', marginBottom: 16 }}>
|
||||||
<div
|
<h2 style={{ fontSize: 16, marginBottom: 0 }}>任务信息</h2>
|
||||||
style={{
|
<div style={{ display: 'flex', gap: 12 }}>
|
||||||
display: 'grid',
|
{assets.length > 0 && (
|
||||||
gridTemplateColumns: '120px 1fr',
|
<button
|
||||||
gap: '8px 16px',
|
className="btn-primary"
|
||||||
fontSize: 13,
|
onClick={() => downloadAssets(assets)}
|
||||||
}}
|
style={{ padding: '8px 20px', fontSize: 13 }}
|
||||||
>
|
>
|
||||||
|
下载全部素材
|
||||||
|
</button>
|
||||||
|
)}
|
||||||
|
<Link
|
||||||
|
to={`/projects/${projectId}/generate`}
|
||||||
|
className="btn-secondary"
|
||||||
|
style={{ padding: '8px 20px', fontSize: 13, borderRadius: 'var(--radius)', display: 'inline-block' }}
|
||||||
|
>
|
||||||
|
继续生成
|
||||||
|
</Link>
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
|
<div style={{ display: 'grid', gridTemplateColumns: '120px 1fr', gap: '8px 16px', fontSize: 13 }}>
|
||||||
<span style={{ color: 'var(--text-secondary)' }}>提示词</span>
|
<span style={{ color: 'var(--text-secondary)' }}>提示词</span>
|
||||||
<span>{task.prompt}</span>
|
<span>{task.prompt}</span>
|
||||||
<span style={{ color: 'var(--text-secondary)' }}>素材类型</span>
|
<span style={{ color: 'var(--text-secondary)' }}>素材类型</span>
|
||||||
<span>{task.assetType}</span>
|
<span>{task.assetType}</span>
|
||||||
<span style={{ color: 'var(--text-secondary)' }}>状态</span>
|
<span style={{ color: 'var(--text-secondary)' }}>状态</span>
|
||||||
<span
|
<span style={{ color: task.status === 'completed' ? 'var(--success)' : 'var(--error)' }}>
|
||||||
style={{
|
|
||||||
color: task.status === 'completed' ? 'var(--success)' : 'var(--error)',
|
|
||||||
}}
|
|
||||||
>
|
|
||||||
{task.status === 'completed' ? '已完成' : '失败'}
|
{task.status === 'completed' ? '已完成' : '失败'}
|
||||||
</span>
|
</span>
|
||||||
<span style={{ color: 'var(--text-secondary)' }}>创建时间</span>
|
<span style={{ color: 'var(--text-secondary)' }}>创建时间</span>
|
||||||
@@ -93,29 +136,29 @@ export default function ResultPage() {
|
|||||||
</div>
|
</div>
|
||||||
</section>
|
</section>
|
||||||
|
|
||||||
{/* 素材预览 */}
|
|
||||||
<section className="card" style={{ marginBottom: 24 }}>
|
<section className="card" style={{ marginBottom: 24 }}>
|
||||||
<div
|
<h2 style={{ fontSize: 16, marginBottom: 16 }}>素材预览</h2>
|
||||||
style={{
|
|
||||||
display: 'flex',
|
|
||||||
justifyContent: 'space-between',
|
|
||||||
alignItems: 'center',
|
|
||||||
marginBottom: 16,
|
|
||||||
}}
|
|
||||||
>
|
|
||||||
<h2 style={{ fontSize: 16 }}>素材预览</h2>
|
|
||||||
{assets.length > 0 && (
|
|
||||||
<button
|
|
||||||
className="btn-primary"
|
|
||||||
onClick={() => alert('下载功能即将上线')}
|
|
||||||
style={{ padding: '8px 20px', fontSize: 13 }}
|
|
||||||
>
|
|
||||||
下载素材
|
|
||||||
</button>
|
|
||||||
)}
|
|
||||||
</div>
|
|
||||||
<AssetPreview assets={assets} />
|
<AssetPreview assets={assets} />
|
||||||
</section>
|
</section>
|
||||||
</div>
|
</div>
|
||||||
)
|
)
|
||||||
}
|
}
|
||||||
|
|
||||||
|
async function downloadAssets(assets: Asset[]) {
|
||||||
|
for (const a of assets) {
|
||||||
|
try {
|
||||||
|
const res = await fetch(a.url)
|
||||||
|
const blob = await res.blob()
|
||||||
|
const blobUrl = URL.createObjectURL(blob)
|
||||||
|
const link = document.createElement('a')
|
||||||
|
link.href = blobUrl
|
||||||
|
link.download = `${a.id}.${a.format}`
|
||||||
|
document.body.appendChild(link)
|
||||||
|
link.click()
|
||||||
|
document.body.removeChild(link)
|
||||||
|
URL.revokeObjectURL(blobUrl)
|
||||||
|
} catch {
|
||||||
|
window.open(a.url, '_blank')
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@@ -1,6 +1,6 @@
|
|||||||
import { create } from 'zustand'
|
import { create } from 'zustand'
|
||||||
import type { Asset, GenerateRequest, GenerateResponse } from '../api/types'
|
import type { Asset, GenerateRequest } from '../api/types'
|
||||||
import { submitGenerate } from '../api/generate'
|
import { submitGenerate, getTask, getAssets } from '../api/generate'
|
||||||
|
|
||||||
type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed'
|
type Status = 'idle' | 'submitting' | 'running' | 'completed' | 'failed'
|
||||||
|
|
||||||
@@ -9,66 +9,99 @@ interface GenerationState {
|
|||||||
projectId: string | null
|
projectId: string | null
|
||||||
progress: number
|
progress: number
|
||||||
status: Status
|
status: Status
|
||||||
|
statusText: string
|
||||||
assets: Asset[]
|
assets: Asset[]
|
||||||
error: string | null
|
error: string | null
|
||||||
submit: (req: GenerateRequest) => Promise<void>
|
submit: (req: GenerateRequest) => Promise<void>
|
||||||
reset: () => void
|
reset: () => void
|
||||||
}
|
}
|
||||||
|
|
||||||
export const useGenerationStore = create<GenerationState>((set) => ({
|
let pollTimer: ReturnType<typeof setInterval> | null = null
|
||||||
|
|
||||||
|
function stopPolling() {
|
||||||
|
if (pollTimer) {
|
||||||
|
clearInterval(pollTimer)
|
||||||
|
pollTimer = null
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
export const useGenerationStore = create<GenerationState>((set, get) => ({
|
||||||
taskId: null,
|
taskId: null,
|
||||||
projectId: null,
|
projectId: null,
|
||||||
progress: 0,
|
progress: 0,
|
||||||
status: 'idle',
|
status: 'idle',
|
||||||
|
statusText: '',
|
||||||
assets: [],
|
assets: [],
|
||||||
error: null,
|
error: null,
|
||||||
|
|
||||||
submit: async (req) => {
|
submit: async (req) => {
|
||||||
set({ status: 'submitting', error: null })
|
stopPolling()
|
||||||
|
set({ status: 'submitting', error: null, statusText: '提交中...' })
|
||||||
try {
|
try {
|
||||||
set({ status: 'running', progress: 30 })
|
const { taskId } = await submitGenerate(req)
|
||||||
const result = await submitGenerate(req)
|
|
||||||
|
|
||||||
set({ progress: 80 })
|
|
||||||
const assets = mapAssets(result)
|
|
||||||
|
|
||||||
set({
|
set({
|
||||||
taskId: result.taskId,
|
taskId,
|
||||||
projectId: req.projectId,
|
projectId: req.projectId,
|
||||||
|
status: 'running',
|
||||||
|
progress: 10,
|
||||||
|
statusText: '任务已提交,等待生成...',
|
||||||
|
})
|
||||||
|
|
||||||
|
// 开始轮询进度
|
||||||
|
pollTimer = setInterval(async () => {
|
||||||
|
try {
|
||||||
|
const task = await getTask(taskId)
|
||||||
|
|
||||||
|
set({
|
||||||
|
progress: task.progress ?? get().progress,
|
||||||
|
statusText:
|
||||||
|
task.status === 'running'
|
||||||
|
? '生成中...'
|
||||||
|
: task.status === 'pending'
|
||||||
|
? '排队中...'
|
||||||
|
: task.status,
|
||||||
|
})
|
||||||
|
|
||||||
|
if (task.status === 'completed') {
|
||||||
|
stopPolling()
|
||||||
|
set({ progress: 90, statusText: '获取结果...' })
|
||||||
|
|
||||||
|
const assets = await getAssets(taskId)
|
||||||
|
set({
|
||||||
status: 'completed',
|
status: 'completed',
|
||||||
progress: 100,
|
progress: 100,
|
||||||
|
statusText: '生成完成',
|
||||||
assets,
|
assets,
|
||||||
})
|
})
|
||||||
|
} else if (task.status === 'failed') {
|
||||||
|
stopPolling()
|
||||||
|
set({
|
||||||
|
status: 'failed',
|
||||||
|
error: task.error || '生成失败',
|
||||||
|
statusText: '生成失败',
|
||||||
|
})
|
||||||
|
}
|
||||||
|
} catch {
|
||||||
|
// 网络错误不中断轮询
|
||||||
|
}
|
||||||
|
}, 2000)
|
||||||
} catch (err) {
|
} catch (err) {
|
||||||
const errMsg = (err as Error).message
|
stopPolling()
|
||||||
set({ status: 'failed', error: errMsg })
|
set({ status: 'failed', error: (err as Error).message, statusText: '提交失败' })
|
||||||
}
|
}
|
||||||
},
|
},
|
||||||
|
|
||||||
reset: () => {
|
reset: () => {
|
||||||
|
stopPolling()
|
||||||
set({
|
set({
|
||||||
taskId: null,
|
taskId: null,
|
||||||
projectId: null,
|
projectId: null,
|
||||||
progress: 0,
|
progress: 0,
|
||||||
status: 'idle',
|
status: 'idle',
|
||||||
|
statusText: '',
|
||||||
assets: [],
|
assets: [],
|
||||||
error: null,
|
error: null,
|
||||||
})
|
})
|
||||||
},
|
},
|
||||||
}))
|
}))
|
||||||
|
|
||||||
function mapAssets(resp: GenerateResponse): Asset[] {
|
|
||||||
return resp.assets.map((a, i) => ({
|
|
||||||
id: `asset-${i}`,
|
|
||||||
url: a.url,
|
|
||||||
format: a.format,
|
|
||||||
width: resp.metadata.frameWidth,
|
|
||||||
height: resp.metadata.frameHeight,
|
|
||||||
metadata: {
|
|
||||||
frameWidth: resp.metadata.frameWidth,
|
|
||||||
frameHeight: resp.metadata.frameHeight,
|
|
||||||
frameCount: resp.metadata.frameCount,
|
|
||||||
directions: resp.metadata.directions,
|
|
||||||
},
|
|
||||||
}))
|
|
||||||
}
|
|
||||||
|
|||||||
Reference in New Issue
Block a user