feat(product-suite): persist current generation rounds
This commit is contained in:
@@ -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",
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user