feat: 配置加载 config(三 yaml + env 优先级 + 脱敏)
This commit is contained in:
@@ -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"))
|
||||
Vendored
+7
@@ -0,0 +1,7 @@
|
||||
server:
|
||||
max_upload_mb: 10
|
||||
session:
|
||||
sqlite_path: "C:/tmp/genesis.db"
|
||||
task_queue:
|
||||
backend: memory
|
||||
timeout_sec: 300
|
||||
Vendored
+8
@@ -0,0 +1,8 @@
|
||||
models:
|
||||
primary:
|
||||
provider: deepseek
|
||||
name: deepseek-chat
|
||||
llm_calls:
|
||||
max_context_tokens: 16000
|
||||
structured_output:
|
||||
max_parse_retry: 3
|
||||
Vendored
+9
@@ -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
|
||||
@@ -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"] == "***"
|
||||
Reference in New Issue
Block a user