diff --git a/__init__.py b/__init__.py index 25d922d..947d4c2 100644 --- a/__init__.py +++ b/__init__.py @@ -20,11 +20,12 @@ _seeder_filter = lambda record: not any( ) logging.getLogger().addFilter(_seeder_filter) -from .nodes import NanoBananaPro, BatchNanoBananaPro, GoogleGemini, LoadFile, ImageStitchPro, BatchCleanMetadata, VideoPreview, GoogleVeo, FluxImageEdit, UniversalLLMChat, KlingVideo, KlingFirstLastFrame, KlingMotionControlTest, AspectRatioPreset, BatchImagesO1key, Seedance, SeedanceMultiModal, StreamPreview, DoubaoImage, O1keyGPTImage, O1keyGrokImage, KVideoFirstLast, KVideoImage2Video +from .nodes import NanoBananaPro, BatchNanoBananaPro, GoogleGemini, LoadFile, ImageStitchPro, BatchCleanMetadata, VideoPreview, GoogleVeo, FluxImageEdit, UniversalLLMChat, KlingVideo, KlingFirstLastFrame, KlingMotionControlTest, AspectRatioPreset, BatchImagesO1key, Seedance, SeedanceMultiModal, StreamPreview, DoubaoImage, O1keyGPTImage, O1keyGPTImageBatch, O1keyGrokImage, KVideoFirstLast, KVideoImage2Video from .nodes import K3Video, K3VideoFirstLast, K3MotionControl, K3MotionVideoCheck, NanoBananaV2, NanoBananaV2Batch, SaveImageFormat from .nodes import O1keySavePSD from .nodes import O1keyRemoveBackground from .nodes import O1keyColorRemoveBG +from .nodes import O1keyGridSplitter # 报错弹框友好文案(不修改原节点代码,仅在外层统一处理) _MSG_TIMEOUT = "API 请求超时,请稍后重试或检查网络。" @@ -85,6 +86,7 @@ NODE_CLASS_MAPPINGS = { "StreamPreview": StreamPreview, "DoubaoImage": DoubaoImage, "O1keyGPTImage": O1keyGPTImage, + "O1keyGPTImageBatch": O1keyGPTImageBatch, "O1keyGrokImage": O1keyGrokImage, "KVideoFirstLast": KVideoFirstLast, "KVideoImage2Video": KVideoImage2Video, @@ -98,6 +100,7 @@ NODE_CLASS_MAPPINGS = { "O1keySavePSD": O1keySavePSD, "O1keyRemoveBackground": O1keyRemoveBackground, "O1keyColorRemoveBG": O1keyColorRemoveBG, + "O1keyGridSplitter": O1keyGridSplitter, } NODE_DISPLAY_NAME_MAPPINGS = { @@ -123,6 +126,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "StreamPreview": "流式文本预览", "DoubaoImage": "豆包生图", "O1keyGPTImage": "o1key GPT Image", + "O1keyGPTImageBatch": "o1key GPT Image(批量)", "O1keyGrokImage": "Grok Image", "KVideoFirstLast": "K26 图生视频(首尾帧)", "KVideoImage2Video": "K26 图生视频", @@ -136,6 +140,7 @@ NODE_DISPLAY_NAME_MAPPINGS = { "O1keySavePSD": "保存 PSD(分层)", "O1keyRemoveBackground": "去背景(rembg)", "O1keyColorRemoveBG": "颜色去背景", + "O1keyGridSplitter": "合并图智能切割", } WEB_DIRECTORY = "./web" @@ -149,6 +154,35 @@ try: import folder_paths from .utils.config import CONFIG_FILE, load_config, NETWORK_ROUTES + def _get_o1key_server_port(): + try: + import comfy.cli_args as _cli_args + args = getattr(_cli_args, "args", None) + port = getattr(args, "port", None) if args else None + port = port or getattr(_cli_args, "server_port", None) or getattr(_cli_args, "port", None) + if port is not None: + return str(int(port)) + except Exception: + pass + try: + import sys as _sys + for idx, arg in enumerate(_sys.argv): + if arg in ("--port", "--listen-port") and idx + 1 < len(_sys.argv): + return str(int(_sys.argv[idx + 1])) + for prefix in ("--port=", "--listen-port="): + if arg.startswith(prefix): + return str(int(arg.split("=", 1)[1])) + except Exception: + pass + return "8188" + + def _get_o1key_history_meta_file(output_dir): + import os as _os_history + return _os_history.path.join( + output_dir, + f".o1key_history_{_get_o1key_server_port()}.json", + ) + @PromptServer.instance.routes.get("/o1key/input_dir") async def get_input_dir(request): import os @@ -220,11 +254,11 @@ try: @PromptServer.instance.routes.get("/o1key/output_history") async def get_output_history(request): """读取 output 目录文件,按执行分组返回 /api/jobs 兼容格式""" - import os, uuid, json as _json + import os, json as _json limit = int(request.query.get("limit", "200")) offset = int(request.query.get("offset", "0")) output_dir = os.path.abspath(folder_paths.get_output_directory()) - meta_file = os.path.join(output_dir, ".o1key_history.json") + meta_file = _get_o1key_history_meta_file(output_dir) meta = {} if os.path.isfile(meta_file): try: @@ -235,7 +269,7 @@ try: supported_ext = {'.png', '.jpg', '.jpeg', '.webp', '.gif', '.mp4', '.webm'} # 收集所有文件并按 workflow_id 分组 all_files = [] - for fname in os.listdir(output_dir): + for fname in meta.keys(): ext = os.path.splitext(fname)[1].lower() if ext not in supported_ext: continue @@ -247,14 +281,11 @@ try: all_files.append({"name": fname, "mtime": mtime, "media": media}) # 按 workflow_id 分组(同一次执行合并为一个 job) groups = {} - ungrouped = [] for f in all_files: m = meta.get(f["name"], {}) wid = m.get("workflow_id") if wid: groups.setdefault(wid, []).append((f, m)) - else: - ungrouped.append((f, {})) # 构建 job 列表 jobs = [] for wid, items in groups.items(): @@ -280,27 +311,6 @@ try: "execution_error": None, "workflow_id": wid, }) - # 无元数据的文件各自作为独立 job - for f, m in ungrouped: - mtime_ms = int(f["mtime"] * 1000) - job_id = str(uuid.uuid5(uuid.NAMESPACE_URL, f["name"])) - jobs.append({ - "id": job_id, - "status": "completed", - "create_time": mtime_ms, - "execution_start_time": mtime_ms, - "execution_end_time": mtime_ms, - "preview_output": { - "filename": f["name"], - "subfolder": "", - "type": "output", - "nodeId": "0", - "mediaType": f["media"], - }, - "outputs_count": 1, - "execution_error": None, - "workflow_id": None, - }) # 按时间倒序排列,分页 jobs.sort(key=lambda x: x["create_time"], reverse=True) total = len(jobs) @@ -349,11 +359,11 @@ try: @PromptServer.instance.routes.get("/o1key/job_detail/{job_id}") async def get_job_detail(request): - """根据 job_id (workflow_id 或 uuid5) 返回含工作流的 job 详情""" - import os, uuid, struct, json as _json + """根据 job_id 返回当前端口持久化历史中的 job 详情""" + import os, struct, json as _json job_id = request.match_info["job_id"] output_dir = os.path.abspath(folder_paths.get_output_directory()) - meta_file = os.path.join(output_dir, ".o1key_history.json") + meta_file = _get_o1key_history_meta_file(output_dir) meta = {} if os.path.isfile(meta_file): try: @@ -362,17 +372,15 @@ try: except Exception: pass supported_ext = {'.png', '.jpg', '.jpeg', '.webp', '.gif', '.mp4', '.webm'} - # 查找属于该 job 的所有文件(按 workflow_id 或 uuid5 匹配) + # 只在当前端口的持久化记录中查找该 job 的文件 matched_files = [] - for fname in os.listdir(output_dir): + for fname in meta.keys(): ext = os.path.splitext(fname)[1].lower() if ext not in supported_ext: continue m = meta.get(fname, {}) if m.get("workflow_id") == job_id: matched_files.append(fname) - elif str(uuid.uuid5(uuid.NAMESPACE_URL, fname)) == job_id: - matched_files.append(fname) if not matched_files: return web.json_response({"error": "not found"}, status=404) # 用最新文件作为代表 @@ -444,13 +452,13 @@ try: @PromptServer.instance.routes.post("/o1key/delete_history") async def delete_history_item(request): """删除持久化历史记录及对应的输出文件""" - import os, uuid, json as _json + import os, json as _json body = await request.json() job_ids = body.get("delete", []) if not job_ids: return web.json_response({"success": False, "error": "missing ids"}, status=400) output_dir = os.path.abspath(folder_paths.get_output_directory()) - meta_file = os.path.join(output_dir, ".o1key_history.json") + meta_file = _get_o1key_history_meta_file(output_dir) meta = {} if os.path.isfile(meta_file): try: @@ -464,8 +472,6 @@ try: for fname, m in list(meta.items()): if m.get("workflow_id") == job_id: files_to_remove.append(fname) - elif str(uuid.uuid5(uuid.NAMESPACE_URL, fname)) == job_id: - files_to_remove.append(fname) for fname in files_to_remove: meta.pop(fname, None) fpath = os.path.join(output_dir, fname) @@ -574,7 +580,7 @@ try: end_time = _time.time() start_time = tracker["start"] output_dir = _os.path.abspath(folder_paths.get_output_directory()) - meta_file = _os.path.join(output_dir, ".o1key_history.json") + meta_file = _get_o1key_history_meta_file(output_dir) meta = {} if _os.path.isfile(meta_file): try: diff --git a/clients/gpt_image_client.py b/clients/gpt_image_client.py index ce9484f..ef129c2 100644 --- a/clients/gpt_image_client.py +++ b/clients/gpt_image_client.py @@ -6,9 +6,8 @@ GPT Image API 客户端 设计原则: - 与 doubao_image_client.py 保持相同的异步 + 同步双入口模式 - - 图像以 multipart/form-data 方式上传(edits 接口) - - generations 接口使用 JSON 请求体,图像以 data URI base64 内联传递 - - 响应支持 url 和 b64_json 两种格式,优先处理 b64_json(避免二次下载) + - generations / edits 接口均使用 multipart/form-data + - 响应兼容 SSE 流式、JSON、url 和 b64_json """ import asyncio @@ -28,11 +27,6 @@ from ..utils.config import get_api_key_or_raise, get_api_base_url from ..utils.image_utils import tensor_to_pil, encode_image_to_base64 from ..utils.http_error import RETRYABLE_STATUS_CODES, HTTP_ERROR_MESSAGES, _compute_delay, DEFAULT_MAX_RETRIES, DEFAULT_BASE_DELAY, DEFAULT_MAX_DELAY, DEFAULT_BACKOFF_FACTOR, get_friendly_message -# GPT Image 专属错误文案 -_GPT_ERROR_MESSAGES = { - 500: "触发内容风控,或服务器繁忙!", -} - try: from comfy.model_management import processing_interrupted, InterruptProcessingException _INTERRUPT_AVAILABLE = True @@ -43,7 +37,7 @@ except ImportError: # ── 接口端点 ────────────────────────────────────────────────────────────────── _ENDPOINT_GENERATIONS = "/v1/images/generations/" -_ENDPOINT_EDITS = "/v1/images/edits/" +_ENDPOINT_EDITS = "/v1/images/edits" # ── 模型名映射(UI 显示名 → API 实际参数名)───────────────────────────────── _MODEL_NAME_MAP = { @@ -60,11 +54,10 @@ class GptImageClient: GPT Image API 客户端 接口说明: - generations:JSON body,支持 quality / size / n / model - edits:multipart/form-data,必须包含 image(PNG),可选 mask(PNG) + generations:multipart/form-data,支持 quality / size / n / model + edits:multipart/form-data,图片和 mask 使用 PNG 文件上传 - 两个接口的响应格式相同: - { "data": [ {"url": "..."} | {"b64_json": "..."} ] } + 响应支持 JSON 和 SSE 流式格式。 """ def __init__(self): @@ -82,6 +75,26 @@ class GptImageClient: "Content-Type": "application/json", } + @staticmethod + def _new_multipart_form() -> aiohttp.FormData: + try: + return aiohttp.FormData(default_to_multipart=True) + except TypeError: + form = aiohttp.FormData() + form._is_multipart = True + return form + + @staticmethod + def _add_form_fields(form: aiohttp.FormData, fields: dict) -> None: + for key, value in fields.items(): + if value is None: + continue + if isinstance(value, bool): + value = "true" if value else "false" + elif isinstance(value, (dict, list)): + value = json.dumps(value, ensure_ascii=False) + form.add_field(key, str(value)) + # ── 图像转换工具 ────────────────────────────────────────────────────────── # ── 请求体大小限制 ──────────────────────────────────────────────────────── @@ -208,6 +221,10 @@ class GptImageClient: raise RuntimeError(f"API 返回错误: {msg}") data_list = resp_json.get("data") + if data_list is None: + data_list = resp_json.get("images") + if data_list is None and (resp_json.get("b64_json") or resp_json.get("url")): + data_list = [resp_json] if not data_list: raise RuntimeError( f"API 响应中未找到 data 字段,完整响应:\n" @@ -247,6 +264,203 @@ class GptImageClient: return images + @staticmethod + def _decode_b64_image(b64: str, label: str) -> Image.Image: + try: + img_bytes = base64.b64decode(b64) + img = Image.open(BytesIO(img_bytes)) + print(f"[o1key GPT Image] {label} base64 解码完成 ({img.size[0]}×{img.size[1]})") + return img + except Exception as e: + raise RuntimeError(f"{label} base64 解码失败: {e}") from None + + async def _append_images_from_payload( + self, + payload: dict, + session: aiohttp.ClientSession, + images: List[Image.Image], + event_name: str = "", + ) -> bool: + if isinstance(payload, dict) and "error" in payload: + err = payload["error"] + msg = ( + err.get("message") or err.get("msg") or json.dumps(err, ensure_ascii=False) + if isinstance(err, dict) + else str(err) + ) + raise RuntimeError(get_friendly_message(500, msg)) from None + + if not isinstance(payload, dict): + return False + + event_type = payload.get("type") or event_name + if "partial_image" in event_type: + return False + + for key in ("data", "images"): + data_list = payload.get(key) + if isinstance(data_list, list): + parsed = await self._parse_response({"data": data_list}, session) + images.extend(parsed) + return True + + if payload.get("b64_json") or payload.get("url"): + parsed = await self._parse_response({"data": [payload]}, session) + images.extend(parsed) + return True + + image_obj = payload.get("image") + if isinstance(image_obj, dict) and (image_obj.get("b64_json") or image_obj.get("url")): + parsed = await self._parse_response({"data": [image_obj]}, session) + images.extend(parsed) + return True + + return False + + async def _parse_edit_stream_response( + self, + resp: aiohttp.ClientResponse, + session: aiohttp.ClientSession, + ) -> List[Image.Image]: + """ + Parse /v1/images/edits SSE events and return final completed images. + Partial images are intentionally ignored so the node output stays unchanged. + """ + images: List[Image.Image] = [] + buffer = "" + event_name = "" + data_lines = [] + partial_count = 0 + debug_body_parts = [] + + async def _handle_event(): + nonlocal event_name, data_lines, partial_count, images + if not data_lines: + event_name = "" + return + + data_str = "\n".join(data_lines).strip() + event_name = event_name.strip() + data_lines = [] + + if not data_str or data_str == "[DONE]": + return + + try: + payload = json.loads(data_str) + except Exception: + raise RuntimeError(get_friendly_message(500, data_str)) from None + + if isinstance(payload, dict): + event_type = payload.get("type") or event_name + else: + event_type = event_name + if "partial_image" in event_type: + partial_count += 1 + return + + await self._append_images_from_payload(payload, session, images, event_name) + + async for raw_chunk in resp.content.iter_any(): + chunk_text = raw_chunk.decode("utf-8", errors="ignore") + debug_body_parts.append(chunk_text) + buffer += chunk_text + while "\n" in buffer: + line, buffer = buffer.split("\n", 1) + line = line.rstrip("\r") + if line == "": + await _handle_event() + event_name = "" + continue + if line.startswith(":"): + continue + if line.startswith("event:"): + event_name = line[len("event:"):].strip() + elif line.startswith("data:"): + data_lines.append(line[len("data:"):].lstrip()) + + if buffer.strip(): + data_lines.append(buffer.strip()) + await _handle_event() + + if partial_count: + print(f"[o1key GPT Image] 流式中间图 {partial_count} 张(已忽略,仅输出最终图)") + if not images: + raise RuntimeError("流式响应结束,但未收到最终图片") + + return images + + async def _parse_stream_text_response( + self, + text: str, + session: aiohttp.ClientSession, + ) -> List[Image.Image]: + images: List[Image.Image] = [] + partial_count = 0 + event_name = "" + data_lines = [] + + async def _handle_event(): + nonlocal event_name, data_lines, partial_count, images + if not data_lines: + event_name = "" + return + data_str = "\n".join(data_lines).strip() + event_name = event_name.strip() + data_lines = [] + if not data_str or data_str == "[DONE]": + return + payload = json.loads(data_str) + if isinstance(payload, dict): + event_type = payload.get("type") or event_name + else: + event_type = event_name + if "partial_image" in event_type: + partial_count += 1 + return + await self._append_images_from_payload(payload, session, images, event_name) + + for raw_line in text.splitlines(): + line = raw_line.rstrip("\r") + if line == "": + await _handle_event() + event_name = "" + continue + if line.startswith(":"): + continue + if line.startswith("event:"): + event_name = line[len("event:"):].strip() + elif line.startswith("data:"): + data_lines.append(line[len("data:"):].lstrip()) + await _handle_event() + + if partial_count: + print(f"[o1key GPT Image] 流式中间图 {partial_count} 张(已忽略,仅输出最终图)") + if not images: + raise RuntimeError("流式响应结束,但未收到最终图片") + return images + + async def _parse_success_response( + self, + resp: aiohttp.ClientResponse, + session: aiohttp.ClientSession, + label: str = "", + ) -> List[Image.Image]: + content_type = resp.headers.get("Content-Type", "").lower() + if "event-stream" in content_type: + return await self._parse_edit_stream_response(resp, session) + + text = await resp.text() + stripped = text.lstrip() + if stripped.startswith("data:") or stripped.startswith("event:"): + return await self._parse_stream_text_response(text, session) + + try: + resp_json = json.loads(text) + except Exception: + raise RuntimeError(f"响应 JSON 解析失败,原始内容:{text[:500]}") + return await self._parse_response(resp_json, session) + # ── 中断轮询 ────────────────────────────────────────────────────────────── @staticmethod @@ -311,13 +525,15 @@ class GptImageClient: "quality": quality, "n": n, "moderation": "low", + "partial_images": 0, } body["size"] = size if size else "auto" - # 图生图:将 tensor 列表转成 data URI 内联 + image_files = [] + + # 图生图:multipart 方式上传参考图 if image_list is not None: - data_urls = [] for idx_img, img_tensor in enumerate(image_list): pil_images = tensor_to_pil(img_tensor) img = pil_images[0] @@ -333,10 +549,8 @@ class GptImageClient: png_budget = int(per_image_budget * 3 / 4) label = f"第{idx_img + 1}张" if len(image_list) > 1 else "" png_bytes = self._shrink_png_to_limit(png_bytes, png_budget, label) - b64 = base64.b64encode(png_bytes).decode("utf-8") - data_urls.append(f"data:image/png;base64,{b64}") - body["image"] = data_urls[0] if len(data_urls) == 1 else data_urls - mode = f"图生图(参考图 {len(data_urls)} 张)" + image_files.append(png_bytes) + mode = f"图生图(参考图 {len(image_files)} 张)" else: mode = "文生图" @@ -347,25 +561,40 @@ class GptImageClient: connector = aiohttp.TCPConnector(ssl=False, force_close=True) timeout = aiohttp.ClientTimeout(total=_REQUEST_TIMEOUT) + def _build_multipart_form() -> aiohttp.FormData: + form = self._new_multipart_form() + self._add_form_fields(form, body) + image_field = "image[]" if len(image_files) > 1 else "image" + for idx_img, png_bytes in enumerate(image_files): + form.add_field( + image_field, + png_bytes, + filename=f"image_{idx_img + 1}.png", + content_type="image/png", + ) + return form + async def _do_request(): async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session: last_status = None for attempt in range(DEFAULT_MAX_RETRIES + 1): t0 = time.time() - async with session.post(url, json=body, headers=self._json_headers()) as resp: + async with session.post( + url, + data=_build_multipart_form(), + headers=self._auth_headers(), + ) as resp: elapsed = time.time() - t0 - text = await resp.text() if resp.status != 200: last_status = resp.status + text = await resp.text() if resp.status in RETRYABLE_STATUS_CODES and attempt < DEFAULT_MAX_RETRIES: - friendly = HTTP_ERROR_MESSAGES.get(resp.status) + friendly = get_friendly_message(resp.status) delay = _compute_delay(attempt, DEFAULT_BASE_DELAY, DEFAULT_MAX_DELAY, DEFAULT_BACKOFF_FACTOR) print(f"[o1key GPT Image] {friendly} {delay:.1f}s 后重试 ({attempt+1}/{DEFAULT_MAX_RETRIES})...") await asyncio.sleep(delay) continue - if resp.status in _GPT_ERROR_MESSAGES: - raise RuntimeError(_GPT_ERROR_MESSAGES[resp.status]) if resp.status in HTTP_ERROR_MESSAGES: raise RuntimeError(HTTP_ERROR_MESSAGES[resp.status]) try: @@ -380,13 +609,8 @@ class GptImageClient: msg = text raise RuntimeError(get_friendly_message(resp.status, msg)) - try: - resp_json = json.loads(text) - except Exception: - raise RuntimeError(f"响应 JSON 解析失败,原始内容:{text[:500]}") - - print(f"[o1key GPT Image] API 响应耗时 {elapsed:.1f}s") - return await self._parse_response(resp_json, session) + print(f"[o1key GPT Image] API 响应耗时 {elapsed:.1f}s") + return await self._parse_success_response(resp, session, "GENERATIONS") if last_status and last_status in HTTP_ERROR_MESSAGES: raise RuntimeError(HTTP_ERROR_MESSAGES[last_status]) @@ -394,7 +618,7 @@ class GptImageClient: return await self._run_with_interrupt(_do_request()) - # ── 图像编辑(edits 接口,multipart/form-data)────────────────────────── + # ── 图像编辑(edits 接口,multipart/form-data)────────────────────── async def _edit_async( self, @@ -408,7 +632,7 @@ class GptImageClient: mask_tensor: Optional[torch.Tensor] = None, ) -> List[Image.Image]: """ - 调用 /v1/images/edits/ 接口(multipart/form-data)。 + 调用 /v1/images/edits 接口(multipart/form-data)。 """ # 模型名映射:UI 显示名 → API 参数名 api_model = _MODEL_NAME_MAP.get(model, model) @@ -421,15 +645,9 @@ class GptImageClient: normalized_tensors.append(t) num_images = len(normalized_tensors) - form = aiohttp.FormData() - form.add_field("model", api_model) - form.add_field("prompt", prompt) - form.add_field("n", str(n)) - form.add_field("quality", quality) + image_files = [] - form.add_field("size", size if size else "auto") - - # 多图:用 image[] 数组字段逐张附加,支持 gpt-image-1.5 最多 16 张 + # 多图:用 multipart image/image[] 字段逐张上传 # 预算:20MB 按图数平摊,蒙版预留 1MB mask_reserve = 1024 * 1024 if mask_tensor is not None else 0 per_image_budget = max( @@ -440,29 +658,51 @@ class GptImageClient: img_bytes = self._tensor_to_png_bytes(frame) label = f"第{i + 1}张" if num_images > 1 else "" img_bytes = self._shrink_png_to_limit(img_bytes, per_image_budget, label) - form.add_field( - "image[]", - img_bytes, - filename=f"image_{i}.png", - content_type="image/png", - ) + image_files.append(img_bytes) # 蒙版尺寸校验以第一张图为基准 first_tensor = normalized_tensors[0] ih, iw = first_tensor.shape[1], first_tensor.shape[2] + mask_png = None if mask_tensor is not None: mask_png = self._mask_tensor_to_rgba_png_bytes(mask_tensor, (ih, iw)) - form.add_field( - "mask", - mask_png, - filename="mask.png", - content_type="image/png", - ) mode = "图像编辑(带蒙版)" else: mode = "图像编辑(无蒙版)" + form_fields = { + "model": api_model, + "prompt": prompt, + "partial_images": 0, + "n": n, + "quality": quality, + "size": size if size else "auto", + "output_format": "png", + "background": "opaque", + "moderation": "low", + } + + def _build_multipart_form() -> aiohttp.FormData: + form = self._new_multipart_form() + self._add_form_fields(form, form_fields) + + for idx_img, img_bytes in enumerate(image_files): + form.add_field( + "image[]", + img_bytes, + filename=f"image_{idx_img + 1}.png", + content_type="image/png", + ) + if mask_png is not None: + form.add_field( + "mask", + mask_png, + filename="mask.png", + content_type="image/png", + ) + return form + url = f"{self.base_url}{_ENDPOINT_EDITS}" print(f"[o1key GPT Image] {mode} | 模型={model} | 参考图={num_images}张 | " f"quality={quality} | size={size} | n={n}") @@ -472,39 +712,45 @@ class GptImageClient: async def _do_request(): async with aiohttp.ClientSession(connector=connector, timeout=timeout) as session: - t0 = time.time() - async with session.post( - url, - data=form, - headers=self._auth_headers(), - ) as resp: - elapsed = time.time() - t0 - text = await resp.text() + last_status = None + for attempt in range(DEFAULT_MAX_RETRIES + 1): + t0 = time.time() + async with session.post( + url, + data=_build_multipart_form(), + headers=self._auth_headers(), + ) as resp: + elapsed = time.time() - t0 - if resp.status != 200: - if resp.status in _GPT_ERROR_MESSAGES: - raise RuntimeError(_GPT_ERROR_MESSAGES[resp.status]) - if resp.status in HTTP_ERROR_MESSAGES: - raise RuntimeError(HTTP_ERROR_MESSAGES[resp.status]) - try: - err_json = json.loads(text) - err_obj = err_json.get("error", {}) - msg = ( - err_obj.get("message") or err_obj.get("msg") or text - if isinstance(err_obj, dict) - else str(err_obj) or text - ) - except Exception: - msg = text - raise RuntimeError(get_friendly_message(resp.status, msg)) + if resp.status != 200: + text = await resp.text() + last_status = resp.status + if resp.status in RETRYABLE_STATUS_CODES and attempt < DEFAULT_MAX_RETRIES: + friendly = get_friendly_message(resp.status) + delay = _compute_delay(attempt, DEFAULT_BASE_DELAY, DEFAULT_MAX_DELAY, DEFAULT_BACKOFF_FACTOR) + print(f"[o1key GPT Image] {friendly} retrying in {delay:.1f}s ({attempt+1}/{DEFAULT_MAX_RETRIES})...") + await asyncio.sleep(delay) + continue + if resp.status in HTTP_ERROR_MESSAGES: + raise RuntimeError(HTTP_ERROR_MESSAGES[resp.status]) + try: + err_json = json.loads(text) + err_obj = err_json.get("error", {}) + msg = ( + err_obj.get("message") or err_obj.get("msg") or text + if isinstance(err_obj, dict) + else str(err_obj) or text + ) + except Exception: + msg = text + raise RuntimeError(get_friendly_message(resp.status, msg)) - try: - resp_json = json.loads(text) - except Exception: - raise RuntimeError(f"响应 JSON 解析失败,原始内容:{text[:500]}") + print(f"[o1key GPT Image] API 响应耗时 {elapsed:.1f}s") + return await self._parse_success_response(resp, session, "EDITS") - print(f"[o1key GPT Image] API 响应耗时 {elapsed:.1f}s") - return await self._parse_response(resp_json, session) + if last_status and last_status in HTTP_ERROR_MESSAGES: + raise RuntimeError(HTTP_ERROR_MESSAGES[last_status]) + raise RuntimeError(f"Request failed after {DEFAULT_MAX_RETRIES} retries") return await self._run_with_interrupt(_do_request()) @@ -525,7 +771,7 @@ class GptImageClient: 同步入口,在独立线程中运行事件循环,避免与 ComfyUI 主循环冲突。 路由逻辑: - - 无 image_tensor → generations 接口(文生图,JSON body) + - 无 image_tensor → generations 接口(文生图,multipart/form-data) - 有 image_tensor → edits 接口(图生图/编辑,multipart/form-data) """ use_edits = (image_tensor is not None) diff --git a/clients/kling_client.py b/clients/kling_client.py index c091c87..d7c03c8 100644 --- a/clients/kling_client.py +++ b/clients/kling_client.py @@ -11,6 +11,16 @@ import aiohttp from ..utils.config import get_api_key_or_raise, get_api_base_url from ..utils.http_error import async_request_with_retry +from ..utils.video_task import ( + check_interrupt, + extract_error_message, + extract_progress, + extract_status, + interruptible_sleep, + is_failure_status, + is_success_status, + run_with_interrupt, +) class KlingClient: @@ -50,9 +60,11 @@ class KlingClient: ) -> Dict[str, Any]: url = f"{self.base_url}{self.ENDPOINTS[endpoint_type]}" - resp = await async_request_with_retry( + check_interrupt() + resp = await run_with_interrupt(async_request_with_retry( session, "POST", url, json=body, headers=self._headers(), prefix="Kling 提交: " - ) + )) + check_interrupt() text = await resp.text() return json.loads(text) @@ -69,6 +81,7 @@ class KlingClient: interval = self.POLL_INITIAL_INTERVAL while True: + check_interrupt() async with session.get(url, headers=self._headers()) as resp: text = await resp.text() if resp.status != 200: @@ -77,36 +90,22 @@ class KlingClient: data = result.get("data", {}) inner_data = data.get("data", {}) if isinstance(data, dict) else {} - status = ( - data.get("status") or - inner_data.get("task_status") or - result.get("status") or - "" - ) - status = status.lower() if status else "" + status = extract_status(result) - progress_str = data.get("progress", "0%") - try: - progress_pct = int(str(progress_str).replace("%", "").strip()) - except (ValueError, AttributeError): - progress_pct = 0 + progress_pct = extract_progress(result) print(f"[视频生成] 生成中 {progress_pct}%") if on_progress: on_progress(progress_pct) - if status in ("success", "completed", "done", "finished", "succeed"): + if is_success_status(status): return result - elif status in ("failed", "fail"): - error_info = result.get("error", {}) - if isinstance(error_info, dict): - error_msg = error_info.get("message", "未知错误") - else: - error_msg = str(error_info) + elif is_failure_status(status, result): + error_msg = extract_error_message(result) raise RuntimeError(f"生成失败:{error_msg}") - await asyncio.sleep(interval) + await interruptible_sleep(interval) interval = min(interval * 1.5, self.POLL_MAX_INTERVAL) # ── 下载视频 ────────────────────────────────────────────────────── @@ -118,12 +117,14 @@ class KlingClient: session: aiohttp.ClientSession, ) -> str: print("[视频生成] 下载视频...") + check_interrupt() async with session.get(video_url, allow_redirects=True) as resp: if resp.status != 200: raise RuntimeError(f"视频下载失败 ({resp.status})") os.makedirs(os.path.dirname(save_path), exist_ok=True) with open(save_path, "wb") as f: async for chunk in resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) return save_path @@ -201,12 +202,14 @@ class KlingClient: async with aiohttp.ClientSession(connector=connector) as session: # 1. 提交 + check_interrupt() if on_stage: on_stage("submitting") create_url = f"{self.base_url}{self.NEW_API_CREATE}" - resp = await async_request_with_retry( + resp = await run_with_interrupt(async_request_with_retry( session, "POST", create_url, json=body, headers=headers, prefix="Kling 动作控制提交: " - ) + )) + check_interrupt() text = await resp.text() create_resp = json.loads(text) @@ -219,6 +222,7 @@ class KlingClient: # 2. 轮询 status_url = f"{self.base_url}{self.NEW_API_STATUS.format(video_id=video_id)}" while True: + check_interrupt() async with session.get(status_url, headers=headers) as resp: text = await resp.text() if resp.status != 200: @@ -230,29 +234,24 @@ class KlingClient: raise RuntimeError(f"状态查询失败 ({resp.status}): {msg}") status_resp = json.loads(text) - status = status_resp.get("status", "").lower() - progress_raw = status_resp.get("progress", 0) - try: - progress_pct = int(str(progress_raw).rstrip("%").strip()) - except (ValueError, AttributeError): - progress_pct = 0 + status = extract_status(status_resp) + progress_pct = extract_progress(status_resp) print(f"[动作控制] 生成中 {progress_pct}%") if on_progress: on_progress(progress_pct) - if status == "completed": + if is_success_status(status): break - if status == "failed": - error_info = status_resp.get("error", {}) - error_msg = (error_info.get("message", "未知错误") - if isinstance(error_info, dict) else str(error_info)) + if is_failure_status(status, status_resp): + error_msg = extract_error_message(status_resp) raise RuntimeError(f"动作控制生成失败:{error_msg}") - await asyncio.sleep(interval) + await interruptible_sleep(interval) interval = min(interval * 1.5, self.POLL_MAX_INTERVAL) # 3. 下载 + check_interrupt() if on_stage: on_stage("downloading") content_url = f"{self.base_url}{self.NEW_API_CONTENT.format(video_id=video_id)}" @@ -272,14 +271,15 @@ class KlingClient: os.makedirs(os.path.dirname(save_path), exist_ok=True) with open(save_path, "wb") as f: async for chunk in dl_resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) else: os.makedirs(os.path.dirname(save_path), exist_ok=True) with open(save_path, "wb") as f: async for chunk in resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) if on_stage: on_stage("done") return save_path - diff --git a/clients/seedance_client.py b/clients/seedance_client.py index b763d86..de1ac45 100644 --- a/clients/seedance_client.py +++ b/clients/seedance_client.py @@ -12,6 +12,16 @@ import aiohttp from ..utils.config import get_api_key_or_raise from ..utils.http_error import async_request_with_retry +from ..utils.video_task import ( + check_interrupt, + extract_error_message, + extract_progress, + extract_status, + interruptible_sleep, + is_failure_status, + is_success_status, + run_with_interrupt, +) class SeedanceClient: @@ -48,9 +58,11 @@ class SeedanceClient: ) -> str: """提交视频生成任务,返回 task_id""" url = f"{self.base_url}{self.CREATE_ENDPOINT}" - resp = await async_request_with_retry( + check_interrupt() + resp = await run_with_interrupt(async_request_with_retry( session, "POST", url, json=body, headers=self._headers(), prefix="Seedance 提交: " - ) + )) + check_interrupt() text = await resp.text() data = json.loads(text) @@ -73,6 +85,7 @@ class SeedanceClient: interval = self.POLL_INITIAL_INTERVAL while True: + check_interrupt() async with session.get(url, headers=self._headers()) as resp: text = await resp.text() if resp.status != 200: @@ -89,20 +102,16 @@ class SeedanceClient: # new-api 包装格式:真实数据在 result["data"] 里 inner = result.get("data") or result - status = (inner.get("status") or "").lower() + status = extract_status(result) # 解析进度 - progress_raw = inner.get("progress", "0") - try: - progress_pct = int(str(progress_raw).rstrip("%").strip()) - except (ValueError, AttributeError): - progress_pct = 0 + progress_pct = extract_progress(result) print(f"[Seedance] 生成中 {progress_pct}%") if on_progress: on_progress(progress_pct) - if status in self.SUCCESS_STATUSES: + if is_success_status(status): # 响应结构:result["data"] = inner,inner["data"] = platform_data # 视频 URL 在 inner["result_url"] 或 inner["data"]["content"]["video_url"] platform_data = inner.get("data") or {} @@ -123,15 +132,11 @@ class SeedanceClient: ) return video_url, last_frame_url - if status in self.FAILURE_STATUSES: - reason = ( - inner.get("fail_reason") - or (inner.get("error") or {}).get("message") - or "未知错误" - ) + if is_failure_status(status, result): + reason = extract_error_message(result) raise RuntimeError(f"视频生成失败:{reason}") - await asyncio.sleep(interval) + await interruptible_sleep(interval) interval = min(interval * 1.5, self.POLL_MAX_INTERVAL) # ── 3. 下载视频 ──────────────────────────────────────────────────── @@ -144,12 +149,14 @@ class SeedanceClient: ) -> str: """下载视频到本地,返回本地路径""" print(f"[Seedance] 下载视频...") + check_interrupt() async with session.get(video_url, allow_redirects=True) as resp: if resp.status != 200: raise RuntimeError(f"视频下载失败 ({resp.status})") os.makedirs(os.path.dirname(save_path), exist_ok=True) with open(save_path, "wb") as f: async for chunk in resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) return save_path @@ -167,6 +174,7 @@ class SeedanceClient: async with aiohttp.ClientSession(connector=connector) as session: # 提交 + check_interrupt() if on_stage: on_stage("submitting") task_id = await self.submit_async(body, session) diff --git a/mockups/note-sidebar-sketch.html b/mockups/note-sidebar-sketch.html new file mode 100644 index 0000000..e3a6995 --- /dev/null +++ b/mockups/note-sidebar-sketch.html @@ -0,0 +1,555 @@ + + + + + + O1Key 笔记侧栏草图 + + + +
+ + + + +
+
ComfyUI 画布区域,笔记列表不再占用右侧空间
+
+
+
+ +
+
+
编辑笔记
+ +
+
+
产品主图:高级玻璃质感
+
+
+ 产品图 × + 玻璃 × + 灯光 × + + 添加标签 +
+
只保留标签作为组织方式;可搜索、筛选、删除。
+
+
Clean studio lighting, translucent glass material, subtle caustics, soft shadow, premium product photography, 85mm lens, minimal background, high detail. + +使用方式: +1. 把产品图作为参考图输入 +2. 保留主体轮廓,只调整材质和灯光 +3. 如果玻璃过亮,降低 “caustics” 权重
+
+
+ + + + +
+
+
+ + diff --git a/nodes/K3_motion_control.py b/nodes/K3_motion_control.py index 5584e91..2b54da1 100644 --- a/nodes/K3_motion_control.py +++ b/nodes/K3_motion_control.py @@ -17,6 +17,17 @@ from ..utils.config import get_api_key_or_raise, get_async_api_base_url, NETWORK from ..utils.r2_uploader import upload_video, upload_image from ..utils.image_utils import tensor_to_pil from ..utils.http_error import async_request_with_retry +from ..utils.video_task import ( + check_interrupt, + extract_error_message, + extract_progress, + extract_status, + extract_video_url, + interruptible_sleep, + is_failure_status, + is_success_status, + run_with_interrupt, +) try: from comfy_api.latest import InputImpl @@ -183,8 +194,10 @@ class K3MotionControl: # ── 图片 & 视频上传 R2 → 获取公网 URL ──────────────────────── _stage("uploading") + check_interrupt() pil_list = tensor_to_pil(参考图片) image_url = await upload_image(pil_list[0].convert("RGB")) + check_interrupt() video_url = await upload_video(参考视频) # ── 构建请求体 ──────────────────────────────────────────────── @@ -207,13 +220,15 @@ class K3MotionControl: async with aiohttp.ClientSession(connector=connector) as session: # 1. 提交任务 + check_interrupt() _stage("submitting") create_url = f"{base_url}{_ENDPOINT_CREATE}" - resp = await async_request_with_retry( + resp = await run_with_interrupt(async_request_with_retry( session, "POST", create_url, data=json.dumps(body, ensure_ascii=False).encode("utf-8"), headers=headers, prefix="K3 动作控制提交: " - ) + )) + check_interrupt() text = await resp.text() create_resp = json.loads(text) @@ -233,7 +248,8 @@ class K3MotionControl: video_result_url = None while True: - await asyncio.sleep(interval) + await interruptible_sleep(interval) + check_interrupt() async with session.get(status_url, headers=headers) as resp: text = await resp.text() if resp.status != 200: @@ -247,30 +263,17 @@ class K3MotionControl: # 兼容扁平结构和 data 嵌套结构 data = sr.get("data", sr) - status = (data.get("status") or sr.get("status") or "").lower() + status = extract_status(sr) - pct_raw = data.get("progress", 0) - try: - pct = int(str(pct_raw).rstrip("%").strip()) - except (ValueError, AttributeError): - pct = 0 + pct = extract_progress(sr) print(f"[K3 动作控制] 生成中 {pct}%") _progress(pct) - if status in ("success", "completed", "done", "finished", "succeed"): - video_result_url = ( - data.get("video_url") - or data.get("result_url") - or data.get("url") - or (data.get("result", {}) or {}).get("url") - or sr.get("video_url") - or sr.get("url") - ) + if is_success_status(status): + video_result_url = extract_video_url(sr) break - elif status in ("failed", "fail"): - err_info = data.get("error") or sr.get("error") or {} - err_msg = (err_info.get("message", "未知错误") - if isinstance(err_info, dict) else str(err_info)) + elif is_failure_status(status, sr): + err_msg = extract_error_message(sr) raise RuntimeError(f"K3 动作控制生成失败:{err_msg}") interval = min(interval * 1.3, _POLL_MAX) @@ -279,6 +282,7 @@ class K3MotionControl: raise RuntimeError(f"API 未返回视频 URL,响应:{sr}") # 3. 下载视频 + check_interrupt() _stage("downloading") async with session.get(video_result_url, allow_redirects=True) as resp: if resp.status != 200: @@ -286,6 +290,7 @@ class K3MotionControl: os.close(tmp_fd) with open(save_path, "wb") as f: async for chunk in resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) _stage("done") diff --git a/nodes/K3_video.py b/nodes/K3_video.py index 9970a0b..b1c71ca 100644 --- a/nodes/K3_video.py +++ b/nodes/K3_video.py @@ -14,6 +14,17 @@ import aiohttp from ..utils.config import get_api_key_or_raise, get_async_api_base_url, NETWORK_ROUTE_OPTIONS, get_base_url_by_route from ..utils.image_utils import tensor_to_pil, encode_image_to_base64 from ..utils.http_error import async_request_with_retry +from ..utils.video_task import ( + check_interrupt, + extract_error_message, + extract_progress, + extract_status, + extract_video_url, + interruptible_sleep, + is_failure_status, + is_success_status, + run_with_interrupt, +) try: from comfy_api.latest import InputImpl @@ -246,11 +257,13 @@ class K3Video: async with aiohttp.ClientSession(connector=connector) as session: # 1. 提交 + check_interrupt() _stage("submitting") create_url = f"{base_url}{_ENDPOINT_CREATE}" - resp = await async_request_with_retry( + resp = await run_with_interrupt(async_request_with_retry( session, "POST", create_url, json=body, headers=headers, prefix="K3 提交: " - ) + )) + check_interrupt() text = await resp.text() create_resp = json.loads(text) @@ -269,6 +282,7 @@ class K3Video: video_url = None while True: + check_interrupt() async with session.get(status_url, headers=headers) as resp: text = await resp.text() if resp.status != 200: @@ -281,39 +295,27 @@ class K3Video: sr = json.loads(text) data = sr.get("data", sr) - status = (data.get("status") or sr.get("status") or "").lower() + status = extract_status(sr) - pct_raw = data.get("progress", 0) - try: - pct = int(str(pct_raw).rstrip("%").strip()) - except (ValueError, AttributeError): - pct = 0 + pct = extract_progress(sr) print(f"[K3 {tag}] 生成中 {pct}%") _progress(pct) - if status in ("success", "completed", "done", "finished", "succeed"): - video_url = ( - data.get("video_url") - or data.get("result_url") - or data.get("url") - or (data.get("result", {}) or {}).get("url") - or sr.get("video_url") - or sr.get("url") - ) + if is_success_status(status): + video_url = extract_video_url(sr) break - if status in ("failed", "fail"): - err_info = data.get("error") or sr.get("error") or {} - err_msg = (err_info.get("message", "未知错误") - if isinstance(err_info, dict) else str(err_info)) + if is_failure_status(status, sr): + err_msg = extract_error_message(sr) raise RuntimeError(f"K3 生成失败:{err_msg}") - await asyncio.sleep(interval) + await interruptible_sleep(interval) interval = min(interval * 1.5, _POLL_MAX) if not video_url: raise RuntimeError(f"API 未返回视频 URL,响应:{sr}") # 3. 下载 + check_interrupt() _stage("downloading") async with session.get(video_url, allow_redirects=True) as resp: if resp.status != 200: @@ -321,6 +323,7 @@ class K3Video: os.close(tmp_fd) with open(save_path, "wb") as f: async for chunk in resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) _stage("done") diff --git a/nodes/K3_video_firstlast.py b/nodes/K3_video_firstlast.py index c0b33f2..f6bc1ae 100644 --- a/nodes/K3_video_firstlast.py +++ b/nodes/K3_video_firstlast.py @@ -13,6 +13,17 @@ import aiohttp from ..utils.config import get_api_key_or_raise, get_async_api_base_url, NETWORK_ROUTE_OPTIONS, get_base_url_by_route from ..utils.image_utils import tensor_to_pil, encode_image_to_base64 from ..utils.http_error import async_request_with_retry +from ..utils.video_task import ( + check_interrupt, + extract_error_message, + extract_progress, + extract_status, + extract_video_url, + interruptible_sleep, + is_failure_status, + is_success_status, + run_with_interrupt, +) try: from comfy_api.latest import InputImpl @@ -185,11 +196,13 @@ class K3VideoFirstLast: async with aiohttp.ClientSession(connector=connector) as session: # 1. 提交 + check_interrupt() _stage("submitting") create_url = f"{base_url}{_ENDPOINT_CREATE}" - resp = await async_request_with_retry( + resp = await run_with_interrupt(async_request_with_retry( session, "POST", create_url, json=body, headers=headers, prefix="K3 首尾帧提交: " - ) + )) + check_interrupt() text = await resp.text() create_resp = json.loads(text) @@ -208,6 +221,7 @@ class K3VideoFirstLast: video_url = None while True: + check_interrupt() async with session.get(status_url, headers=headers) as resp: text = await resp.text() if resp.status != 200: @@ -220,39 +234,27 @@ class K3VideoFirstLast: sr = json.loads(text) data = sr.get("data", sr) - status = (data.get("status") or sr.get("status") or "").lower() + status = extract_status(sr) - pct_raw = data.get("progress", 0) - try: - pct = int(str(pct_raw).rstrip("%").strip()) - except (ValueError, AttributeError): - pct = 0 + pct = extract_progress(sr) print(f"[K3 首尾帧] 生成中 {pct}%") _progress(pct) - if status in ("success", "completed", "done", "finished", "succeed"): - video_url = ( - data.get("video_url") - or data.get("result_url") - or data.get("url") - or (data.get("result", {}) or {}).get("url") - or sr.get("video_url") - or sr.get("url") - ) + if is_success_status(status): + video_url = extract_video_url(sr) break - if status in ("failed", "fail"): - err_info = data.get("error") or sr.get("error") or {} - err_msg = (err_info.get("message", "未知错误") - if isinstance(err_info, dict) else str(err_info)) + if is_failure_status(status, sr): + err_msg = extract_error_message(sr) raise RuntimeError(f"K3 首尾帧生成失败:{err_msg}") - await asyncio.sleep(interval) + await interruptible_sleep(interval) interval = min(interval * 1.5, _POLL_MAX) if not video_url: raise RuntimeError(f"API 未返回视频 URL,响应:{sr}") # 3. 下载 + check_interrupt() _stage("downloading") async with session.get(video_url, allow_redirects=True) as resp: if resp.status != 200: @@ -260,6 +262,7 @@ class K3VideoFirstLast: os.close(tmp_fd) with open(save_path, "wb") as f: async for chunk in resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) _stage("done") diff --git a/nodes/K_video_firstlast.py b/nodes/K_video_firstlast.py index 261b402..5594b99 100644 --- a/nodes/K_video_firstlast.py +++ b/nodes/K_video_firstlast.py @@ -13,6 +13,17 @@ import aiohttp from ..utils.config import get_api_key_or_raise, get_api_base_url, NETWORK_ROUTE_OPTIONS, get_base_url_by_route from ..utils.image_utils import tensor_to_pil, encode_image_to_base64 from ..utils.http_error import async_request_with_retry +from ..utils.video_task import ( + check_interrupt, + extract_error_message, + extract_progress, + extract_status, + extract_video_url, + interruptible_sleep, + is_failure_status, + is_success_status, + run_with_interrupt, +) try: from comfy_api.latest import InputImpl @@ -154,11 +165,13 @@ class KVideoFirstLast: async with aiohttp.ClientSession(connector=connector) as session: # 1. 提交 + check_interrupt() _stage("submitting") create_url = f"{base_url}{_ENDPOINT_CREATE}" - resp = await async_request_with_retry( + resp = await run_with_interrupt(async_request_with_retry( session, "POST", create_url, json=body, headers=headers, prefix="K26 提交: " - ) + )) + check_interrupt() text = await resp.text() create_resp = json.loads(text) @@ -177,6 +190,7 @@ class KVideoFirstLast: video_url = None while True: + check_interrupt() async with session.get(status_url, headers=headers) as resp: text = await resp.text() if resp.status != 200: @@ -189,40 +203,28 @@ class KVideoFirstLast: sr = json.loads(text) data = sr.get("data", sr) - status = (data.get("status") or sr.get("status") or "").lower() + status = extract_status(sr) - pct_raw = data.get("progress", 0) - try: - pct = int(str(pct_raw).rstrip("%").strip()) - except (ValueError, AttributeError): - pct = 0 + pct = extract_progress(sr) print(f"[K26 图生视频] 生成中 {pct}%") _progress(pct) - if status in ("success", "completed", "done", "finished", "succeed"): + if is_success_status(status): # 提取视频 URL - video_url = ( - data.get("video_url") - or data.get("result_url") - or data.get("url") - or (data.get("result", {}) or {}).get("url") - or sr.get("video_url") - or sr.get("url") - ) + video_url = extract_video_url(sr) break - if status in ("failed", "fail"): - err_info = data.get("error") or sr.get("error") or {} - err_msg = (err_info.get("message", "未知错误") - if isinstance(err_info, dict) else str(err_info)) + if is_failure_status(status, sr): + err_msg = extract_error_message(sr) raise RuntimeError(f"K26 生成失败:{err_msg}") - await asyncio.sleep(interval) + await interruptible_sleep(interval) interval = min(interval * 1.5, _POLL_MAX) if not video_url: raise RuntimeError(f"API 未返回视频 URL,响应:{sr}") # 3. 下载 + check_interrupt() _stage("downloading") async with session.get(video_url, allow_redirects=True) as resp: if resp.status != 200: @@ -230,6 +232,7 @@ class KVideoFirstLast: os.close(tmp_fd) with open(save_path, "wb") as f: async for chunk in resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) _stage("done") diff --git a/nodes/K_video_image2video.py b/nodes/K_video_image2video.py index d784cda..6b5f976 100644 --- a/nodes/K_video_image2video.py +++ b/nodes/K_video_image2video.py @@ -14,6 +14,17 @@ import aiohttp from ..utils.config import get_api_key_or_raise, get_api_base_url, NETWORK_ROUTE_OPTIONS, get_base_url_by_route from ..utils.image_utils import tensor_to_pil, encode_image_to_base64 from ..utils.http_error import async_request_with_retry +from ..utils.video_task import ( + check_interrupt, + extract_error_message, + extract_progress, + extract_status, + extract_video_url, + interruptible_sleep, + is_failure_status, + is_success_status, + run_with_interrupt, +) try: from comfy_api.latest import InputImpl @@ -150,11 +161,13 @@ class KVideoImage2Video: async with aiohttp.ClientSession(connector=connector) as session: # 1. 提交 + check_interrupt() _stage("submitting") create_url = f"{base_url}{_ENDPOINT_CREATE}" - resp = await async_request_with_retry( + resp = await run_with_interrupt(async_request_with_retry( session, "POST", create_url, json=body, headers=headers, prefix="K26 图生视频提交: " - ) + )) + check_interrupt() sr = await resp.json() task_id = sr.get("task_id") or sr.get("id") @@ -169,8 +182,9 @@ class KVideoImage2Video: video_url = None while True: - await asyncio.sleep(interval) + await interruptible_sleep(interval) + check_interrupt() async with session.get(status_url, headers=headers) as resp: if resp.status != 200: err_text = await resp.text() @@ -178,40 +192,27 @@ class KVideoImage2Video: sr = await resp.json() data = sr.get("data", {}) or {} - status = (sr.get("status") or data.get("status") or "").lower() + status = extract_status(sr) - pct_raw = str(data.get("progress", 0)).strip().rstrip('%') - try: - pct = max(0, min(100, int(float(pct_raw)))) - except (ValueError, TypeError): - pct = 0 + pct = extract_progress(sr) print(f"[K26 图生视频] 生成中 {pct}%") _progress(pct) - if status in ("success", "completed", "done", "finished", "succeed"): + if is_success_status(status): # 提取视频 URL - video_url = ( - data.get("video_url") - or data.get("result_url") - or data.get("url") - or (data.get("result", {}) or {}).get("url") - or sr.get("video_url") - or sr.get("url") - ) + video_url = extract_video_url(sr) break - if status in ("failed", "fail"): - err_info = data.get("error") or sr.get("error") or {} - err_msg = (err_info.get("message", "未知错误") - if isinstance(err_info, dict) else str(err_info)) + if is_failure_status(status, sr): + err_msg = extract_error_message(sr) raise RuntimeError(f"K26 生成失败:{err_msg}") - await asyncio.sleep(interval) interval = min(interval * 1.5, _POLL_MAX) if not video_url: raise RuntimeError(f"API 未返回视频 URL,响应:{sr}") # 3. 下载 + check_interrupt() _stage("downloading") async with session.get(video_url, allow_redirects=True) as resp: if resp.status != 200: @@ -219,6 +220,7 @@ class KVideoImage2Video: os.close(tmp_fd) with open(save_path, "wb") as f: async for chunk in resp.content.iter_chunked(8192): + check_interrupt() f.write(chunk) _stage("done") diff --git a/nodes/__init__.py b/nodes/__init__.py index 2c5608c..023eeda 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -20,7 +20,7 @@ from .batch_images_o1key import BatchImagesO1key from .seedance_video import Seedance, SeedanceMultiModal from .nano_banana_v2 import NanoBananaV2, NanoBananaV2Batch, AsyncImageGenerator, BatchAsyncImageGenerator from .doubao_image import DoubaoImage -from .gpt_image import O1keyGPTImage +from .gpt_image import O1keyGPTImage, O1keyGPTImageBatch from .grok_image import O1keyGrokImage from .K_video_firstlast import KVideoFirstLast from .K_video_image2video import KVideoImage2Video @@ -31,5 +31,6 @@ from .save_image_format import SaveImageFormat from .save_psd import O1keySavePSD from .remove_bg import O1keyRemoveBackground from .color_remove_bg import O1keyColorRemoveBG +from .grid_splitter import O1keyGridSplitter -__all__ = ['NanoBananaV2', 'NanoBananaV2Batch', 'NanoBanana', 'BatchNanoBananaPro', 'GoogleGemini', 'LoadFile', 'ImageStitchPro', 'BatchCleanMetadata', 'VideoPreview', 'KlingVideo', 'KlingFirstLastFrame', 'KlingMotionControlTest', 'AspectRatioPreset', 'GoogleVeo', 'FluxImageEdit', 'UniversalLLMChat', 'BatchImagesO1key', 'Seedance', 'SeedanceMultiModal', 'StreamPreview', 'DoubaoImage', 'O1keyGPTImage', 'O1keyGrokImage', 'KVideoFirstLast', 'KVideoImage2Video', 'K3Video', 'K3VideoFirstLast', 'K3MotionControl', 'K3MotionVideoCheck', 'AsyncImageGenerator', 'BatchAsyncImageGenerator', 'SaveImageFormat', 'O1keySavePSD', 'O1keyRemoveBackground', 'O1keyColorRemoveBG'] +__all__ = ['NanoBananaV2', 'NanoBananaV2Batch', 'NanoBanana', 'BatchNanoBananaPro', 'GoogleGemini', 'LoadFile', 'ImageStitchPro', 'BatchCleanMetadata', 'VideoPreview', 'KlingVideo', 'KlingFirstLastFrame', 'KlingMotionControlTest', 'AspectRatioPreset', 'GoogleVeo', 'FluxImageEdit', 'UniversalLLMChat', 'BatchImagesO1key', 'Seedance', 'SeedanceMultiModal', 'StreamPreview', 'DoubaoImage', 'O1keyGPTImage', 'O1keyGPTImageBatch', 'O1keyGrokImage', 'KVideoFirstLast', 'KVideoImage2Video', 'K3Video', 'K3VideoFirstLast', 'K3MotionControl', 'K3MotionVideoCheck', 'AsyncImageGenerator', 'BatchAsyncImageGenerator', 'SaveImageFormat', 'O1keySavePSD', 'O1keyRemoveBackground', 'O1keyColorRemoveBG', 'O1keyGridSplitter'] diff --git a/nodes/batch_nano_banana.py b/nodes/batch_nano_banana.py index c7c90b7..583c5dd 100644 --- a/nodes/batch_nano_banana.py +++ b/nodes/batch_nano_banana.py @@ -20,7 +20,7 @@ from PIL import Image import torch import numpy as np -from ..utils.image_utils import tensor_to_pil, pil_to_tensor, parse_batch_prompts, encode_image_to_base64, encode_image_to_base64_limited +from ..utils.image_utils import tensor_to_pil, pil_to_tensor, parse_batch_prompts, encode_images_for_request_body_limit from ..utils.file_utils import ( ImageInfo, load_images_from_folder, @@ -95,22 +95,6 @@ def _build_request_body( images: Optional[List[Image.Image]] = None, enable_grounding: bool = False, ) -> dict: - content_parts = [{"type": "text", "text": prompt}] - - if images: - for img in images: - b64 = encode_image_to_base64_limited(img, format="PNG") - content_parts.append({ - "type": "image_url", - "image_url": {"url": f"data:image/png;base64,{b64}"} - }) - - body = { - "model": model, - "stream": True, - "messages": [{"role": "user", "content": content_parts}], - } - google_config = { "image_config": { "image_size": resolution, @@ -118,12 +102,32 @@ def _build_request_body( } if aspect_ratio and aspect_ratio != "智能": google_config["image_config"]["aspect_ratio"] = aspect_ratio - body["extra_body"] = {"google": google_config} - if enable_grounding: - body["extra_body"]["google_search"] = True + def _make_body(encoded_images: Optional[List[tuple]] = None) -> dict: + content_parts = [{"type": "text", "text": prompt}] + if encoded_images: + for mime_type, b64 in encoded_images: + content_parts.append({ + "type": "image_url", + "image_url": {"url": f"data:{mime_type};base64,{b64}"} + }) - return body + body = { + "model": model, + "stream": True, + "messages": [{"role": "user", "content": content_parts}], + "extra_body": {"google": google_config}, + } + + if enable_grounding: + body["extra_body"]["google_search"] = True + return body + + encoded_images = None + if images: + encoded_images = encode_images_for_request_body_limit(images, _make_body) + + return _make_body(encoded_images) async def _generate_single_openai( @@ -1159,4 +1163,4 @@ class BatchNanoBananaPro: # 最终内存清理 import gc gc.collect() - print(f"BatchNanoBananaPro: 最终内存清理完成") \ No newline at end of file + print(f"BatchNanoBananaPro: 最终内存清理完成") diff --git a/nodes/gpt_image.py b/nodes/gpt_image.py index 8c7393d..68b7a4e 100644 --- a/nodes/gpt_image.py +++ b/nodes/gpt_image.py @@ -3,10 +3,23 @@ o1key GPT Image 节点 支持 gpt-image-1 / gpt-image-1.5 模型的文生图、图生图、图像编辑(带蒙版) """ +import os import time +from typing import List, Optional, Tuple + +from PIL import Image + from ..clients.gpt_image_client import GptImageClient -from ..utils.image_utils import parse_batch_prompts +from ..utils.image_utils import parse_batch_prompts, pil_to_tensor, tensor_to_pil from ..utils.config import NETWORK_ROUTE_OPTIONS, get_base_url_by_route +from ..utils.file_utils import ( + ImageInfo, + generate_timestamp_filename, + load_images_from_folder, + pair_images_by_name, + pair_images_cartesian, + save_image, +) try: from comfy.model_management import processing_interrupted, InterruptProcessingException @@ -16,6 +29,18 @@ except ImportError: processing_interrupted = lambda: False InterruptProcessingException = RuntimeError +try: + from comfy.utils import ProgressBar + _PROGRESS_BAR_AVAILABLE = True +except ImportError: + _PROGRESS_BAR_AVAILABLE = False + +try: + import folder_paths + _FOLDER_PATHS_AVAILABLE = True +except ImportError: + _FOLDER_PATHS_AVAILABLE = False + class O1keyGPTImage: """ @@ -270,3 +295,455 @@ class O1keyGPTImage: print(f"[o1key GPT Image] {balance_info}") except Exception: pass + + +class O1keyGPTImageBatch: + """ + o1key GPT Image 批量节点 + + 复用 BatchNanoBananaPro 的批量思路: + - 从文件夹批量加载图片 + - 按文件名同名 / 1*N / 不配对 三种模式创建任务 + - 可追加节点手动输入参考图 + - prompt 支持用独占一行 --- 展开为多提示词任务 + - 每个任务调用 GPT Image 客户端并保存到磁盘 + """ + + PAIRING_MODES = ["按相同图片命名", "1*N", "不配对"] + IMAGE_FORMATS = ["原始", "JPEG", "PNG", "WebP"] + MODEL_OPTIONS = ["gpt-image-2-按量", "gpt-image-2-次卡"] + QUALITY_OPTIONS = ["高", "中", "低", "自动"] + RESOLUTION_OPTIONS = [ + "智能", + "1024x1024(1K 正方形 1:1)", + "1536x1024(1K 横版 3:2)", + "1024x1536(1K 竖版 2:3)", + "1360x1024(1K 横版 4:3)", + "1024x1360(1K 竖版 3:4)", + "1824x1024(1K 横版 16:9)", + "1024x1824(1K 竖版 9:16)", + "2048x2048(2K 正方形 1:1)", + "3072x2048(2K 横版 3:2)", + "2048x3072(2K 竖版 2:3)", + "2736x2048(2K 横版 4:3)", + "2048x2736(2K 竖版 3:4)", + "3648x2048(2K 横版 16:9)", + "2048x3648(2K 竖版 9:16)", + "2880x2880(4K 正方形 1:1)", + "3504x2336(4K 横版 3:2)", + "2336x3504(4K 竖版 2:3)", + "3264x2448(4K 横版 4:3)", + "2448x3264(4K 竖版 3:4)", + "3840x2160(4K 横版 16:9)", + "2160x3840(4K 竖版 9:16)", + ] + + @classmethod + def INPUT_TYPES(cls): + optional_inputs = {} + for image_index in range(1, 10): + optional_inputs[f"参考图{image_index}"] = ("IMAGE", { + "tooltip": "追加到每个批量任务末尾的固定参考图。", + }) + + optional_inputs["遮罩"] = ("MASK", { + "tooltip": "可选蒙版,会应用到每个任务的第一张参考图;请确保尺寸一致。", + }) + optional_inputs["图片配对模式"] = (cls.PAIRING_MODES, { + "default": "不配对", + "tooltip": "文件夹图片的组合方式;手动参考图只追加,不参与配对。", + }) + + return { + "required": { + "prompt": ("STRING", { + "default": "", + "multiline": True, + "tooltip": "提示词;可用独占一行的 --- 分隔多条批量提示词。", + }), + "模型": (cls.MODEL_OPTIONS, { + "default": "gpt-image-2-次卡", + }), + "网络": (NETWORK_ROUTE_OPTIONS, { + "default": "全球加速", + }), + "分辨率": (cls.RESOLUTION_OPTIONS, { + "default": "智能", + }), + "生图数量": ("INT", { + "default": 1, + "min": 1, + "max": 8, + "step": 1, + "display": "number", + }), + "质量": (cls.QUALITY_OPTIONS, { + "default": "自动", + }), + "seed": ("INT", { + "default": 0, + "min": 0, + "max": 2**31 - 1, + "step": 1, + "display": "number", + "control_after_generate": True, + }), + "图片格式": (cls.IMAGE_FORMATS, { + "default": "原始", + }), + "文件夹1": ("STRING", { + "default": "", + "multiline": False, + }), + "文件夹2": ("STRING", { + "default": "", + "multiline": False, + }), + "文件夹3": ("STRING", { + "default": "", + "multiline": False, + }), + "文件夹4": ("STRING", { + "default": "", + "multiline": False, + }), + "文件夹5": ("STRING", { + "default": "", + "multiline": False, + }), + "保存路径": ("STRING", { + "default": "", + "multiline": False, + "tooltip": "为空时优先使用 ComfyUI 默认 output 目录。", + }), + }, + "optional": optional_inputs, + } + + RETURN_TYPES = ("IMAGE",) + RETURN_NAMES = ("IMAGE",) + FUNCTION = "process_batch" + CATEGORY = "o1key/image" + OUTPUT_NODE = False + + def _load_folders(self, folders: List[str]) -> List[List[ImageInfo]]: + image_lists = [] + for folder_index, folder in enumerate(folders, 1): + if not folder or not folder.strip(): + continue + try: + loaded_images = load_images_from_folder(folder) + if loaded_images: + image_lists.append(loaded_images) + except ValueError as error: + print(f"[o1key GPT Image Batch] 文件夹{folder_index} 加载失败 - {error}") + return image_lists + + def _create_pairs( + self, + image_lists: List[List[ImageInfo]], + pairing_mode: str, + manual_images: Optional[List[ImageInfo]] = None, + ) -> List[Tuple[ImageInfo, ...]]: + if pairing_mode == "不配对": + if len(image_lists) > 1: + raise ValueError("「不配对」模式只支持单个文件夹,请清空其他文件夹路径") + + if image_lists and manual_images: + return [ + (folder_image,) + tuple(manual_images) + for folder_image in image_lists[0] + ] + if image_lists: + return [(folder_image,) for folder_image in image_lists[0]] + return [] + + if not image_lists: + return [] + + if len(image_lists) == 1: + base_pairs = [(folder_image,) for folder_image in image_lists[0]] + elif pairing_mode == "按相同图片命名": + base_pairs = list(pair_images_by_name(*image_lists)) + else: + base_pairs = list(pair_images_cartesian(*image_lists)) + + if manual_images: + manual_tuple = tuple(manual_images) + base_pairs = [pair + manual_tuple for pair in base_pairs] + + return base_pairs + + def _collect_manual_images(self, kwargs) -> List[ImageInfo]: + manual_images = [] + for image_index in range(1, 10): + key = f"参考图{image_index}" + if key not in kwargs or kwargs[key] is None: + continue + for tensor_index, image in enumerate(tensor_to_pil(kwargs[key])): + manual_images.append(ImageInfo( + image=image, + filename=f"manual_{image_index}_{tensor_index}", + extension=".png", + source_path="", + )) + return manual_images + + @staticmethod + def _pair_to_tensors(pair: Tuple[ImageInfo, ...]) -> List: + return [pil_to_tensor([image_info.image]) for image_info in pair] + + @staticmethod + def _resolve_size(分辨率: str) -> str: + return "auto" if 分辨率 == "智能" else 分辨率.split("(")[0].strip() + + @staticmethod + def _resolve_model(模型: str) -> str: + model_map = { + "gpt-image-2-次卡": "gpt-image-2-c", + "gpt-image-2-按量": "gpt-image-2", + } + return model_map.get(模型, 模型) + + @staticmethod + def _resolve_quality(质量: str) -> str: + quality_map = {"高": "high", "中": "medium", "低": "low", "自动": "auto"} + return quality_map.get(质量, "auto") + + @staticmethod + def _ensure_output_folder(保存路径: str) -> str: + output_folder = (保存路径 or "").strip() + if not output_folder and _FOLDER_PATHS_AVAILABLE: + output_folder = folder_paths.get_output_directory() + print(f"[o1key GPT Image Batch] 未设置保存路径,使用 ComfyUI 默认 output 目录: {output_folder}") + + if not output_folder: + raise ValueError("未设置保存路径,且当前环境无法获取 ComfyUI 默认 output 目录") + + os.makedirs(output_folder, exist_ok=True) + test_path = os.path.join(output_folder, ".write_test") + with open(test_path, "w", encoding="utf-8") as test_file: + test_file.write("test") + os.remove(test_path) + return output_folder + + @staticmethod + def _save_images( + images: List[Image.Image], + output_folder: str, + image_format: str, + base_filename: Optional[str] = None, + ) -> List[str]: + format_ext_map = {"JPEG": ".jpg", "PNG": ".png", "WebP": ".webp"} + save_ext = format_ext_map.get(image_format, ".png") + saved_files = [] + + for image in images: + if base_filename: + counter = 0 + while True: + suffix = "" if counter == 0 else f"+{counter}" + filename = f"{base_filename}{suffix}{save_ext}" + output_path = os.path.join(output_folder, filename) + if not os.path.exists(output_path): + break + counter += 1 + else: + output_path = generate_timestamp_filename( + output_folder=output_folder, + extension=save_ext, + ) + + if image_format == "JPEG": + if image.mode != "RGB": + image = image.convert("RGB") + image.save(output_path, quality=100) + elif image_format == "WebP": + image.save(output_path, lossless=True) + else: + save_image(image, output_path) + + saved_files.append(output_path) + + return saved_files + + def process_batch( + self, + prompt: str, + 模型: str, + 网络: str, + 分辨率: str, + 生图数量: int, + 质量: str, + seed: int, + 图片格式: str, + 文件夹1: str, + 文件夹2: str, + 文件夹3: str, + 文件夹4: str, + 文件夹5: str, + 保存路径: str = "", + 图片配对模式: str = "不配对", + 遮罩=None, + **kwargs, + ): + start_time = time.time() + client = None + + try: + if not prompt or not prompt.strip(): + raise ValueError("提示词不能为空") + + folders = [文件夹1, 文件夹2, 文件夹3, 文件夹4, 文件夹5] + if not any(folder and folder.strip() for folder in folders): + raise ValueError("请至少填写一个文件夹路径,该节点专为批量文件夹处理设计") + + image_lists = self._load_folders(folders) + total_folder_images = sum(len(image_list) for image_list in image_lists) + if total_folder_images == 0: + raise ValueError("文件夹中未找到任何图片,请检查文件夹路径是否正确") + + manual_images = self._collect_manual_images(kwargs) + pairs = self._create_pairs( + image_lists=image_lists, + pairing_mode=图片配对模式, + manual_images=manual_images if manual_images else None, + ) + if not pairs: + raise ValueError("配对结果为空,请检查输入") + + batch_prompts = parse_batch_prompts(prompt) + prompts_per_task = None + if batch_prompts: + expanded_pairs = [] + expanded_prompts = [] + for pair in pairs: + for batch_prompt in batch_prompts: + expanded_pairs.append(pair) + expanded_prompts.append(batch_prompt) + pairs = expanded_pairs + prompts_per_task = expanded_prompts + + total_tasks = len(pairs) + if batch_prompts: + print( + f"[o1key GPT Image Batch] 批量任务 | {图片配对模式} × " + f"{len(batch_prompts)} 个提示词 | 共 {total_tasks} 任务" + ) + else: + print(f"[o1key GPT Image Batch] 批量任务 | {图片配对模式} | 共 {total_tasks} 任务") + + output_folder = self._ensure_output_folder(保存路径) + size = self._resolve_size(分辨率) + model = self._resolve_model(模型) + quality = self._resolve_quality(质量) + + client = GptImageClient() + client.base_url = get_base_url_by_route(网络) + + progress_bar = ProgressBar(total_tasks) if _PROGRESS_BAR_AVAILABLE else None + results = [] + all_saved_files = [] + + for task_index, pair in enumerate(pairs, 1): + if _INTERRUPT_AVAILABLE and processing_interrupted(): + print("[o1key GPT Image Batch] 用户取消,已中断批量生成") + raise InterruptProcessingException() + + task_prompt = prompts_per_task[task_index - 1] if prompts_per_task else prompt + base_filename = pair[0].filename if pair else None + result = { + "task_index": task_index, + "success": False, + "generated_count": 0, + "saved_files": [], + "error": None, + } + + try: + pil_images = client.run_sync( + prompt=task_prompt, + model=model, + quality=quality, + size=size, + n=生图数量, + seed=seed, + image_tensor=self._pair_to_tensors(pair), + mask_tensor=遮罩, + ) + saved_files = self._save_images( + images=pil_images, + output_folder=output_folder, + image_format=图片格式, + base_filename=base_filename, + ) + result["success"] = bool(pil_images) + result["generated_count"] = len(pil_images) + result["saved_files"] = saved_files + all_saved_files.extend(saved_files) + print(f"[o1key GPT Image Batch] [{task_index}/{total_tasks}] ✓ {base_filename or 'task'}") + except InterruptProcessingException: + raise + except Exception as error: + error_msg = str(error).split("\n")[0] + result["error"] = error_msg + print(f"[o1key GPT Image Batch] [{task_index}/{total_tasks}] ❌ {base_filename or 'task'} → {error_msg}") + + results.append(result) + if progress_bar is not None: + progress_bar.update(1) + + success_count = sum(1 for result in results if result.get("success", False)) + total_generated = sum(result.get("generated_count", 0) for result in results) + if success_count == 0: + raise RuntimeError("所有批量任务均生成失败,无可用图像输出") + + output_images = [] + for file_path in all_saved_files[-10:]: + try: + loaded_image = Image.open(file_path) + loaded_image.load() + output_images.append(loaded_image) + except Exception as error: + print(f"[o1key GPT Image Batch] 无法加载输出图片 {file_path} - {error}") + + if not output_images: + output_images = [Image.new("RGBA", (512, 512), (128, 128, 128, 255))] + + output_tensor = GptImageClient._pil_list_to_tensor(output_images) + elapsed = time.time() - start_time + print("=" * 60) + print( + f"[o1key GPT Image Batch] 完成!耗时 {elapsed:.1f}s | " + f"成功 {success_count}/{total_tasks} | 生成 {total_generated} 张" + ) + print(f"[o1key GPT Image Batch] 保存路径: {output_folder}") + if all_saved_files: + print(f"[o1key GPT Image Batch] 最新保存文件: {all_saved_files[-1]}") + + failed_results = [result for result in results if not result.get("success", False)] + if failed_results: + print(f"[o1key GPT Image Batch] 失败任务: {len(failed_results)} 个") + for failed_result in failed_results[:3]: + print( + f" - #{failed_result.get('task_index')}: " + f"{failed_result.get('error', '未知错误')}" + ) + + return (output_tensor,) + + except ValueError as error: + if str(error) == "未授权!": + print("[o1key GPT Image Batch] 请联系作者授权后方可使用!") + raise ValueError("未授权!") from None + raise ValueError(str(error)) from None + except RuntimeError as error: + raise RuntimeError(str(error)) from None + finally: + if client is not None: + try: + balance_data = client.query_balance_sync() + balance_info = client.format_balance_info(balance_data) + print(f"[o1key GPT Image Batch] {balance_info}") + except Exception: + pass diff --git a/nodes/grid_splitter.py b/nodes/grid_splitter.py new file mode 100644 index 0000000..abd239e --- /dev/null +++ b/nodes/grid_splitter.py @@ -0,0 +1,398 @@ +""" +Merged grid image splitter. + +This node is designed for AI-generated contact sheets such as 3x3 or 2x3 +grids. Auto mode scores common layouts by looking for strong seams or flat +separator bands near the expected grid lines, then crops each cell. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import List, Sequence, Tuple + +import numpy as np +import torch +from PIL import Image + +from ..utils.image_utils import pil_to_tensor, tensor_to_pil + + +_AUTO_LAYOUTS: Sequence[Tuple[int, int]] = ( + (3, 3), + (2, 3), + (3, 2), + (2, 2), + (1, 2), + (2, 1), + (1, 3), + (3, 1), + (4, 4), + (3, 4), + (4, 3), +) + +_LAYOUTS = [ + "auto", + "1x2", + "2x1", + "1x3", + "3x1", + "2x2", + "2x3", + "3x2", + "3x3", + "3x4", + "4x3", + "4x4", + "custom", +] + + +@dataclass(frozen=True) +class _AxisCut: + seam: int + span_start: int + span_end: int + score: float + + +@dataclass(frozen=True) +class _AxisPlan: + intervals: List[Tuple[int, int]] + cuts: List[_AxisCut] + score: float + + +def _to_float_array(image: Image.Image) -> np.ndarray: + if image.mode != "RGB": + image = image.convert("RGB") + return np.asarray(image).astype(np.float32) / 255.0 + + +def _axis_texture(arr: np.ndarray, axis: str) -> np.ndarray: + if axis == "x": + profile = arr.std(axis=(0, 2)) + else: + profile = arr.std(axis=(1, 2)) + high = np.percentile(profile, 95) + 1e-6 + return np.clip(profile / high, 0.0, 1.0) + + +def _axis_edge(arr: np.ndarray, axis: str) -> np.ndarray: + if axis == "x": + diff = np.abs(np.diff(arr, axis=1)).mean(axis=(0, 2)) + length = arr.shape[1] + else: + diff = np.abs(np.diff(arr, axis=0)).mean(axis=(1, 2)) + length = arr.shape[0] + + padded = np.zeros(length, dtype=np.float32) + if diff.size: + padded[1:] = diff + high = np.percentile(padded, 95) + 1e-6 + return np.clip(padded / high, 0.0, 1.5) + + +def _smooth(profile: np.ndarray, radius: int = 2) -> np.ndarray: + if radius <= 0 or profile.size < radius * 2 + 1: + return profile + kernel = np.ones(radius * 2 + 1, dtype=np.float32) / float(radius * 2 + 1) + return np.convolve(profile, kernel, mode="same") + + +def _separator_span( + texture: np.ndarray, + seam: int, + search_px: int, + min_separator_px: int, +) -> Tuple[int, int]: + length = texture.size + if length <= 1: + return 0, length + + limit = max(1, min(search_px, length // 8)) + threshold = max(0.08, min(0.28, float(np.percentile(texture, 12)) * 1.8)) + + left = seam + while left > 0 and seam - left < limit and texture[left - 1] <= threshold: + left -= 1 + + right = seam + while right < length and right - seam < limit and texture[right] <= threshold: + right += 1 + + if right - left >= max(1, min_separator_px): + return left, right + + return seam, seam + + +def _edge_trim(texture: np.ndarray, search_px: int, min_cell: int) -> Tuple[int, int]: + length = texture.size + if length <= 2: + return 0, length + + max_trim = max(0, min(search_px * 2, min_cell // 3, length // 6)) + if max_trim <= 0: + return 0, length + + threshold = max(0.08, min(0.24, float(np.percentile(texture, 12)) * 1.6)) + + start = 0 + while start < max_trim and texture[start] <= threshold: + start += 1 + + end = length + while length - end < max_trim and end > start + min_cell and texture[end - 1] <= threshold: + end -= 1 + + return start, end + + +def _axis_plan( + arr: np.ndarray, + cells: int, + axis: str, + search_px: int, + crop_separators: bool, + trim_outer: bool, + min_separator_px: int, +) -> _AxisPlan: + length = arr.shape[1] if axis == "x" else arr.shape[0] + if cells <= 1: + return _AxisPlan(intervals=[(0, length)], cuts=[], score=0.0) + + raw_texture = _axis_texture(arr, axis) + raw_edge = _axis_edge(arr, axis) + texture = _smooth(raw_texture, radius=2) + edge = _smooth(raw_edge, radius=1) + evidence = np.maximum(edge, (1.0 - texture) * 0.75) + exact_evidence = np.maximum(raw_edge, (1.0 - raw_texture) * 0.75) + + cuts: List[_AxisCut] = [] + scores: List[float] = [] + for idx in range(1, cells): + expected = round(length * idx / cells) + start = max(1, expected - search_px) + end = min(length - 1, expected + search_px) + if start >= end: + seam = expected + score = 0.0 + else: + window = evidence[start:end + 1] + offset = int(window.argmax()) + coarse = start + offset + fine_start = max(start, coarse - 2) + fine_end = min(end, coarse + 2) + fine_window = exact_evidence[fine_start:fine_end + 1] + seam = fine_start + int(fine_window.argmax()) + score = float(window[offset]) + + span_start, span_end = _separator_span( + raw_texture, + seam, + search_px=search_px, + min_separator_px=min_separator_px, + ) + cuts.append(_AxisCut(seam=seam, span_start=span_start, span_end=span_end, score=score)) + scores.append(score) + + min_cell = max(1, length // cells) + outer_start, outer_end = _edge_trim(raw_texture, search_px, min_cell) if trim_outer else (0, length) + + intervals: List[Tuple[int, int]] = [] + cursor = outer_start + for cut in cuts: + split_start = cut.span_start if crop_separators else cut.seam + split_end = cut.span_end if crop_separators else cut.seam + intervals.append((cursor, split_start)) + cursor = split_end + intervals.append((cursor, outer_end)) + + cleaned: List[Tuple[int, int]] = [] + for start, end in intervals: + start = max(0, min(length - 1, int(start))) + end = max(start + 1, min(length, int(end))) + cleaned.append((start, end)) + + return _AxisPlan( + intervals=cleaned, + cuts=cuts, + score=float(np.mean(scores)) if scores else 0.0, + ) + + +def _parse_layout(layout: str, custom_rows: int, custom_cols: int) -> Tuple[int, int]: + if layout == "custom": + return max(1, int(custom_rows)), max(1, int(custom_cols)) + rows_text, cols_text = layout.split("x", 1) + return int(rows_text), int(cols_text) + + +def _fallback_layout(width: int, height: int) -> Tuple[int, int]: + aspect = width / max(1, height) + if 0.82 <= aspect <= 1.22: + return 3, 3 + if aspect > 1.22: + return 2, 3 + return 3, 2 + + +def _choose_auto_layout( + arr: np.ndarray, + search_px: int, + crop_separators: bool, + trim_outer: bool, + min_separator_px: int, +) -> Tuple[int, int, _AxisPlan, _AxisPlan, float, bool]: + height, width = arr.shape[:2] + best = None + + for rows, cols in _AUTO_LAYOUTS: + x_plan = _axis_plan(arr, cols, "x", search_px, crop_separators, trim_outer, min_separator_px) + y_plan = _axis_plan(arr, rows, "y", search_px, crop_separators, trim_outer, min_separator_px) + score = (x_plan.score + y_plan.score) / 2.0 + + # Prefer common 3x3 / 2x3 / 3x2 layouts when the image gives weak signals. + if (rows, cols) in ((3, 3), (2, 3), (3, 2)): + score += 0.025 + + if best is None or score > best[0]: + best = (score, rows, cols, x_plan, y_plan) + + assert best is not None + score, rows, cols, x_plan, y_plan = best + confident = score >= 0.22 + + if confident: + return rows, cols, x_plan, y_plan, score, True + + rows, cols = _fallback_layout(width, height) + x_plan = _axis_plan(arr, cols, "x", search_px, crop_separators, trim_outer, min_separator_px) + y_plan = _axis_plan(arr, rows, "y", search_px, crop_separators, trim_outer, min_separator_px) + return rows, cols, x_plan, y_plan, score, False + + +def _normalize_sizes(crops: List[Image.Image]) -> List[Image.Image]: + min_w = min(crop.width for crop in crops) + min_h = min(crop.height for crop in crops) + normalized = [] + for crop in crops: + left = max(0, (crop.width - min_w) // 2) + top = max(0, (crop.height - min_h) // 2) + normalized.append(crop.crop((left, top, left + min_w, top + min_h))) + return normalized + + +def _split_one( + image: Image.Image, + layout: str, + custom_rows: int, + custom_cols: int, + search_px: int, + crop_separators: bool, + trim_outer: bool, + min_separator_px: int, +) -> Tuple[List[Image.Image], str]: + arr = _to_float_array(image) + + if layout == "auto": + rows, cols, x_plan, y_plan, confidence, confident = _choose_auto_layout( + arr, + search_px=search_px, + crop_separators=crop_separators, + trim_outer=trim_outer, + min_separator_px=min_separator_px, + ) + mode_note = "auto" if confident else "auto-low-confidence-fallback" + else: + rows, cols = _parse_layout(layout, custom_rows, custom_cols) + x_plan = _axis_plan(arr, cols, "x", search_px, crop_separators, trim_outer, min_separator_px) + y_plan = _axis_plan(arr, rows, "y", search_px, crop_separators, trim_outer, min_separator_px) + confidence = (x_plan.score + y_plan.score) / 2.0 + mode_note = "manual" + + crops: List[Image.Image] = [] + for y0, y1 in y_plan.intervals: + for x0, x1 in x_plan.intervals: + crops.append(image.crop((x0, y0, x1, y1))) + + crops = _normalize_sizes(crops) + info = ( + f"{mode_note}: {rows}x{cols}, cells={len(crops)}, " + f"confidence={confidence:.3f}, " + f"x={x_plan.intervals}, y={y_plan.intervals}" + ) + return crops, info + + +class O1keyGridSplitter: + """Split AI-generated grid/contact-sheet images into individual cells.""" + + @classmethod + def INPUT_TYPES(cls): + return { + "required": { + "图像": ("IMAGE",), + "布局": (_LAYOUTS, {"default": "auto"}), + "自定义行数": ("INT", {"default": 3, "min": 1, "max": 12, "step": 1}), + "自定义列数": ("INT", {"default": 3, "min": 1, "max": 12, "step": 1}), + "搜索范围px": ("INT", {"default": 32, "min": 0, "max": 256, "step": 1}), + "裁掉分隔线": ("BOOLEAN", {"default": True}), + "裁掉外边距": ("BOOLEAN", {"default": True}), + "最小分隔线px": ("INT", {"default": 2, "min": 0, "max": 64, "step": 1}), + "最大输出张数": ("INT", {"default": 16, "min": 1, "max": 144, "step": 1}), + } + } + + RETURN_TYPES = ("IMAGE", "STRING") + RETURN_NAMES = ("切割图像", "检测信息") + FUNCTION = "split_grid" + CATEGORY = "o1key/image" + DESCRIPTION = ( + "智能切割 AI 生成的九宫格、六宫格等合并图。" + "自动模式会检测常见布局;没有明显分隔线时建议手动选择布局。" + ) + + def split_grid( + self, + 图像: torch.Tensor, + 布局: str = "auto", + 自定义行数: int = 3, + 自定义列数: int = 3, + 搜索范围px: int = 32, + 裁掉分隔线: bool = True, + 裁掉外边距: bool = True, + 最小分隔线px: int = 2, + 最大输出张数: int = 16, + ): + source_images = tensor_to_pil(图像) + all_crops: List[Image.Image] = [] + info_lines: List[str] = [] + + for batch_index, image in enumerate(source_images, start=1): + crops, info = _split_one( + image=image, + layout=布局, + custom_rows=自定义行数, + custom_cols=自定义列数, + search_px=搜索范围px, + crop_separators=裁掉分隔线, + trim_outer=裁掉外边距, + min_separator_px=最小分隔线px, + ) + if len(crops) > 最大输出张数: + raise ValueError( + f"合并图切割:检测到 {len(crops)} 张,超过最大输出张数 {最大输出张数}。" + "请调大最大输出张数,或检查布局设置。" + ) + all_crops.extend(crops) + info_lines.append(f"batch {batch_index}: {info}") + + if not all_crops: + raise ValueError("合并图切割:没有生成任何切片。") + + all_crops = _normalize_sizes(all_crops) + print("[o1key 合并图切割] " + " | ".join(info_lines)) + return (pil_to_tensor(all_crops), "\n".join(info_lines)) diff --git a/nodes/nano_banana.py b/nodes/nano_banana.py index 1668ea1..3e16679 100644 --- a/nodes/nano_banana.py +++ b/nodes/nano_banana.py @@ -22,7 +22,7 @@ from PIL import Image from comfy_api.latest import io -from ..utils.image_utils import tensor_to_pil, pil_to_tensor, parse_batch_prompts, encode_image_to_base64, encode_image_to_base64_limited +from ..utils.image_utils import tensor_to_pil, pil_to_tensor, parse_batch_prompts, encode_images_for_image_size_limit from ..utils.config import ( NETWORK_ROUTE_OPTIONS, get_base_url_by_route, @@ -43,6 +43,14 @@ try: except ImportError: PROGRESS_BAR_AVAILABLE = False +try: + from comfy.model_management import processing_interrupted, InterruptProcessingException + INTERRUPT_AVAILABLE = True +except ImportError: + INTERRUPT_AVAILABLE = False + InterruptProcessingException = RuntimeError + processing_interrupted = lambda: False + try: import psutil MEMORY_MONITOR_AVAILABLE = True @@ -54,6 +62,8 @@ REQUEST_LOG_ENABLED = False _NODE = "Nano Banana" _ENDPOINT = "/v1/chat/completions" +_REQUEST_TIMEOUT = 900 +_INTERRUPT_CHECK_INTERVAL = 0.2 _client_instance = None @@ -65,6 +75,43 @@ def _get_client(): return _client_instance +async def _poll_interrupt(): + while True: + await asyncio.sleep(_INTERRUPT_CHECK_INTERVAL) + if INTERRUPT_AVAILABLE and processing_interrupted(): + return + + +async def _run_with_interrupt(coro): + if not INTERRUPT_AVAILABLE: + return await coro + + request_task = asyncio.ensure_future(coro) + interrupt_task = asyncio.ensure_future(_poll_interrupt()) + + done, pending = await asyncio.wait( + [request_task, interrupt_task], + return_when=asyncio.FIRST_COMPLETED, + ) + + for task in pending: + task.cancel() + try: + await task + except (asyncio.CancelledError, Exception): + pass + + if interrupt_task in done and request_task not in done: + raise InterruptProcessingException() + + return request_task.result() + + +def _check_interrupt(): + if INTERRUPT_AVAILABLE and processing_interrupted(): + raise InterruptProcessingException() + + def _images_to_tensor_safe(images: List[Image.Image], node_label: str) -> torch.Tensor: if not images: placeholder = Image.new('RGB', (512, 512), color=(128, 128, 128)) @@ -143,22 +190,6 @@ def _build_request_body( enable_grounding: bool = False, thinking_level: Optional[str] = None, ) -> dict: - content_parts = [{"type": "text", "text": prompt}] - - if images: - for img in images: - b64 = encode_image_to_base64_limited(img, format="PNG") - content_parts.append({ - "type": "image_url", - "image_url": {"url": f"data:image/png;base64,{b64}"} - }) - - body = { - "model": model, - "stream": True, - "messages": [{"role": "user", "content": content_parts}], - } - google_config = { "image_config": { "image_size": resolution, @@ -171,12 +202,32 @@ def _build_request_body( "thinking_level": thinking_level.lower(), "include_thoughts": True, } - body["extra_body"] = {"google": google_config} - if enable_grounding: - body["extra_body"]["google_search"] = True + def _make_body(encoded_images: Optional[List[tuple]] = None) -> dict: + content_parts = [{"type": "text", "text": prompt}] + if encoded_images: + for mime_type, b64 in encoded_images: + content_parts.append({ + "type": "image_url", + "image_url": {"url": f"data:{mime_type};base64,{b64}"} + }) - return body + body = { + "model": model, + "stream": True, + "messages": [{"role": "user", "content": content_parts}], + "extra_body": {"google": google_config}, + } + + if enable_grounding: + body["extra_body"]["google_search"] = True + return body + + encoded_images = None + if images: + encoded_images = encode_images_for_image_size_limit(images) + + return _make_body(encoded_images) async def _generate_single( @@ -208,8 +259,11 @@ async def _generate_single( print(f"[请求] POST {url} | model={model} | extra_body={extra}") last_status = None + resp = None + timeout = aiohttp.ClientTimeout(total=_REQUEST_TIMEOUT, connect=30, sock_read=_REQUEST_TIMEOUT) for attempt in range(DEFAULT_MAX_RETRIES + 1): - resp = await session.post(url, headers=headers, json=body) + _check_interrupt() + resp = await session.post(url, headers=headers, json=body, timeout=timeout) if resp.status == 200: break last_status = resp.status @@ -239,27 +293,36 @@ async def _generate_single( buffer = "" t_request = time.time() t_first_token = None - async for raw_chunk in resp.content.iter_any(): - if t_first_token is None: - t_first_token = time.time() - buffer += raw_chunk.decode("utf-8") - while "\n" in buffer: - line_str, buffer = buffer.split("\n", 1) - line_str = line_str.strip() - if not line_str or not line_str.startswith("data:"): - continue - data_str = line_str[5:].strip() - if data_str == "[DONE]": - break - try: - chunk = json.loads(data_str) - delta = chunk.get("choices", [{}])[0].get("delta", {}) - if "content" in delta: - full_content += delta["content"] - except (json.JSONDecodeError, IndexError): - continue + try: + async for raw_chunk in resp.content.iter_any(): + _check_interrupt() + if t_first_token is None: + t_first_token = time.time() + buffer += raw_chunk.decode("utf-8") + while "\n" in buffer: + line_str, buffer = buffer.split("\n", 1) + line_str = line_str.strip() + if not line_str or not line_str.startswith("data:"): + continue + data_str = line_str[5:].strip() + if data_str == "[DONE]": + break + try: + chunk = json.loads(data_str) + delta = chunk.get("choices", [{}])[0].get("delta", {}) + if "content" in delta: + full_content += delta["content"] + except (json.JSONDecodeError, IndexError): + continue + except aiohttp.ClientPayloadError as e: + if full_content and _IMAGE_RE.search(full_content): + print(f"Nano Banana: 响应流提前结束,但已收到完整图片,继续解析 ({e})") + else: + raise RuntimeError(f"响应流下载中断,请重试或检查网络/代理: {e}") from None + finally: + if resp is not None: + resp.close() t_done = time.time() - resp.close() if not full_content: raise RuntimeError("API 未返回有效内容") @@ -316,6 +379,8 @@ async def _generate_single_task( result["output_images"] = gen_images result["success"] = True result["generated_count"] = len(gen_images) + except InterruptProcessingException: + raise except Exception as e: result["error"] = str(e) return result @@ -352,11 +417,13 @@ async def _process_batch_async( async with aiohttp.ClientSession(connector=connector) as session: for batch_idx in range(num_batches): + _check_interrupt() start_idx = batch_idx * max_concurrent end_idx = min(start_idx + max_concurrent, total_tasks) tasks = [] for i in range(start_idx, end_idx): + _check_interrupt() _, _, prompt = tasks_def[i] task = asyncio.create_task( _generate_single_task( @@ -377,6 +444,7 @@ async def _process_batch_async( batch_results = [] for coro in asyncio.as_completed(tasks): + _check_interrupt() result_data = None try: result = await coro @@ -384,6 +452,11 @@ async def _process_batch_async( result_data = {"success": False, "error": str(result), "generated_count": 0, "output_images": [], "prompt": ""} else: result_data = result + except InterruptProcessingException: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + raise except Exception as e: result_data = {"success": False, "error": str(e), "generated_count": 0, "output_images": [], "prompt": ""} @@ -472,6 +545,7 @@ class NanoBanana(io.ComfyNode): @classmethod def execute(cls, prompt, 模型, 生图数量, 计费, 网络, 谷歌搜索, seed, **kwargs) -> io.NodeOutput: start_time = time.time() + was_interrupted = False model_name = 模型["模型"] 宽高比 = 模型["宽高比"] @@ -533,7 +607,7 @@ class NanoBanana(io.ComfyNode): asyncio.set_event_loop(loop) try: return loop.run_until_complete( - _process_batch_async( + _run_with_interrupt(_process_batch_async( base_url=base_url, api_key=api_key, prompts=prompts, @@ -545,7 +619,7 @@ class NanoBanana(io.ComfyNode): pbar=pbar, enable_grounding=enable_grounding, thinking_level=thinking_level, - ) + )) ) finally: loop.close() @@ -553,9 +627,9 @@ class NanoBanana(io.ComfyNode): with ThreadPoolExecutor(max_workers=1) as executor: future = executor.submit(run_async_in_thread) try: - results = future.result(timeout=900) + results = future.result(timeout=_REQUEST_TIMEOUT) except TimeoutError: - raise RuntimeError("任务执行超时(900秒)") + raise RuntimeError(f"任务执行超时({_REQUEST_TIMEOUT}秒)") success_count = sum(1 for r in results if r.get("success", False)) fail_count = len(results) - success_count @@ -594,13 +668,13 @@ class NanoBanana(io.ComfyNode): enable_grounding=enable_grounding, thinking_level=thinking_level, ) - return loop.run_until_complete(_do()) + return loop.run_until_complete(_run_with_interrupt(_do())) finally: loop.close() with ThreadPoolExecutor(max_workers=1) as executor: future = executor.submit(run_single) - generated_images, first_token_ms, download_ms = future.result(timeout=900) + generated_images, first_token_ms, download_ms = future.result(timeout=_REQUEST_TIMEOUT) if pbar is not None: pbar.update(1) @@ -615,6 +689,10 @@ class NanoBanana(io.ComfyNode): import gc; gc.collect() return io.NodeOutput(output_tensor) + except InterruptProcessingException: + was_interrupted = True + print("Nano Banana: 用户取消") + raise except ValueError as e: if str(e) == "未授权!": print("请联系作者授权后方可使用!") @@ -625,12 +703,13 @@ class NanoBanana(io.ComfyNode): except Exception as e: raise type(e)(str(e)) from None finally: - try: - client = _get_client() - client.base_url = base_url - balance_data = client.query_balance_sync() - balance_info = client.format_balance_info(balance_data) - print(f"Nano Banana: {balance_info}") - except Exception: - pass - import gc; gc.collect() \ No newline at end of file + if not was_interrupted: + try: + client = _get_client() + client.base_url = base_url + balance_data = client.query_balance_sync() + balance_info = client.format_balance_info(balance_data) + print(f"Nano Banana: {balance_info}") + except Exception: + pass + import gc; gc.collect() diff --git a/utils/http_error.py b/utils/http_error.py index 548f379..dfcb11b 100644 --- a/utils/http_error.py +++ b/utils/http_error.py @@ -28,12 +28,15 @@ HTTP_ERROR_MESSAGES = { # 错误内容关键词 → 用户友好文案(优先于状态码匹配) ERROR_CONTENT_MESSAGES = { + "Your request was rejected by the safety system": "请求被安全系统拦截:请调整提示词,避免敏感、违规、血腥、色情、仇恨、未成年人或真实人物等高风险内容。", + "safety system": "请求被安全系统拦截:请调整提示词,避免敏感、违规、血腥、色情、仇恨、未成年人或真实人物等高风险内容。", + "unexpected end of JSON input": "通常重试能解决;反复出现就降低分辨率、数量或换网络线路。", "The current model has a high load": "模型过载,请稍后重试!", "system error": "系统错误,请稍后重试。", } # 可退避重试的状态码 -RETRYABLE_STATUS_CODES = {429, 502, 503, 504} +RETRYABLE_STATUS_CODES = {429, 502, 503, 504, 524} # 退避重试默认参数 DEFAULT_MAX_RETRIES = 3 @@ -44,10 +47,15 @@ DEFAULT_BACKOFF_FACTOR = 2.0 # 指数退避因子 def get_friendly_message(status_code: int, raw_message: str = "") -> str: """根据状态码/错误内容返回友好文案,未匹配则返回原始信息""" + if status_code == 524: + return "Gateway timed out while waiting for upstream image generation. Please retry, lower resolution/count, or switch network route." if raw_message: + raw_message_lower = raw_message.lower() for keyword, friendly_msg in ERROR_CONTENT_MESSAGES.items(): - if keyword in raw_message: + if keyword.lower() in raw_message_lower: return friendly_msg + if status_code == 500: + return "服务器返回 500:上游生成失败或服务端临时异常。请稍后重试;如果多次出现,请降低分辨率/数量,或调整提示词。" friendly = HTTP_ERROR_MESSAGES.get(status_code) if friendly: return friendly @@ -88,7 +96,7 @@ async def async_request_with_retry( """ 带退避重试的 aiohttp 请求。 - 仅对 RETRYABLE_STATUS_CODES (429/503/504) 进行重试。 + 仅对 RETRYABLE_STATUS_CODES (429/502/503/504/524) 进行重试。 超过最大重试次数后抛出友好 RuntimeError。 成功时返回 response 对象(调用者需在 async with 外自行处理 body)。 diff --git a/utils/image_utils.py b/utils/image_utils.py index 5b7f8e7..93742d9 100644 --- a/utils/image_utils.py +++ b/utils/image_utils.py @@ -5,7 +5,8 @@ import base64 from io import BytesIO -from typing import List +import json +from typing import Callable, List, Tuple import numpy as np import torch @@ -110,7 +111,130 @@ def encode_image_to_base64(image: Image.Image, format: str = "PNG") -> str: return base64.b64encode(img_bytes).decode('utf-8') -_MAX_IMAGE_BYTES = 10 * 1024 * 1024 # 10MB base64 上限 +_MAX_REQUEST_BODY_BYTES = 50 * 1024 * 1024 # 50MB 请求体上限 + + +def _encode_image_to_base64_with_quality(image: Image.Image, quality: int) -> str: + buffered = BytesIO() + working = image + if working.mode != 'RGB': + working = working.convert('RGB') + + working.save( + buffered, + format="JPEG", + quality=quality, + optimize=True, + subsampling=2, + ) + return base64.b64encode(buffered.getvalue()).decode('utf-8') + + +def encode_images_for_request_body_limit( + images: List[Image.Image], + build_body: Callable[[List[Tuple[str, str]]], dict], + max_body_bytes: int = _MAX_REQUEST_BODY_BYTES, +) -> List[Tuple[str, str]]: + """ + 为请求体编码图片,并保证完整 JSON 请求体不超过 max_body_bytes。 + + 策略: + - 先按原始 PNG 编码估算完整请求体; + - 若超过限制,改用 JPEG 质量压缩,逐步降低 quality; + - 全程不缩放图片尺寸。 + + Returns: + [(mime_type, base64), ...] + """ + encoded = [("image/png", encode_image_to_base64(img, format="PNG")) for img in images] + body_size = len(json.dumps(build_body(encoded)).encode("utf-8")) + if body_size <= max_body_bytes: + return encoded + + for quality in [95, 90, 85, 80, 75, 70, 65, 60, 55, 50, 45, 40, 35, 30, 25, 20, 15, 10, 5, 1]: + encoded = [ + ("image/jpeg", _encode_image_to_base64_with_quality(img, quality)) + for img in images + ] + body_size = len(json.dumps(build_body(encoded)).encode("utf-8")) + if body_size <= max_body_bytes: + print( + f"输入图片已通过 JPEG 质量压缩控制请求体积: " + f"quality={quality}, 请求体积={body_size / 1024 / 1024:.2f}MB " + f"(限制 {max_body_bytes / 1024 / 1024:.0f}MB)" + ) + return encoded + + raise ValueError( + f"请求体超过 {max_body_bytes / 1024 / 1024:.0f}MB," + "即使压缩到最低图片质量仍无法满足限制;请减少参考图数量或输入图片内容复杂度" + ) + + +_MAX_IMAGE_BYTES = 10 * 1024 * 1024 # 10MB 单张图片上限 + + +def _encode_image_to_bytes(image: Image.Image, format: str = "PNG", quality: int = None) -> bytes: + buffered = BytesIO() + working = image + + if format.upper() == "JPEG" and working.mode != 'RGB': + working = working.convert('RGB') + elif working.mode == 'RGBA': + working = working.convert('RGB') + + save_kwargs = {"format": format} + if quality is not None: + save_kwargs.update({ + "quality": quality, + "optimize": True, + "subsampling": 2, + }) + + working.save(buffered, **save_kwargs) + return buffered.getvalue() + + +def encode_images_for_image_size_limit( + images: List[Image.Image], + max_image_bytes: int = _MAX_IMAGE_BYTES, +) -> List[Tuple[str, str]]: + """ + 将图片编码为 base64,并保证每张编码前的图片文件体积不超过 max_image_bytes。 + + 策略: + - 先尝试 PNG 原图尺寸编码; + - 单张超过限制时,改用 JPEG 质量压缩; + - 全程不缩放图片尺寸。 + + Returns: + [(mime_type, base64), ...] + """ + encoded = [] + + for idx, img in enumerate(images, start=1): + png_bytes = _encode_image_to_bytes(img, format="PNG") + if len(png_bytes) <= max_image_bytes: + encoded.append(("image/png", base64.b64encode(png_bytes).decode('utf-8'))) + continue + + for quality in [95, 90, 85, 80, 75, 70, 65, 60, 55, 50, 45, 40, 35, 30, 25, 20, 15, 10, 5, 1]: + jpg_bytes = _encode_image_to_bytes(img, format="JPEG", quality=quality) + if len(jpg_bytes) <= max_image_bytes: + print( + f"输入图片 {idx} 已通过 JPEG 质量压缩控制单图体积: " + f"quality={quality}, 图片体积={len(jpg_bytes) / 1024 / 1024:.2f}MB " + f"(限制 {max_image_bytes / 1024 / 1024:.0f}MB),尺寸保持 {img.width}x{img.height}" + ) + encoded.append(("image/jpeg", base64.b64encode(jpg_bytes).decode('utf-8'))) + break + else: + raise ValueError( + f"输入图片 {idx} 超过 {max_image_bytes / 1024 / 1024:.0f}MB," + "即使压缩到最低图片质量仍无法满足限制;请减少图片内容复杂度或手动处理图片" + ) + + return encoded def encode_image_to_base64_limited( @@ -241,4 +365,4 @@ def parse_batch_prompts(prompt: str) -> List[str]: if not filtered_prompts: raise ValueError("批量提示词模式下,所有提示词都为空,请至少提供一个有效的提示词") - return filtered_prompts \ No newline at end of file + return filtered_prompts diff --git a/utils/rembg_utils.py b/utils/rembg_utils.py index 783634d..64bf9cf 100644 --- a/utils/rembg_utils.py +++ b/utils/rembg_utils.py @@ -3,6 +3,7 @@ 基于 rembg 库实现,支持 CPU 推理 """ +import os import numpy as np import torch from PIL import Image @@ -14,6 +15,13 @@ def _get_session(): """懒加载 rembg session,避免启动时加载模型""" global _session if _session is None: + try: + import folder_paths + models_dir = os.path.join(folder_paths.models_dir, "rembg") + os.makedirs(models_dir, exist_ok=True) + os.environ["U2NET_HOME"] = models_dir + except Exception: + pass try: from rembg import new_session _session = new_session("isnet-general-use") diff --git a/utils/video_task.py b/utils/video_task.py new file mode 100644 index 0000000..60f175b --- /dev/null +++ b/utils/video_task.py @@ -0,0 +1,181 @@ +import asyncio +from typing import Any, Dict + + +try: + from comfy.model_management import processing_interrupted, InterruptProcessingException + INTERRUPT_AVAILABLE = True +except Exception: + INTERRUPT_AVAILABLE = False + processing_interrupted = lambda: False + + class InterruptProcessingException(Exception): + pass + + +SUCCESS_STATUSES = {"succeed", "succeeded", "success", "completed", "done", "finished"} +FAILURE_STATUSES = { + "fail", + "failed", + "failure", + "error", + "expired", + "timeout", + "timed_out", + "cancel", + "canceled", + "cancelled", + "rejected", +} + + +def check_interrupt() -> None: + if INTERRUPT_AVAILABLE and processing_interrupted(): + raise InterruptProcessingException() + + +async def interruptible_sleep(seconds: float, step: float = 0.2) -> None: + elapsed = 0.0 + while elapsed < seconds: + check_interrupt() + delay = min(step, seconds - elapsed) + await asyncio.sleep(delay) + elapsed += delay + check_interrupt() + + +async def run_with_interrupt(coro, step: float = 0.2): + task = asyncio.ensure_future(coro) + try: + while not task.done(): + check_interrupt() + await asyncio.wait({task}, timeout=step) + check_interrupt() + return await task + except InterruptProcessingException: + task.cancel() + try: + await task + except BaseException: + pass + raise + + +def _as_dict(value: Any) -> Dict[str, Any]: + return value if isinstance(value, dict) else {} + + +def _nested_payloads(payload: Dict[str, Any]): + root = _as_dict(payload) + data = _as_dict(root.get("data")) + inner = _as_dict(data.get("data")) + return root, data, inner + + +def extract_status(payload: Dict[str, Any]) -> str: + root, data, inner = _nested_payloads(payload) + keys = ("status", "task_status", "state", "task_state") + statuses = [] + for source in (data, inner, root): + for key in keys: + value = source.get(key) + if value is not None and str(value).strip(): + statuses.append(str(value).strip().lower()) + for status in statuses: + if status in FAILURE_STATUSES or any( + token in status for token in ("fail", "error", "reject", "timeout", "cancel") + ): + return status + for status in statuses: + if status in SUCCESS_STATUSES: + return status + return statuses[0] if statuses else "" + + +def extract_progress(payload: Dict[str, Any]) -> int: + root, data, inner = _nested_payloads(payload) + for source in (data, inner, root): + value = source.get("progress") + if value is None: + continue + try: + return max(0, min(100, int(float(str(value).strip().rstrip("%"))))) + except (TypeError, ValueError): + return 0 + return 0 + + +def extract_error_message(payload: Dict[str, Any], default: str = "未知错误") -> str: + root, data, inner = _nested_payloads(payload) + keys = ( + "fail_reason", + "failure_reason", + "task_status_msg", + "status_msg", + "error_message", + "message", + "msg", + "reason", + "detail", + "details", + ) + for source in (data, inner, root): + error = source.get("error") + if isinstance(error, dict): + for key in ("message", "msg", "detail", "reason", "code"): + value = error.get(key) + if value: + return str(value) + elif error: + return str(error) + + for key in keys: + value = source.get(key) + if value: + return str(value) + return default + + +def extract_video_url(payload: Dict[str, Any]) -> str | None: + root, data, inner = _nested_payloads(payload) + for source in (data, inner, root): + for key in ("video_url", "result_url", "url", "download_url"): + value = source.get(key) + if value: + return str(value) + + result = _as_dict(source.get("result")) + for key in ("video_url", "result_url", "url", "download_url"): + value = result.get(key) + if value: + return str(value) + + content = _as_dict(source.get("content")) + value = content.get("video_url") or content.get("url") + if value: + return str(value) + + task_result = _as_dict(source.get("task_result")) + videos = task_result.get("videos") + if isinstance(videos, list) and videos: + first = _as_dict(videos[0]) + value = first.get("url") or first.get("video_url") + if value: + return str(value) + return None + + +def is_success_status(status: str) -> bool: + return status in SUCCESS_STATUSES + + +def is_failure_status(status: str, payload: Dict[str, Any] | None = None) -> bool: + if status in FAILURE_STATUSES: + return True + if any(token in status for token in ("fail", "error", "reject", "timeout", "cancel")): + return True + if payload is None: + return False + root, data, inner = _nested_payloads(payload) + failure_keys = ("error", "fail_reason", "failure_reason", "task_status_msg", "error_message") + return any(any(source.get(key) for key in failure_keys) for source in (data, inner, root)) diff --git a/web/js/notePanel.js b/web/js/notePanel.js new file mode 100644 index 0000000..6a0e58a --- /dev/null +++ b/web/js/notePanel.js @@ -0,0 +1,701 @@ +import { app } from "../../../scripts/app.js"; + +const STORAGE_KEY = "o1key-notes"; +const SEEDED_KEY = "o1key-notes-seeded-v2"; +const STYLE_ID = "o1key-notes-styles"; + +let notes = []; +let searchText = ""; +let activeFilter = "all"; +let noteContainer = null; +let editingNoteId = null; +let draftNote = null; +let pendingDelete = false; + +const SAMPLE_NOTES = [ + { + title: "产品主图:高级玻璃质感", + tags: ["产品图", "玻璃", "灯光"], + content: `Clean studio lighting, translucent glass material, subtle caustics, soft shadow, premium product photography, 85mm lens, minimal background, high detail. + +使用方式: +1. 把产品图作为参考图输入 +2. 保留主体轮廓,只调整材质和灯光 +3. 如果玻璃过亮,降低 "caustics" 权重` + }, + { + title: "Nano Banana 参考图经验", + tags: ["Nano Banana", "参考图"], + content: `参考图越多越容易跑偏,主体一致性优先用 1-3 张图。 + +复杂场景建议分两步: +1. 先生成主体和构图 +2. 再用局部或参考图做材质、背景、文字等精修 + +如果提示词和参考图冲突,模型通常会优先参考图。` + }, + { + title: "电商模特换装模板", + tags: ["电商", "模特", "换装"], + content: `Keep face identity and original pose. Replace the outfit with: [服装描述]. +Realistic fabric texture, natural folds, accurate seams, studio e-commerce photography, clean background, consistent lighting. + +Avoid changing body shape, face, hairstyle, camera angle, or hand position.` + }, + { + title: "常用负向词", + tags: ["负向词", "通用"], + content: "low quality, blurry, deformed hands, extra fingers, bad anatomy, distorted text, watermark, logo, oversaturated, plastic skin, broken geometry" + }, + { + title: "图片批量命名经验", + tags: ["批量", "工作流"], + content: `批量生图前先确认保存路径和命名规则。 + +推荐流程: +1. 小批量跑 2-3 张确认风格 +2. 固定提示词和参考图 +3. 再放大批量数量 + +这样失败成本最低,也更容易定位是哪一环导致跑偏。` + } +]; + +const CSS = ` +#o1key-notes-root{position:absolute;inset:0;display:flex;flex-direction:column;min-height:0;color:#ddd;background:var(--comfy-menu-bg,#202020);overflow:hidden;font-family:inherit;--o1n-blue:#4f8cff;--o1n-blue-2:#7eb8f7;--o1n-blue-soft:rgba(79,140,255,.14)} +#o1key-notes-header{padding:12px 14px 10px;border-bottom:1px solid rgba(255,255,255,.06);flex-shrink:0} +.o1n-title-row{display:flex;align-items:center;justify-content:space-between;margin-bottom:10px} +.o1n-title{font-size:14px;font-weight:700;color:#eee;letter-spacing:.2px} +.o1n-head-actions{display:flex;gap:5px} +.o1n-icon-btn{width:28px;height:28px;border:1px solid rgba(255,255,255,.09);border-radius:7px;background:rgba(255,255,255,.035);color:#888;display:flex;align-items:center;justify-content:center;cursor:pointer;transition:all .15s;flex-shrink:0} +.o1n-icon-btn:hover{color:#ddd;background:rgba(255,255,255,.08);border-color:rgba(255,255,255,.16)} +.o1n-icon-btn.primary{background:var(--o1n-blue);color:#fff;border-color:transparent} +.o1n-icon-btn.primary:hover{background:#6ca0ff;color:#fff} +#o1n-search-box{height:34px;border:1px solid rgba(255,255,255,.1);border-radius:8px;background:rgba(255,255,255,.045);display:flex;align-items:center;gap:8px;padding:0 10px;transition:border-color .15s} +#o1n-search-box:focus-within{border-color:rgba(79,140,255,.42)} +#o1n-search{flex:1;background:transparent;border:0;outline:0;color:#ddd;font-size:12px;min-width:0} +#o1n-search::placeholder{color:#666} +#o1n-panel-status{height:16px;margin-top:7px;color:#777;font-size:11px;line-height:16px;white-space:nowrap;overflow:hidden;text-overflow:ellipsis} +.o1n-tag-strip{display:flex;gap:6px;margin-top:10px;overflow-x:auto;padding-bottom:1px} +.o1n-tag-strip::-webkit-scrollbar{display:none} +.o1n-tag-filter{height:26px;padding:0 9px;border-radius:7px;border:1px solid rgba(255,255,255,.09);background:transparent;color:#929292;font-size:12px;display:flex;align-items:center;gap:5px;cursor:pointer;white-space:nowrap;transition:all .15s} +.o1n-tag-filter:hover{color:#cfcfcf;border-color:rgba(255,255,255,.16)} +.o1n-tag-filter.active{color:var(--o1n-blue-2);border-color:rgba(79,140,255,.36);background:var(--o1n-blue-soft)} +#o1key-notes-list{flex:1;min-height:0;overflow:auto;padding:8px 10px 12px;display:flex;flex-direction:column;gap:6px} +#o1key-notes-list::-webkit-scrollbar,#o1n-content::-webkit-scrollbar{width:4px} +#o1key-notes-list::-webkit-scrollbar-thumb,#o1n-content::-webkit-scrollbar-thumb{background:rgba(255,255,255,.1);border-radius:2px} +.o1n-item{border:1px solid transparent;border-radius:8px;padding:10px;background:transparent;cursor:pointer;transition:all .12s} +.o1n-item:hover{background:rgba(255,255,255,.04);border-color:rgba(255,255,255,.07)} +.o1n-item.active{background:rgba(255,255,255,.065);border-color:rgba(255,255,255,.12);box-shadow:inset 2px 0 0 var(--o1n-blue)} +.o1n-item-top{display:flex;align-items:center;gap:8px;margin-bottom:6px} +.o1n-item-title{font-size:13px;color:#e0e0e0;white-space:nowrap;overflow:hidden;text-overflow:ellipsis;flex:1;min-width:0} +.o1n-item-edit{font-size:11px;color:#777;flex-shrink:0} +.o1n-item-text{font-size:12px;line-height:1.45;color:#8b8b8b;display:-webkit-box;-webkit-line-clamp:3;-webkit-box-orient:vertical;overflow:hidden;word-break:break-word} +.o1n-item-meta{margin-top:8px;display:flex;align-items:center;justify-content:space-between;gap:8px;color:#666;font-size:10px} +.o1n-item-tags{display:flex;gap:4px;overflow:hidden;min-width:0} +.o1n-tag{color:var(--o1n-blue-2);background:rgba(79,140,255,.1);border:1px solid rgba(79,140,255,.2);border-radius:5px;padding:2px 5px;white-space:nowrap;max-width:120px;overflow:hidden;text-overflow:ellipsis} +.o1n-empty{height:100%;display:flex;align-items:center;justify-content:center;text-align:center;color:#666;font-size:13px;line-height:1.6;padding:20px} +.o1n-editor-backdrop{position:absolute;inset:0;background:rgba(0,0,0,.34);display:flex;align-items:stretch;justify-content:flex-end;z-index:30;animation:o1n-fade .12s ease} +.o1n-editor{width:100%;height:100%;background:var(--comfy-menu-bg,#202020);border-left:1px solid rgba(79,140,255,.28);box-shadow:0 22px 70px rgba(0,0,0,.46);display:flex;flex-direction:column;min-height:0} +.o1n-editor-top{height:48px;padding:0 14px;border-bottom:1px solid rgba(255,255,255,.08);display:flex;align-items:center;justify-content:space-between;flex-shrink:0} +.o1n-editor-title{font-size:13px;font-weight:700;color:#eee} +.o1n-editor-body{padding:12px;display:flex;flex-direction:column;gap:9px;flex:1;min-height:0} +#o1n-title-input,#o1n-content{width:100%;border:1px solid rgba(255,255,255,.1);border-radius:8px;background:rgba(255,255,255,.045);color:#ddd;outline:0;font-family:inherit;transition:border-color .15s} +#o1n-title-input{height:34px;padding:0 10px;font-size:13px;font-weight:700;flex-shrink:0} +#o1n-content{flex:1;min-height:160px;resize:none;padding:11px 12px;font-size:13px;line-height:1.58} +#o1n-title-input:focus,#o1n-content:focus{border-color:rgba(79,140,255,.42)} +.o1n-tag-editor{min-height:74px;border:1px solid rgba(255,255,255,.1);border-radius:8px;background:rgba(255,255,255,.035);padding:7px;flex-shrink:0} +.o1n-tag-row{display:flex;align-items:center;gap:6px;flex-wrap:wrap} +.o1n-tag-token{height:24px;padding:0 7px;border-radius:6px;border:1px solid rgba(79,140,255,.25);background:rgba(79,140,255,.1);color:var(--o1n-blue-2);display:inline-flex;align-items:center;gap:6px;font-size:12px} +.o1n-tag-remove{color:#8baee8;font-size:13px;line-height:1;cursor:pointer} +.o1n-tag-input{height:24px;min-width:104px;flex:1;border:1px dashed rgba(255,255,255,.14);border-radius:6px;background:transparent;color:#ddd;padding:0 7px;font-size:12px;outline:0} +.o1n-tag-input::placeholder{color:#777} +.o1n-tag-hint{margin-top:7px;color:#666;font-size:11px;line-height:1.35} +.o1n-editor-actions{display:grid;grid-template-columns:1fr 1fr 1fr;gap:7px;padding:12px;border-top:1px solid rgba(255,255,255,.08);flex-shrink:0} +.o1n-action{height:32px;border:1px solid rgba(255,255,255,.1);border-radius:7px;background:rgba(255,255,255,.035);color:#aaa;font-size:12px;display:flex;align-items:center;justify-content:center;gap:6px;cursor:pointer;transition:all .15s} +.o1n-action:hover{color:#e5e5e5;background:rgba(255,255,255,.075);border-color:rgba(255,255,255,.16)} +.o1n-action.primary{background:var(--o1n-blue);color:#fff;border-color:transparent;font-weight:700} +.o1n-action.primary:hover{background:#6ca0ff;color:#fff} +.o1n-action.accent{background:#d8b45b;color:#171717;border-color:transparent;font-weight:700} +.o1n-action.accent:hover{background:#e3c474;color:#111} +.o1n-action.danger:hover{background:rgba(215,101,101,.18);border-color:rgba(215,101,101,.32);color:#f0a0a0} +.o1n-delete-confirm{grid-column:1 / -1;display:none;align-items:center;gap:7px;padding:7px 8px;border:1px solid rgba(215,101,101,.24);border-radius:7px;background:rgba(215,101,101,.08);color:#ccc;font-size:12px} +.o1n-delete-confirm.show{display:flex} +.o1n-delete-confirm span{flex:1;min-width:0} +.o1n-mini-btn{height:24px;border:1px solid rgba(255,255,255,.1);border-radius:6px;background:rgba(255,255,255,.06);color:#bbb;padding:0 8px;font-size:11px;cursor:pointer} +.o1n-mini-btn.danger{background:rgba(215,101,101,.5);border-color:rgba(215,101,101,.3);color:#fff} +#o1n-status{grid-column:1 / -1;height:16px;color:#777;font-size:11px;line-height:16px;white-space:nowrap;overflow:hidden;text-overflow:ellipsis} +#o1n-import-file{display:none} +@keyframes o1n-fade{from{opacity:0}to{opacity:1}} +`; + +function genId() { + return Date.now().toString(36) + Math.random().toString(36).slice(2, 8); +} + +function now() { + return Date.now(); +} + +function parseTags(value) { + if (Array.isArray(value)) return value.map(String).map(s => s.trim()).filter(Boolean); + return String(value || "") + .split(/[,,、\s]+/) + .map(s => s.trim()) + .filter(Boolean); +} + +function makeNote(data = {}) { + const ts = now(); + return { + id: data.id || genId(), + title: String(data.title || "未命名笔记"), + tags: parseTags(data.tags || data.category || ""), + content: String(data.content || ""), + createdAt: data.createdAt || ts, + updatedAt: data.updatedAt || ts, + }; +} + +function cloneNote(note) { + return { + ...note, + tags: [...(note.tags || [])], + }; +} + +function formatDate(ts) { + const d = new Date(ts || now()); + const today = new Date(); + const pad = n => String(n).padStart(2, "0"); + if (d.toDateString() === today.toDateString()) { + return `今天 ${pad(d.getHours())}:${pad(d.getMinutes())}`; + } + return `${pad(d.getMonth() + 1)}/${pad(d.getDate())}`; +} + +function escapeHtml(value) { + return String(value ?? "") + .replace(/&/g, "&") + .replace(//g, ">") + .replace(/"/g, """) + .replace(/'/g, "'"); +} + +function icon(name, size = 14) { + const icons = { + plus: ``, + search: ``, + copy: ``, + insert: ``, + save: ``, + trash: ``, + download: ``, + upload: ``, + node: ``, + close: ``, + }; + return `${icons[name] || icons.plus}`; +} + +function injectStyles() { + if (document.getElementById(STYLE_ID)) return; + const el = document.createElement("style"); + el.id = STYLE_ID; + el.textContent = CSS; + document.head.appendChild(el); +} + +function loadNotes() { + try { + const raw = localStorage.getItem(STORAGE_KEY); + notes = raw ? JSON.parse(raw).map(makeNote) : []; + } catch { + notes = []; + } + + if (!localStorage.getItem(SEEDED_KEY)) { + const existingTitles = new Set(notes.map(n => n.title)); + const samples = SAMPLE_NOTES.map(makeNote).filter(n => !existingTitles.has(n.title)); + notes = [...samples, ...notes]; + localStorage.setItem(SEEDED_KEY, "1"); + saveNotes(); + } +} + +function saveNotes() { + try { + localStorage.setItem(STORAGE_KEY, JSON.stringify(notes)); + } catch {} +} + +function allTags() { + const counts = new Map(); + for (const note of notes) { + for (const tag of note.tags || []) counts.set(tag, (counts.get(tag) || 0) + 1); + } + return [...counts.entries()].sort((a, b) => b[1] - a[1]).slice(0, 12); +} + +function filteredNotes() { + const q = searchText.trim().toLowerCase(); + return notes.filter(note => { + if (activeFilter.startsWith("tag:") && !(note.tags || []).includes(activeFilter.slice(4))) return false; + if (!q) return true; + return [note.title, note.content, ...(note.tags || [])].join("\n").toLowerCase().includes(q); + }).sort((a, b) => (b.updatedAt || 0) - (a.updatedAt || 0)); +} + +app.registerExtension({ + name: "o1key.notePanel", + async setup() { + loadNotes(); + app.extensionManager.registerSidebarTab({ + id: "o1key-notes", + title: "笔记", + icon: "pi pi-pencil", + type: "custom", + render: (container) => { + injectStyles(); + renderNotePanel(container); + }, + }); + }, +}); + +function renderNotePanel(container) { + container.innerHTML = ""; + container.style.position = "relative"; + container.style.height = "100%"; + container.style.overflow = "hidden"; + + const root = document.createElement("div"); + root.id = "o1key-notes-root"; + root.innerHTML = ` +
+
+
笔记
+
+ + + +
+
+ +
+
+
+
+ + `; + container.appendChild(root); + noteContainer = root; + bindPanelEvents(root); + renderTagFilters(); + renderList(); + renderEditor(); +} + +function bindPanelEvents(root) { + root.querySelector("#o1n-new").addEventListener("click", createNote); + root.querySelector("#o1n-export").addEventListener("click", exportNotes); + root.querySelector("#o1n-import").addEventListener("click", () => root.querySelector("#o1n-import-file").click()); + root.querySelector("#o1n-import-file").addEventListener("change", importNotes); + root.querySelector("#o1n-search").addEventListener("input", (e) => { + searchText = e.target.value; + renderList(); + }); +} + +function renderTagFilters() { + const row = noteContainer?.querySelector("#o1n-tag-strip"); + if (!row) return; + const tags = allTags(); + row.innerHTML = `` + + tags.map(([tag, count]) => { + const filter = `tag:${tag}`; + return ``; + }).join(""); + + row.querySelectorAll("[data-filter]").forEach(btn => { + btn.addEventListener("click", () => { + activeFilter = btn.dataset.filter; + renderTagFilters(); + renderList(); + }); + }); +} + +function renderList() { + const list = noteContainer.querySelector("#o1key-notes-list"); + const visible = filteredNotes(); + + if (!visible.length) { + list.innerHTML = `
没有匹配的笔记
`; + return; + } + + list.innerHTML = visible.map(note => { + const tags = (note.tags || []).slice(0, 3).map(tag => `#${escapeHtml(tag)}`).join(""); + return `
+
+
${escapeHtml(note.title || "未命名笔记")}
+
编辑
+
+
${escapeHtml(note.content || "空笔记")}
+
${tags || `#未分类`}
${formatDate(note.updatedAt)}
+
`; + }).join(""); + + list.querySelectorAll(".o1n-item").forEach(item => { + item.addEventListener("click", () => openEditor(item.dataset.id)); + }); +} + +function renderEditor() { + noteContainer.querySelector(".o1n-editor-backdrop")?.remove(); + if (!draftNote) return; + pendingDelete = false; + + const overlay = document.createElement("div"); + overlay.className = "o1n-editor-backdrop"; + overlay.innerHTML = ` +
+
+
${editingNoteId ? "编辑笔记" : "新建笔记"}
+ +
+
+ +
+
+ ${(draftNote.tags || []).map(tag => `${escapeHtml(tag)} ×`).join("")} + +
+
输入标签后按 Enter,或用逗号分隔多个标签。
+
+ +
+
+ + + + + + +
+ 确定删除这条笔记? + + +
+
+
+
+ `; + noteContainer.appendChild(overlay); + bindEditorEvents(overlay); +} + +function bindEditorEvents(overlay) { + overlay.querySelector("#o1n-close-editor").addEventListener("click", closeEditor); + overlay.querySelector("#o1n-cancel").addEventListener("click", closeEditor); + overlay.querySelector("#o1n-title-input").addEventListener("input", updateDraftFromEditor); + overlay.querySelector("#o1n-content").addEventListener("input", updateDraftFromEditor); + overlay.querySelectorAll(".o1n-tag-remove").forEach(btn => { + btn.addEventListener("click", () => removeDraftTag(btn.dataset.tag)); + }); + overlay.querySelector("#o1n-tag-input").addEventListener("keydown", handleTagInputKeydown); + overlay.querySelector("#o1n-tag-input").addEventListener("blur", commitTagInput); + overlay.querySelector("#o1n-copy").addEventListener("click", copyDraftContent); + overlay.querySelector("#o1n-insert").addEventListener("click", insertDraftContent); + overlay.querySelector("#o1n-save-from-node").addEventListener("click", saveFromCurrentNode); + overlay.querySelector("#o1n-delete").addEventListener("click", requestDeleteEditingNote); + overlay.querySelector("#o1n-delete-yes").addEventListener("click", deleteEditingNote); + overlay.querySelector("#o1n-delete-no").addEventListener("click", hideDeleteConfirm); + overlay.querySelector("#o1n-save").addEventListener("click", saveDraft); +} + +function openEditor(id) { + const note = notes.find(n => n.id === id); + if (!note) return; + editingNoteId = id; + draftNote = cloneNote(note); + renderList(); + renderEditor(); +} + +function closeEditor() { + editingNoteId = null; + draftNote = null; + pendingDelete = false; + renderList(); + renderEditor(); +} + +function createNote() { + editingNoteId = null; + draftNote = makeNote({ title: "新笔记", tags: ["未分类"], content: "" }); + renderList(); + renderEditor(); + noteContainer.querySelector("#o1n-title-input")?.focus(); +} + +function updateDraftFromEditor() { + if (!draftNote) return; + draftNote.title = noteContainer.querySelector("#o1n-title-input")?.value.trim() || "未命名笔记"; + draftNote.content = noteContainer.querySelector("#o1n-content")?.value || ""; + draftNote.updatedAt = now(); +} + +function addTagsFromText(value) { + if (!draftNote) return false; + const incoming = parseTags(value); + if (!incoming.length) return false; + const existing = new Set(draftNote.tags || []); + for (const tag of incoming) existing.add(tag); + draftNote.tags = [...existing]; + draftNote.updatedAt = now(); + renderDraftTags(); + return true; +} + +function commitTagInput() { + const input = noteContainer?.querySelector("#o1n-tag-input"); + if (!input) return; + if (addTagsFromText(input.value)) input.value = ""; +} + +function handleTagInputKeydown(e) { + if (e.key !== "Enter" && e.key !== "," && e.key !== ",") return; + e.preventDefault(); + commitTagInput(); +} + +function removeDraftTag(tag) { + if (!draftNote) return; + draftNote.tags = (draftNote.tags || []).filter(t => t !== tag); + draftNote.updatedAt = now(); + renderDraftTags(); +} + +function renderDraftTags() { + const row = noteContainer?.querySelector(".o1n-tag-row"); + if (!row || !draftNote) return; + row.innerHTML = ` + ${(draftNote.tags || []).map(tag => `${escapeHtml(tag)} ×`).join("")} + + `; + row.querySelectorAll(".o1n-tag-remove").forEach(btn => { + btn.addEventListener("click", () => removeDraftTag(btn.dataset.tag)); + }); + const input = row.querySelector("#o1n-tag-input"); + input.addEventListener("keydown", handleTagInputKeydown); + input.addEventListener("blur", commitTagInput); + input.focus(); +} + +function saveDraft() { + if (!draftNote) return; + updateDraftFromEditor(); + draftNote.tags = draftNote.tags?.length ? draftNote.tags : ["未分类"]; + if (editingNoteId) { + const idx = notes.findIndex(n => n.id === editingNoteId); + if (idx >= 0) notes[idx] = cloneNote(draftNote); + } else { + draftNote.id = genId(); + draftNote.createdAt = now(); + draftNote.updatedAt = now(); + notes.unshift(cloneNote(draftNote)); + editingNoteId = draftNote.id; + } + saveNotes(); + renderTagFilters(); + closeEditor(); +} + +function requestDeleteEditingNote() { + if (!editingNoteId) { + closeEditor(); + return; + } + pendingDelete = true; + const box = noteContainer?.querySelector("#o1n-delete-confirm"); + box?.classList.add("show"); + setStatus(""); +} + +function hideDeleteConfirm() { + pendingDelete = false; + noteContainer?.querySelector("#o1n-delete-confirm")?.classList.remove("show"); +} + +function deleteEditingNote() { + if (!editingNoteId) { + closeEditor(); + return; + } + const note = notes.find(n => n.id === editingNoteId); + if (!note) return; + notes = notes.filter(n => n.id !== editingNoteId); + saveNotes(); + renderTagFilters(); + closeEditor(); +} + +async function copyDraftContent() { + updateDraftFromEditor(); + try { + await navigator.clipboard.writeText(draftNote?.content || ""); + setStatus("已复制"); + } catch { + setStatus("复制失败,请手动选择内容"); + } +} + +function insertDraftContent() { + updateDraftFromEditor(); + const text = draftNote?.content || ""; + if (!text.trim()) { + setStatus("当前笔记内容为空"); + return; + } + if (insertIntoFocusedInput(text) || insertIntoSelectedNode(text)) { + setStatus("已插入"); + return; + } + navigator.clipboard?.writeText(text).catch(() => {}); + setStatus("未找到可插入位置,已复制内容"); +} + +function setStatus(message) { + const el = noteContainer?.querySelector("#o1n-status"); + if (!el) return; + el.textContent = message; + clearTimeout(setStatus._timer); + setStatus._timer = setTimeout(() => { + if (el.textContent === message) el.textContent = ""; + }, 2200); +} + +function setPanelStatus(message) { + const el = noteContainer?.querySelector("#o1n-panel-status"); + if (!el) return; + el.textContent = message; + clearTimeout(setPanelStatus._timer); + setPanelStatus._timer = setTimeout(() => { + if (el.textContent === message) el.textContent = ""; + }, 2600); +} + +function insertIntoFocusedInput(text) { + const el = document.activeElement; + if (!el || noteContainer.contains(el)) return false; + if (!(el instanceof HTMLTextAreaElement || el instanceof HTMLInputElement)) return false; + const start = el.selectionStart ?? el.value.length; + const end = el.selectionEnd ?? el.value.length; + const before = el.value.slice(0, start); + const after = el.value.slice(end); + const spacer = before && !before.endsWith("\n") ? "\n" : ""; + el.value = before + spacer + text + after; + const pos = (before + spacer + text).length; + el.setSelectionRange(pos, pos); + el.dispatchEvent(new Event("input", { bubbles: true })); + el.dispatchEvent(new Event("change", { bubbles: true })); + return true; +} + +function findSelectedNode() { + const selectedNodes = app.canvas?.selected_nodes; + if (selectedNodes) { + const values = Array.isArray(selectedNodes) ? selectedNodes : Object.values(selectedNodes); + if (values.length) return values[0]; + } + return app.canvas?.selected_node || null; +} + +function isPromptWidget(widget) { + const name = String(widget?.name || "").toLowerCase(); + const value = widget?.value; + if (typeof value !== "string") return false; + return ( + name.includes("prompt") || + name.includes("提示词") || + name.includes("正向") || + name === "text" || + name === "文本" + ); +} + +function insertIntoSelectedNode(text) { + const node = findSelectedNode(); + if (!node?.widgets?.length) return false; + const widget = node.widgets.find(isPromptWidget) || node.widgets.find(w => typeof w.value === "string"); + if (!widget) return false; + const current = widget.value || ""; + const spacer = current && !current.endsWith("\n") ? "\n" : ""; + widget.value = current + spacer + text; + widget.callback?.(widget.value, app.canvas, node, widget); + app.graph?.setDirtyCanvas?.(true, true); + return true; +} + +function getPromptWidgetsFromSelectedNode() { + const node = findSelectedNode(); + if (!node?.widgets?.length) return []; + return node.widgets.filter(w => typeof w.value === "string" && String(w.value).trim()); +} + +function saveFromCurrentNode() { + const widgets = getPromptWidgetsFromSelectedNode(); + if (!widgets.length) { + setStatus("当前没有选中包含文本的节点"); + return; + } + + const preferred = widgets.find(isPromptWidget) || widgets[0]; + const node = findSelectedNode(); + draftNote = makeNote({ + title: `${node?.title || node?.type || "节点"}:${preferred.name || "提示词"}`, + tags: [node?.title || node?.type || "节点"], + content: preferred.value, + }); + editingNoteId = null; + renderEditor(); + setStatus("已读取当前节点内容,保存后写入笔记"); +} + +function exportNotes() { + const blob = new Blob([JSON.stringify(notes, null, 2)], { type: "application/json" }); + const url = URL.createObjectURL(blob); + const a = document.createElement("a"); + a.href = url; + a.download = `o1key-notes-${new Date().toISOString().slice(0, 10)}.json`; + document.body.appendChild(a); + a.click(); + a.remove(); + URL.revokeObjectURL(url); +} + +async function importNotes(e) { + const file = e.target.files?.[0]; + e.target.value = ""; + if (!file) return; + try { + const imported = JSON.parse(await file.text()); + const list = Array.isArray(imported) ? imported : imported.notes; + if (!Array.isArray(list)) throw new Error("Invalid notes file"); + const existing = new Set(notes.map(n => n.id)); + const normalized = list.map(makeNote).map(n => { + if (existing.has(n.id)) n.id = genId(); + return n; + }); + notes = [...normalized, ...notes]; + saveNotes(); + renderTagFilters(); + renderList(); + setPanelStatus(`已导入 ${normalized.length} 条笔记`); + } catch { + setPanelStatus("导入失败,请确认 JSON 格式"); + } +} diff --git a/web/js/paintBrush.js b/web/js/paintBrush.js index d055793..7b83657 100644 --- a/web/js/paintBrush.js +++ b/web/js/paintBrush.js @@ -25,6 +25,7 @@ const STYLES = ` .pb-toolbar .pb-sep { width:1px; height:24px; background:#555; margin:0 4px; } .pb-toolbar input[type=color] { width:32px; height:28px; border:none; padding:0; cursor:pointer; border-radius:4px; } .pb-toolbar input[type=range] { width:80px; accent-color:#0066ff; } +.pb-toolbar input.pb-mosaic-size { width:90px; } .pb-toolbar .pb-label { color:#aaa; font-size:12px; } .pb-canvas-wrap { border:2px solid #444; border-radius:4px; overflow:hidden; } .pb-actions { display:flex; gap:10px; margin-top:10px; } @@ -41,6 +42,64 @@ function injectStyles() { document.head.appendChild(el); } +function clamp(value, min, max) { + return Math.max(min, Math.min(max, value)); +} + +const PB_EXPORT_PROPS = ["pbTool"]; + +function clampPointer(canvas, pointer) { + return { + x: clamp(pointer.x, 0, canvas.width), + y: clamp(pointer.y, 0, canvas.height), + }; +} + +function makeCircleCursor(size) { + const cursorSize = clamp(Math.round(size), 18, 80); + const center = cursorSize / 2; + const hotspot = Math.round(center); + const radius = Math.max(3, center - 2); + const svg = ` + + + + + `.trim(); + return `url("data:image/svg+xml;charset=UTF-8,${encodeURIComponent(svg)}") ${hotspot} ${hotspot}, crosshair`; +} + +function applyMosaicCursor(canvas, state) { + const cursor = makeCircleCursor(state.mosaicSize); + canvas.defaultCursor = cursor; + canvas.hoverCursor = cursor; +} + +function configureMosaicObject(obj) { + obj.set({ + selectable: false, + hasControls: false, + hasBorders: false, + lockMovementX: true, + lockMovementY: true, + lockScalingX: true, + lockScalingY: true, + lockRotation: true, + perPixelTargetFind: true, + objectCaching: false, + }); + obj.pbTool = "mosaic"; + return obj; +} + +function normalizeMosaicObjects(canvas) { + canvas.getObjects().forEach((obj) => { + if (obj.pbTool === "mosaic" || obj.type === "image") { + configureMosaicObject(obj); + } + }); +} + // --- Get current image URL from node --- function getImageUrl(node) { if (node.imgs && node.imgs.length > 0) { @@ -71,7 +130,7 @@ class HistoryManager { if (this.locked) return; this.index++; this.stack.length = this.index; - this.stack.push(this.canvas.toJSON()); + this.stack.push(this.canvas.toJSON(PB_EXPORT_PROPS)); } undo() { if (this.index <= 0) return; @@ -86,6 +145,7 @@ class HistoryManager { _restore() { this.locked = true; this.canvas.loadFromJSON(this.stack[this.index], () => { + normalizeMosaicObjects(this.canvas); this.canvas.renderAll(); this.locked = false; }); @@ -97,7 +157,7 @@ function setupShapeDrawing(canvas, state) { let startX, startY, shape; canvas.on("mouse:down", (opt) => { - if (state.tool === "select" || state.tool === "brush") return; + if (state.tool === "select" || state.tool === "brush" || state.tool === "eraser" || state.tool === "mosaic") return; const ptr = canvas.getPointer(opt.e); startX = ptr.x; startY = ptr.y; @@ -132,12 +192,247 @@ function setupShapeDrawing(canvas, state) { }); canvas.on("mouse:up", () => { - if (!state.drawing) return; + if (!state.drawing || !shape) return; state.drawing = false; shape = null; }); } +// --- Mosaic brush handler --- +function createMosaicSource(sourceImg, canvas, blockSize) { + const scaleX = canvas.width / sourceImg.width; + const scaleY = canvas.height / sourceImg.height; + const sourceBlockSize = Math.max(1, Math.round(blockSize / Math.min(scaleX, scaleY))); + const smallW = Math.max(1, Math.ceil(sourceImg.width / sourceBlockSize)); + const smallH = Math.max(1, Math.ceil(sourceImg.height / sourceBlockSize)); + + const smallCanvas = document.createElement("canvas"); + smallCanvas.width = smallW; + smallCanvas.height = smallH; + const smallCtx = smallCanvas.getContext("2d"); + smallCtx.imageSmoothingEnabled = true; + smallCtx.drawImage(sourceImg, 0, 0, sourceImg.width, sourceImg.height, 0, 0, smallW, smallH); + + const pixelCanvas = document.createElement("canvas"); + pixelCanvas.width = sourceImg.width; + pixelCanvas.height = sourceImg.height; + const pixelCtx = pixelCanvas.getContext("2d"); + pixelCtx.imageSmoothingEnabled = false; + pixelCtx.drawImage(smallCanvas, 0, 0, smallW, smallH, 0, 0, sourceImg.width, sourceImg.height); + + return pixelCanvas; +} + +function getMosaicBrushBounds(sourceImg, canvas, centerX, centerY, size) { + const scaleX = canvas.width / sourceImg.width; + const scaleY = canvas.height / sourceImg.height; + const sourceX = centerX / scaleX; + const sourceY = centerY / scaleY; + const radiusX = Math.max(1, (size / 2) / scaleX); + const radiusY = Math.max(1, (size / 2) / scaleY); + + return { + centerX: sourceX, + centerY: sourceY, + radius: Math.max(radiusX, radiusY), + left: clamp(Math.floor(sourceX - radiusX), 0, sourceImg.width), + top: clamp(Math.floor(sourceY - radiusY), 0, sourceImg.height), + right: clamp(Math.ceil(sourceX + radiusX), 0, sourceImg.width), + bottom: clamp(Math.ceil(sourceY + radiusY), 0, sourceImg.height), + }; +} + +function expandMosaicBounds(bounds, brushBounds) { + if (!bounds.left && bounds.left !== 0) { + bounds.left = brushBounds.left; + bounds.top = brushBounds.top; + bounds.right = brushBounds.right; + bounds.bottom = brushBounds.bottom; + return; + } + bounds.left = Math.min(bounds.left, brushBounds.left); + bounds.top = Math.min(bounds.top, brushBounds.top); + bounds.right = Math.max(bounds.right, brushBounds.right); + bounds.bottom = Math.max(bounds.bottom, brushBounds.bottom); +} + +function paintMosaicStamp(sourceImg, canvas, mosaicSource, strokeCtx, bounds, point) { + const brushBounds = getMosaicBrushBounds(sourceImg, canvas, point.x, point.y, point.size); + if (brushBounds.right <= brushBounds.left || brushBounds.bottom <= brushBounds.top) return false; + + strokeCtx.save(); + strokeCtx.beginPath(); + strokeCtx.arc(brushBounds.centerX, brushBounds.centerY, brushBounds.radius, 0, Math.PI * 2); + strokeCtx.clip(); + strokeCtx.drawImage( + mosaicSource, + brushBounds.left, + brushBounds.top, + brushBounds.right - brushBounds.left, + brushBounds.bottom - brushBounds.top, + brushBounds.left, + brushBounds.top, + brushBounds.right - brushBounds.left, + brushBounds.bottom - brushBounds.top + ); + strokeCtx.restore(); + + expandMosaicBounds(bounds, brushBounds); + return true; +} + +function paintMosaicLine(sourceImg, canvas, mosaicSource, strokeCtx, bounds, from, to, size) { + const dx = to.x - from.x; + const dy = to.y - from.y; + const distance = Math.hypot(dx, dy); + const step = Math.max(2, size * 0.25); + const steps = Math.max(1, Math.ceil(distance / step)); + let painted = false; + + for (let i = 1; i <= steps; i++) { + const t = i / steps; + painted = paintMosaicStamp(sourceImg, canvas, mosaicSource, strokeCtx, bounds, { + x: from.x + dx * t, + y: from.y + dy * t, + size, + }) || painted; + } + + return painted; +} + +function createMosaicStrokeObject(strokeCanvas, bounds, canvas, sourceImg) { + const width = bounds.right - bounds.left; + const height = bounds.bottom - bounds.top; + if (width < 1 || height < 1) return Promise.resolve(null); + + const cropCanvas = document.createElement("canvas"); + cropCanvas.width = width; + cropCanvas.height = height; + cropCanvas.getContext("2d").drawImage( + strokeCanvas, + bounds.left, + bounds.top, + width, + height, + 0, + 0, + width, + height + ); + + const scaleX = canvas.width / sourceImg.width; + const scaleY = canvas.height / sourceImg.height; + const dataUrl = cropCanvas.toDataURL("image/png"); + return new Promise(resolve => { + fabric.Image.fromURL(dataUrl, (imgObj) => { + configureMosaicObject(imgObj).set({ + left: bounds.left * scaleX, + top: bounds.top * scaleY, + scaleX, + scaleY, + }); + resolve(imgObj); + }); + }); +} + +function setupMosaicDrawing(canvas, state, sourceImg, history) { + let lastPoint, mosaicSource, strokeCanvas, strokeCtx, strokePreview, strokeBounds, strokePainted; + + const refreshPreview = () => { + if (!strokePreview) return; + strokePreview.dirty = true; + canvas.requestRenderAll(); + }; + + canvas.on("mouse:down", (opt) => { + if (state.tool !== "mosaic") return; + const ptr = clampPointer(canvas, canvas.getPointer(opt.e)); + lastPoint = ptr; + state.drawing = true; + strokePainted = false; + strokeBounds = {}; + mosaicSource = createMosaicSource(sourceImg, canvas, state.mosaicSize); + + strokeCanvas = document.createElement("canvas"); + strokeCanvas.width = sourceImg.width; + strokeCanvas.height = sourceImg.height; + strokeCtx = strokeCanvas.getContext("2d"); + strokeCtx.imageSmoothingEnabled = false; + + strokePreview = new fabric.Image(strokeCanvas, { + left: 0, + top: 0, + scaleX: canvas.width / sourceImg.width, + scaleY: canvas.height / sourceImg.height, + selectable: false, + evented: false, + excludeFromExport: true, + objectCaching: false, + }); + + history.locked = true; + canvas.add(strokePreview); + history.locked = false; + + strokePainted = paintMosaicStamp(sourceImg, canvas, mosaicSource, strokeCtx, strokeBounds, { + x: ptr.x, + y: ptr.y, + size: state.mosaicSize, + }) || strokePainted; + refreshPreview(); + }); + + canvas.on("mouse:move", (opt) => { + if (state.tool !== "mosaic" || !state.drawing || !strokeCtx || !lastPoint) return; + const ptr = clampPointer(canvas, canvas.getPointer(opt.e)); + strokePainted = paintMosaicLine( + sourceImg, + canvas, + mosaicSource, + strokeCtx, + strokeBounds, + lastPoint, + ptr, + state.mosaicSize + ) || strokePainted; + lastPoint = ptr; + refreshPreview(); + }); + + canvas.on("mouse:up", async () => { + if (!state.drawing || !strokePreview) return; + + history.locked = true; + canvas.remove(strokePreview); + history.locked = false; + state.drawing = false; + lastPoint = null; + + if (!strokePainted) { + strokeCanvas = null; + strokeCtx = null; + strokePreview = null; + mosaicSource = null; + canvas.renderAll(); + return; + } + + const mosaicObj = await createMosaicStrokeObject(strokeCanvas, strokeBounds, canvas, sourceImg); + mosaicSource = null; + strokeCanvas = null; + strokeCtx = null; + strokePreview = null; + + if (mosaicObj) { + canvas.add(mosaicObj); + canvas.discardActiveObject(); + } + canvas.renderAll(); + }); +} + // --- Open Paint Modal --- async function openPaintModal(node) { injectStyles(); @@ -192,7 +487,7 @@ async function openPaintModal(node) { const cw = Math.round(img.width * scale); const ch = Math.round(img.height * scale); - const state = { tool: "brush", color: "#ff0000", width: 4, drawing: false }; + const state = { tool: "brush", color: "#ff0000", width: 4, mosaicSize: 14, drawing: false }; const toolbar = buildToolbar(state); overlay.appendChild(toolbar); @@ -222,6 +517,7 @@ async function openPaintModal(node) { if (savedState) { await new Promise(resolve => { canvas.loadFromJSON(savedState, () => { + normalizeMosaicObjects(canvas); // Re-apply background since loadFromJSON may clear it canvas.setBackgroundImage(bgUrl, () => { canvas.renderAll(); resolve(); }, { scaleX: cw / img.width, scaleY: ch / img.height, crossOrigin: "anonymous" @@ -238,6 +534,7 @@ async function openPaintModal(node) { // Shape drawing setupShapeDrawing(canvas, state); + setupMosaicDrawing(canvas, state, img, history); wireToolbar(toolbar, canvas, state, history); // Actions buttons @@ -265,7 +562,7 @@ async function openPaintModal(node) { // Confirm - save painted image and store canvas state for re-editing actions.querySelector(".pb-confirm").onclick = async () => { // Save canvas objects (without background) for future re-editing - node.properties.paintBrushCanvas = canvas.toJSON(); + node.properties.paintBrushCanvas = canvas.toJSON(PB_EXPORT_PROPS); node.graph?.change?.(); await savePaintedImage(canvas, node, img.width, img.height); close(); @@ -326,12 +623,15 @@ function buildToolbar(state) { + 颜色 线宽 + 块大小 + @@ -351,6 +651,8 @@ function wireToolbar(toolbar, canvas, state, history) { if (state.tool === "brush") { canvas.isDrawingMode = true; canvas.selection = false; + canvas.defaultCursor = "default"; + canvas.hoverCursor = "move"; canvas.freeDrawingBrush.color = state.color; canvas.freeDrawingBrush.width = state.width; } else if (state.tool === "eraser") { @@ -359,6 +661,11 @@ function wireToolbar(toolbar, canvas, state, history) { canvas.selection = true; canvas.defaultCursor = "crosshair"; canvas.hoverCursor = "pointer"; + } else if (state.tool === "mosaic") { + canvas.isDrawingMode = false; + canvas.selection = false; + canvas.discardActiveObject(); + applyMosaicCursor(canvas, state); } else { canvas.isDrawingMode = false; canvas.selection = false; @@ -396,6 +703,13 @@ function wireToolbar(toolbar, canvas, state, history) { } }; + toolbar.querySelector(".pb-mosaic-size").oninput = (e) => { + state.mosaicSize = parseInt(e.target.value); + if (state.tool === "mosaic") { + applyMosaicCursor(canvas, state); + } + }; + // Action buttons toolbar.querySelector("[data-action=undo]").onclick = () => history.undo(); toolbar.querySelector("[data-action=clear]").onclick = () => {