diff --git a/core/backend/apps/ai/migrations/0006_seed_tokenssr_models.py b/core/backend/apps/ai/migrations/0006_seed_tokenssr_models.py index 13ed19e..732e7cf 100644 --- a/core/backend/apps/ai/migrations/0006_seed_tokenssr_models.py +++ b/core/backend/apps/ai/migrations/0006_seed_tokenssr_models.py @@ -36,6 +36,11 @@ def seed(apps, schema_editor): if not tokenssr.base_url: tokenssr.base_url = TOKENSSR_BASE_URL tokenssr.save(update_fields=["base_url"]) + # down→up 循环后 get_or_create 命中已存在(可能 disabled)的 provider,不会用 defaults 覆盖 status, + # 必须显式确保 active,否则 get_default_model(要求 provider__status=active)选不到 tokenssr 模型。 + if tokenssr.status != "active": + tokenssr.status = "active" + tokenssr.save(update_fields=["status"]) for name, display, meta in TEXT_MODELS: ModelConfig.objects.update_or_create( @@ -64,8 +69,10 @@ def seed(apps, schema_editor): }, ) - # 图像主力切到 tokenssr:gpt-image-2;停用只能纯文生图的 yunqi:gpt-image-2(参考图分镜要靠 tokenssr) - ModelConfig.objects.filter(provider__name="yunqi", capability="image").update(status="disabled") + # 图像主力切到 tokenssr:gpt-image-2(唯一 active 图像模型)。 + # 必须停用**所有**非 tokenssr 图像模型 —— 否则 catalog bootstrap 种的火山 Seedream(active 且 created_at 更早) + # 会被 get_default_model(order_by created_at)选中,参考图分镜/模特库(需 image_edit)全部被架空。 + ModelConfig.objects.filter(capability="image").exclude(provider__name="tokenssr").update(status="disabled") def unseed(apps, schema_editor): @@ -73,7 +80,8 @@ def unseed(apps, schema_editor): ModelProvider = apps.get_model("ai", "ModelProvider") ModelConfig = apps.get_model("ai", "ModelConfig") ModelConfig.objects.filter(provider__name="tokenssr").update(status="disabled") - ModelConfig.objects.filter(provider__name="yunqi", capability="image").update(status="active") + # 恢复被本迁移停用的非 tokenssr 图像模型(火山 Seedream / yunqi) + ModelConfig.objects.filter(capability="image").exclude(provider__name="tokenssr").update(status="active") ModelProvider.objects.filter(name="tokenssr").update(status="disabled") diff --git a/core/backend/apps/ai/script_agent.py b/core/backend/apps/ai/script_agent.py index aa38a81..3b8a69f 100644 --- a/core/backend/apps/ai/script_agent.py +++ b/core/backend/apps/ai/script_agent.py @@ -137,15 +137,59 @@ def build_agent_messages( # --------------------------------------------------------------------------- # # JSON 抽取 + 契约规范化(模型无关,后端兜底) # --------------------------------------------------------------------------- # -def _extract_json(text: str) -> str | None: - fenced = re.search(r"```(?:json)?\s*(.+?)```", text, re.DOTALL) - candidate = fenced.group(1) if fenced else text - start, end = candidate.find("{"), candidate.rfind("}") - if start != -1 and end != -1 and end > start: - return candidate[start : end + 1] +def _balanced_object(text: str) -> str | None: + """从首个 '{' 起按括号深度扫描,返回第一个配平的 {...}(忽略字符串内的括号)。 + 避免 rfind('}') 在 JSON 后还有含花括号的散文时越界截出非法片段。""" + start = text.find("{") + if start == -1: + return None + depth = 0 + in_str = False + esc = False + for i in range(start, len(text)): + c = text[i] + if in_str: + if esc: + esc = False + elif c == "\\": + esc = True + elif c == '"': + in_str = False + continue + if c == '"': + in_str = True + elif c == "{": + depth += 1 + elif c == "}": + depth -= 1 + if depth == 0: + return text[start : i + 1] return None +def _looks_like_draft(blob: str) -> bool: + try: + d = json.loads(blob) + except (ValueError, TypeError): + return False + return isinstance(d, dict) and ("segments" in d or "hook" in d) + + +def _extract_json(text: str) -> str | None: + """抽取 ScriptDraft JSON。容错:模型可能先给示例 ```json 块再给正式块, + 故取**最后一个**含 segments/hook 的合法围栏块;都不像草稿再退而取末个配平对象;无围栏再裸扫。""" + fences = re.findall(r"```(?:json)?\s*(.+?)```", text, re.DOTALL) + for block in reversed(fences): + obj = _balanced_object(block) + if obj and _looks_like_draft(obj): + return obj + for block in reversed(fences): + obj = _balanced_object(block) + if obj: + return obj + return _balanced_object(text) + + def _nearest_duration(value) -> int: try: value = int(value) @@ -381,69 +425,81 @@ def stream_script_agent( yield _sse({"type": "error", "detail": f"任务创建失败(可能额度不足):{exc}"}) return reservation = task.credit_reservation - - yield _sse({"type": "tool", "id": "generate", "label": "按黄金结构生成分镜", "status": "running"}) - full: list[str] = [] - shown = 0 - forwarding = True + # 额度是否已结算(charge 成功 / release 失败)。客户端中途断连时,生成器被 .close() 抛 + # GeneratorExit —— 它是 BaseException 不是 Exception,普通 except 抓不到,会让预扣额度冻结。 + # 故用 try/finally 兜底:任何未结算路径(含断连)都释放预扣。 + settled = False try: - task.status = AITask.Status.SUBMITTED - task.submitted_at = timezone.now() - task.save(update_fields=["status", "submitted_at", "updated_at"]) - provider = build_provider(model_config) - for ev in provider.chat_completion_stream( - model=model_config.name, - endpoint=model_config.endpoint, - messages=messages, - temperature=0.85, - ): - if ev.get("type") == "delta": - full.append(ev["text"]) - if forwarding: - text = "".join(full) - cut = _visible_cut(text) - if cut < len(text): - forwarding = False - visible = text[:cut] - if len(visible) > shown: - piece = visible[shown:] - shown = len(visible) - if piece.strip(): - yield _sse({"type": "delta", "text": piece}) - elif ev.get("type") == "done": - break - raw = "".join(full) - draft = normalize_draft(raw, aspect_ratio=aspect_ratio, total_duration=total_duration) - except Exception as exc: # noqa: BLE001 - _fail_task(task, reservation, str(exc)) - yield _sse({"type": "tool", "id": "generate", "status": "error"}) - yield _sse({"type": "error", "detail": f"脚本生成失败:{exc}"}) - return + yield _sse({"type": "tool", "id": "generate", "label": "按黄金结构生成分镜", "status": "running"}) + full: list[str] = [] + shown = 0 + forwarding = True + try: + task.status = AITask.Status.SUBMITTED + task.submitted_at = timezone.now() + task.save(update_fields=["status", "submitted_at", "updated_at"]) + provider = build_provider(model_config) + for ev in provider.chat_completion_stream( + model=model_config.name, + endpoint=model_config.endpoint, + messages=messages, + temperature=0.85, + ): + if ev.get("type") == "delta": + full.append(ev["text"]) + if forwarding: + text = "".join(full) + cut = _visible_cut(text) + if cut < len(text): + forwarding = False + visible = text[:cut] + if len(visible) > shown: + piece = visible[shown:] + shown = len(visible) + if piece.strip(): + yield _sse({"type": "delta", "text": piece}) + elif ev.get("type") == "done": + break + raw = "".join(full) + draft = normalize_draft(raw, aspect_ratio=aspect_ratio, total_duration=total_duration) + except Exception as exc: # noqa: BLE001 + _fail_task(task, reservation, str(exc)) + settled = True + yield _sse({"type": "tool", "id": "generate", "status": "error"}) + yield _sse({"type": "error", "detail": f"脚本生成失败:{exc}"}) + return - yield _sse({"type": "tool", "id": "generate", "status": "done"}) - yield _sse( - { - "type": "tool", - "id": "extract", - "label": f"提取实体 {len(draft['entities'])} 个 · {len(draft['segments'])} 镜", - "status": "done", - } - ) - yield _sse({"type": "tool", "id": "check", "label": "自检:镜数 / ≤55字 / 违规词", "status": "done"}) - yield _sse({"type": "draft", "draft": draft}) + yield _sse({"type": "tool", "id": "generate", "status": "done"}) + yield _sse( + { + "type": "tool", + "id": "extract", + "label": f"提取实体 {len(draft['entities'])} 个 · {len(draft['segments'])} 镜", + "status": "done", + } + ) + yield _sse({"type": "tool", "id": "check", "label": "自检:镜数 / ≤55字 / 违规词", "status": "done"}) + yield _sse({"type": "draft", "draft": draft}) - try: from django.db import transaction - with transaction.atomic(): - task.status = AITask.Status.SUCCEEDED - task.response_payload = {"raw": raw[:8000]} - task.actual_cost = task.estimated_cost - task.completed_at = timezone.now() - task.save(update_fields=["status", "response_payload", "actual_cost", "completed_at", "updated_at"]) - charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost) - source = "revise" if mode == "revise" else ("theme" if mode == "theme" else "ai") - script = persist_script_draft(project=project, user=user, task=task, draft=draft, source=source) + try: + with transaction.atomic(): + task.status = AITask.Status.SUCCEEDED + task.response_payload = {"raw": raw[:8000]} + task.actual_cost = task.estimated_cost + task.completed_at = timezone.now() + task.save(update_fields=["status", "response_payload", "actual_cost", "completed_at", "updated_at"]) + charge_reserved_credit(reservation=reservation, actual_amount=task.actual_cost) + source = "revise" if mode == "revise" else ("theme" if mode == "theme" else "ai") + script = persist_script_draft(project=project, user=user, task=task, draft=draft, source=source) + settled = True # charge 已提交 + except Exception as exc: # noqa: BLE001 — 落库失败:atomic 已回滚 charge,补释放预留 + _fail_task(task, reservation, f"保存脚本失败:{exc}") + settled = True + yield _sse({"type": "error", "detail": f"保存脚本失败:{exc}"}) + return + from apps.projects.serializers import ScriptVersionSerializer yield _sse( @@ -453,12 +509,11 @@ def stream_script_agent( "version": ScriptVersionSerializer(script).data, } ) - except Exception as exc: # noqa: BLE001 — 落库失败:回滚已撤销扣费,补释放预留 - _fail_task(task, reservation, f"保存脚本失败:{exc}") - yield _sse({"type": "error", "detail": f"保存脚本失败:{exc}"}) - return - - yield _sse({"type": "done"}) + yield _sse({"type": "done"}) + finally: + # 断连(GeneratorExit)或任何 settled=False 的退出路径:释放预扣,避免额度冻结 + if not settled: + _fail_task(task, reservation, "stream aborted (client disconnected)") def _fail_task(task, reservation, message: str) -> None: diff --git a/core/backend/apps/ai/services.py b/core/backend/apps/ai/services.py index 01b6c09..7ee2354 100644 --- a/core/backend/apps/ai/services.py +++ b/core/backend/apps/ai/services.py @@ -60,11 +60,12 @@ def resolve_provider_credentials(provider) -> tuple[str | None, str | None]: def build_provider(model_config: ModelConfig): - """按 provider.name 分流:火山官方直连 → VolcanoArkProvider;其余 → 通用 OpenAICompatibleProvider。""" + """按 provider.name 分流:火山官方直连 → VolcanoArkProvider;其余 → 通用 OpenAICompatibleProvider。 + 两条路都走 resolve_provider_credentials,统一 DB→.env 优先级(官方直连的 None 再由 __post_init__ 回退 settings.VOLCANO)。""" provider = model_config.provider - if provider.name in OFFICIAL_DIRECT_PROVIDERS: - return VolcanoArkProvider(base_url=provider.base_url or None) base_url, api_key = resolve_provider_credentials(provider) + if provider.name in OFFICIAL_DIRECT_PROVIDERS: + return VolcanoArkProvider(base_url=base_url, api_key=api_key) return OpenAICompatibleProvider(base_url=base_url, api_key=api_key) @@ -768,10 +769,9 @@ def _storyboard_frame_worker(task_id, version_id, segment_id, user_id) -> None: refs = _storyboard_reference_images(project, segment) ref_urls = [r["url"] for r in refs] if ref_urls and hasattr(provider, "image_edit"): - # gpt-image-2 多图参考:把角色/场景/商品合成进本镜(@图1@图2@图3),锁脸锁外观保一致 - frame_prompt = task.request_payload.get("prompt") or build_storyboard_frame_prompt_refs( - project, version, segment, refs - ) + # gpt-image-2 多图参考:必须用 refs 版提示词(点名「参考图N=角色/场景/商品」+锁脸锁商品), + # 不能复用 request_payload['prompt'](那是建任务时写死的基础提示词,恒为真值会架空一致性约束)。 + frame_prompt = build_storyboard_frame_prompt_refs(project, version, segment, refs) response = provider.image_edit( model=model_config.name, prompt=frame_prompt, diff --git a/core/frontend/src/routes/pipeline.tsx b/core/frontend/src/routes/pipeline.tsx index 3a5df4d..e168ed6 100644 --- a/core/frontend/src/routes/pipeline.tsx +++ b/core/frontend/src/routes/pipeline.tsx @@ -655,6 +655,7 @@ export function PipelinePage(props: { const agentMode = mode ?? mapSourceToMode(source ?? chatMode); const baseVersionId = agentMode === "revise" ? currentScript?.id : undefined; let ok = false; + let receivedEvent = false; try { await api.agentScriptStream( project.id, @@ -667,6 +668,7 @@ export function PipelinePage(props: { total_duration: 60 }, (evt) => { + receivedEvent = true; if (evt.type === "tool") { // 工具卡:running 时把 label 滚进进度流(done/error 暂只用于结束态) if (evt.status === "running" && typeof evt.label === "string") { @@ -685,9 +687,16 @@ export function PipelinePage(props: { } ); } catch { - // 流式不可用(网关不支持 SSE 等)→ 退回旧同步端点,保证可用性 - const res = await onGenerateScript(prompt, source ?? chatMode).catch(() => null); - ok = !!res; + if (ok) { + // 已收到 saved:流其实成功了(只是收尾抖动),绝不回退重发,避免重复生成+重复扣费 + } else if (!receivedEvent) { + // 一个事件都没收到 = 流式根本没跑起来(网关不支持 SSE)→ 退回旧同步端点 + const res = await onGenerateScript(prompt, source ?? chatMode).catch(() => null); + ok = !!res; + } else { + // 流中途断了(后端已通过 finally 释放预扣额度):不重试,提示用户 + pushMsg("ai", "生成中断了,请重试。"); + } } setChatMsgs((list) => list.map((m) => (m.id === progressId ? { ...m, done: true } : m))); if (ok) {