feat: snapshot product suite reference assets
This commit is contained in:
+51
-3
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user