File size: 4,091 Bytes
9dbb7e3
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
"""
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