use pydantic instead

This commit is contained in:
Terry Jia 2025-06-01 08:37:40 -04:00
parent a21e49c18b
commit 37fcf3e8e8
3 changed files with 61 additions and 105 deletions

View File

@ -1,23 +1,17 @@
import os import os
from typing import Optional from typing import Optional
import tomlkit import tomllib
import tomlkit.exceptions
from pydantic import ValidationError
import logging import logging
from comfy_config.types import ( from comfy_config.types import (
ComfyConfig,
License,
Model,
ProjectConfig, ProjectConfig,
PyProjectConfig, PyProjectConfig,
URLs,
) )
""" """
Original implementation comes from https://github.com/Comfy-Org/comfy-cli/blob/2e36f33dd39ef43b5acf7d1fc5acc5e01be92360/comfy_cli/registry/config_parser.py#L146
Extract configuration from a custom node directory's pyproject.toml file. Extract configuration from a custom node directory's pyproject.toml file.
This function reads and parses the pyproject.toml file in the specified directory This function reads and parses the pyproject.toml file in the specified directory
@ -46,70 +40,42 @@ Example:
>>> print(project_config.project.name) # "my_custom_node" or name from pyproject.toml >>> print(project_config.project.name) # "my_custom_node" or name from pyproject.toml
>>> nodes.EXTENSION_WEB_DIRS[project_config.project.name] = js_dir >>> nodes.EXTENSION_WEB_DIRS[project_config.project.name] = js_dir
""" """
def extract_node_configuration(
path,
) -> Optional[PyProjectConfig]: def extract_node_configuration(path) -> Optional[PyProjectConfig]:
folder_name = os.path.basename(path) folder_name = os.path.basename(path)
toml_path = os.path.join(path, "pyproject.toml")
path = os.path.join(path, "pyproject.toml") if not os.path.isfile(toml_path):
logging.warning(
if not os.path.isfile(path): "No pyproject.toml file found in the current directory, will use custom node folder name as project name as default.")
logging.warning("No pyproject.toml file found in the current directory, will use custom node folder name as project name as default.")
project = ProjectConfig(
name=folder_name,
)
try:
project = ProjectConfig(name=folder_name)
return PyProjectConfig(project=project) return PyProjectConfig(project=project)
except ValidationError as e:
logging.error(f"Failed to create default configuration: {e}")
return None
with open(path, "r") as file: try:
data = tomlkit.load(file) with open(toml_path, "rb") as f:
data = tomllib.load(f)
except Exception as e:
logging.error(f"Failed to read pyproject.toml: {e}")
return None
project_data = data.get("project", {}) try:
urls_data = project_data.get("urls", {}) config_data = {
comfy_data = data.get("tool", {}).get("comfy", {}) "project": data.get("project", {}),
"tool_comfy": data.get("tool", {}).get("comfy", {})
}
license_data = project_data.get("license", {}) return PyProjectConfig(**config_data)
if isinstance(license_data, str):
license = License(text=license_data)
logging.warning(
'Warning: License should be in one of these two formats: license = {file = "LICENSE"} OR license = {text = "MIT License"}. Please check the documentation: https://docs.comfy.org/registry/specifications.'
)
elif isinstance(license_data, dict):
if "file" in license_data or "text" in license_data:
license = License(file=license_data.get("file", ""), text=license_data.get("text", ""))
else:
logging.warning(
'Warning: License should be in one of these two formats: license = {file = "LICENSE"} OR license = {text = "MIT License"}. Please check the documentation: https://docs.comfy.org/registry/specifications.'
)
license = License()
else:
license = License()
logging.warning(
'Warning: License should be in one of these two formats: license = {file = "LICENSE"} OR license = {text = "MIT License"}. Please check the documentation: https://docs.comfy.org/registry/specifications.'
)
project = ProjectConfig( except ValidationError as e:
name=project_data.get("name", ""), logging.error(f"Validation error while parsing configuration: {e}")
description=project_data.get("description", ""), logging.error(f"Validation details: {e.errors()}")
version=project_data.get("version", ""), return None
requires_python=project_data.get("requires-python", ""), except Exception as e:
dependencies=project_data.get("dependencies", []), logging.error(f"Unexpected error while parsing configuration: {e}")
license=license, return None
urls=URLs(
homepage=urls_data.get("Homepage", ""),
documentation=urls_data.get("Documentation", ""),
repository=urls_data.get("Repository", ""),
issues=urls_data.get("Issues", ""),
),
)
comfy = ComfyConfig(
publisher_id=comfy_data.get("PublisherId", ""),
display_name=comfy_data.get("DisplayName", ""),
icon=comfy_data.get("Icon", ""),
models=[Model(location=m["location"], model_url=m["model_url"]) for m in comfy_data.get("Models", [])],
includes=comfy_data.get("includes", []),
)
return PyProjectConfig(project=project, tool_comfy=comfy)

View File

@ -1,12 +1,11 @@
from dataclasses import dataclass, field from pydantic import BaseModel, Field
from typing import List, Optional from typing import List, Optional
# IMPORTANT: The type definitions specified in pyproject.toml for custom nodes # IMPORTANT: The type definitions specified in pyproject.toml for custom nodes
# must remain synchronized with the corresponding files in the https://github.com/Comfy-Org/comfy-cli/blob/main/comfy_cli/registry/types.py. # must remain synchronized with the corresponding files in the https://github.com/Comfy-Org/comfy-cli/blob/main/comfy_cli/registry/types.py.
# Any changes to one must be reflected in the other to maintain consistency. # Any changes to one must be reflected in the other to maintain consistency.
@dataclass class NodeVersion(BaseModel):
class NodeVersion:
changelog: str changelog: str
dependencies: List[str] dependencies: List[str]
deprecated: bool deprecated: bool
@ -15,8 +14,7 @@ class NodeVersion:
download_url: str download_url: str
@dataclass class Node(BaseModel):
class Node:
id: str id: str
name: str name: str
description: str description: str
@ -24,57 +22,50 @@ class Node:
license: Optional[str] = None license: Optional[str] = None
icon: Optional[str] = None icon: Optional[str] = None
repository: Optional[str] = None repository: Optional[str] = None
tags: List[str] = field(default_factory=list) tags: List[str] = Field(default_factory=list)
latest_version: Optional[NodeVersion] = None latest_version: Optional[NodeVersion] = None
@dataclass class PublishNodeVersionResponse(BaseModel):
class PublishNodeVersionResponse:
node_version: NodeVersion node_version: NodeVersion
signedUrl: str signedUrl: str
@dataclass class URLs(BaseModel):
class URLs: homepage: str = Field(default="", alias="Homepage")
homepage: str = "" documentation: str = Field(default="", alias="Documentation")
documentation: str = "" repository: str = Field(default="", alias="Repository")
repository: str = "" issues: str = Field(default="", alias="Issues")
issues: str = ""
@dataclass class Model(BaseModel):
class Model:
location: str location: str
model_url: str model_url: str
@dataclass class ComfyConfig(BaseModel):
class ComfyConfig: publisher_id: str = Field(default="", alias="PublisherId")
publisher_id: str = "" display_name: str = Field(default="", alias="DisplayName")
display_name: str = "" icon: str = Field(default="", alias="Icon")
icon: str = "" models: List[Model] = Field(default_factory=list, alias="Models")
models: List[Model] = field(default_factory=list) includes: List[str] = Field(default_factory=list)
includes: List[str] = field(default_factory=list)
@dataclass class License(BaseModel):
class License:
file: str = "" file: str = ""
text: str = "" text: str = ""
@dataclass class ProjectConfig(BaseModel):
class ProjectConfig:
name: str = "" name: str = ""
description: str = "" description: str = ""
version: str = "1.0.0" version: str = "1.0.0"
requires_python: str = ">= 3.9" requires_python: str = Field(default=">= 3.9", alias="requires-python")
dependencies: List[str] = field(default_factory=list) dependencies: List[str] = Field(default_factory=list)
license: License = field(default_factory=License) license: License = Field(default_factory=License)
urls: URLs = field(default_factory=URLs) urls: URLs = Field(default_factory=URLs)
@dataclass class PyProjectConfig(BaseModel):
class PyProjectConfig: project: ProjectConfig = Field(default_factory=ProjectConfig)
project: ProjectConfig = field(default_factory=ProjectConfig) tool_comfy: ComfyConfig = Field(default_factory=ComfyConfig)
tool_comfy: ComfyConfig = field(default_factory=ComfyConfig)

View File

@ -24,4 +24,3 @@ spandrel
soundfile soundfile
av>=14.2.0 av>=14.2.0
pydantic~=2.0 pydantic~=2.0
tomlkit