diff --git a/comfy_api_nodes/apis/recraft.py b/comfy_api_nodes/apis/recraft.py index 64780d73b1f..26834f7767d 100644 --- a/comfy_api_nodes/apis/recraft.py +++ b/comfy_api_nodes/apis/recraft.py @@ -250,6 +250,26 @@ class RecraftImageSize(str, Enum): "1536x2688", ] +RECRAFT_V4_STYLES_MODELS = frozenset( + { + "recraftv4_styles", + "recraftv4_styles_vector", + "recraftv4_styles_pro", + "recraftv4_styles_pro_vector", + } +) + +RECRAFT_V4_VECTOR_MODEL_FOR_STYLE = { + "recraftv4": "recraftv4_vector", + "recraftv4_pro": "recraftv4_pro_vector", +} + +RECRAFT_STYLE_MATCH_OPTIONS = ["precise", "flexible"] + +RECRAFT_STYLE_REFERENCES_MAX = 10 + +RECRAFT_STYLE_REFERENCES_MAX_BYTES = 10 * 1024 * 1024 + class RecraftColorObject(BaseModel): rgb: list[int] = Field(..., description='An array of 3 integer values in range of 0...255 defining RGB Color Model') @@ -272,6 +292,8 @@ class RecraftImageGenerationRequest(BaseModel): substyle: str | None = Field(None, description='The substyle to apply to the generated image, depending on the style input') controls: RecraftControlsObject | None = Field(None, description='A set of custom parameters to tweak generation process') style_id: str | None = Field(None, description='Use a previously uploaded style as a reference; UUID') + style_match: str | None = Field(None, description='How closely to follow the referenced style: "precise" or "flexible" for V4 models') + style_reference_urls: list[str] | None = Field(None, description='URLs or data URLs of style reference images; a private style is created from them and returned as style_id') strength: float | None = Field(None, description='Defines the difference with the original image, should lie in [0, 1], where 0 means almost identical, and 1 means miserable similarity') random_seed: int | None = Field(None, description="Seed for video generation") @@ -286,10 +308,12 @@ class RecraftImageGenerationResponse(BaseModel): credits: int = Field(..., description='Number of credits used for the generation') data: list[RecraftReturnedObject] | None = Field(None, description='Array of generated image information') image: RecraftReturnedObject | None = Field(None, description='Single generated image') + style_id: str | None = Field(None, description='The style applied to the generation, including one auto-created from style references') class RecraftCreateStyleRequest(BaseModel): - style: str = Field(..., description="realistic_image, digital_illustration, vector_illustration, or icon") + style: str = Field(..., description="any, realistic_image, digital_illustration, vector_illustration, or icon") + model: str | None = Field(None, description="The model family the style is created for, e.g. recraftv4_styles") class RecraftCreateStyleResponse(BaseModel): diff --git a/comfy_api_nodes/nodes_recraft.py b/comfy_api_nodes/nodes_recraft.py index 9f1823426f1..d2784557a0e 100644 --- a/comfy_api_nodes/nodes_recraft.py +++ b/comfy_api_nodes/nodes_recraft.py @@ -1,3 +1,5 @@ +import base64 +import uuid from io import BytesIO import aiohttp @@ -8,8 +10,13 @@ from comfy.utils import ProgressBar from comfy_api.latest import IO, ComfyExtension from comfy_api_nodes.apis.recraft import ( + RECRAFT_STYLE_MATCH_OPTIONS, + RECRAFT_STYLE_REFERENCES_MAX, + RECRAFT_STYLE_REFERENCES_MAX_BYTES, RECRAFT_V4_PRO_SIZES, RECRAFT_V4_SIZES, + RECRAFT_V4_STYLES_MODELS, + RECRAFT_V4_VECTOR_MODEL_FOR_STYLE, RecraftColor, RecraftColorChain, RecraftControls, @@ -173,6 +180,60 @@ def __exit__(self, exc_type, exc_val, exc_tb): ) +def style_reference_images(images: IO.Autogrow.Type | None) -> list[torch.Tensor]: + references = [] + for batch in (images or {}).values(): + if batch is None: + continue + if batch.ndim == 3: + batch = batch.unsqueeze(0) + references.extend(batch[i] for i in range(batch.shape[0])) + return references + + +def encode_style_references(references: list[torch.Tensor]) -> list[bytes]: + if not references: + raise ValueError("At least one style reference image is required.") + if len(references) > RECRAFT_STYLE_REFERENCES_MAX: + raise ValueError( + f"At most {RECRAFT_STYLE_REFERENCES_MAX} style reference images are allowed; got {len(references)}." + ) + encoded = [] + total_size = 0 + for reference in references: + data = tensor_to_bytesio(reference, total_pixels=2048 * 2048, mime_type="image/webp").read() + total_size += len(data) + if total_size > RECRAFT_STYLE_REFERENCES_MAX_BYTES: + raise ValueError("Total size of style reference images exceeds the 10 MB limit.") + encoded.append(data) + return encoded + + +def resolve_v4_style( + model: str, style_id: str, style_references: IO.Autogrow.Type | None +) -> tuple[str | None, list[str] | None]: + style_id = style_id.strip() + references = style_reference_images(style_references) + if style_id and references: + raise ValueError("Provide either a style_id or style reference images, not both.") + if style_id: + try: + uuid.UUID(style_id) + except ValueError: + raise ValueError(f"style_id must be a UUID; got '{style_id}'.") from None + return style_id, None + if references: + return None, [ + f"data:image/webp;base64,{base64.b64encode(data).decode()}" + for data in encode_style_references(references) + ] + if model in RECRAFT_V4_STYLES_MODELS: + raise ValueError( + f"Model '{model}' always requires a style: connect style reference images or provide a style_id." + ) + return None, None + + class RecraftColorRGBNode(IO.ComfyNode): @classmethod def define_schema(cls): @@ -272,9 +333,9 @@ class RecraftStyleV3VectorIllustrationNode(RecraftStyleV3RealisticImageNode): def define_schema(cls): return IO.Schema( node_id="RecraftStyleV3VectorIllustrationNode", - display_name="Recraft Style - Realistic Image", + display_name="Recraft Style - Vector Illustration", category="partner/image/Recraft", - description="Select realistic_image style and optional substyle.", + description="Select vector_illustration style and optional substyle.", inputs=[ IO.Combo.Input("substyle", options=get_v3_substyles(cls.RECRAFT_STYLE)), ], @@ -293,7 +354,7 @@ def define_schema(cls): node_id="RecraftStyleV3LogoRaster", display_name="Recraft Style - Logo Raster", category="partner/image/Recraft", - description="Select realistic_image style and optional substyle.", + description="Select logo_raster style and substyle.", inputs=[ IO.Combo.Input("substyle", options=get_v3_substyles(cls.RECRAFT_STYLE, include_none=False)), ], @@ -395,6 +456,74 @@ async def execute( return IO.NodeOutput(response.id) +class RecraftV4CreateStyleNode(IO.ComfyNode): + @classmethod + def define_schema(cls): + return IO.Schema( + node_id="RecraftV4CreateStyleNode", + display_name="Recraft V4 Create Style", + category="partner/image/Recraft", + description="Create a reusable Recraft V4 style from 1-10 reference images. " + "The returned style_id works with every Recraft V4 and V4.1 model of the same output type. " + "Total size of all images is limited to 10 MB.", + inputs=[ + IO.Combo.Input( + "model", + options=["recraftv4_styles", "recraftv4_styles_vector"], + tooltip="Output type the style is created for: recraftv4_styles for raster images, " + "recraftv4_styles_vector for SVG.", + ), + IO.Autogrow.Input( + "images", + template=IO.Autogrow.TemplatePrefix( + IO.Image.Input("image"), + prefix="image", + min=1, + max=RECRAFT_STYLE_REFERENCES_MAX, + ), + tooltip="Reference images defining the style. Similar references sharpen the match, " + "varied references widen it.", + ), + ], + outputs=[ + IO.String.Output(display_name="style_id"), + ], + hidden=[ + IO.Hidden.auth_token_comfy_org, + IO.Hidden.api_key_comfy_org, + IO.Hidden.unique_id, + ], + is_api_node=True, + price_badge=IO.PriceBadge( + expr="""{"type":"usd","usd": 0.005}""", + ), + ) + + @classmethod + async def execute( + cls, + model: str, + images: IO.Autogrow.Type, + ) -> IO.NodeOutput: + files = [ + (f"file{i + 1}", data) + for i, data in enumerate(encode_style_references(style_reference_images(images))) + ] + response = await sync_op( + cls, + endpoint=ApiEndpoint(path="/proxy/recraft/styles", method="POST"), + response_model=RecraftCreateStyleResponse, + files=files, + data=RecraftCreateStyleRequest( + style="vector_illustration" if model.endswith("_vector") else "any", + model=model, + ), + content_type="multipart/form-data", + max_retries=1, + ) + return IO.NodeOutput(response.id) + + class RecraftTextToImageNode(IO.ComfyNode): @classmethod def define_schema(cls): @@ -1170,8 +1299,31 @@ def define_schema(cls): ), ], ), + IO.DynamicCombo.Option( + "recraftv4_styles", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_SIZES, + default="1024x1024", + tooltip="The size of the generated image.", + ), + ], + ), + IO.DynamicCombo.Option( + "recraftv4_styles_pro", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_PRO_SIZES, + default="2048x2048", + tooltip="The size of the generated image.", + ), + ], + ), ], - tooltip="The model to use for generation.", + tooltip="The model to use for generation. The recraftv4_styles models are built for " + "style-consistent generation and always require a style_id or style_references.", ), IO.Int.Input( "n", @@ -1194,9 +1346,36 @@ def define_schema(cls): tooltip="Optional additional controls over the generation via the Recraft Controls node.", optional=True, ), + IO.String.Input( + "style_id", + default="", + optional=True, + tooltip="UUID of a Recraft V4 style to apply, e.g. from the Recraft V4 Create Style node " + "or the style_id output of a previous run. Cannot be combined with style_references.", + ), + IO.Combo.Input( + "style_match", + options=RECRAFT_STYLE_MATCH_OPTIONS, + optional=True, + tooltip="How closely to follow the style: precise reproduces it in detail, " + "flexible matches the general look. Only used when a style is provided.", + ), + IO.Autogrow.Input( + "style_references", + template=IO.Autogrow.TemplatePrefix( + IO.Image.Input("style_reference"), + prefix="style_reference", + min=0, + max=RECRAFT_STYLE_REFERENCES_MAX, + ), + optional=True, + tooltip="Reference images to create a style from on the fly, billed on top of the generation. " + "The created style is returned as style_id for reuse. Cannot be combined with style_id.", + ), ], outputs=[ IO.Image.Output(), + IO.String.Output(id="style_id", display_name="style_id"), ], hidden=[ IO.Hidden.auth_token_comfy_org, @@ -1205,7 +1384,7 @@ def define_schema(cls): ], is_api_node=True, price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model", "n"]), + depends_on=IO.PriceBadgeDepends(widgets=["model", "n"], input_groups=["style_references"]), expr=""" ( $prices := { @@ -1214,9 +1393,13 @@ def define_schema(cls): "recraftv4_1_pro": 0.21, "recraftv4_1_utility_pro": 0.21, "recraftv4": 0.04, - "recraftv4_pro": 0.25 + "recraftv4_pro": 0.25, + "recraftv4_styles": 0.035, + "recraftv4_styles_pro": 0.10 }; - {"type":"usd","usd": $lookup($prices, widgets.model) * widgets.n} + $references := $lookup(inputGroups, "style_references"); + $style := ($references ? $references : 0) > 0 ? 0.005 : 0; + {"type":"usd","usd": $lookup($prices, widgets.model) * widgets.n + $style} ) """, ), @@ -1231,8 +1414,12 @@ async def execute( n: int, seed: int, recraft_controls: RecraftControls | None = None, + style_id: str = "", + style_match: str = "precise", + style_references: IO.Autogrow.Type | None = None, ) -> IO.NodeOutput: validate_string(prompt, strip_whitespace=True, min_length=1, max_length=10000) + style_id, style_reference_urls = resolve_v4_style(model["model"], style_id, style_references) response = await sync_op( cls, ApiEndpoint(path="/proxy/recraft/image_generation", method="POST"), @@ -1242,6 +1429,9 @@ async def execute( model=model["model"], size=model["size"], n=n, + style_id=style_id, + style_match=style_match if style_id or style_reference_urls else None, + style_reference_urls=style_reference_urls, controls=recraft_controls.create_api_model() if recraft_controls else None, ), max_retries=1, @@ -1253,7 +1443,7 @@ async def execute( if len(image.shape) < 4: image = image.unsqueeze(0) images.append(image) - return IO.NodeOutput(torch.cat(images, dim=0)) + return IO.NodeOutput(torch.cat(images, dim=0), response.style_id or "") class RecraftV4TextToVectorNode(IO.ComfyNode): @@ -1345,8 +1535,31 @@ def define_schema(cls): ), ], ), + IO.DynamicCombo.Option( + "recraftv4_styles_vector", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_SIZES, + default="1024x1024", + tooltip="The size of the generated image.", + ), + ], + ), + IO.DynamicCombo.Option( + "recraftv4_styles_pro_vector", + [ + IO.Combo.Input( + "size", + options=RECRAFT_V4_PRO_SIZES, + default="2048x2048", + tooltip="The size of the generated image.", + ), + ], + ), ], - tooltip="The model to use for generation.", + tooltip="The model to use for generation. The recraftv4_styles models are built for " + "style-consistent generation and always require a style_id or style_references.", ), IO.Int.Input( "n", @@ -1369,9 +1582,36 @@ def define_schema(cls): tooltip="Optional additional controls over the generation via the Recraft Controls node.", optional=True, ), + IO.String.Input( + "style_id", + default="", + optional=True, + tooltip="UUID of a Recraft V4 vector style to apply, e.g. from the Recraft V4 Create Style node " + "or the style_id output of a previous run. Cannot be combined with style_references.", + ), + IO.Combo.Input( + "style_match", + options=RECRAFT_STYLE_MATCH_OPTIONS, + optional=True, + tooltip="How closely to follow the style: precise reproduces it in detail, " + "flexible matches the general look. Only used when a style is provided.", + ), + IO.Autogrow.Input( + "style_references", + template=IO.Autogrow.TemplatePrefix( + IO.Image.Input("style_reference"), + prefix="style_reference", + min=0, + max=RECRAFT_STYLE_REFERENCES_MAX, + ), + optional=True, + tooltip="Reference images to create a vector style from on the fly, billed on top of the generation. " + "The created style is returned as style_id for reuse. Cannot be combined with style_id.", + ), ], outputs=[ IO.SVG.Output(), + IO.String.Output(id="style_id", display_name="style_id"), ], hidden=[ IO.Hidden.auth_token_comfy_org, @@ -1380,7 +1620,7 @@ def define_schema(cls): ], is_api_node=True, price_badge=IO.PriceBadge( - depends_on=IO.PriceBadgeDepends(widgets=["model", "n"]), + depends_on=IO.PriceBadgeDepends(widgets=["model", "n"], input_groups=["style_references"]), expr=""" ( $prices := { @@ -1389,9 +1629,13 @@ def define_schema(cls): "recraftv4_1_pro_vector": 0.30, "recraftv4_1_utility_pro_vector": 0.30, "recraftv4": 0.08, - "recraftv4_pro": 0.30 + "recraftv4_pro": 0.30, + "recraftv4_styles_vector": 0.05, + "recraftv4_styles_pro_vector": 0.12 }; - {"type":"usd","usd": $lookup($prices, widgets.model) * widgets.n} + $references := $lookup(inputGroups, "style_references"); + $style := ($references ? $references : 0) > 0 ? 0.005 : 0; + {"type":"usd","usd": $lookup($prices, widgets.model) * widgets.n + $style} ) """, ), @@ -1406,19 +1650,30 @@ async def execute( n: int, seed: int, recraft_controls: RecraftControls | None = None, + style_id: str = "", + style_match: str = "precise", + style_references: IO.Autogrow.Type | None = None, ) -> IO.NodeOutput: validate_string(prompt, strip_whitespace=True, min_length=1, max_length=10000) + model_id = model["model"] + style_id, style_reference_urls = resolve_v4_style(model_id, style_id, style_references) + has_style = bool(style_id or style_reference_urls) + if has_style: + model_id = RECRAFT_V4_VECTOR_MODEL_FOR_STYLE.get(model_id, model_id) response = await sync_op( cls, ApiEndpoint(path="/proxy/recraft/image_generation", method="POST"), response_model=RecraftImageGenerationResponse, data=RecraftImageGenerationRequest( prompt=prompt, - model=model["model"], + model=model_id, size=model["size"], n=n, - style=None if model["model"].endswith("_vector") else "vector_illustration", + style=None if has_style or model_id.endswith("_vector") else "vector_illustration", substyle=None, + style_id=style_id, + style_match=style_match if has_style else None, + style_reference_urls=style_reference_urls, controls=recraft_controls.create_api_model() if recraft_controls else None, ), max_retries=1, @@ -1426,7 +1681,7 @@ async def execute( svg_data = [] for data in response.data: svg_data.append(await download_url_as_bytesio(data.url, timeout=1024)) - return IO.NodeOutput(SVG(svg_data)) + return IO.NodeOutput(SVG(svg_data), response.style_id or "") class RecraftExtension(ComfyExtension): @@ -1451,6 +1706,7 @@ async def get_node_list(self) -> list[type[IO.ComfyNode]]: RecraftControlsNode, RecraftV4TextToImageNode, RecraftV4TextToVectorNode, + RecraftV4CreateStyleNode, ]