from __future__ import annotations
import os
import sys
import json
import threading
from typing import TYPE_CHECKING
from modules import cmd_args, errors
from modules.json_helpers import readfile, writefile
from modules.shared_legacy import LegacyOption
from modules.logger import log


if TYPE_CHECKING:
    from collections.abc import Callable
    from modules.options import OptionInfo
    from typing import Any

cmd_opts = cmd_args.parse_args()
compatibility_opts = ['clip_skip', 'uni_pc_lower_order_final', 'uni_pc_order']
secrets_pattern = ['_version', '_token', '_key', '_secret', '_password']


class Options:
    data_labels: dict[str, OptionInfo | LegacyOption]
    data: dict[str, Any]
    secrets: dict[str, Any]
    typemap = {int: float}
    debug = os.environ.get('SD_CONFIG_DEBUG', None) is not None
    secrets_debug = os.environ.get("SD_SECRETS_DEBUG", None) is not None

    def __init__(self, options_templates: dict[str, OptionInfo | LegacyOption] | None = None, restricted: set[str] | None = None, *, filename = '', secrets = ''):
        if options_templates is None:
            options_templates = {}
        if restricted is None:
            restricted = set()
        super().__setattr__('data_labels', options_templates)
        super().__setattr__('data', {k: v.default for k, v in options_templates.items()})
        super().__setattr__('secrets', {})
        self.filename: str = filename or cmd_opts.config
        self.secretsfn: str = secrets or cmd_opts.secrets
        self.restricted: set[str] = restricted
        self.legacy = [k for k, v in options_templates.items() if isinstance(v, LegacyOption)]
        self.load()

    def __getattr__(self, item):
        if item == 'secrets':
            return super().__getattribute__('secrets')
        if item == 'data':
            return super().__getattribute__('data')
        if item in self.secrets:
            if self.secrets_debug:
                fn = f"{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}"  # pylint: disable=protected-access
                log.trace(f"Secret: get={item} fn={fn}")
            return self.secrets[item]
        if item in self.data:
            return self.data[item]
        if item in self.data_labels:
            return self.data_labels[item].default
        return super().__getattribute__(item)  # pylint: disable=super-with-arguments

    def get(self, item):
        if item in self.secrets:
            if self.secrets_debug:
                fn = f'{sys._getframe(2).f_code.co_name}:{sys._getframe(1).f_code.co_name}' # pylint: disable=protected-access
                log.trace(f"Secret: get={item} fn={fn}")
            return self.secrets[item]
        if item in self.data:
            return self.data[item]
        if item in self.data_labels:
            return self.data_labels[item].default
        return super().__getattribute__(item)  # pylint: disable=super-with-arguments

    def __setattr__(self, key, value):  # pylint: disable=inconsistent-return-statements
        if (key in self.data_labels) or (key in self.data) or (key in self.secrets):
            if cmd_opts.freeze:
                log.warning(f"Settings are frozen: {key}")
                return
            if cmd_opts.hide_ui_dir_config and key in self.restricted:
                log.warning(f"Settings key is restricted: {key}")
                return
            if self.debug:
                log.trace(f"Settings set: {key}={value}")
            if key in self.legacy:
                log.warning(f"Settings set: {key}={value} legacy")
            if any(key.endswith(pattern) for pattern in secrets_pattern):
                if self.secrets_debug:
                    log.trace(f"Secret: set={key}")
                self.secrets[key] = value
            else:
                self.data[key] = value
            return
        return super().__setattr__(key, value)  # pylint: disable=super-with-arguments

    def set(self, key, value, force=False):
        """sets an option and calls its onchange callback, returning True if the option changed and False otherwise"""
        if key in self.secrets:
            oldval = self.secrets.get(key, None)
        else:
            oldval = self.data.get(key, None)
        if oldval is None:
            if key in self.data_labels:
                oldval = self.data_labels[key].default
            else:
                log.warning(f'Settings: key={key} value={value} unknown')
                return False
        if oldval == value and not force:
            return False
        try:
            setattr(self, key, value)
        except RuntimeError:
            return False
        # compatibility_opts (e.g. clip_skip) live in data without a data_labels entry
        func = self.data_labels[key].onchange if key in self.data_labels else None
        if func is not None:
            try:
                func()
            except Exception as err:
                log.error(f'Error in onchange callback: {key} {value} {err}')
                errors.display(err, 'Error in onchange callback')
                setattr(self, key, oldval)
                return False
        return True

    def get_default(self, key):
        """returns the default value for the key"""
        data_label = self.data_labels.get(key)
        return data_label.default if data_label is not None else None

    def list(self):
        """list all visible options"""
        components = [k for k, v in self.data_labels.items() if v.visible]
        return components

    def save_atomic(self, silent=False):
        if self.debug:
            log.debug(f'Settings: save settings="{self.filename}" secrets="{self.secretsfn}" cmd="{cmd_opts.config}" cwd="{os.getcwd()}"')
        filename = os.path.abspath(self.filename)
        secretsfn = os.path.abspath(self.secretsfn)
        if cmd_opts.freeze:
            log.warning(f'Setting: fn="{filename}" save disabled')
            return
        try:
            unused_settings = []
            if self.debug:
                log.debug(f'Settings: total={len(self.data.keys())} secrets={len(self.secrets.keys())} known={len(self.data_labels.keys())}')

            all_options = self.data | self.secrets
            diff = {}
            for k, v in all_options.items():
                if k in self.data_labels:
                    default = self.data_labels[k].default
                    if isinstance(v, list):
                        if (len(default) != len(v) or set(default) != set(v)): # list order is non-deterministic
                            diff[k] = v
                            if self.debug:
                                log.trace(f'Settings changed: {k}={v} default={default}')
                    elif self.data_labels[k].default != v:
                        diff[k] = v
                        if self.debug:
                            log.trace(f'Settings changed: {k}={v} default={default}')
                else:
                    if k not in compatibility_opts:
                        diff[k] = v
                        if not k.startswith('uiux_'):
                            unused_settings.append(k)
                        if self.debug:
                            log.trace(f'Settings unknown: {k}={v}')
            options = {}
            secrets = {}
            for k, v in diff.items():
                if any(k.endswith(pattern) for pattern in secrets_pattern):
                    secrets[k] = v
                else:
                    options[k] = v
            writefile(options, filename, silent=silent)
            writefile(secrets, secretsfn, silent=silent)
            if self.debug:
                log.trace(f'Settings save: count={len(diff.keys())} {diff}')
            if len(unused_settings) > 0:
                log.debug(f"Settings: unused={unused_settings}")
        except Exception as err:
            log.error(f'Settings: config="{filename}" secrets="{secretsfn}" {err}')

    def save(self, silent=False):
        threading.Thread(target=self.save_atomic, args=(silent,)).start()

    def same_type(self, x, y):
        if x is None or y is None:
            return True
        type_x = self.typemap.get(type(x), type(x))
        type_y = self.typemap.get(type(y), type(y))
        return type_x == type_y

    def load(self):
        filename = os.path.abspath(self.filename)
        secretsfn = os.path.abspath(self.secretsfn)
        if not os.path.isfile(filename):
            log.debug(f'Settings: config="{filename}" secrets="{secretsfn}" created')
            self.save()
            return
        self.data = readfile(filename, lock=True, as_type="dict")
        self.secrets = readfile(secretsfn, lock=True, as_type="dict")
        if self.data.get('quicksettings') is not None and self.data.get('quicksettings_list') is None:
            self.data['quicksettings_list'] = [i.strip() for i in self.data.get('quicksettings', '').split(',')]
        unknown_settings = []
        for k, v in self.data.items():
            info = self.data_labels.get(k, None)
            if info is not None:
                if not info.validate(k, v):
                    self.data[k] = info.default
            if info is not None and not self.same_type(info.default, v):
                log.warning(f"Setting validation: {k}={v} ({type(v).__name__} expected={type(info.default).__name__})")
                self.data[k] = info.default
            if info is None and k not in compatibility_opts and not k.startswith('uiux_'):
                unknown_settings.append(k)
        if len(unknown_settings) > 0:
            log.warning(f"Setting validation: unknown={unknown_settings}")

    def onchange(self, key, func: Callable, call=True):
        item = self.data_labels.get(key)
        item.onchange = func
        if call:
            func()

    def dumpjson(self):
        d = {k: self.data.get(k, self.data_labels.get(k).default) for k in self.data_labels.keys()}
        metadata = {
            k: {
                "is_stored": k in self.data and self.data[k] != self.data_labels[k].default, # pylint: disable=unnecessary-dict-index-lookup
                "tab_name": v.section[0]
            } for k, v in self.data_labels.items()
        }
        return json.dumps({"values": d, "metadata": metadata})

    def add_option(self, key, info):
        self.data_labels[key] = info

    def reorder(self):
        """reorder settings so that all items related to section always go together"""
        section_ids = {}
        settings_items = self.data_labels.items()
        for _k, item in settings_items:
            if item.section not in section_ids:
                section_ids[item.section] = len(section_ids)
        self.data_labels = dict(sorted(settings_items, key=lambda x: section_ids[x[1].section]))

    def cast_value(self, key, value):
        """casts an arbitrary to the same type as this setting's value with key
        Example: cast_value("eta_noise_seed_delta", "12") -> returns 12 (an int rather than str)
        """
        if value is None:
            return None
        default_value = self.data_labels[key].default
        if default_value is None:
            default_value = getattr(self, key, None)
        if default_value is None:
            return None
        expected_type = type(default_value)
        if expected_type == bool and value == "False":
            value = False
        elif expected_type == type(value):
            pass
        else:
            value = expected_type(value)
        return value