feat: snapshot product suite reference assets

This commit is contained in:
chengma
2026-07-17 17:03:40 +08:00
parent d11dc752b5
commit 1fabd8afcf
12 changed files with 252 additions and 5 deletions
+51 -3
View File
@@ -78,6 +78,7 @@ class ImageStudioJob:
id: int
project_id: int
source_asset_id: Optional[int]
reference_asset_ids: Optional[str]
output_asset_id: Optional[int]
generation_source: str
provider: str
@@ -344,6 +345,41 @@ def _ensure_assets_belong_to_project(database, project_id, asset_ids):
raise db.DbError("AI工场资产不属于当前项目")
def _parse_reference_asset_ids(value, source_asset_id=None):
if value is None or str(value).strip() == "":
return []
if isinstance(value, str):
try:
value = json.loads(value)
except (TypeError, ValueError, json.JSONDecodeError) as exc:
raise db.DbError("AI工场参考图快照格式无效") from exc
if not isinstance(value, (list, tuple)):
raise db.DbError("AI工场参考图快照必须是图片ID列表")
source_id = int(source_asset_id) if source_asset_id is not None else None
normalized = []
for asset_id in value:
try:
parsed = int(asset_id)
except (TypeError, ValueError) as exc:
raise db.DbError("AI工场参考图快照包含无效图片ID") from exc
if parsed <= 0:
raise db.DbError("AI工场参考图快照包含无效图片ID")
if source_id is not None and parsed == source_id:
raise db.DbError("AI工场参考图不能包含主图")
if parsed in normalized:
raise db.DbError("AI工场参考图不能重复")
normalized.append(parsed)
return normalized
def job_reference_asset_ids(job):
"""Return a validated, ordered reference asset snapshot for one job."""
return _parse_reference_asset_ids(
getattr(job, "reference_asset_ids", None),
getattr(job, "source_asset_id", None),
)
def get_project(project_id, path=None, conn=None, include_deleted=False):
clauses = ["id = ?"]
params = [int(project_id)]
@@ -1114,6 +1150,7 @@ def create_job(
project_id,
*,
source_asset_id=None,
reference_asset_ids=None,
job_type="main",
prompt="",
task_key=None,
@@ -1126,6 +1163,12 @@ def create_job(
):
now = _now()
task_key = str(task_key or _task_key(project_id))
reference_ids = _parse_reference_asset_ids(reference_asset_ids, source_asset_id)
reference_json = None if reference_asset_ids is None else json.dumps(
reference_ids,
ensure_ascii=True,
separators=(",", ":"),
)
if generation_round_key is not None:
generation_round_key = str(generation_round_key).strip()
if not generation_round_key:
@@ -1140,18 +1183,23 @@ def create_job(
with _connection(conn, path) as database:
try:
with database:
_ensure_assets_belong_to_project(database, project_id, [source_asset_id])
_ensure_assets_belong_to_project(
database,
project_id,
[source_asset_id, *reference_ids],
)
cursor = database.execute(
"""
INSERT INTO image_studio_jobs
(project_id, source_asset_id, generation_source, provider,
(project_id, source_asset_id, reference_asset_ids, generation_source, provider,
job_type, task_key, status, prompt, generation_round_key,
generation_slot_index, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, 'pending', ?, ?, ?, ?, ?)
VALUES (?, ?, ?, ?, ?, ?, ?, 'pending', ?, ?, ?, ?, ?)
""",
(
int(project_id),
source_asset_id,
reference_json,
str(generation_source or "cmhub"),
str(provider or "cmhub"),
str(job_type),