feat(channel/config): 新增完整的配置模块实现
新增了channel目录下的配置相关模块,包含默认配置、配置迁移、配置差异计算、文件加载器、重载计划、重载器以及配置校验功能,完善了配置管理体系
This commit is contained in:
parent
5d098c3423
commit
4962d1f6f8
3
backend/package/yuxi/channel/config/__init__.py
Normal file
3
backend/package/yuxi/channel/config/__init__.py
Normal file
@ -0,0 +1,3 @@
|
||||
from yuxi.channel.config.defaults import TIMEOUT, TimeoutConfig
|
||||
|
||||
__all__ = ["TIMEOUT", "TimeoutConfig"]
|
||||
66
backend/package/yuxi/channel/config/defaults.py
Normal file
66
backend/package/yuxi/channel/config/defaults.py
Normal file
@ -0,0 +1,66 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
# ── Streaming defaults ────────────────────────────────────
|
||||
|
||||
STREAMING_PREVIEW_THROTTLE_MS: int = 160
|
||||
STREAMING_PREVIEW_MIN_INITIAL_CHARS: int = 18
|
||||
|
||||
BLOCK_STREAMING_BREAK: str = "text_end"
|
||||
BLOCK_STREAMING_CHUNK_MIN_CHARS: int = 800
|
||||
BLOCK_STREAMING_CHUNK_MAX_CHARS: int = 1200
|
||||
BLOCK_STREAMING_CHUNK_BREAK_PREFERENCE: str = "paragraph"
|
||||
|
||||
BLOCK_STREAMING_COALESCE_MIN_CHARS: int | None = None
|
||||
BLOCK_STREAMING_COALESCE_MAX_CHARS: int | None = None
|
||||
BLOCK_STREAMING_COALESCE_IDLE_MS: int = 1000
|
||||
BLOCK_STREAMING_COALESCE_JOINER: str = "\n\n"
|
||||
|
||||
SSE_MAX_REASONING_CHARS: int = 4096
|
||||
|
||||
|
||||
class HttpTimeoutConfig(BaseModel):
|
||||
short: float = 10.0
|
||||
normal: float = 30.0
|
||||
long: float = 60.0
|
||||
|
||||
|
||||
class WebhookTimeoutConfig(BaseModel):
|
||||
process: float = 120.0
|
||||
body_read: float = 5.0
|
||||
|
||||
|
||||
class GatewayTimeoutConfig(BaseModel):
|
||||
send: float = 15.0
|
||||
stream_rpc: float = 30.0
|
||||
connect: float = 5.0
|
||||
heartbeat: float = 90.0
|
||||
heartbeat_interval: float = 30.0
|
||||
negotiation: float = 5.0
|
||||
|
||||
|
||||
class CronTimeoutConfig(BaseModel):
|
||||
gather: float = 15.0
|
||||
task: float = 10.0
|
||||
delivery: float = 300.0
|
||||
|
||||
|
||||
class DoctorTimeoutConfig(BaseModel):
|
||||
check: float = 30.0
|
||||
|
||||
|
||||
class TimeoutConfig(BaseModel):
|
||||
http: HttpTimeoutConfig = HttpTimeoutConfig()
|
||||
webhook: WebhookTimeoutConfig = WebhookTimeoutConfig()
|
||||
gateway: GatewayTimeoutConfig = GatewayTimeoutConfig()
|
||||
cron: CronTimeoutConfig = CronTimeoutConfig()
|
||||
doctor: DoctorTimeoutConfig = DoctorTimeoutConfig()
|
||||
|
||||
stream_finalize: float = 30.0
|
||||
stream_queue_poll: float = 0.5
|
||||
connection: float = 10.0
|
||||
job_timeout: int = 900
|
||||
|
||||
|
||||
TIMEOUT = TimeoutConfig()
|
||||
135
backend/package/yuxi/channel/config/diff.py
Normal file
135
backend/package/yuxi/channel/config/diff.py
Normal file
@ -0,0 +1,135 @@
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from .snapshot import AccountConfigEntry, is_sensitive_config_path
|
||||
|
||||
_SENTINEL = object()
|
||||
|
||||
|
||||
def _redact_path(path: str) -> str:
|
||||
"""对敏感配置路径进行脱敏处理。"""
|
||||
if not is_sensitive_config_path(path):
|
||||
return path
|
||||
parts = path.rsplit(".", 1)
|
||||
if len(parts) == 2:
|
||||
return f"{parts[0]}.<redacted>"
|
||||
return "<redacted>"
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConfigDiff:
|
||||
added: list[AccountConfigEntry] = field(default_factory=list)
|
||||
removed: list[AccountConfigEntry] = field(default_factory=list)
|
||||
changed: list[AccountConfigEntry] = field(default_factory=list)
|
||||
unchanged: list[AccountConfigEntry] = field(default_factory=list)
|
||||
routes_changed: bool = False
|
||||
changed_paths: list[str] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def has_changes(self) -> bool:
|
||||
return bool(self.added or self.removed or self.changed or self.routes_changed or self.changed_paths)
|
||||
|
||||
@property
|
||||
def summary(self) -> str:
|
||||
parts = []
|
||||
if self.added:
|
||||
parts.append(f"added={len(self.added)}")
|
||||
if self.removed:
|
||||
parts.append(f"removed={len(self.removed)}")
|
||||
if self.changed:
|
||||
parts.append(f"changed={len(self.changed)}")
|
||||
if self.routes_changed:
|
||||
parts.append("routes_changed")
|
||||
if self.changed_paths:
|
||||
redacted_paths = [_redact_path(p) for p in self.changed_paths[:5]]
|
||||
paths_preview = ", ".join(redacted_paths)
|
||||
if len(self.changed_paths) > 5:
|
||||
paths_preview += f" (+{len(self.changed_paths) - 5} more)"
|
||||
parts.append(f"paths=[{paths_preview}]")
|
||||
return f"ConfigDiff({', '.join(parts)})"
|
||||
|
||||
|
||||
def compute_diff(
|
||||
prev: list[AccountConfigEntry],
|
||||
next_: list[AccountConfigEntry],
|
||||
) -> ConfigDiff:
|
||||
prev = prev or []
|
||||
next_ = next_ or []
|
||||
|
||||
prev_map = {f"{e.channel_type}:{e.account_id}": e for e in prev}
|
||||
next_map = {f"{e.channel_type}:{e.account_id}": e for e in next_}
|
||||
|
||||
diff = ConfigDiff()
|
||||
|
||||
for key, entry in next_map.items():
|
||||
if key not in prev_map:
|
||||
diff.added.append(entry)
|
||||
|
||||
for key, entry in prev_map.items():
|
||||
if key not in next_map:
|
||||
diff.removed.append(entry)
|
||||
|
||||
for key, entry in next_map.items():
|
||||
prev_entry = prev_map.get(key)
|
||||
if prev_entry is None:
|
||||
continue
|
||||
if (
|
||||
entry.enabled != prev_entry.enabled
|
||||
or entry.configured != prev_entry.configured
|
||||
or entry.config_hash != prev_entry.config_hash
|
||||
):
|
||||
diff.changed.append(entry)
|
||||
else:
|
||||
diff.unchanged.append(entry)
|
||||
|
||||
return diff
|
||||
|
||||
|
||||
def _deep_equal(a, b) -> bool:
|
||||
if a is b:
|
||||
return True
|
||||
if type(a) is not type(b):
|
||||
return False
|
||||
if isinstance(a, dict):
|
||||
if len(a) != len(b):
|
||||
return False
|
||||
return all(k in b and _deep_equal(a[k], b[k]) for k in a)
|
||||
if isinstance(a, (list, tuple)):
|
||||
if len(a) != len(b):
|
||||
return False
|
||||
return all(_deep_equal(x, y) for x, y in zip(a, b, strict=False))
|
||||
return a == b
|
||||
|
||||
|
||||
def diff_config_paths(prev, next_, prefix: str = "") -> list[str]:
|
||||
if prev is next_:
|
||||
return []
|
||||
|
||||
if isinstance(prev, dict) and isinstance(next_, dict):
|
||||
keys = set(prev.keys()) | set(next_.keys())
|
||||
paths: list[str] = []
|
||||
for key in keys:
|
||||
prev_value = prev.get(key, _SENTINEL)
|
||||
next_value = next_.get(key, _SENTINEL)
|
||||
if prev_value is _SENTINEL and next_value is _SENTINEL:
|
||||
continue
|
||||
child_prefix = f"{prefix}.{key}" if prefix else key
|
||||
child_paths = diff_config_paths(prev_value, next_value, child_prefix)
|
||||
if child_paths:
|
||||
paths.extend(child_paths)
|
||||
return paths
|
||||
|
||||
if isinstance(prev, (list, tuple)) and isinstance(next_, (list, tuple)):
|
||||
if _deep_equal(prev, next_):
|
||||
return []
|
||||
max_len = max(len(prev), len(next_))
|
||||
paths: list[str] = []
|
||||
for i in range(max_len):
|
||||
prev_item = prev[i] if i < len(prev) else _SENTINEL
|
||||
next_item = next_[i] if i < len(next_) else _SENTINEL
|
||||
item_prefix = f"{prefix}[{i}]"
|
||||
child_paths = diff_config_paths(prev_item, next_item, item_prefix)
|
||||
if child_paths:
|
||||
paths.extend(child_paths)
|
||||
return paths
|
||||
|
||||
return [prefix or "<root>"]
|
||||
225
backend/package/yuxi/channel/config/file_loader.py
Normal file
225
backend/package/yuxi/channel/config/file_loader.py
Normal file
@ -0,0 +1,225 @@
|
||||
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
|
||||
164
backend/package/yuxi/channel/config/migration.py
Normal file
164
backend/package/yuxi/channel/config/migration.py
Normal file
@ -0,0 +1,164 @@
|
||||
"""配置版本迁移 — schema_version + 迁移规则 + 自动执行 + 回滚"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import copy
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
CURRENT_SCHEMA_VERSION = 1
|
||||
|
||||
ConfigDict = dict[str, Any]
|
||||
MigrationFn = Callable[[ConfigDict], ConfigDict]
|
||||
|
||||
|
||||
@dataclass
|
||||
class MigrationRule:
|
||||
from_version: int
|
||||
to_version: int
|
||||
description: str
|
||||
migrate: MigrationFn
|
||||
|
||||
|
||||
_migration_registry: dict[int, MigrationRule] = {}
|
||||
|
||||
|
||||
def register_migration(from_version: int, to_version: int, description: str):
|
||||
"""装饰器:注册配置迁移规则"""
|
||||
|
||||
if to_version != from_version + 1:
|
||||
raise ValueError(
|
||||
f"迁移规则版本不连续: v{from_version} → v{to_version}, "
|
||||
f"期望 to_version={from_version + 1}"
|
||||
)
|
||||
|
||||
def decorator(fn: MigrationFn) -> MigrationFn:
|
||||
existing = _migration_registry.get(from_version)
|
||||
if existing is not None:
|
||||
logger.warning(
|
||||
"重复注册迁移规则 v%d → v%d: 旧规则 (%s) 将被覆盖为新规则 (%s)",
|
||||
from_version,
|
||||
to_version,
|
||||
existing.description,
|
||||
description,
|
||||
)
|
||||
_migration_registry[from_version] = MigrationRule(
|
||||
from_version=from_version,
|
||||
to_version=to_version,
|
||||
description=description,
|
||||
migrate=fn,
|
||||
)
|
||||
logger.info(
|
||||
"注册配置迁移规则: v%d → v%d (%s)",
|
||||
from_version,
|
||||
to_version,
|
||||
description,
|
||||
)
|
||||
return fn
|
||||
|
||||
return decorator
|
||||
|
||||
|
||||
def get_migration_path(current_version: int) -> list[MigrationRule]:
|
||||
if current_version >= CURRENT_SCHEMA_VERSION:
|
||||
return []
|
||||
path: list[MigrationRule] = []
|
||||
version = current_version
|
||||
while version < CURRENT_SCHEMA_VERSION:
|
||||
rule = _migration_registry.get(version)
|
||||
if rule is None:
|
||||
logger.warning("缺少迁移规则: v%d → v%d", version, version + 1)
|
||||
break
|
||||
path.append(rule)
|
||||
version = rule.to_version
|
||||
return path
|
||||
|
||||
|
||||
def migrate_config(
|
||||
config: ConfigDict,
|
||||
current_version: int | None = None,
|
||||
*,
|
||||
dry_run: bool = False,
|
||||
) -> tuple[ConfigDict, int, list[MigrationRule]]:
|
||||
"""执行配置迁移,返回 (新配置, 新版本号, 已应用的迁移规则列表)。
|
||||
|
||||
失败时自动恢复到原始配置(回滚)。
|
||||
|
||||
dry_run=True 时只返回迁移后的配置预览,不实际修改任何数据。
|
||||
"""
|
||||
if current_version is None:
|
||||
version = config.get("schema_version", 0)
|
||||
else:
|
||||
version = current_version
|
||||
|
||||
path = get_migration_path(version)
|
||||
if not path:
|
||||
return config, version, []
|
||||
|
||||
original = copy.deepcopy(config)
|
||||
working = copy.deepcopy(config)
|
||||
applied: list[MigrationRule] = []
|
||||
last_version = version
|
||||
|
||||
for rule in path:
|
||||
try:
|
||||
working = rule.migrate(working)
|
||||
working["schema_version"] = rule.to_version
|
||||
last_version = rule.to_version
|
||||
applied.append(rule)
|
||||
logger.info(
|
||||
"配置迁移成功: v%d → v%d (%s)%s",
|
||||
rule.from_version,
|
||||
rule.to_version,
|
||||
rule.description,
|
||||
" [dry-run]" if dry_run else "",
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"配置迁移失败: v%d → v%d (%s), 回滚到 v%d",
|
||||
rule.from_version,
|
||||
rule.to_version,
|
||||
rule.description,
|
||||
version,
|
||||
)
|
||||
return original, version, applied
|
||||
|
||||
if dry_run:
|
||||
logger.info(
|
||||
"迁移 dry-run 完成: v%d → v%d, 共 %d 步",
|
||||
version,
|
||||
last_version,
|
||||
len(applied),
|
||||
)
|
||||
return working, last_version, applied
|
||||
|
||||
return working, last_version, applied
|
||||
|
||||
|
||||
def backup_config(config: ConfigDict, reason: str = "迁移前备份") -> ConfigDict:
|
||||
logger.debug("备份配置: %s", reason)
|
||||
return copy.deepcopy(config)
|
||||
|
||||
|
||||
# ── 内置迁移规则 ────────────────────────────────────────
|
||||
|
||||
|
||||
@register_migration(
|
||||
from_version=0,
|
||||
to_version=1,
|
||||
description="初始化 schema_version 字段",
|
||||
)
|
||||
def migration_v0_to_v1(config: ConfigDict) -> ConfigDict:
|
||||
result = dict(config)
|
||||
|
||||
if "approval" not in result:
|
||||
result["approval"] = {
|
||||
"enabled": False,
|
||||
"approvers": [],
|
||||
"auto_approve_internal": True,
|
||||
}
|
||||
return result
|
||||
337
backend/package/yuxi/channel/config/reload_plan.py
Normal file
337
backend/package/yuxi/channel/config/reload_plan.py
Normal file
@ -0,0 +1,337 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from enum import StrEnum
|
||||
from functools import lru_cache
|
||||
|
||||
from .snapshot import is_sensitive_config_path
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _redact_path(path: str) -> str:
|
||||
"""对敏感配置路径进行脱敏处理。"""
|
||||
if not is_sensitive_config_path(path):
|
||||
return path
|
||||
parts = path.rsplit(".", 1)
|
||||
if len(parts) == 2:
|
||||
return f"{parts[0]}.<redacted>"
|
||||
return "<redacted>"
|
||||
|
||||
|
||||
class ReloadMode(StrEnum):
|
||||
HOT = "hot"
|
||||
RESTART = "restart"
|
||||
HYBRID = "hybrid"
|
||||
OFF = "off"
|
||||
|
||||
|
||||
class ReloadAction(StrEnum):
|
||||
RELOAD_HOOKS = "reload_hooks"
|
||||
RESTART_GMAIL_WATCHER = "restart_gmail_watcher"
|
||||
RESTART_CRON = "restart_cron"
|
||||
RESTART_HEARTBEAT = "restart_heartbeat"
|
||||
RESTART_HEALTH_MONITOR = "restart_health_monitor"
|
||||
RELOAD_PLUGINS = "reload_plugins"
|
||||
DISPOSE_MCP_RUNTIMES = "dispose_mcp_runtimes"
|
||||
RESTART_CHANNEL = "restart_channel"
|
||||
RESTART_GATEWAY = "restart_gateway"
|
||||
|
||||
@classmethod
|
||||
def _missing_(cls, value: object) -> "ReloadAction | None":
|
||||
if isinstance(value, str) and value.startswith("restart_channel:"):
|
||||
return cls._create_dynamic(value)
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def _create_dynamic(cls, value: str) -> "ReloadAction":
|
||||
obj = str.__new__(cls, value)
|
||||
obj._name_ = value
|
||||
obj._value_ = value
|
||||
return obj
|
||||
|
||||
|
||||
class ReloadRuleKind(StrEnum):
|
||||
RESTART = "restart"
|
||||
HOT = "hot"
|
||||
NONE = "none"
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True)
|
||||
class ReloadRule:
|
||||
prefix: str
|
||||
kind: ReloadRuleKind
|
||||
actions: tuple[ReloadAction, ...] = ()
|
||||
|
||||
def matches(self, path: str) -> bool:
|
||||
return path == self.prefix or path.startswith(f"{self.prefix}.")
|
||||
|
||||
|
||||
BASE_RELOAD_RULES: tuple[ReloadRule, ...] = (
|
||||
ReloadRule("gateway.remote", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("gateway.reload", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule(
|
||||
"gateway.channel_health_check_minutes",
|
||||
kind=ReloadRuleKind.HOT,
|
||||
actions=(ReloadAction.RESTART_HEALTH_MONITOR,),
|
||||
),
|
||||
ReloadRule(
|
||||
"gateway.channel_stale_event_threshold_minutes",
|
||||
kind=ReloadRuleKind.HOT,
|
||||
actions=(ReloadAction.RESTART_HEALTH_MONITOR,),
|
||||
),
|
||||
ReloadRule(
|
||||
"gateway.channel_max_restarts_per_hour",
|
||||
kind=ReloadRuleKind.HOT,
|
||||
actions=(ReloadAction.RESTART_HEALTH_MONITOR,),
|
||||
),
|
||||
ReloadRule("diagnostics.stuck_session_warn_ms", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("diagnostics.stuck_session_abort_ms", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("hooks.gmail", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RESTART_GMAIL_WATCHER,)),
|
||||
ReloadRule("hooks", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RELOAD_HOOKS,)),
|
||||
ReloadRule("agents.defaults.heartbeat", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RESTART_HEARTBEAT,)),
|
||||
ReloadRule("agents.defaults.models", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RESTART_HEARTBEAT,)),
|
||||
ReloadRule("agents.defaults.model", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RESTART_HEARTBEAT,)),
|
||||
ReloadRule("models.pricing", kind=ReloadRuleKind.RESTART),
|
||||
ReloadRule("models", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RESTART_HEARTBEAT,)),
|
||||
ReloadRule("agents.list", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RESTART_HEARTBEAT,)),
|
||||
ReloadRule("agent.heartbeat", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RESTART_HEARTBEAT,)),
|
||||
ReloadRule("cron", kind=ReloadRuleKind.HOT, actions=(ReloadAction.RESTART_CRON,)),
|
||||
ReloadRule("mcp", kind=ReloadRuleKind.HOT, actions=(ReloadAction.DISPOSE_MCP_RUNTIMES,)),
|
||||
ReloadRule("plugins.load", kind=ReloadRuleKind.RESTART),
|
||||
ReloadRule("plugins.installs", kind=ReloadRuleKind.RESTART),
|
||||
)
|
||||
|
||||
BASE_RELOAD_RULES_TAIL: tuple[ReloadRule, ...] = (
|
||||
ReloadRule("meta", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("identity", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("wizard", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("logging", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("agents", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("tools", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("bindings", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("audio", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("agent", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("routing", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("messages", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("session", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("talk", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("skills", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("secrets", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule(
|
||||
"plugins",
|
||||
kind=ReloadRuleKind.HOT,
|
||||
actions=(ReloadAction.RELOAD_PLUGINS, ReloadAction.DISPOSE_MCP_RUNTIMES),
|
||||
),
|
||||
ReloadRule("ui", kind=ReloadRuleKind.NONE),
|
||||
ReloadRule("gateway", kind=ReloadRuleKind.RESTART),
|
||||
ReloadRule("discovery", kind=ReloadRuleKind.RESTART),
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class GatewayReloadPlan:
|
||||
mode: ReloadMode = ReloadMode.HYBRID
|
||||
changed_paths: list[str] = field(default_factory=list)
|
||||
restart_gateway: bool = False
|
||||
restart_reasons: list[str] = field(default_factory=list)
|
||||
hot_reasons: list[str] = field(default_factory=list)
|
||||
reload_hooks: bool = False
|
||||
restart_gmail_watcher: bool = False
|
||||
restart_cron: bool = False
|
||||
restart_heartbeat: bool = False
|
||||
restart_health_monitor: bool = False
|
||||
reload_plugins: bool = False
|
||||
restart_channels: set[str] = field(default_factory=set)
|
||||
dispose_mcp_runtimes: bool = False
|
||||
noop_paths: list[str] = field(default_factory=list)
|
||||
|
||||
@property
|
||||
def has_channel_restarts(self) -> bool:
|
||||
return bool(self.restart_channels)
|
||||
|
||||
@property
|
||||
def requires_full_restart(self) -> bool:
|
||||
return self.restart_gateway
|
||||
|
||||
@property
|
||||
def is_noop(self) -> bool:
|
||||
return (
|
||||
not self.restart_gateway
|
||||
and not self.hot_reasons
|
||||
and not self.reload_hooks
|
||||
and not self.restart_gmail_watcher
|
||||
and not self.restart_cron
|
||||
and not self.restart_heartbeat
|
||||
and not self.restart_health_monitor
|
||||
and not self.reload_plugins
|
||||
and not self.dispose_mcp_runtimes
|
||||
and not self.restart_channels
|
||||
)
|
||||
|
||||
@property
|
||||
def summary(self) -> str:
|
||||
parts = [f"mode={self.mode.value}"]
|
||||
if self.restart_gateway:
|
||||
redacted_reasons = [_redact_path(p) for p in self.restart_reasons[:3]]
|
||||
parts.append(f"restart({', '.join(redacted_reasons)})")
|
||||
if self.hot_reasons:
|
||||
redacted_hot = [_redact_path(p) for p in self.hot_reasons[:5]]
|
||||
hot_preview = ", ".join(redacted_hot)
|
||||
parts.append(f"hot=[{hot_preview}]")
|
||||
flags = []
|
||||
if self.reload_hooks:
|
||||
flags.append("reload_hooks")
|
||||
if self.restart_cron:
|
||||
flags.append("restart_cron")
|
||||
if self.restart_heartbeat:
|
||||
flags.append("restart_heartbeat")
|
||||
if self.restart_health_monitor:
|
||||
flags.append("restart_health_monitor")
|
||||
if self.reload_plugins:
|
||||
flags.append("reload_plugins")
|
||||
if self.dispose_mcp_runtimes:
|
||||
flags.append("dispose_mcp")
|
||||
if self.restart_channels:
|
||||
flags.append(f"channels={{{','.join(sorted(self.restart_channels))}}}")
|
||||
if self.noop_paths:
|
||||
flags.append(f"noop={len(self.noop_paths)}")
|
||||
if flags:
|
||||
parts.append(", ".join(flags))
|
||||
return f"GatewayReloadPlan({'; '.join(parts)})"
|
||||
|
||||
|
||||
_plugin_rules: dict[str, tuple[ReloadRule, ...]] = {}
|
||||
|
||||
|
||||
def clear_plugin_reload_rules() -> None:
|
||||
"""清除所有已注册的插件重载规则(主要用于测试)。"""
|
||||
_plugin_rules.clear()
|
||||
_build_merged_rules.cache_clear()
|
||||
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def _build_merged_rules() -> tuple[ReloadRule, ...]:
|
||||
rules: list[ReloadRule] = list(BASE_RELOAD_RULES)
|
||||
for plugin_rules in _plugin_rules.values():
|
||||
rules.extend(plugin_rules)
|
||||
rules.extend(BASE_RELOAD_RULES_TAIL)
|
||||
return tuple(rules)
|
||||
|
||||
|
||||
def register_plugin_reload_rules(
|
||||
channel_id: str,
|
||||
config_prefixes: list[str] | None = None,
|
||||
noop_prefixes: list[str] | None = None,
|
||||
) -> None:
|
||||
rules: list[ReloadRule] = []
|
||||
for prefix in config_prefixes or []:
|
||||
rules.append(
|
||||
ReloadRule(
|
||||
prefix,
|
||||
kind=ReloadRuleKind.HOT,
|
||||
actions=(ReloadAction(f"restart_channel:{channel_id}"),),
|
||||
)
|
||||
)
|
||||
for prefix in noop_prefixes or []:
|
||||
rules.append(ReloadRule(prefix, kind=ReloadRuleKind.NONE))
|
||||
rules.append(
|
||||
ReloadRule(
|
||||
f"plugins.entries.{channel_id}",
|
||||
kind=ReloadRuleKind.HOT,
|
||||
actions=(
|
||||
ReloadAction.RELOAD_PLUGINS,
|
||||
ReloadAction.DISPOSE_MCP_RUNTIMES,
|
||||
ReloadAction(f"restart_channel:{channel_id}"),
|
||||
),
|
||||
),
|
||||
)
|
||||
_plugin_rules[channel_id] = tuple(rules)
|
||||
_build_merged_rules.cache_clear()
|
||||
|
||||
|
||||
def unregister_plugin_reload_rules(channel_id: str) -> None:
|
||||
"""注销指定插件的重载规则。"""
|
||||
if channel_id in _plugin_rules:
|
||||
del _plugin_rules[channel_id]
|
||||
_build_merged_rules.cache_clear()
|
||||
|
||||
|
||||
def _match_rule(path: str) -> ReloadRule | None:
|
||||
rules = _build_merged_rules()
|
||||
for rule in rules:
|
||||
if rule.matches(path):
|
||||
return rule
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_action(action: ReloadAction, plan: GatewayReloadPlan) -> None:
|
||||
action_str = action.value
|
||||
if action_str.startswith("restart_channel:"):
|
||||
channel = action_str[len("restart_channel:"):]
|
||||
plan.restart_channels.add(channel)
|
||||
return
|
||||
match action:
|
||||
case ReloadAction.RELOAD_HOOKS:
|
||||
plan.reload_hooks = True
|
||||
case ReloadAction.RESTART_GMAIL_WATCHER:
|
||||
plan.restart_gmail_watcher = True
|
||||
case ReloadAction.RESTART_CRON:
|
||||
plan.restart_cron = True
|
||||
case ReloadAction.RESTART_HEARTBEAT:
|
||||
plan.restart_heartbeat = True
|
||||
case ReloadAction.RESTART_HEALTH_MONITOR:
|
||||
plan.restart_health_monitor = True
|
||||
case ReloadAction.RELOAD_PLUGINS:
|
||||
plan.reload_plugins = True
|
||||
case ReloadAction.DISPOSE_MCP_RUNTIMES:
|
||||
plan.dispose_mcp_runtimes = True
|
||||
case ReloadAction.RESTART_CHANNEL | ReloadAction.RESTART_GATEWAY:
|
||||
logger.warning("Unhandled reload action: %s", action.value)
|
||||
case _:
|
||||
logger.warning("Unknown reload action: %s", action.value)
|
||||
|
||||
|
||||
def build_gateway_reload_plan(
|
||||
changed_paths: list[str],
|
||||
mode: ReloadMode = ReloadMode.HYBRID,
|
||||
) -> GatewayReloadPlan:
|
||||
if mode == ReloadMode.OFF:
|
||||
return GatewayReloadPlan(mode=mode, changed_paths=list(changed_paths))
|
||||
|
||||
plan = GatewayReloadPlan(mode=mode, changed_paths=list(changed_paths))
|
||||
|
||||
if mode == ReloadMode.RESTART:
|
||||
plan.restart_gateway = True
|
||||
plan.restart_reasons = list(changed_paths)
|
||||
return plan
|
||||
|
||||
for path in changed_paths:
|
||||
rule = _match_rule(path)
|
||||
if rule is None:
|
||||
plan.restart_gateway = True
|
||||
plan.restart_reasons.append(path)
|
||||
continue
|
||||
if rule.kind == ReloadRuleKind.RESTART:
|
||||
plan.restart_gateway = True
|
||||
plan.restart_reasons.append(path)
|
||||
continue
|
||||
if rule.kind == ReloadRuleKind.NONE:
|
||||
plan.noop_paths.append(path)
|
||||
continue
|
||||
plan.hot_reasons.append(path)
|
||||
for action in rule.actions:
|
||||
_resolve_action(action, plan)
|
||||
|
||||
if plan.restart_gmail_watcher:
|
||||
plan.reload_hooks = True
|
||||
|
||||
if mode == ReloadMode.HOT and plan.restart_gateway:
|
||||
plan.restart_gateway = False
|
||||
plan.restart_reasons.clear()
|
||||
for path in changed_paths:
|
||||
if path not in plan.noop_paths:
|
||||
plan.hot_reasons.append(path)
|
||||
|
||||
return plan
|
||||
424
backend/package/yuxi/channel/config/reloader.py
Normal file
424
backend/package/yuxi/channel/config/reloader.py
Normal file
@ -0,0 +1,424 @@
|
||||
import asyncio
|
||||
import inspect
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from .diff import ConfigDiff, compute_diff, diff_config_paths
|
||||
from .snapshot import ConfigSnapshot, RevisionCounter
|
||||
from .reload_plan import (
|
||||
ReloadMode,
|
||||
build_gateway_reload_plan,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from yuxi.channel.events.bus import ChannelEventBus
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
MISSING_CONFIG_RETRY_DELAY = 0.15
|
||||
MISSING_CONFIG_MAX_RETRIES = 2
|
||||
|
||||
|
||||
async def _maybe_await(result):
|
||||
if inspect.isawaitable(result):
|
||||
return await result
|
||||
return result
|
||||
|
||||
|
||||
class ConfigWriteNotification:
|
||||
def __init__(
|
||||
self,
|
||||
event_type: str,
|
||||
config_id: str | None = None,
|
||||
diff: ConfigDiff | None = None,
|
||||
after_write_mode: str = "auto",
|
||||
after_write_reason: str = "",
|
||||
persisted_hash: str | None = None,
|
||||
):
|
||||
self.event_type = event_type
|
||||
self.config_id = config_id
|
||||
self.diff = diff
|
||||
self.revision = RevisionCounter.next()
|
||||
self.after_write_mode = after_write_mode
|
||||
self.after_write_reason = after_write_reason
|
||||
self.persisted_hash = persisted_hash
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"ConfigWriteNotification(event={self.event_type}, config_id={self.config_id}, revision={self.revision})"
|
||||
|
||||
|
||||
class ConfigReloader:
|
||||
def __init__(
|
||||
self,
|
||||
poll_interval_seconds: float = 5.0,
|
||||
debounce_seconds: float = 0.3,
|
||||
reload_mode: ReloadMode = ReloadMode.HYBRID,
|
||||
event_bus: "ChannelEventBus | None" = None,
|
||||
):
|
||||
self._poll_interval = poll_interval_seconds
|
||||
self._debounce = debounce_seconds
|
||||
self._reload_mode = reload_mode
|
||||
self._event_bus = event_bus
|
||||
self._prev_snapshot: ConfigSnapshot | None = None
|
||||
self._prev_source_config: dict | None = None
|
||||
self._running = False
|
||||
self._task: asyncio.Task | None = None
|
||||
self._on_diff: list = []
|
||||
self._on_reload_plan: list = []
|
||||
self._write_subscribers: list = []
|
||||
self._debounce_timer: asyncio.Task | None = None
|
||||
self._pending = False
|
||||
self._reloading = False
|
||||
self._missing_retries = 0
|
||||
self._pending_in_process: ConfigWriteNotification | None = None
|
||||
self._last_applied_write_hash: str | None = None
|
||||
self._get_current_config = None
|
||||
self._lock = asyncio.Lock()
|
||||
|
||||
def on_diff(self, callback):
|
||||
self._on_diff.append(callback)
|
||||
return callback
|
||||
|
||||
def on_reload_plan(self, callback):
|
||||
self._on_reload_plan.append(callback)
|
||||
return callback
|
||||
|
||||
def subscribe_to_writes(self, callback):
|
||||
self._write_subscribers.append(callback)
|
||||
return callback
|
||||
|
||||
async def _publish_event(self, topic: str, *args, **kwargs) -> None:
|
||||
if self._event_bus is None:
|
||||
return
|
||||
try:
|
||||
await self._event_bus.publish(topic, *args, **kwargs)
|
||||
except Exception:
|
||||
logger.exception("EventBus publish failed: topic=%s", topic)
|
||||
|
||||
def _notify_write_subscribers(self, notification: ConfigWriteNotification) -> None:
|
||||
for subscriber in self._write_subscribers:
|
||||
try:
|
||||
subscriber(notification)
|
||||
except Exception:
|
||||
logger.exception("ConfigReloader: write subscriber failed")
|
||||
|
||||
def set_reload_mode(self, mode: ReloadMode) -> None:
|
||||
self._reload_mode = mode
|
||||
|
||||
async def notify_write(
|
||||
self,
|
||||
event_type: str,
|
||||
config_id: str | None = None,
|
||||
after_write_mode: str = "auto",
|
||||
after_write_reason: str = "",
|
||||
persisted_hash: str | None = None,
|
||||
) -> None:
|
||||
if persisted_hash and persisted_hash == self._last_applied_write_hash:
|
||||
logger.debug(
|
||||
"ConfigReloader: skipping duplicate write notification (hash=%s)",
|
||||
persisted_hash,
|
||||
)
|
||||
return
|
||||
notification = ConfigWriteNotification(
|
||||
event_type=event_type,
|
||||
config_id=config_id,
|
||||
after_write_mode=after_write_mode,
|
||||
after_write_reason=after_write_reason,
|
||||
persisted_hash=persisted_hash,
|
||||
)
|
||||
self._pending_in_process = notification
|
||||
self._last_applied_write_hash = persisted_hash
|
||||
await self._schedule_after(0)
|
||||
|
||||
async def start(self, get_current_config):
|
||||
self._running = True
|
||||
snapshot = await get_current_config()
|
||||
if snapshot is None:
|
||||
raise RuntimeError("ConfigReloader.start: get_current_config() returned None")
|
||||
self._get_current_config = get_current_config
|
||||
self._prev_snapshot = snapshot
|
||||
self._prev_source_config = _extract_source_config(snapshot)
|
||||
self._task = asyncio.create_task(self._poll_loop(get_current_config))
|
||||
logger.info(
|
||||
"ConfigReloader started (poll_interval=%.1fs, debounce=%.1fs, mode=%s)",
|
||||
self._poll_interval,
|
||||
self._debounce,
|
||||
self._reload_mode.value,
|
||||
)
|
||||
|
||||
async def stop(self):
|
||||
self._running = False
|
||||
if self._debounce_timer and not self._debounce_timer.done():
|
||||
self._debounce_timer.cancel()
|
||||
if self._task:
|
||||
self._task.cancel()
|
||||
try:
|
||||
await self._task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
logger.info("ConfigReloader stopped")
|
||||
|
||||
async def _schedule(self) -> None:
|
||||
await self._schedule_after(self._debounce)
|
||||
|
||||
async def _schedule_after(self, wait: float) -> None:
|
||||
if not self._running:
|
||||
return
|
||||
async with self._lock:
|
||||
if self._debounce_timer and not self._debounce_timer.done():
|
||||
self._debounce_timer.cancel()
|
||||
self._debounce_timer = asyncio.create_task(self._delayed_reload(wait))
|
||||
|
||||
async def _delayed_reload(self, wait: float) -> None:
|
||||
await asyncio.sleep(wait)
|
||||
await self._run_reload()
|
||||
|
||||
async def _poll_loop(self, get_current_config):
|
||||
while self._running:
|
||||
await asyncio.sleep(self._poll_interval)
|
||||
try:
|
||||
next_snapshot = await get_current_config()
|
||||
except Exception:
|
||||
logger.exception("ConfigReloader: failed to fetch config")
|
||||
continue
|
||||
|
||||
if next_snapshot is None:
|
||||
logger.warning("ConfigReloader: get_current_config() returned None, skipping poll cycle")
|
||||
continue
|
||||
|
||||
if next_snapshot.version == self._prev_snapshot.version:
|
||||
continue
|
||||
|
||||
next_snapshot.revision = RevisionCounter.next()
|
||||
|
||||
prev_accounts = self._prev_snapshot.accounts or []
|
||||
next_accounts = next_snapshot.accounts or []
|
||||
diff = compute_diff(prev_accounts, next_accounts)
|
||||
|
||||
routes_changed = next_snapshot.route_bindings_hash != self._prev_snapshot.route_bindings_hash
|
||||
diff.routes_changed = routes_changed
|
||||
|
||||
prev_source = self._prev_source_config or {}
|
||||
next_source = _extract_source_config(next_snapshot)
|
||||
config_paths = diff_config_paths(prev_source, next_source)
|
||||
diff.changed_paths = config_paths
|
||||
|
||||
if diff.has_changes:
|
||||
logger.info("ConfigReloader: detected changes %s", diff.summary)
|
||||
|
||||
reload_mode = _resolve_reload_mode(next_source, self._reload_mode)
|
||||
reload_plan = build_gateway_reload_plan(config_paths, mode=reload_mode)
|
||||
if not reload_plan.is_noop:
|
||||
logger.info("ConfigReloader: reload plan %s", reload_plan.summary)
|
||||
|
||||
for cb in self._on_reload_plan:
|
||||
try:
|
||||
await _maybe_await(cb(reload_plan))
|
||||
except Exception:
|
||||
logger.exception("ConfigReloader: on_reload_plan callback failed")
|
||||
|
||||
await self._publish_event("config.changed", diff=diff, reload_plan=reload_plan)
|
||||
|
||||
notification = ConfigWriteNotification(
|
||||
event_type="config_changed",
|
||||
diff=diff,
|
||||
)
|
||||
self._notify_write_subscribers(notification)
|
||||
|
||||
for cb in self._on_diff:
|
||||
try:
|
||||
await _maybe_await(cb(diff))
|
||||
except Exception:
|
||||
logger.exception("ConfigReloader: on_diff callback failed")
|
||||
|
||||
self._prev_snapshot = next_snapshot
|
||||
self._prev_source_config = next_source
|
||||
|
||||
async def _run_reload(self) -> None:
|
||||
if not self._running:
|
||||
return
|
||||
async with self._lock:
|
||||
if self._reloading:
|
||||
self._pending = True
|
||||
return
|
||||
self._reloading = True
|
||||
if self._debounce_timer and not self._debounce_timer.done():
|
||||
self._debounce_timer.cancel()
|
||||
self._debounce_timer = None
|
||||
|
||||
try:
|
||||
if self._pending_in_process:
|
||||
notification = self._pending_in_process
|
||||
self._pending_in_process = None
|
||||
diff = await self._compute_write_diff()
|
||||
if diff is not None and diff.has_changes:
|
||||
notification.diff = diff
|
||||
await self._apply_diff(diff)
|
||||
if self._prev_snapshot:
|
||||
self._prev_source_config = _extract_source_config(self._prev_snapshot)
|
||||
self._missing_retries = 0
|
||||
else:
|
||||
await self._poll_once()
|
||||
except Exception:
|
||||
logger.exception("ConfigReloader: reload failed")
|
||||
finally:
|
||||
async with self._lock:
|
||||
self._reloading = False
|
||||
pending = self._pending
|
||||
self._pending = False
|
||||
if pending:
|
||||
await self._schedule()
|
||||
|
||||
async def _compute_write_diff(self) -> ConfigDiff | None:
|
||||
if self._get_current_config is None or self._prev_snapshot is None:
|
||||
return None
|
||||
try:
|
||||
next_snapshot = await self._get_current_config()
|
||||
except Exception:
|
||||
logger.exception("ConfigReloader: failed to fetch config in _compute_write_diff")
|
||||
return None
|
||||
|
||||
if next_snapshot is None:
|
||||
for attempt in range(1, MISSING_CONFIG_MAX_RETRIES + 1):
|
||||
self._missing_retries += 1
|
||||
logger.warning(
|
||||
"ConfigReloader: get_current_config() returned None (retry %d/%d), will retry after %.1fs",
|
||||
attempt,
|
||||
MISSING_CONFIG_MAX_RETRIES,
|
||||
MISSING_CONFIG_RETRY_DELAY,
|
||||
)
|
||||
await asyncio.sleep(MISSING_CONFIG_RETRY_DELAY)
|
||||
try:
|
||||
next_snapshot = await self._get_current_config()
|
||||
except Exception:
|
||||
logger.exception("ConfigReloader: retry fetch config failed")
|
||||
return None
|
||||
if next_snapshot is not None:
|
||||
break
|
||||
else:
|
||||
logger.warning(
|
||||
"ConfigReloader: get_current_config() returned None after %d retries, giving up",
|
||||
MISSING_CONFIG_MAX_RETRIES,
|
||||
)
|
||||
return None
|
||||
|
||||
if next_snapshot is None:
|
||||
return None
|
||||
|
||||
next_snapshot.revision = RevisionCounter.next()
|
||||
|
||||
prev_accounts = self._prev_snapshot.accounts or []
|
||||
next_accounts = next_snapshot.accounts or []
|
||||
diff = compute_diff(prev_accounts, next_accounts)
|
||||
|
||||
routes_changed = next_snapshot.route_bindings_hash != self._prev_snapshot.route_bindings_hash
|
||||
diff.routes_changed = routes_changed
|
||||
|
||||
prev_source = self._prev_source_config or {}
|
||||
next_source = _extract_source_config(next_snapshot)
|
||||
config_paths = diff_config_paths(prev_source, next_source)
|
||||
diff.changed_paths = config_paths
|
||||
|
||||
if diff.has_changes:
|
||||
logger.info("ConfigReloader: write-triggered diff %s", diff.summary)
|
||||
self._prev_snapshot = next_snapshot
|
||||
self._prev_source_config = next_source
|
||||
|
||||
return diff
|
||||
|
||||
async def _poll_once(self) -> None:
|
||||
pass
|
||||
|
||||
async def _apply_diff(self, diff: ConfigDiff) -> None:
|
||||
for cb in self._on_diff:
|
||||
try:
|
||||
await _maybe_await(cb(diff))
|
||||
except Exception:
|
||||
logger.exception("ConfigReloader: on_diff callback failed")
|
||||
|
||||
self._notify_write_subscribers(ConfigWriteNotification(event_type="config_reloaded", diff=diff))
|
||||
|
||||
await self._publish_event("config.reloaded", diff=diff)
|
||||
|
||||
|
||||
def _extract_source_config(snapshot: ConfigSnapshot) -> dict:
|
||||
try:
|
||||
return getattr(snapshot, "source_config", None) or {}
|
||||
except Exception:
|
||||
return {}
|
||||
|
||||
|
||||
def _resolve_reload_mode(source_config: dict, fallback: ReloadMode) -> ReloadMode:
|
||||
try:
|
||||
gateway = source_config.get("gateway", {})
|
||||
reload_cfg = gateway.get("reload", {}) if isinstance(gateway, dict) else {}
|
||||
raw_mode = reload_cfg.get("mode")
|
||||
if raw_mode in ("off", "restart", "hot", "hybrid"):
|
||||
return ReloadMode(raw_mode)
|
||||
except Exception:
|
||||
pass
|
||||
return fallback
|
||||
|
||||
|
||||
async def reload_channel_config(
|
||||
channel_type: str,
|
||||
account_id: str,
|
||||
config_update: dict,
|
||||
) -> None:
|
||||
logger.info(
|
||||
"reload_channel_config: %s:%s with keys=%s",
|
||||
channel_type,
|
||||
account_id,
|
||||
list(config_update.keys()),
|
||||
)
|
||||
|
||||
try:
|
||||
from yuxi.channel.runtime.manager import gateway
|
||||
from yuxi.channel.protocols import ConfigProtocol
|
||||
from yuxi.channel.plugins.registry import ChannelPluginRegistry
|
||||
|
||||
plugin = ChannelPluginRegistry.get(channel_type)
|
||||
if plugin is None:
|
||||
logger.warning("reload_channel_config: plugin not found for %s", channel_type)
|
||||
return
|
||||
|
||||
if isinstance(plugin, ConfigProtocol):
|
||||
try:
|
||||
account = await plugin.resolve_account(account_id)
|
||||
if not plugin.is_configured(account):
|
||||
logger.warning(
|
||||
"reload_channel_config: %s:%s not configured, skipping restart",
|
||||
channel_type,
|
||||
account_id,
|
||||
)
|
||||
return
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"reload_channel_config: resolve_account failed for %s:%s",
|
||||
channel_type,
|
||||
account_id,
|
||||
)
|
||||
return
|
||||
|
||||
snapshot = gateway.get_snapshot(channel_type, account_id)
|
||||
if snapshot is None or snapshot.state.value == "stopped":
|
||||
logger.info(
|
||||
"reload_channel_config: %s:%s is not running, skipping restart",
|
||||
channel_type,
|
||||
account_id,
|
||||
)
|
||||
return
|
||||
|
||||
await gateway.stop_channel(channel_type, account_id)
|
||||
await gateway.start_channel(channel_type, account_id, gateway.global_config)
|
||||
logger.info(
|
||||
"reload_channel_config: %s:%s restarted with updated config",
|
||||
channel_type,
|
||||
account_id,
|
||||
)
|
||||
except Exception:
|
||||
logger.exception(
|
||||
"reload_channel_config: unexpected error for %s:%s",
|
||||
channel_type,
|
||||
account_id,
|
||||
)
|
||||
551
backend/package/yuxi/channel/config/snapshot.py
Normal file
551
backend/package/yuxi/channel/config/snapshot.py
Normal file
@ -0,0 +1,551 @@
|
||||
import hashlib
|
||||
import json
|
||||
import logging
|
||||
import re
|
||||
import threading
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import UTC, datetime
|
||||
from itertools import count
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_SENSITIVE_KEY_PATTERNS = [
|
||||
re.compile(r"token$", re.IGNORECASE),
|
||||
re.compile(r"password", re.IGNORECASE),
|
||||
re.compile(r"secret", re.IGNORECASE),
|
||||
re.compile(r"api[\W_]?key", re.IGNORECASE),
|
||||
re.compile(r"encrypt[\W_]?key", re.IGNORECASE),
|
||||
re.compile(r"private[\W_]?key", re.IGNORECASE),
|
||||
re.compile(r"serviceaccount(?:ref)?$", re.IGNORECASE),
|
||||
]
|
||||
|
||||
_SENSITIVE_KEY_WHITELIST_SUFFIXES = [
|
||||
"maxtokens",
|
||||
"maxoutputtokens",
|
||||
"maxinputtokens",
|
||||
"maxcompletiontokens",
|
||||
"contexttokens",
|
||||
"totaltokens",
|
||||
"tokencount",
|
||||
"tokenlimit",
|
||||
"tokenbudget",
|
||||
"passwordfile",
|
||||
]
|
||||
|
||||
REDACTED_SENTINEL = "__FORCEPILOT_REDACTED__"
|
||||
|
||||
_ENV_VAR_PLACEHOLDER_PATTERN = re.compile(r"^\$\{[^}]*\}$")
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# AccountConfigEntry – 渠道账户配置与运行时状态
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class AccountConfigEntry:
|
||||
channel_type: str
|
||||
account_id: str
|
||||
enabled: bool
|
||||
configured: bool
|
||||
config_hash: str
|
||||
|
||||
name: str | None = None
|
||||
connected: bool | None = None
|
||||
running: bool | None = None
|
||||
linked: bool | None = None
|
||||
restart_pending: bool | None = None
|
||||
status_state: str | None = None
|
||||
health_state: str | None = None
|
||||
mode: str | None = None
|
||||
dm_policy: str | None = None
|
||||
allow_from: list[str] | None = None
|
||||
reconnect_attempts: int | None = None
|
||||
last_connected_at: float | None = None
|
||||
last_message_at: float | None = None
|
||||
last_inbound_at: float | None = None
|
||||
last_outbound_at: float | None = None
|
||||
last_event_at: float | None = None
|
||||
last_transport_activity_at: float | None = None
|
||||
last_error: str | None = None
|
||||
last_start_at: float | None = None
|
||||
last_stop_at: float | None = None
|
||||
busy: bool | None = None
|
||||
active_runs: int | None = None
|
||||
last_run_activity_at: float | None = None
|
||||
token_status: str | None = None
|
||||
bot_token_status: str | None = None
|
||||
app_token_status: str | None = None
|
||||
signing_secret_status: str | None = None
|
||||
user_token_status: str | None = None
|
||||
base_url: str | None = None
|
||||
port: int | None = None
|
||||
allow_unmentioned_groups: bool | None = None
|
||||
cli_path: str | None = None
|
||||
db_path: str | None = None
|
||||
|
||||
|
||||
_SENSITIVE_ACCOUNT_FIELDS: set[str] = {
|
||||
"webhook_url",
|
||||
"webhook_path",
|
||||
"audience",
|
||||
"public_key",
|
||||
"channel_access_token",
|
||||
"channel_secret",
|
||||
"token",
|
||||
"bot_token",
|
||||
"app_token",
|
||||
"signing_secret",
|
||||
"user_token",
|
||||
}
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 安全投影 – 剥离敏感字段
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def project_safe_account_fields(account: dict[str, Any] | Any) -> dict[str, Any]:
|
||||
"""从账户对象中提取对外安全的快照字段,剥离敏感信息"""
|
||||
if not isinstance(account, dict):
|
||||
return {}
|
||||
|
||||
record: dict[str, Any] = account
|
||||
result: dict[str, Any] = {}
|
||||
|
||||
for key in ("name", "status_state", "health_state", "mode", "dm_policy"):
|
||||
val = record.get(key)
|
||||
if isinstance(val, str) and val:
|
||||
result[key] = val
|
||||
|
||||
for key in ("linked", "running", "connected", "restart_pending", "busy"):
|
||||
val = record.get(key)
|
||||
if isinstance(val, bool):
|
||||
result[key] = val
|
||||
|
||||
for key in (
|
||||
"reconnect_attempts",
|
||||
"active_runs",
|
||||
"last_connected_at",
|
||||
"last_message_at",
|
||||
"last_inbound_at",
|
||||
"last_outbound_at",
|
||||
"last_event_at",
|
||||
"last_transport_activity_at",
|
||||
"last_start_at",
|
||||
"last_stop_at",
|
||||
"last_run_activity_at",
|
||||
):
|
||||
val = record.get(key)
|
||||
if isinstance(val, (int, float)):
|
||||
result[key] = val
|
||||
elif val is None and key in record:
|
||||
result[key] = None
|
||||
|
||||
for key in ("token_status", "bot_token_status", "app_token_status", "signing_secret_status", "user_token_status"):
|
||||
val = record.get(key)
|
||||
if val in ("available", "configured_unavailable", "missing"):
|
||||
result[key] = val
|
||||
|
||||
base_url = record.get("base_url")
|
||||
if isinstance(base_url, str) and base_url:
|
||||
result["base_url"] = _strip_url_userinfo(base_url)
|
||||
|
||||
allow_from = record.get("allow_from")
|
||||
if isinstance(allow_from, list):
|
||||
normalized = [str(e).strip() for e in allow_from if isinstance(e, (str, int, float)) and str(e).strip()]
|
||||
if normalized:
|
||||
result["allow_from"] = normalized
|
||||
|
||||
cli_path = record.get("cli_path")
|
||||
if isinstance(cli_path, str) and cli_path:
|
||||
result["cli_path"] = cli_path
|
||||
|
||||
db_path = record.get("db_path")
|
||||
if isinstance(db_path, str) and db_path:
|
||||
result["db_path"] = db_path
|
||||
|
||||
port = record.get("port")
|
||||
if isinstance(port, (int, float)):
|
||||
result["port"] = int(port)
|
||||
elif port is None and "port" in record:
|
||||
result["port"] = None
|
||||
|
||||
allow_unmentioned = record.get("allow_unmentioned_groups")
|
||||
if isinstance(allow_unmentioned, bool):
|
||||
result["allow_unmentioned_groups"] = allow_unmentioned
|
||||
|
||||
account_id = record.get("account_id")
|
||||
if isinstance(account_id, str):
|
||||
result["account_id"] = account_id
|
||||
|
||||
enabled = record.get("enabled")
|
||||
if isinstance(enabled, bool):
|
||||
result["enabled"] = enabled
|
||||
|
||||
configured = record.get("configured")
|
||||
if isinstance(configured, bool):
|
||||
result["configured"] = configured
|
||||
|
||||
return result
|
||||
|
||||
|
||||
def _strip_url_userinfo(url: str) -> str:
|
||||
"""剥离 URL 中的用户信息部分 (https://user:pass@host -> https://host)"""
|
||||
import re as _re
|
||||
|
||||
return _re.sub(r"://[^@/]+@", "://", url)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 敏感路径检测
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _is_whitelisted_sensitive_path(path: str) -> bool:
|
||||
lower = path.lower()
|
||||
return any(lower.endswith(suffix) for suffix in _SENSITIVE_KEY_WHITELIST_SUFFIXES)
|
||||
|
||||
|
||||
def is_sensitive_config_path(path: str) -> bool:
|
||||
"""检测配置路径是否指向敏感字段 (token/password/secret/api_key 等)"""
|
||||
if _is_whitelisted_sensitive_path(path):
|
||||
return False
|
||||
return any(pattern.search(path) for pattern in _SENSITIVE_KEY_PATTERNS)
|
||||
|
||||
|
||||
def _is_env_var_placeholder(value: str) -> bool:
|
||||
return bool(_ENV_VAR_PLACEHOLDER_PATTERN.match(value.strip()))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 脱敏 & 还原 (Redact / Restore)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _is_object_record(value: Any) -> bool:
|
||||
return isinstance(value, dict)
|
||||
|
||||
|
||||
def _collect_sensitive_strings(value: Any, values: list[str]) -> None:
|
||||
if isinstance(value, str):
|
||||
if not _is_env_var_placeholder(value):
|
||||
values.append(value)
|
||||
return
|
||||
if isinstance(value, list):
|
||||
for item in value:
|
||||
_collect_sensitive_strings(item, values)
|
||||
return
|
||||
if _is_object_record(value):
|
||||
for item in value.values():
|
||||
_collect_sensitive_strings(item, values)
|
||||
|
||||
|
||||
def redact_config_object(obj: Any, ui_hints: dict[str, Any] | None = None) -> Any:
|
||||
"""深度遍历对象,将敏感路径上的字符串值替换为 REDACTED_SENTINEL"""
|
||||
if ui_hints:
|
||||
lookup = _build_redaction_lookup(ui_hints)
|
||||
if "" in lookup:
|
||||
return _redact_with_lookup(obj, lookup, "")
|
||||
return _redact_guessing(obj, "", ui_hints)
|
||||
return _redact_guessing(obj, "")
|
||||
|
||||
|
||||
def restore_redacted_values(incoming: Any, original: Any, ui_hints: dict[str, Any] | None = None) -> Any:
|
||||
"""深度遍历 incoming,将 REDACTED_SENTINEL 替换回 original 中的真实值"""
|
||||
if incoming is None:
|
||||
return incoming
|
||||
if not isinstance(incoming, dict):
|
||||
return incoming
|
||||
|
||||
if ui_hints:
|
||||
lookup = _build_redaction_lookup(ui_hints)
|
||||
if "" in lookup:
|
||||
return _restore_with_lookup(incoming, original, lookup, "", ui_hints)
|
||||
return _restore_guessing(incoming, original, "", ui_hints)
|
||||
return _restore_guessing(incoming, original, "")
|
||||
|
||||
|
||||
def _build_redaction_lookup(hints: dict[str, Any]) -> set[str]:
|
||||
result: set[str] = set()
|
||||
for path, hint in hints.items():
|
||||
if not hint.get("sensitive"):
|
||||
continue
|
||||
parts = path.split(".")
|
||||
joined = parts[0] if parts else ""
|
||||
result.add(joined)
|
||||
if joined.endswith("[]"):
|
||||
result.add(joined[:-2])
|
||||
for part in parts[1:]:
|
||||
if part.endswith("[]"):
|
||||
result.add(f"{joined}.{part[:-2]}")
|
||||
joined = f"{joined}.{part}"
|
||||
result.add(joined)
|
||||
if result:
|
||||
result.add("")
|
||||
return result
|
||||
|
||||
|
||||
def _redact_with_lookup(obj: Any, lookup: set[str], prefix: str) -> Any:
|
||||
if obj is None:
|
||||
return obj
|
||||
if isinstance(obj, list):
|
||||
path = f"{prefix}[]"
|
||||
if path not in lookup:
|
||||
return [_redact_guessing(item, prefix) for item in obj]
|
||||
return [
|
||||
REDACTED_SENTINEL
|
||||
if isinstance(item, str) and not _is_env_var_placeholder(item)
|
||||
else _redact_with_lookup(item, lookup, path)
|
||||
for item in obj
|
||||
]
|
||||
if _is_object_record(obj):
|
||||
result: dict[str, Any] = {}
|
||||
for key, value in obj.items():
|
||||
path = f"{prefix}.{key}" if prefix else key
|
||||
wildcard_path = f"{prefix}.*" if prefix else "*"
|
||||
matched = False
|
||||
for candidate in (path, wildcard_path):
|
||||
if candidate in lookup:
|
||||
matched = True
|
||||
if isinstance(value, str) and not _is_env_var_placeholder(value):
|
||||
result[key] = REDACTED_SENTINEL
|
||||
elif isinstance(value, dict):
|
||||
result[key] = _redact_with_lookup(value, lookup, candidate)
|
||||
elif isinstance(value, list):
|
||||
result[key] = _redact_with_lookup(value, lookup, candidate)
|
||||
else:
|
||||
result[key] = value
|
||||
break
|
||||
if not matched:
|
||||
result[key] = _redact_guessing(value, path)
|
||||
return result
|
||||
return obj
|
||||
|
||||
|
||||
def _redact_guessing(obj: Any, prefix: str, hints: dict[str, Any] | None = None) -> Any:
|
||||
if obj is None:
|
||||
return obj
|
||||
if isinstance(obj, list):
|
||||
return [_redact_guessing(item, f"{prefix}[]", hints) for item in obj]
|
||||
if _is_object_record(obj):
|
||||
result: dict[str, Any] = {}
|
||||
for key, value in obj.items():
|
||||
dot_path = f"{prefix}.{key}" if prefix else key
|
||||
if isinstance(value, str) and not _is_env_var_placeholder(value) and is_sensitive_config_path(dot_path):
|
||||
result[key] = REDACTED_SENTINEL
|
||||
elif isinstance(value, dict):
|
||||
result[key] = _redact_guessing(value, dot_path, hints)
|
||||
elif isinstance(value, list):
|
||||
result[key] = _redact_guessing(value, dot_path, hints)
|
||||
else:
|
||||
result[key] = value
|
||||
return result
|
||||
return obj
|
||||
|
||||
|
||||
def _restore_with_lookup(incoming: Any, original: Any, lookup: set[str], prefix: str, hints: dict[str, Any]) -> Any:
|
||||
if incoming is None or not isinstance(incoming, (dict, list)):
|
||||
return incoming
|
||||
|
||||
if isinstance(incoming, list):
|
||||
path = f"{prefix}[]"
|
||||
if path not in lookup:
|
||||
return [_restore_guessing(item, original, prefix, hints) for item in incoming]
|
||||
orig_list = original if isinstance(original, list) else []
|
||||
return [
|
||||
orig_list[i]
|
||||
if item == REDACTED_SENTINEL and i < len(orig_list)
|
||||
else _restore_with_lookup(item, orig_list[i] if i < len(orig_list) else None, lookup, path, hints)
|
||||
for i, item in enumerate(incoming)
|
||||
]
|
||||
|
||||
if _is_object_record(incoming):
|
||||
orig = original if _is_object_record(original) else {}
|
||||
result: dict[str, Any] = {}
|
||||
for key, value in incoming.items():
|
||||
path = f"{prefix}.{key}" if prefix else key
|
||||
wildcard_path = f"{prefix}.*" if prefix else "*"
|
||||
matched = False
|
||||
for candidate in (path, wildcard_path):
|
||||
if candidate in lookup:
|
||||
matched = True
|
||||
if value == REDACTED_SENTINEL:
|
||||
result[key] = orig.get(key, value)
|
||||
elif isinstance(value, (dict, list)):
|
||||
result[key] = _restore_with_lookup(value, orig.get(key), lookup, candidate, hints)
|
||||
else:
|
||||
result[key] = value
|
||||
break
|
||||
if not matched:
|
||||
result[key] = _restore_guessing(value, orig, prefix, hints)
|
||||
return result
|
||||
return incoming
|
||||
|
||||
|
||||
def _restore_guessing(incoming: Any, original: Any, prefix: str, hints: dict[str, Any] | None = None) -> Any:
|
||||
if incoming is None or not isinstance(incoming, (dict, list)):
|
||||
return incoming
|
||||
|
||||
if isinstance(incoming, list):
|
||||
orig_list = original if isinstance(original, list) else []
|
||||
path = f"{prefix}[]"
|
||||
return [
|
||||
orig_list[i]
|
||||
if item == REDACTED_SENTINEL and is_sensitive_config_path(path) and i < len(orig_list)
|
||||
else _restore_guessing(item, orig_list[i] if i < len(orig_list) else None, path, hints)
|
||||
for i, item in enumerate(incoming)
|
||||
]
|
||||
|
||||
if _is_object_record(incoming):
|
||||
orig = original if _is_object_record(original) else {}
|
||||
result: dict[str, Any] = {}
|
||||
for key, value in incoming.items():
|
||||
dot_path = f"{prefix}.{key}" if prefix else key
|
||||
if value == REDACTED_SENTINEL and is_sensitive_config_path(dot_path):
|
||||
result[key] = orig.get(key, value)
|
||||
elif isinstance(value, (dict, list)):
|
||||
result[key] = _restore_guessing(value, orig.get(key), dot_path, hints)
|
||||
else:
|
||||
result[key] = value
|
||||
return result
|
||||
return incoming
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ConfigFileSnapshot – 配置文件全量快照
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConfigFileSnapshot:
|
||||
path: str
|
||||
exists: bool
|
||||
raw: str | None
|
||||
parsed: Any
|
||||
source_config: dict[str, Any] = field(default_factory=dict)
|
||||
resolved: dict[str, Any] = field(default_factory=dict)
|
||||
runtime_config: dict[str, Any] = field(default_factory=dict)
|
||||
valid: bool = False
|
||||
hash: str = ""
|
||||
issues: list[dict[str, Any]] = field(default_factory=list)
|
||||
warnings: list[dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
def __post_init__(self):
|
||||
if not self.hash and self.runtime_config:
|
||||
self.hash = hash_config_value(self.runtime_config)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# ConfigSnapshot – 渠道配置运行时快照
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class ConfigSnapshot:
|
||||
version: str
|
||||
accounts: list[AccountConfigEntry]
|
||||
route_bindings_hash: str
|
||||
created_at: datetime = field(default_factory=lambda: datetime.now(UTC))
|
||||
revision: int = 0
|
||||
source_fingerprint: str = ""
|
||||
runtime_fingerprint: str = ""
|
||||
|
||||
@property
|
||||
def stable_hash(self) -> str:
|
||||
if self.source_fingerprint:
|
||||
return self.source_fingerprint
|
||||
return _compute_snapshot_stable_hash(self.accounts, self.route_bindings_hash)
|
||||
|
||||
|
||||
def _compute_snapshot_stable_hash(accounts: list[AccountConfigEntry], routes_hash: str) -> str:
|
||||
sorted_accounts = sorted(
|
||||
[
|
||||
{
|
||||
"t": a.channel_type,
|
||||
"id": a.account_id,
|
||||
"e": a.enabled,
|
||||
"c": a.configured,
|
||||
"h": a.config_hash,
|
||||
}
|
||||
for a in accounts
|
||||
],
|
||||
key=lambda x: (x["t"], x["id"]),
|
||||
)
|
||||
payload = {"accounts": sorted_accounts, "routes_hash": routes_hash}
|
||||
raw = json.dumps(payload, sort_keys=True, ensure_ascii=False)
|
||||
return hashlib.sha256(raw.encode()).hexdigest()[:16]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# 通用哈希 & 快照比较
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _stable_config_stringify(value: Any) -> str:
|
||||
if value is None:
|
||||
return "null"
|
||||
if isinstance(value, bool):
|
||||
return "true" if value else "false"
|
||||
if isinstance(value, (int, float)):
|
||||
return json.dumps(value)
|
||||
if isinstance(value, str):
|
||||
return json.dumps(value, ensure_ascii=False)
|
||||
if isinstance(value, (list, tuple)):
|
||||
return f"[{','.join(_stable_config_stringify(v) for v in value)}]"
|
||||
if isinstance(value, dict):
|
||||
keys = sorted(value.keys())
|
||||
return (
|
||||
"{"
|
||||
+ ",".join(f"{json.dumps(k, ensure_ascii=False)}:{_stable_config_stringify(value[k])}" for k in keys)
|
||||
+ "}"
|
||||
)
|
||||
return json.dumps(value, default=str, ensure_ascii=False)
|
||||
|
||||
|
||||
def hash_config_value(value: dict[str, Any]) -> str:
|
||||
return hashlib.sha256(_stable_config_stringify(value).encode()).hexdigest()[:16]
|
||||
|
||||
|
||||
def compute_source_fingerprint(config: dict) -> str:
|
||||
raw = json.dumps(config, sort_keys=True, default=str, ensure_ascii=False)
|
||||
return hashlib.sha256(raw.encode()).hexdigest()[:16]
|
||||
|
||||
|
||||
def config_snapshots_match(left: dict[str, Any], right: dict[str, Any]) -> bool:
|
||||
if left is right:
|
||||
return True
|
||||
try:
|
||||
return _stable_config_stringify(left) == _stable_config_stringify(right)
|
||||
except Exception as e:
|
||||
logger.warning("Config snapshot comparison failed: %s", e)
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# RevisionCounter – 全局递增计数器
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class RevisionCounter:
|
||||
_revision = count(1)
|
||||
_current: int = 0
|
||||
_lock = threading.Lock()
|
||||
|
||||
@classmethod
|
||||
def next(cls) -> int:
|
||||
with cls._lock:
|
||||
cls._current = next(cls._revision)
|
||||
return cls._current
|
||||
|
||||
@classmethod
|
||||
def current(cls) -> int:
|
||||
with cls._lock:
|
||||
return cls._current
|
||||
|
||||
@classmethod
|
||||
def reset(cls) -> None:
|
||||
with cls._lock:
|
||||
cls._revision = count(1)
|
||||
cls._current = 0
|
||||
416
backend/package/yuxi/channel/config/validation.py
Normal file
416
backend/package/yuxi/channel/config/validation.py
Normal file
@ -0,0 +1,416 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@dataclass
|
||||
class ValidationIssue:
|
||||
path: str
|
||||
message: str
|
||||
severity: str = "error"
|
||||
allowed_values: list[str] | None = None
|
||||
allowed_values_hidden_count: int | None = None
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
result = {"path": self.path, "message": self.message, "severity": self.severity}
|
||||
if self.allowed_values is not None:
|
||||
result["allowedValues"] = self.allowed_values
|
||||
if self.allowed_values_hidden_count is not None:
|
||||
result["allowedValuesHiddenCount"] = self.allowed_values_hidden_count
|
||||
return result
|
||||
|
||||
|
||||
@dataclass
|
||||
class ValidationResult:
|
||||
ok: bool = True
|
||||
issues: list[ValidationIssue] = field(default_factory=list)
|
||||
warnings: list[ValidationIssue] = field(default_factory=list)
|
||||
legacy_issues: list[ValidationIssue] = field(default_factory=list)
|
||||
|
||||
def add_issue(
|
||||
self,
|
||||
path: str,
|
||||
message: str,
|
||||
severity: str = "error",
|
||||
allowed_values: list[str] | None = None,
|
||||
allowed_values_hidden_count: int | None = None,
|
||||
) -> None:
|
||||
issue = ValidationIssue(
|
||||
path=path,
|
||||
message=message,
|
||||
severity=severity,
|
||||
allowed_values=allowed_values,
|
||||
allowed_values_hidden_count=allowed_values_hidden_count,
|
||||
)
|
||||
if severity == "warning":
|
||||
self.warnings.append(issue)
|
||||
elif severity == "legacy":
|
||||
self.legacy_issues.append(issue)
|
||||
else:
|
||||
self.issues.append(issue)
|
||||
self.ok = False
|
||||
|
||||
def merge(self, other: ValidationResult) -> None:
|
||||
self.issues.extend(other.issues)
|
||||
self.warnings.extend(other.warnings)
|
||||
self.legacy_issues.extend(other.legacy_issues)
|
||||
if not other.ok:
|
||||
self.ok = False
|
||||
|
||||
def to_dict(self) -> dict:
|
||||
return {
|
||||
"ok": self.ok,
|
||||
"issues": [i.to_dict() for i in self.issues],
|
||||
"warnings": [w.to_dict() for w in self.warnings],
|
||||
"legacy_issues": [li.to_dict() for li in self.legacy_issues],
|
||||
}
|
||||
|
||||
|
||||
# ── Schema Validation ────────────────────────────────────
|
||||
|
||||
|
||||
def _collect_allowed_values_from_schema_node(node: dict) -> tuple[list[str], bool]:
|
||||
if "const" in node:
|
||||
return [str(node["const"])], False
|
||||
if "enum" in node and isinstance(node["enum"], list):
|
||||
return [str(v) for v in node["enum"]], False
|
||||
node_type = node.get("type")
|
||||
if node_type == "boolean" or (isinstance(node_type, list) and "boolean" in node_type):
|
||||
return ["true", "false"], False
|
||||
|
||||
union_branches = node.get("anyOf") or node.get("oneOf")
|
||||
if not isinstance(union_branches, list):
|
||||
return [], False
|
||||
|
||||
collected: list[str] = []
|
||||
for branch in union_branches:
|
||||
if not isinstance(branch, dict):
|
||||
return [], True
|
||||
branch_values, incomplete = _collect_allowed_values_from_schema_node(branch)
|
||||
if incomplete or not branch_values:
|
||||
return [], True
|
||||
collected.extend(branch_values)
|
||||
return collected, False
|
||||
|
||||
|
||||
def _lookup_schema_node(schema: dict, path_parts: list[str]) -> dict | None:
|
||||
current: Any = schema
|
||||
for part in path_parts:
|
||||
if not isinstance(current, dict):
|
||||
return None
|
||||
properties = current.get("properties")
|
||||
if isinstance(properties, dict):
|
||||
current = properties.get(part)
|
||||
elif "additionalProperties" in current:
|
||||
current = current["additionalProperties"]
|
||||
else:
|
||||
return None
|
||||
return current if isinstance(current, dict) else None
|
||||
|
||||
|
||||
def _check_type(entry_value: Any, expected_type: str) -> bool:
|
||||
if expected_type == "string":
|
||||
return isinstance(entry_value, str)
|
||||
if expected_type == "number":
|
||||
return isinstance(entry_value, (int, float)) and not isinstance(entry_value, bool)
|
||||
if expected_type == "integer":
|
||||
return isinstance(entry_value, int) and not isinstance(entry_value, bool)
|
||||
if expected_type == "boolean":
|
||||
return isinstance(entry_value, bool)
|
||||
if expected_type == "array":
|
||||
return isinstance(entry_value, list)
|
||||
if expected_type == "object":
|
||||
return isinstance(entry_value, dict)
|
||||
if expected_type == "null":
|
||||
return entry_value is None
|
||||
return True
|
||||
|
||||
|
||||
def _matches_type(entry_value: Any, prop_type: str | list[str] | None) -> bool:
|
||||
if prop_type is None:
|
||||
return True
|
||||
if isinstance(prop_type, list):
|
||||
return any(_check_type(entry_value, t) for t in prop_type)
|
||||
return _check_type(entry_value, prop_type)
|
||||
|
||||
|
||||
def _validate_numeric_bounds(
|
||||
entry_value: Any,
|
||||
prop_schema: dict,
|
||||
field_path: str,
|
||||
key: str,
|
||||
result: ValidationResult,
|
||||
) -> None:
|
||||
if not isinstance(entry_value, (int, float)) or isinstance(entry_value, bool):
|
||||
return
|
||||
|
||||
minimum = prop_schema.get("minimum")
|
||||
maximum = prop_schema.get("maximum")
|
||||
exclusive_minimum = prop_schema.get("exclusiveMinimum")
|
||||
exclusive_maximum = prop_schema.get("exclusiveMaximum")
|
||||
|
||||
if minimum is not None and entry_value < minimum:
|
||||
result.add_issue(field_path, f"'{key}' must be >= {minimum}")
|
||||
if maximum is not None and entry_value > maximum:
|
||||
result.add_issue(field_path, f"'{key}' must be <= {maximum}")
|
||||
if exclusive_minimum is not None and entry_value <= exclusive_minimum:
|
||||
result.add_issue(field_path, f"'{key}' must be > {exclusive_minimum}")
|
||||
if exclusive_maximum is not None and entry_value >= exclusive_maximum:
|
||||
result.add_issue(field_path, f"'{key}' must be < {exclusive_maximum}")
|
||||
|
||||
|
||||
def _validate_string_constraints(
|
||||
entry_value: Any,
|
||||
prop_schema: dict,
|
||||
field_path: str,
|
||||
key: str,
|
||||
result: ValidationResult,
|
||||
) -> None:
|
||||
if not isinstance(entry_value, str):
|
||||
return
|
||||
|
||||
min_length = prop_schema.get("minLength")
|
||||
max_length = prop_schema.get("maxLength")
|
||||
pattern = prop_schema.get("pattern")
|
||||
format_ = prop_schema.get("format")
|
||||
|
||||
if min_length is not None and len(entry_value) < min_length:
|
||||
result.add_issue(field_path, f"'{key}' must be at least {min_length} characters")
|
||||
if max_length is not None and len(entry_value) > max_length:
|
||||
result.add_issue(field_path, f"'{key}' must be at most {max_length} characters")
|
||||
if pattern is not None:
|
||||
import re
|
||||
|
||||
try:
|
||||
if not re.search(pattern, entry_value):
|
||||
result.add_issue(field_path, f"'{key}' does not match pattern '{pattern}'")
|
||||
except re.error:
|
||||
logger.warning("Invalid pattern in schema for '%s': %s", key, pattern)
|
||||
if format_ == "email" and "@" not in entry_value:
|
||||
result.add_issue(field_path, f"'{key}' must be a valid email address")
|
||||
if format_ == "uri" and not (entry_value.startswith("http://") or entry_value.startswith("https://")):
|
||||
result.add_issue(field_path, f"'{key}' must be a valid URI")
|
||||
|
||||
|
||||
def _validate_array_constraints(
|
||||
entry_value: Any,
|
||||
prop_schema: dict,
|
||||
field_path: str,
|
||||
key: str,
|
||||
result: ValidationResult,
|
||||
) -> None:
|
||||
if not isinstance(entry_value, list):
|
||||
return
|
||||
|
||||
min_items = prop_schema.get("minItems")
|
||||
max_items = prop_schema.get("maxItems")
|
||||
items_schema = prop_schema.get("items")
|
||||
|
||||
if min_items is not None and len(entry_value) < min_items:
|
||||
result.add_issue(field_path, f"'{key}' must have at least {min_items} items")
|
||||
if max_items is not None and len(entry_value) > max_items:
|
||||
result.add_issue(field_path, f"'{key}' must have at most {max_items} items")
|
||||
|
||||
if isinstance(items_schema, dict):
|
||||
for i, item in enumerate(entry_value):
|
||||
item_path = f"{field_path}[{i}]"
|
||||
item_type = items_schema.get("type")
|
||||
if not _matches_type(item, item_type):
|
||||
type_label = item_type if isinstance(item_type, str) else " | ".join(item_type) if isinstance(item_type, list) else "any"
|
||||
result.add_issue(item_path, f"item expects a {type_label} value")
|
||||
if isinstance(item, dict) and isinstance(items_schema.get("properties"), dict):
|
||||
_validate_entry_against_schema(item, items_schema, item_path, result)
|
||||
|
||||
|
||||
def _validate_entry_against_schema(
|
||||
entry: dict,
|
||||
schema: dict,
|
||||
prefix: str,
|
||||
result: ValidationResult,
|
||||
) -> None:
|
||||
if schema.get("type") != "object":
|
||||
return
|
||||
|
||||
properties = schema.get("properties")
|
||||
required_fields: list[str] = schema.get("required", [])
|
||||
if not isinstance(properties, dict):
|
||||
return
|
||||
|
||||
seen_keys: set[str] = set()
|
||||
|
||||
for key, entry_value in entry.items():
|
||||
if key in ("channel_type", "account_id"):
|
||||
continue
|
||||
seen_keys.add(key)
|
||||
|
||||
prop_schema = properties.get(key) if isinstance(properties, dict) else None
|
||||
if prop_schema is None:
|
||||
continue
|
||||
|
||||
field_path = f"{prefix}.{key}"
|
||||
prop_type = prop_schema.get("type")
|
||||
|
||||
if not _matches_type(entry_value, prop_type):
|
||||
help_text = prop_schema.get("description", "")
|
||||
if isinstance(prop_type, list):
|
||||
type_label = " | ".join(prop_type)
|
||||
elif isinstance(prop_type, str):
|
||||
type_label = prop_type
|
||||
else:
|
||||
type_label = "any"
|
||||
msg = f"'{key}' expects a {type_label} value. {help_text}".strip()
|
||||
result.add_issue(field_path, msg)
|
||||
continue
|
||||
|
||||
if isinstance(entry_value, dict) and isinstance(prop_schema.get("properties"), dict):
|
||||
_validate_entry_against_schema(entry_value, prop_schema, field_path, result)
|
||||
|
||||
if isinstance(entry_value, str) and not entry_value.strip() and key in required_fields:
|
||||
help_text = prop_schema.get("description", "")
|
||||
msg = f"'{key}' is required but empty. {help_text}".strip()
|
||||
result.add_issue(field_path, msg)
|
||||
continue
|
||||
|
||||
if isinstance(entry_value, str) and entry_value.strip():
|
||||
allowed, incomplete = _collect_allowed_values_from_schema_node(prop_schema)
|
||||
if not incomplete and allowed and entry_value not in allowed:
|
||||
result.add_issue(
|
||||
field_path,
|
||||
f"'{key}' has invalid value '{entry_value}'",
|
||||
allowed_values=allowed,
|
||||
)
|
||||
|
||||
_validate_numeric_bounds(entry_value, prop_schema, field_path, key, result)
|
||||
_validate_string_constraints(entry_value, prop_schema, field_path, key, result)
|
||||
_validate_array_constraints(entry_value, prop_schema, field_path, key, result)
|
||||
|
||||
for req_key in required_fields:
|
||||
if req_key not in seen_keys:
|
||||
prop_schema = properties.get(req_key, {})
|
||||
help_text = prop_schema.get("description", "")
|
||||
msg = f"missing required field '{req_key}'. {help_text}".strip()
|
||||
result.add_issue(f"{prefix}.{req_key}", msg)
|
||||
|
||||
|
||||
# ── Plugin Validation ────────────────────────────────────
|
||||
|
||||
|
||||
def validate_channel_plugin(plugin: Any, result: ValidationResult) -> None:
|
||||
from yuxi.channel.protocols import ConfigProtocol, ConfigSchemaProtocol
|
||||
|
||||
if not isinstance(plugin, ConfigProtocol):
|
||||
return
|
||||
|
||||
plugin_id = getattr(plugin, "id", None)
|
||||
if not plugin_id or not isinstance(plugin_id, str) or not plugin_id.strip():
|
||||
result.add_issue("plugin", "channel plugin missing id")
|
||||
return
|
||||
|
||||
prefix = f"plugin.{plugin_id}"
|
||||
|
||||
if not callable(getattr(plugin, "is_configured", None)):
|
||||
result.add_issue(prefix, "missing is_configured method")
|
||||
|
||||
if not callable(getattr(plugin, "list_account_ids", None)):
|
||||
result.add_issue(prefix, "missing list_account_ids method")
|
||||
|
||||
if not callable(getattr(plugin, "resolve_account", None)):
|
||||
result.add_issue(prefix, "missing resolve_account method")
|
||||
|
||||
if isinstance(plugin, ConfigSchemaProtocol):
|
||||
try:
|
||||
schema = plugin.config_schema()
|
||||
if not isinstance(schema, dict) or not schema:
|
||||
result.add_issue(prefix, "config_schema() returned empty or invalid schema", "warning")
|
||||
except Exception as e:
|
||||
result.add_issue(prefix, f"config_schema() raised: {e}", "warning")
|
||||
|
||||
|
||||
# ── Main Validation ──────────────────────────────────────
|
||||
|
||||
|
||||
async def validate_all_channel_configs(
|
||||
plugins: list,
|
||||
config_entries: list[dict],
|
||||
*,
|
||||
validate_schema: bool = True,
|
||||
validate_plugins: bool = False,
|
||||
) -> ValidationResult:
|
||||
result = ValidationResult()
|
||||
|
||||
from yuxi.channel.protocols import ConfigProtocol, ConfigSchemaProtocol
|
||||
from yuxi.channel.secrets.models import is_secret_ref
|
||||
|
||||
configured_plugins = [p for p in plugins if isinstance(p, ConfigProtocol)]
|
||||
if not configured_plugins:
|
||||
result.add_issue("global", "No ConfigProtocol plugins registered")
|
||||
return result
|
||||
|
||||
if validate_plugins:
|
||||
for plugin in configured_plugins:
|
||||
validate_channel_plugin(plugin, result)
|
||||
|
||||
plugin_ids = {p.id for p in configured_plugins}
|
||||
|
||||
for entry in config_entries:
|
||||
channel_type = entry.get("channel_type", "unknown")
|
||||
account_id = entry.get("account_id", "unknown")
|
||||
prefix = f"channels.{channel_type}.{account_id}"
|
||||
|
||||
if channel_type not in plugin_ids:
|
||||
result.add_issue(prefix, f"No plugin registered for channel_type '{channel_type}'")
|
||||
continue
|
||||
|
||||
plugin = next((p for p in configured_plugins if p.id == channel_type), None)
|
||||
if plugin is None:
|
||||
result.add_issue(prefix, f"Plugin '{channel_type}' disappeared during validation")
|
||||
continue
|
||||
|
||||
try:
|
||||
configured = plugin.is_configured(entry)
|
||||
except Exception as e:
|
||||
result.add_issue(prefix, f"is_configured check failed: {e}")
|
||||
continue
|
||||
|
||||
if not configured:
|
||||
reason = ""
|
||||
if callable(getattr(plugin, "unconfigured_reason", None)):
|
||||
try:
|
||||
reason = plugin.unconfigured_reason(entry, {}) or ""
|
||||
except Exception:
|
||||
pass
|
||||
msg = "Account not fully configured"
|
||||
if reason:
|
||||
msg = f"{msg}: {reason}"
|
||||
result.add_issue(prefix, msg, "warning")
|
||||
continue
|
||||
|
||||
if validate_schema and isinstance(plugin, ConfigSchemaProtocol):
|
||||
try:
|
||||
schema = plugin.config_schema()
|
||||
if isinstance(schema, dict) and schema:
|
||||
_validate_entry_against_schema(entry, schema, prefix, result)
|
||||
except Exception as e:
|
||||
logger.warning("Schema validation skipped for %s: %s", channel_type, e)
|
||||
|
||||
has_secret_refs = any(
|
||||
is_secret_ref(v) for v in entry.values() if isinstance(v, dict)
|
||||
)
|
||||
if has_secret_refs:
|
||||
try:
|
||||
from yuxi.channel.secrets import SecretResolver
|
||||
|
||||
resolver = SecretResolver()
|
||||
resolved = await resolver.resolve(entry)
|
||||
for issue in resolved.issues:
|
||||
result.add_issue(f"{prefix}.secrets", str(issue))
|
||||
for warning in resolved.warnings:
|
||||
result.add_issue(f"{prefix}.secrets", str(warning), "warning")
|
||||
except Exception as e:
|
||||
result.add_issue(f"{prefix}.secrets", f"Secret resolution failed: {e}")
|
||||
|
||||
return result
|
||||
Loading…
Reference in New Issue
Block a user