feat: 配置加载 config(三 yaml + env 优先级 + 脱敏)

This commit is contained in:
lhl
2026-08-08 15:33:36 +08:00
parent 9d69bb3cd7
commit dc62670bb2
5 changed files with 281 additions and 0 deletions
+206
View File
@@ -0,0 +1,206 @@
from __future__ import annotations
import os
from pathlib import Path
from typing import Any
import yaml
from pydantic import BaseModel, Field
from pydantic_settings import BaseSettings, SettingsConfigDict
SECRET_KEYWORDS = ("key", "secret", "token")
ENV_PREFIX = "GENESIS_"
# ---------- 各 yaml 对应的 pydantic 模型 ----------
class ServerConfig(BaseModel):
max_upload_mb: int = 100
allowed_extensions: list[str] = Field(
default_factory=lambda: [".xlsx", ".xls", ".docx", ".pptx", ".java", ".xml", ".yml"]
)
class AppConfig(BaseModel):
name: str = "genesis"
version: str = "0.1.0"
timezone: str = "Asia/Tokyo"
server: ServerConfig = Field(default_factory=ServerConfig)
session: dict[str, Any] = Field(default_factory=lambda: {
"sqlite_path": "/data/db/genesis.db",
"snapshot_dir": "/data/db/snapshots",
})
paths: dict[str, Any] = Field(default_factory=lambda: {
"user_root": "/data/users",
"shared_root": "/data/shared",
})
task_queue: dict[str, Any] = Field(default_factory=lambda: {
"backend": "memory",
"redis_url": "",
"timeout_sec": 600,
"retry_default": 2,
})
class ModelSpec(BaseModel):
provider: str = "deepseek"
name: str = "deepseek-chat"
temperature: float = 0.2
max_tokens: int = 4096
timeout_sec: int = 60
retry_backoff: list[float] = Field(default_factory=lambda: [1.0, 3.0, 7.0])
class InferenceModels(BaseModel):
primary: ModelSpec = Field(default_factory=ModelSpec)
fallback: ModelSpec = Field(default_factory=lambda: ModelSpec(provider="qwen", name="qwen-max"))
vision: ModelSpec = Field(default_factory=lambda: ModelSpec(name="deepseek-vl", timeout_sec=90))
class LlmCallsConfig(BaseModel):
token_estimation: str = "tiktoken"
max_context_tokens: int = 32000
truncation_policy: dict[str, Any] = Field(default_factory=lambda: {
"priority": ["shrink_rule_chunks", "summarize_history", "truncate_data"],
})
class StructuredOutputConfig(BaseModel):
max_parse_retry: int = 2
class PromptRegistryConfig(BaseModel):
prompts_dir: str = "./prompts"
default_version: str = "latest"
class InferenceConfig(BaseModel):
models: InferenceModels = Field(default_factory=InferenceModels)
llm_calls: LlmCallsConfig = Field(default_factory=LlmCallsConfig)
structured_output: StructuredOutputConfig = Field(default_factory=StructuredOutputConfig)
prompt_registry: PromptRegistryConfig = Field(default_factory=PromptRegistryConfig)
class EmbeddingConfig(BaseModel):
model: str = "BAAI/bge-small-zh-v1.5"
device: str = "cpu"
max_batch_size: int = 32
cache_dir: str = "/data/shared/models"
class ChromaStoreConfig(BaseModel):
persist_dir: str = "/data/shared/rules-handbook/chroma"
class QdrantStoreConfig(BaseModel):
url: str = "http://qdrant:6333"
api_key: str = ""
class VectorStoreConfig(BaseModel):
adapter: str = "chroma"
chroma: ChromaStoreConfig = Field(default_factory=ChromaStoreConfig)
qdrant: QdrantStoreConfig = Field(default_factory=QdrantStoreConfig)
class ChunkingConfig(BaseModel):
word_max_tokens: int = 512
excel_rule_block_rows: int = 10
ppt_pages_per_chunk: int = 2
min_tokens: int = 30
class RetrievalConfig(BaseModel):
channel_top_k: int = 10
rrf_k: int = 60
default_top_k: int = 5
contextual_enrichment: bool = True
class RagConfig(BaseModel):
embedding: EmbeddingConfig = Field(default_factory=EmbeddingConfig)
vector_store: VectorStoreConfig = Field(default_factory=VectorStoreConfig)
chunking: ChunkingConfig = Field(default_factory=ChunkingConfig)
retrieval: RetrievalConfig = Field(default_factory=RetrievalConfig)
# ---------- 加载辅助 ----------
def _expand_env(data: Any) -> Any:
"""递归展开 ${VAR} 占位(读环境变量,缺失→空串)"""
if isinstance(data, dict):
return {k: _expand_env(v) for k, v in data.items()}
if isinstance(data, list):
return [_expand_env(v) for v in data]
if isinstance(data, str) and data.startswith("${") and data.endswith("}"):
return os.environ.get(data[2:-1], "")
return data
def _deep_merge(base: dict, override: dict) -> dict:
"""递归合并:override 覆盖 base;非 dict 值直接取 override 存在者"""
out = dict(base)
for k, v in override.items():
if isinstance(v, dict) and isinstance(out.get(k), dict):
out[k] = _deep_merge(out[k], v)
else:
out[k] = v
return out
def _env_overrides() -> dict:
"""收集 GENESIS_ 前缀的条目为嵌套 dict,__ 为嵌套分隔(键统一小写以匹配 yaml)"""
result: dict[str, Any] = {}
for key, value in os.environ.items():
if key.startswith(ENV_PREFIX):
parts = key[len(ENV_PREFIX):].split("__")
node = result
for part in parts[:-1]:
node = node.setdefault(part.lower(), {})
node[parts[-1].lower()] = value
return result
def _load_yaml(config_dir: Path, name: str) -> dict:
path = config_dir / f"{name}.yaml"
if not path.exists():
return {}
with path.open("r", encoding="utf-8") as f:
return yaml.safe_load(f) or {}
def _redact(data: dict) -> dict:
out = {}
for k, v in data.items():
if any(kw in str(k).lower() for kw in SECRET_KEYWORDS):
out[k] = "***"
elif isinstance(v, dict):
out[k] = _redact(v)
else:
out[k] = v
return out
# ---------- 根 Settings ----------
class Settings(BaseSettings):
model_config = SettingsConfigDict(env_prefix=ENV_PREFIX, env_file=".env", extra="ignore")
app: AppConfig = Field(default_factory=AppConfig)
inference: InferenceConfig = Field(default_factory=InferenceConfig)
rag: RagConfig = Field(default_factory=RagConfig)
@classmethod
def from_dir(cls, config_dir: Path | str) -> "Settings":
config_dir = Path(config_dir)
raw = {
"app": _load_yaml(config_dir, "app"),
"inference": _load_yaml(config_dir, "inference"),
"rag": _load_yaml(config_dir, "rag"),
}
env = _env_overrides()
merged = {k: _deep_merge(raw[k], env.get(k, {})) for k in raw}
return cls(**{k: _expand_env(v) for k, v in merged.items()})
def get_redacted(self) -> dict:
return _redact(self.model_dump(mode="json"))
+7
View File
@@ -0,0 +1,7 @@
server:
max_upload_mb: 10
session:
sqlite_path: "C:/tmp/genesis.db"
task_queue:
backend: memory
timeout_sec: 300
+8
View File
@@ -0,0 +1,8 @@
models:
primary:
provider: deepseek
name: deepseek-chat
llm_calls:
max_context_tokens: 16000
structured_output:
max_parse_retry: 3
+9
View File
@@ -0,0 +1,9 @@
embedding:
model: BAAI/bge-small-zh-v1.5
vector_store:
adapter: chroma
qdrant:
url: http://qdrant:6333
api_key: ${QDRANT_API_KEY}
retrieval:
rrf_k: 42
+51
View File
@@ -0,0 +1,51 @@
from pathlib import Path
from genesis.config import Settings
FIXTURES = Path(__file__).parent / "fixtures"
def test_from_dir_maps_yaml_fields():
s = Settings.from_dir(FIXTURES)
assert s.app.server.max_upload_mb == 10
assert s.app.session["sqlite_path"] == "C:/tmp/genesis.db"
assert s.app.task_queue["timeout_sec"] == 300
assert s.inference.models.primary.name == "deepseek-chat"
assert s.inference.llm_calls.max_context_tokens == 16000
assert s.inference.structured_output.max_parse_retry == 3
assert s.rag.embedding.model == "BAAI/bge-small-zh-v1.5"
assert s.rag.retrieval.rrf_k == 42
def test_defaults_when_dir_empty(tmp_path):
s = Settings.from_dir(tmp_path)
assert s.app.name == "genesis"
assert s.app.server.max_upload_mb == 100
assert s.app.task_queue["backend"] == "memory"
assert s.inference.models.primary.name == "deepseek-chat"
assert s.rag.embedding.model == "BAAI/bge-small-zh-v1.5"
assert s.rag.retrieval.rrf_k == 60
def test_env_override_yaml(monkeypatch):
monkeypatch.setenv("GENESIS_APP__SERVER__MAX_UPLOAD_MB", "25")
s = Settings.from_dir(FIXTURES)
assert s.app.server.max_upload_mb == 25
def test_env_create_missing_key(monkeypatch):
monkeypatch.setenv("GENESIS_RAG__RETRIEVAL__DEFAULT_TOP_K", "7")
s = Settings.from_dir(FIXTURES)
assert s.rag.retrieval.default_top_k == 7
def test_env_placeholder_expansion(monkeypatch):
monkeypatch.setenv("QDRANT_API_KEY", "sk-test-xyz")
s = Settings.from_dir(FIXTURES)
assert s.rag.vector_store.qdrant.api_key == "sk-test-xyz"
def test_redacted_hides_secrets():
s = Settings.from_dir(FIXTURES)
red = s.get_redacted()
assert red["rag"]["vector_store"]["qdrant"]["api_key"] == "***"