ImageGen / mcp_tools /get_model_features.py
RioShiina's picture
Add High-Level API/MCP Tools
9dbb7e3 verified
Raw
History Blame Contribute Delete
4.09 kB
"""
MCP Tool: get_model_features
Query metadata for a specified model, including supported task types, extended features, and default inference parameters.
"""
from .common import (
_load_yaml,
_MODEL_LIST_PATH,
_MODEL_DEFAULTS_PATH,
_IMAGE_GEN_FEATURES_PATH,
_MODEL_ARCHITECTURES_PATH,
_CHAIN_FEATURES_PATH,
_TASK_DEFINITIONS,
)
from .error_schema import make_validation_error, make_not_found_error
def handle_get_model_features(model: str) -> dict:
"""Query metadata for a specified model: supported task types, extended features, and default inference parameters."""
if not model:
return make_validation_error(
"Parameter 'model' is required.",
missing_fields=["model"],
)
model_list = _load_yaml(_MODEL_LIST_PATH)
model_defaults = _load_yaml(_MODEL_DEFAULTS_PATH)
features_config = _load_yaml(_IMAGE_GEN_FEATURES_PATH)
arch_config = _load_yaml(_MODEL_ARCHITECTURES_PATH)
chain_features = _load_yaml(_CHAIN_FEATURES_PATH)
found_arch = None
checkpoints = model_list.get("Checkpoint", {})
for arch_name, arch_data in checkpoints.items():
if not isinstance(arch_data, dict):
continue
for m in arch_data.get("models", []):
if m.get("display_name") == model:
found_arch = arch_name
break
if found_arch:
break
if not found_arch:
return make_not_found_error("model", model)
architectures = arch_config.get("architectures", {})
arch_info = architectures.get(found_arch, {})
model_type = arch_info.get("model_type", found_arch.lower())
arch_features = features_config.get(model_type, features_config.get("default", {}))
enabled_chains = arch_features.get("enabled_chains", [])
supported_features = []
for chain_name in enabled_chains:
if chain_name in chain_features:
chain_data = chain_features[chain_name]
visibility = chain_data.get("visibility", "public")
if visibility == "public":
supported_features.append(chain_name)
else:
generic_mapping = {
"krea2_controlnet": "controlnet",
"anima_controlnet_lllite": "controlnet",
"controlnet_model_patch": "controlnet",
"flux1_ipadapter": "ipadapter",
"sd3_ipadapter": "ipadapter",
"hidream_o1_reference": "reference_latent",
}
generic_name = generic_mapping.get(chain_name)
if generic_name and generic_name not in supported_features:
supported_features.append(generic_name)
arch_defaults_section = model_defaults.get(found_arch, {})
arch_level_defaults = arch_defaults_section.get("_defaults", {})
model_specific_defaults = arch_defaults_section.get(model, {})
global_defaults = model_defaults.get("Default", {})
merged_defaults = {**global_defaults, **arch_level_defaults, **model_specific_defaults}
default_parameter = {
"sampler": merged_defaults.get("sampler_name", "euler"),
"scheduler": merged_defaults.get("scheduler", "simple"),
"steps": merged_defaults.get("steps", 20),
"cfg": merged_defaults.get("cfg", 1.0),
}
supported_tasks = [t["task_type"] for t in _TASK_DEFINITIONS]
result = {
"name": model,
"model_architecture": found_arch,
"supported_tasks": supported_tasks,
"supported_features": supported_features,
"default_parameter": default_parameter,
}
default_pos = model_specific_defaults.get(
"positive_prompt",
arch_level_defaults.get("positive_prompt", ""),
)
default_neg = model_specific_defaults.get(
"negative_prompt",
arch_level_defaults.get("negative_prompt", ""),
)
if default_pos:
result["default_positive_prompt"] = default_pos
if default_neg:
result["default_negative_prompt"] = default_neg
return result