226 lines
6.6 KiB
Python
226 lines
6.6 KiB
Python
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
import os
|
|
import re
|
|
from pathlib import Path
|
|
from typing import Any
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
_SAFE_ROOT = Path.cwd()
|
|
|
|
_ENV_VAR_PATTERN = re.compile(r"\$\{([A-Za-z_][A-Za-z0-9_]*)\}")
|
|
_ESCAPED_ENV_PATTERN = re.compile(r"\$\$\{\}")
|
|
|
|
|
|
def set_safe_config_root(path: str | Path) -> None:
|
|
global _SAFE_ROOT
|
|
_SAFE_ROOT = Path(path).resolve()
|
|
|
|
|
|
def merge_patch(target: dict, patch: dict) -> dict:
|
|
"""RFC 7396 JSON Merge Patch."""
|
|
result = dict(target)
|
|
|
|
for key, value in patch.items():
|
|
if value is None:
|
|
result.pop(key, None)
|
|
elif isinstance(value, dict) and isinstance(result.get(key), dict):
|
|
result[key] = merge_patch(result[key], value)
|
|
else:
|
|
result[key] = value
|
|
|
|
return result
|
|
|
|
|
|
def resolve_env_vars(config: dict) -> dict:
|
|
"""Replace ${ENV_VAR} patterns with environment variable values.
|
|
$${} is preserved as literal ${} (escape sequence).
|
|
"""
|
|
|
|
def _resolve(value: Any) -> Any:
|
|
if isinstance(value, str):
|
|
value = _ESCAPED_ENV_PATTERN.sub("\x00ESCAPED\x00", value)
|
|
value = _ENV_VAR_PATTERN.sub(
|
|
lambda m: os.getenv(m.group(1), ""),
|
|
value,
|
|
)
|
|
value = value.replace("\x00ESCAPED\x00", "${}")
|
|
return value
|
|
if isinstance(value, dict):
|
|
return {k: _resolve(v) for k, v in value.items()}
|
|
if isinstance(value, list):
|
|
return [_resolve(v) for v in value]
|
|
return value
|
|
|
|
return _resolve(config)
|
|
|
|
|
|
def resolve_includes(
|
|
config_path: str | Path,
|
|
base_dir: str | Path | None = None,
|
|
max_depth: int = 10,
|
|
) -> dict:
|
|
"""Load a config file and resolve $include directives recursively.
|
|
|
|
$include can be:
|
|
- A string path to another config file (deep merged)
|
|
- An array of paths (each deep merged in order)
|
|
"""
|
|
config_path = Path(config_path)
|
|
base_dir = Path(base_dir) if base_dir else config_path.parent
|
|
base_dir = base_dir.resolve()
|
|
|
|
return _load_and_merge(config_path, base_dir, max_depth, _visited=set())
|
|
|
|
|
|
def _load_and_merge(
|
|
file_path: Path,
|
|
base_dir: Path,
|
|
max_depth: int,
|
|
_visited: set,
|
|
_depth: int = 0,
|
|
) -> dict:
|
|
if _depth > max_depth:
|
|
raise RecursionError(f"$include depth exceeded {max_depth} for {file_path}")
|
|
|
|
resolved = file_path.resolve()
|
|
if not str(resolved).startswith(str(_SAFE_ROOT)):
|
|
raise ValueError(f"$include path outside safe root: {resolved}")
|
|
|
|
if resolved in _visited:
|
|
logger.warning("Circular $include detected: %s", resolved)
|
|
return {}
|
|
|
|
_visited.add(resolved)
|
|
|
|
config = _load_config_file(resolved)
|
|
result = {}
|
|
|
|
for key, value in config.items():
|
|
if key == "$include":
|
|
result = _process_includes(value, base_dir, max_depth, _visited, _depth + 1)
|
|
elif isinstance(value, dict):
|
|
if "$include" in value:
|
|
inner = _process_includes(value.pop("$include"), base_dir, max_depth, _visited, _depth + 1)
|
|
inner.update({k: v for k, v in value.items() if k != "$include"})
|
|
result[key] = inner
|
|
else:
|
|
result[key] = value
|
|
else:
|
|
result[key] = value
|
|
|
|
return result
|
|
|
|
|
|
def _process_includes(
|
|
include_value: Any,
|
|
base_dir: Path,
|
|
max_depth: int,
|
|
_visited: set,
|
|
depth: int,
|
|
) -> dict:
|
|
result: dict = {}
|
|
|
|
if isinstance(include_value, str):
|
|
paths = [include_value]
|
|
elif isinstance(include_value, list):
|
|
paths = include_value
|
|
else:
|
|
logger.warning("$include value must be a string or list, got %s", type(include_value))
|
|
return result
|
|
|
|
for rel_path in paths:
|
|
included_path = base_dir / rel_path
|
|
included = _load_and_merge(included_path, included_path.parent, max_depth, _visited, depth)
|
|
result = merge_patch(result, included)
|
|
|
|
return result
|
|
|
|
|
|
def _load_config_file(file_path: Path) -> dict:
|
|
suffix = file_path.suffix.lower()
|
|
if suffix == ".json":
|
|
return json.loads(file_path.read_text(encoding="utf-8"))
|
|
if suffix == ".json5":
|
|
return _load_json5(file_path)
|
|
if suffix in (".yaml", ".yml"):
|
|
return _load_yaml(file_path)
|
|
if suffix == ".toml":
|
|
return _load_toml(file_path)
|
|
return json.loads(file_path.read_text(encoding="utf-8"))
|
|
|
|
|
|
def _load_json5(file_path: Path) -> dict:
|
|
raw = file_path.read_text(encoding="utf-8")
|
|
cleaned = _strip_json5_comments(raw)
|
|
return json.loads(cleaned)
|
|
|
|
|
|
def _strip_json5_comments(raw: str) -> str:
|
|
lines = []
|
|
in_multiline = False
|
|
for line in raw.split("\n"):
|
|
stripped = line.strip()
|
|
if in_multiline:
|
|
if "*/" in stripped:
|
|
in_multiline = False
|
|
after = line[line.index("*/") + 2 :]
|
|
if after.strip():
|
|
lines.append(after)
|
|
continue
|
|
if "/*" in stripped and "*/" not in stripped:
|
|
in_multiline = True
|
|
before = line[: line.index("/*")]
|
|
lines.append(before)
|
|
continue
|
|
if "//" in stripped:
|
|
idx = line.index("//")
|
|
candidate = line[:idx]
|
|
if not _is_inside_string(candidate):
|
|
lines.append(candidate)
|
|
continue
|
|
lines.append(line)
|
|
|
|
result = "\n".join(lines)
|
|
result = re.sub(r",\s*([}\]])", r"\1", result)
|
|
return result
|
|
|
|
|
|
def _is_inside_string(text: str) -> bool:
|
|
dq_count = text.count('"') - text.count('\\"')
|
|
sq_count = text.count("'") - text.count("\\'")
|
|
return dq_count % 2 != 0 or sq_count % 2 != 0
|
|
|
|
|
|
def _load_yaml(file_path: Path) -> dict:
|
|
try:
|
|
import yaml
|
|
|
|
return yaml.safe_load(file_path.read_text(encoding="utf-8")) or {}
|
|
except ImportError:
|
|
logger.warning("PyYAML not installed, falling back to JSON for %s", file_path)
|
|
return json.loads(file_path.read_text(encoding="utf-8"))
|
|
|
|
|
|
def _load_toml(file_path: Path) -> dict:
|
|
try:
|
|
import tomllib
|
|
|
|
return tomllib.loads(file_path.read_text(encoding="utf-8"))
|
|
except ImportError:
|
|
logger.warning("tomllib not available, falling back to JSON for %s", file_path)
|
|
return json.loads(file_path.read_text(encoding="utf-8"))
|
|
|
|
|
|
def load_file_config(file_path: str | Path) -> dict:
|
|
"""Load a configuration file with full pipeline:
|
|
File → JSON5 parse → $include resolution → ${ENV} replacement.
|
|
"""
|
|
config_path = Path(file_path)
|
|
config = resolve_includes(config_path, config_path.parent)
|
|
config = resolve_env_vars(config)
|
|
return config
|