summaryrefslogtreecommitdiff
path: root/py_modules/lsfg_vk/config_schema.py
blob: 8b1f29c7af9ff9111f0a9707d03ea8ce35862fbe (plain)
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
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
"""Small adapter for the upstream lsfg-vk v2 configuration format."""

import json
import tomllib
from typing import Any, Dict, TypedDict

ConfigurationData = Dict[str, Any]


class ProfileData(TypedDict):
    profiles: Dict[str, Dict[str, Any]]
    global_config: Dict[str, Any]


class UnsupportedConfigurationVersion(ValueError):
    def __init__(self, version: Any):
        self.version = version
        super().__init__("unsupported lsfg-vk configuration version")


PROFILE_DEFAULTS: Dict[str, Any] = {
    "active_in": [],
    "pacing_mode": "vsync",
    "multiplier": 2,
    "flow_scale": 0.8,
    "performance_mode": False,
    "override_present_mode": True,
    "preserve_swapchain_image_count": False,
}
GLOBAL_DEFAULTS: Dict[str, Any] = {"dll": "", "no_fp16": False}


def _toml_value(value: Any) -> str:
    if isinstance(value, bool):
        return str(value).lower()
    if isinstance(value, str):
        return json.dumps(value)
    if isinstance(value, list):
        return "[ " + ", ".join(_toml_value(item) for item in value) + " ]"
    return str(value)


def _normalize_active_in(value: Any) -> list[str]:
    if value in (None, ""):
        return []
    if isinstance(value, str):
        return [value]
    if not isinstance(value, list):
        raise ValueError("active_in must be a string or list of strings")
    return [str(item) for item in value if str(item)]


class ConfigurationManager:
    @staticmethod
    def is_discardable_legacy_version(version: Any) -> bool:
        return version is None or version == "1" or (
            isinstance(version, (int, float))
            and not isinstance(version, bool)
            and version < 2
        )

    @staticmethod
    def get_defaults() -> Dict[str, Any]:
        return {**GLOBAL_DEFAULTS, **PROFILE_DEFAULTS}

    @staticmethod
    def validate_config(config: Dict[str, Any]) -> Dict[str, Any]:
        result = {**PROFILE_DEFAULTS, **GLOBAL_DEFAULTS}
        result.update({key: value for key, value in config.items() if key in result})
        result["active_in"] = _normalize_active_in(result["active_in"])
        result["pacing_mode"] = str(result["pacing_mode"]).lower()
        if result["pacing_mode"] != "vsync":
            raise ValueError("pacing_mode must be vsync")
        result["multiplier"] = int(result["multiplier"])
        if result["multiplier"] < 1:
            raise ValueError("multiplier must be 1 or greater")
        result["flow_scale"] = float(result["flow_scale"])
        if not 0.25 <= result["flow_scale"] <= 1.0:
            raise ValueError("flow_scale must be between 0.25 and 1.0")
        for name in (
            "no_fp16",
            "performance_mode",
            "override_present_mode",
            "preserve_swapchain_image_count",
        ):
            result[name] = bool(result[name])
        result["dll"] = str(result["dll"] or "")
        return result

    @staticmethod
    def generate_toml_content_multi_profile(profile_data: ProfileData) -> str:
        global_config = {**GLOBAL_DEFAULTS, **profile_data.get("global_config", {})}
        lines = ["version = 2", "", "[global]"]
        if global_config["dll"]:
            lines.append(f"dll = {_toml_value(global_config['dll'])}")
        lines.append(f"allow_fp16 = {_toml_value(not bool(global_config['no_fp16']))}")
        profiles = sorted(profile_data["profiles"].items()) or [("", {})]
        for name, raw in profiles:
            config = ConfigurationManager.validate_config({**raw, **global_config})
            lines.extend(["", "[[profile]]", f"name = {_toml_value(name)}"])
            if config["active_in"]:
                lines.append(f"active_in = {_toml_value(config['active_in'])}")
            lines.extend([
                f"pacing_mode = {_toml_value(config['pacing_mode'])}",
                f"multiplier = {config['multiplier']}",
                f"flow_scale = {config['flow_scale']}",
                f"performance_mode = {_toml_value(config['performance_mode'])}",
                f"override_present_mode = {_toml_value(config['override_present_mode'])}",
                f"preserve_swapchain_image_count = {_toml_value(config['preserve_swapchain_image_count'])}",
            ])
        return "\n".join(lines) + "\n"

    @staticmethod
    def parse_toml_content_multi_profile(content: str) -> ProfileData:
        data = tomllib.loads(content)
        version = data.get("version")
        if version != 2:
            raise UnsupportedConfigurationVersion(version)
        raw_global = data.get("global", {})
        global_config = {
            "dll": str(raw_global.get("dll", "") or ""),
            "no_fp16": not bool(raw_global.get("allow_fp16", True)),
        }
        profiles: Dict[str, Dict[str, Any]] = {}
        for profile in data.get("profile", []):
            name = str(profile.get("name", ""))
            config = ConfigurationManager.validate_config({
                **profile,
                **global_config,
            })
            if name or config["active_in"]:
                profiles[name] = config
        return {"profiles": profiles, "global_config": global_config}