feat(product-suite): persist current generation rounds

This commit is contained in:
chengma
2026-07-16 23:27:25 +08:00
parent c146c0b41d
commit f3defdeb95
11 changed files with 806 additions and 16 deletions
+79
View File
@@ -696,10 +696,12 @@ class SuiteTaskState:
last_saved_prompt: str = ""
settings: dict = field(default_factory=product_suite.default_suite_settings)
current_job_ids: list = field(default_factory=list)
current_generation_round_key: str = ""
show_history: bool = False
generation_job_ids: list = field(default_factory=list)
generation_mode: str = "batch"
generation_retry_job_id: int = None
generation_round_key: str = ""
worker: object = None
thread: object = None
generation_run_token: str = ""
@@ -1457,6 +1459,7 @@ class ProductSuiteTab(QWidget):
def _load_state(self, state):
self._sync_state_project_binding(state)
self._restore_current_generation_results(state)
self._loading = True
try:
self.custom_category_edit.hide()
@@ -1723,6 +1726,8 @@ class ProductSuiteTab(QWidget):
return None
state.project_id = int(project.id)
state.project_binding_state = project.binding_state
if previous_id != state.project_id:
self._restore_current_generation_results(state, force=True)
stored_prompt = str(project.draft_prompt or "")
if load_existing and previous_id != state.project_id:
state.prompt = stored_prompt
@@ -3101,6 +3106,25 @@ class ProductSuiteTab(QWidget):
if not specs:
self._message("生成数量为0", "请至少把一个套图分类的数量设为1。")
return False
generation_round_key = ""
if retrying:
try:
original_job = image_studio.get_job(retry_job_id, path=self.db_path)
except Exception as exc:
self._message("读取重试图片失败", _user_error(exc))
return False
if original_job is None or int(original_job.project_id) != int(state.project_id):
self._message("重试图片无效", "该图片不属于当前商品套图任务。")
return False
for spec in specs:
spec["generation_round_key"] = original_job.generation_round_key
spec["generation_slot_index"] = original_job.generation_slot_index
generation_round_key = str(original_job.generation_round_key or "")
else:
generation_round_key = uuid.uuid4().hex
for slot_index, spec in enumerate(specs):
spec["generation_round_key"] = generation_round_key
spec["generation_slot_index"] = slot_index
if confirm_batch and not self._confirm(
"确认生成商品套图",
self._generation_confirmation_message(
@@ -3118,6 +3142,7 @@ class ProductSuiteTab(QWidget):
state.project_id,
specs,
run_token=run_token,
generation_round_key=generation_round_key or None,
aspect_ratio=state.settings["ratio"],
db_path=self.db_path,
config=self.config,
@@ -3130,6 +3155,7 @@ class ProductSuiteTab(QWidget):
state.generation_job_ids = []
state.generation_mode = "retry" if retrying else "batch"
state.generation_retry_job_id = retry_job_id
state.generation_round_key = generation_round_key
state.done = 0
state.failed = 0
state.total = len(specs)
@@ -3406,6 +3432,30 @@ class ProductSuiteTab(QWidget):
current.append(job_id)
state.current_job_ids = current
def _restore_current_generation_results(self, state, *, force=False, allow_running=False):
if state.project_id is None or (
state.generation_running() and not allow_running
):
return False
if state.current_job_ids and not force:
return False
try:
round_key = image_studio.get_current_generation_round(
state.project_id,
path=self.db_path,
)
jobs = image_studio.list_generation_round_current_jobs(
state.project_id,
round_key,
path=self.db_path,
)
except Exception as exc:
self._status("当前生成结果恢复失败:%s" % _user_error(exc), "danger")
return False
state.current_generation_round_key = str(round_key or "")
state.current_job_ids = [int(job.id) for job in jobs]
return True
def _generation_job_snapshot(self, state):
job_ids = self._generation_job_ids(state)
counts = {
@@ -3501,12 +3551,40 @@ class ProductSuiteTab(QWidget):
if state.started_at
else 0
)
promoted = False
if not retrying and state.generation_round_key and success:
try:
promoted = image_studio.promote_generation_round_if_success(
state.project_id,
state.generation_round_key,
path=self.db_path,
)
except Exception as exc:
self._status("当前生成轮次保存失败:%s" % _user_error(exc), "danger")
self._restore_current_generation_results(
state,
force=True,
allow_running=True,
)
elif not retrying:
self._restore_current_generation_results(
state,
force=True,
allow_running=True,
)
elif state.generation_round_key:
self._restore_current_generation_results(
state,
force=True,
allow_running=True,
)
self._generation_run_states.pop(run_token, None)
state.generation_run_token = ""
state.generation_stop_requested = False
state.generation_terminal_streak = 0
state.generation_job_ids = []
state.generation_retry_job_id = None
state.generation_round_key = ""
state.worker = None
state.thread = None
state.done = success + failed + cancelled
@@ -3530,6 +3608,7 @@ class ProductSuiteTab(QWidget):
"active": active,
"elapsed_seconds": elapsed,
"mode": "retry" if retrying else "batch",
"current_round_promoted": promoted,
},
level="WARNING" if active or result.get("ok") is False else "INFO",
)
+5
View File
@@ -334,6 +334,7 @@ class ProductSuiteGenerateWorker(BaseWorker):
job_specs,
*,
run_token="",
generation_round_key=None,
aspect_ratio="1:1",
db_path=None,
config=None,
@@ -343,6 +344,7 @@ class ProductSuiteGenerateWorker(BaseWorker):
self.project_id = int(project_id)
self.job_specs = [dict(spec) for spec in (job_specs or [])]
self.run_token = str(run_token or "")
self.generation_round_key = str(generation_round_key or "").strip()
self.aspect_ratio = str(aspect_ratio or "1:1")
self.db_path = db_path
self.config = config
@@ -366,6 +368,8 @@ class ProductSuiteGenerateWorker(BaseWorker):
prompt=spec.get("prompt") or "",
generation_source="cmhub",
provider="cmhub",
generation_round_key=spec.get("generation_round_key") or self.generation_round_key or None,
generation_slot_index=spec.get("generation_slot_index"),
path=self.db_path,
)
)
@@ -410,6 +414,7 @@ class ProductSuiteGenerateWorker(BaseWorker):
summary["job_ids"] = list(self.job_ids)
summary["cancelled_count"] = int(summary.get("cancelled", 0) or 0)
summary["run_token"] = self.run_token
summary["generation_round_key"] = self.generation_round_key or None
return summary