diff --git a/src/genesis/config.py b/src/genesis/config.py new file mode 100644 index 0000000..0d53a1f --- /dev/null +++ b/src/genesis/config.py @@ -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")) \ No newline at end of file diff --git a/tests/fixtures/app.yaml b/tests/fixtures/app.yaml new file mode 100644 index 0000000..e4e5e7a --- /dev/null +++ b/tests/fixtures/app.yaml @@ -0,0 +1,7 @@ +server: + max_upload_mb: 10 +session: + sqlite_path: "C:/tmp/genesis.db" +task_queue: + backend: memory + timeout_sec: 300 \ No newline at end of file diff --git a/tests/fixtures/inference.yaml b/tests/fixtures/inference.yaml new file mode 100644 index 0000000..e362ad3 --- /dev/null +++ b/tests/fixtures/inference.yaml @@ -0,0 +1,8 @@ +models: + primary: + provider: deepseek + name: deepseek-chat +llm_calls: + max_context_tokens: 16000 +structured_output: + max_parse_retry: 3 \ No newline at end of file diff --git a/tests/fixtures/rag.yaml b/tests/fixtures/rag.yaml new file mode 100644 index 0000000..4d50bbb --- /dev/null +++ b/tests/fixtures/rag.yaml @@ -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 \ No newline at end of file diff --git a/tests/test_config.py b/tests/test_config.py new file mode 100644 index 0000000..3f28bc7 --- /dev/null +++ b/tests/test_config.py @@ -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"] == "***" \ No newline at end of file