Files
comfyui_o1key/nodes/prompt_multi_function.py
T
Jony ba920f2b66 Publish current ComfyUI O1Key code baseline
Replace the prior release tree with the current plugin, frontend, tests, and documentation. Document retired node IDs and the public Gitea update source.
2026-09-24 19:56:48 +08:00

199 lines
7.6 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
提示词(多功能)节点
支持以「第一套---第二套---第三套」格式填入多套提示词,并选择处理方式:
- 全部使用:保留 --- 分隔符输出全部套数,交给下游批量节点并发跑
- 随机抽取n套:按指定数量不重复抽取,按原始顺序输出;数量为 1 时只抽 1 套
- 指定序号:按填写顺序输出一套或多套提示词
下游需连接支持批量提示词(按单独行 --- 分割并发执行)的节点,
如「Nano Banana 批量跑图」等。
"""
import random
import re
from typing import List
from ..utils.image_utils import parse_batch_prompts
_MODE_ALL = "全部使用"
_MODE_RANDOM = "随机抽取n套"
_MODE_SELECTED = "指定序号"
_LEGACY_MODE_RANDOM_ONE = "随机抽取1套"
_LEGACY_MODE_RANDOM_MANY = "随机抽取多套"
_RANDOM_MODES = {
_MODE_RANDOM,
_LEGACY_MODE_RANDOM_ONE,
_LEGACY_MODE_RANDOM_MANY,
}
_MODES = [_MODE_ALL, _MODE_RANDOM, _MODE_SELECTED]
_DEFAULT_SAMPLE_COUNT = 3
_DEFAULT_SELECTED_INDICES = "1,2,3"
_INDEX_SEPARATOR_RE = re.compile(r"[\s,;]+")
_INDEX_RANGE_RE = re.compile(r"^(\d+)[-~~—–](\d+)$")
def _split_prompt_sets(text: str) -> List[str]:
"""按单独行 --- 切分多套提示词;无分隔符时整段视为 1 套。"""
stripped = (text or "").strip()
if not stripped:
return []
sets = parse_batch_prompts(text)
if not sets:
return [stripped]
return sets
def _parse_prompt_indices(value: str, total: int) -> List[int]:
"""解析从 1 开始的序号列表,拒绝重复并保留填写顺序。"""
raw_value = str(value or "").strip()
if not raw_value:
raise ValueError("提示词(多功能):指定序号为空,请填写如 1,3,5。")
tokens = [token for token in _INDEX_SEPARATOR_RE.split(raw_value) if token]
selected: List[int] = []
seen = set()
def append_index(index: int) -> None:
if index < 1 or index > total:
raise ValueError(
f"提示词(多功能):序号 {index} 超出范围,当前共有 {total} 套提示词。"
)
if index in seen:
raise ValueError(f"提示词(多功能):序号 {index} 重复,请勿重复填写。")
seen.add(index)
selected.append(index)
for token in tokens:
if token.isdigit():
append_index(int(token))
continue
match = _INDEX_RANGE_RE.fullmatch(token)
if match:
start, end = (int(part) for part in match.groups())
if end < start:
raise ValueError(
f"提示词(多功能):区间 {token} 必须从小到大填写。"
)
if start < 1 or start > total:
append_index(start)
if end < 1 or end > total:
append_index(end)
for index in range(start, end + 1):
append_index(index)
continue
raise ValueError(
f"提示词(多功能):无法识别序号“{token}”,请填写如 1,3,5 或 2-4。"
)
return selected
class O1keyPromptMultiFunction:
"""多功能提示词节点:全部使用、随机抽取或按序号选择。"""
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"提示词": ("STRING", {
"multiline": True,
"default": "第一套提示词\n---\n第二套提示词\n---\n第三套提示词",
"placeholder": "多套提示词请用单独一行的 --- 分隔",
}),
"功能": (_MODES, {"default": _MODE_ALL}),
"抽取数量": ("INT", {
"default": _DEFAULT_SAMPLE_COUNT,
"min": 1,
"max": 1000,
"step": 1,
"tooltip": "仅“随机抽取n套”生效;从全部提示词中不重复抽取。",
}),
"指定序号": ("STRING", {
"default": _DEFAULT_SELECTED_INDICES,
"placeholder": "例如:1,3,5 或 2-4",
"tooltip": "仅“指定序号”生效;序号从 1 开始,按填写顺序输出。",
}),
}
}
RETURN_TYPES = ("STRING",)
RETURN_NAMES = ("提示词",)
FUNCTION = "process"
CATEGORY = "o1key/prompt"
DESCRIPTION = (
"多套提示词用单独一行的 --- 分隔。\n"
"全部使用:保留 --- 输出全部,交给下游批量节点并发跑每一套。\n"
"随机抽取n套:按“抽取数量”不重复随机选择,并按原始顺序输出。\n"
"指定序号:支持 1,3,5、中文逗号、空格和 2-4 区间,按填写顺序输出。"
)
def process(
self,
提示词: str,
功能: str = _MODE_ALL,
抽取数量: int = _DEFAULT_SAMPLE_COUNT,
指定序号: str = _DEFAULT_SELECTED_INDICES,
):
sets = _split_prompt_sets(提示词)
if not sets:
raise ValueError("提示词(多功能):提示词为空,请至少填写 1 套。")
# 旧 API 工作流可能绕过前端迁移直接提交原模式值,继续保持只抽 1 套。
if 功能 == _LEGACY_MODE_RANDOM_ONE:
chosen = random.choice(sets)
print(f"[o1key 提示词多功能] 随机抽取 1/{len(sets)} 套")
return (chosen,)
if 功能 in {_MODE_RANDOM, _LEGACY_MODE_RANDOM_MANY}:
try:
sample_count = int(抽取数量)
except (TypeError, ValueError) as exc:
raise ValueError("提示词(多功能):抽取数量必须是整数。") from exc
if sample_count < 1:
raise ValueError("提示词(多功能):抽取数量必须至少为 1。")
if sample_count > len(sets):
raise ValueError(
f"提示词(多功能):抽取数量 {sample_count} 超过当前提示词总数 {len(sets)}。"
)
chosen_indices = sorted(random.sample(range(len(sets)), sample_count))
chosen = [sets[index] for index in chosen_indices]
display_indices = ",".join(str(index + 1) for index in chosen_indices)
print(
f"[o1key 提示词多功能] 随机抽取 {sample_count}/{len(sets)} 套,"
f"序号:{display_indices}"
)
return ("\n---\n".join(chosen),)
if 功能 == _MODE_SELECTED:
selected_indices = _parse_prompt_indices(指定序号, len(sets))
chosen = [sets[index - 1] for index in selected_indices]
display_indices = ",".join(str(index) for index in selected_indices)
print(
f"[o1key 提示词多功能] 指定使用 {len(chosen)}/{len(sets)} 套,"
f"序号:{display_indices}"
)
return ("\n---\n".join(chosen),)
# 全部使用:保留单独行 --- 分隔符,下游批量节点可并发分割执行
joined = "\n---\n".join(sets)
print(f"[o1key 提示词多功能] 全部使用,共 {len(sets)} 套")
return (joined,)
@classmethod
def IS_CHANGED(
cls,
提示词,
功能=_MODE_ALL,
抽取数量=_DEFAULT_SAMPLE_COUNT,
指定序号=_DEFAULT_SELECTED_INDICES,
):
# 随机模式每次都重新抽取
if 功能 in _RANDOM_MODES:
return float("nan")
if 功能 == _MODE_SELECTED:
return f"{功能}|{指定序号}|{提示词}"
return f"{功能}|{提示词}"