feat(generate): confirm product status scope
This commit is contained in:
+36
-1
@@ -678,6 +678,10 @@ class GenerateWorker(BaseWorker):
|
||||
db_path=None,
|
||||
config=None,
|
||||
diagnostic_log_dir=None,
|
||||
generation_scope="all",
|
||||
product_status_counts=None,
|
||||
status_scope_excluded=0,
|
||||
generation_plan_fingerprint=None,
|
||||
):
|
||||
super().__init__()
|
||||
self.tasks = list(tasks)
|
||||
@@ -685,6 +689,15 @@ class GenerateWorker(BaseWorker):
|
||||
self.db_path = db_path
|
||||
self.config = config
|
||||
self.diagnostic_log_dir = diagnostic_log_dir
|
||||
self.generation_scope = product_status.normalize_scope(generation_scope)
|
||||
self.product_status_counts = {
|
||||
status: int((product_status_counts or {}).get(status, 0) or 0)
|
||||
for status in product_status.VALID_PRODUCT_STATUSES
|
||||
}
|
||||
self.status_scope_excluded = max(0, int(status_scope_excluded or 0))
|
||||
self.generation_plan_fingerprint = (
|
||||
str(generation_plan_fingerprint or "") or None
|
||||
)
|
||||
self._run_id = None
|
||||
self._account_by_alias = {}
|
||||
self._task_positions = {}
|
||||
@@ -713,6 +726,12 @@ class GenerateWorker(BaseWorker):
|
||||
task for task in self.tasks
|
||||
if ai.is_generatable_task(task, generate_mode=generate_mode)
|
||||
]
|
||||
if self.generation_scope == product_status.SCOPE_NORMAL_ONLY:
|
||||
eligible = [
|
||||
task
|
||||
for task in eligible
|
||||
if product_status.is_normal(getattr(task, "product_status", None))
|
||||
]
|
||||
component_totals = ai.generation_component_totals(
|
||||
eligible,
|
||||
generate_mode=generate_mode,
|
||||
@@ -757,10 +776,18 @@ class GenerateWorker(BaseWorker):
|
||||
title_concurrency=ai_cfg.get("title_concurrency", 1),
|
||||
started_at=self._run_started_at_text,
|
||||
)
|
||||
start_message += ";生成范围:{scope};按范围排除{excluded}条".format(
|
||||
scope=(
|
||||
"仅状态正常"
|
||||
if self.generation_scope == product_status.SCOPE_NORMAL_ONLY
|
||||
else "所有状态"
|
||||
),
|
||||
excluded=self.status_scope_excluded,
|
||||
)
|
||||
self._log_run_event(start_message)
|
||||
try:
|
||||
summary = ai.generate_batch(
|
||||
self.tasks,
|
||||
eligible,
|
||||
self.prompt_values,
|
||||
ai_cfg={
|
||||
"config": self.config,
|
||||
@@ -810,6 +837,10 @@ class GenerateWorker(BaseWorker):
|
||||
summary["billing_error"] = dict(self._billing_error)
|
||||
summary["ok"] = False
|
||||
summary["cancelled"] = True
|
||||
summary["generation_scope"] = self.generation_scope
|
||||
summary["product_status_counts"] = dict(self.product_status_counts)
|
||||
summary["status_scope_excluded"] = self.status_scope_excluded
|
||||
summary["generation_plan_fingerprint"] = self.generation_plan_fingerprint
|
||||
summary["run_id"] = self._run_id
|
||||
summary["batch_ids"] = batch_ids
|
||||
status = "failed" if summary.get("billing_error") or summary.get("error") else ("cancelled" if summary.get("cancelled") else "done")
|
||||
@@ -1131,6 +1162,10 @@ class GenerateWorker(BaseWorker):
|
||||
"image_concurrency": ai_cfg.get("image_concurrency"),
|
||||
"generate_cover": ai_cfg.get("generate_cover", False),
|
||||
"backend": ai_cfg.get("backend", "direct"),
|
||||
"generation_scope": self.generation_scope,
|
||||
"product_status_counts": dict(self.product_status_counts),
|
||||
"status_scope_excluded": self.status_scope_excluded,
|
||||
"generation_plan_fingerprint": self.generation_plan_fingerprint,
|
||||
},
|
||||
path=self.db_path,
|
||||
)
|
||||
|
||||
Reference in New Issue
Block a user