Coverage for src\genesis\config.py: 100%

132 statements  

« prev     ^ index     » next       coverage.py v7.15.4, created at 2026-08-26 14:20 +0800

1from __future__ import annotations 

2 

3import os 

4from pathlib import Path 

5from typing import Any, Literal 

6 

7import yaml 

8from pydantic import BaseModel, Field 

9from pydantic_settings import BaseSettings, SettingsConfigDict 

10 

11SECRET_KEYWORDS = ("key", "secret", "token") 

12ENV_PREFIX = "GENESIS_" 

13 

14 

15# ---------- 各 yaml 对应的 pydantic 模型 ---------- 

16 

17class ServerConfig(BaseModel): 

18 max_upload_mb: int = 100 

19 allowed_extensions: list[str] = Field( 

20 default_factory=lambda: [".xlsx", ".xls", ".docx", ".pptx", ".java", ".xml", ".yml", 

21 ".py", ".ts", ".go", ".cs"] 

22 ) 

23 

24 

25class AppConfig(BaseModel): 

26 name: str = "genesis" 

27 version: str = "0.1.0" 

28 timezone: str = "Asia/Tokyo" 

29 server: ServerConfig = Field(default_factory=ServerConfig) 

30 session: dict[str, Any] = Field(default_factory=lambda: { 

31 "sqlite_path": "/data/db/genesis.db", 

32 "snapshot_dir": "/data/db/snapshots", 

33 }) 

34 paths: dict[str, Any] = Field(default_factory=lambda: { 

35 "user_root": "/data/users", 

36 "shared_root": "/data/shared", 

37 }) 

38 task_queue: dict[str, Any] = Field(default_factory=lambda: { 

39 "backend": "memory", 

40 "timeout_sec": 600, 

41 "retry_default": 2, 

42 }) 

43 

44 

45class ModelSpec(BaseModel): 

46 provider: str = "deepseek" 

47 name: str = "deepseek-chat" 

48 temperature: float = 0.2 

49 max_tokens: int = 4096 

50 timeout_sec: int = 60 

51 retry_backoff: list[float] = Field(default_factory=lambda: [1.0, 3.0, 7.0]) 

52 

53 

54class InferenceModels(BaseModel): 

55 primary: ModelSpec = Field(default_factory=ModelSpec) 

56 fallback: ModelSpec = Field(default_factory=lambda: ModelSpec(provider="qwen", name="qwen-max")) 

57 vision: ModelSpec = Field(default_factory=lambda: ModelSpec(name="deepseek-vl", timeout_sec=90)) 

58 

59 

60class LlmCallsConfig(BaseModel): 

61 token_estimation: str = "tiktoken" 

62 max_context_tokens: int = 32000 

63 truncation_policy: dict[str, Any] = Field(default_factory=lambda: { 

64 "priority": ["shrink_rule_chunks", "summarize_history", "truncate_data"], 

65 }) 

66 

67 

68class StructuredOutputConfig(BaseModel): 

69 max_parse_retry: int = 2 

70 

71 

72class PromptRegistryConfig(BaseModel): 

73 prompts_dir: str = "./prompts" 

74 default_version: str = "latest" 

75 

76 

77class InferenceConfig(BaseModel): 

78 models: InferenceModels = Field(default_factory=InferenceModels) 

79 llm_calls: LlmCallsConfig = Field(default_factory=LlmCallsConfig) 

80 structured_output: StructuredOutputConfig = Field(default_factory=StructuredOutputConfig) 

81 prompt_registry: PromptRegistryConfig = Field(default_factory=PromptRegistryConfig) 

82 

83 

84class EmbeddingConfig(BaseModel): 

85 # OV2/T11:实际语料为日文,bge-small-zh 面向中文 → 默认多语言 bge-m3(中/日/英) 

86 model: str = "BAAI/bge-m3" 

87 device: str = "cpu" 

88 max_batch_size: int = 32 

89 cache_dir: str = "/data/shared/models" 

90 

91 

92class ChromaStoreConfig(BaseModel): 

93 persist_dir: str = "/data/shared/rules-handbook/chroma" 

94 

95 

96class VectorStoreConfig(BaseModel): 

97 adapter: str = "chroma" 

98 chroma: ChromaStoreConfig = Field(default_factory=ChromaStoreConfig) 

99 

100 

101class ChunkingConfig(BaseModel): 

102 word_max_tokens: int = 512 

103 excel_rule_block_rows: int = 10 

104 ppt_pages_per_chunk: int = 2 

105 min_tokens: int = 30 

106 

107 

108class RetrievalConfig(BaseModel): 

109 channel_top_k: int = 10 

110 rrf_k: int = 60 

111 default_top_k: int = 5 

112 contextual_enrichment: bool = True 

113 

114 

115class RerankConfig(BaseModel): 

116 # I6/T6:v1 引入 rerank 精排(2026 主流实践:向量→rerank→精排) 

117 enabled: bool = True 

118 model: str = "BAAI/bge-reranker-v2-m3" 

119 device: str = "cpu" 

120 

121 

122class RagConfig(BaseModel): 

123 embedding: EmbeddingConfig = Field(default_factory=EmbeddingConfig) 

124 vector_store: VectorStoreConfig = Field(default_factory=VectorStoreConfig) 

125 chunking: ChunkingConfig = Field(default_factory=ChunkingConfig) 

126 retrieval: RetrievalConfig = Field(default_factory=RetrievalConfig) 

127 rerank: RerankConfig = Field(default_factory=RerankConfig) 

128 

129 

130class WriterConfig(BaseModel): 

131 """Writer 子系统配置(步骤 1:输出语言参数)。 

132 

133 output_language: 生成概要设计书正文的自然语言 

134 - "auto":与章节标题所用语言保持一致(默认,向后兼容既有日文文档) 

135 - "zh":强制简体中文 

136 - "ja":强制日文 

137 表格数据始终照抄源 Excel 原文(不翻译),见 design.md §7.2。 

138 """ 

139 

140 output_language: Literal["auto", "zh", "ja"] = "auto" 

141 

142 

143# ---------- 加载辅助 ---------- 

144 

145def _expand_env(data: Any) -> Any: 

146 """递归展开 ${VAR} 占位(读环境变量,缺失→空串)""" 

147 if isinstance(data, dict): 

148 return {k: _expand_env(v) for k, v in data.items()} 

149 if isinstance(data, list): 

150 return [_expand_env(v) for v in data] 

151 if isinstance(data, str) and data.startswith("${") and data.endswith("}"): 

152 return os.environ.get(data[2:-1], "") 

153 return data 

154 

155 

156def _deep_merge(base: dict, override: dict) -> dict: 

157 """递归合并:override 覆盖 base;非 dict 值直接取 override 存在者""" 

158 out = dict(base) 

159 for k, v in override.items(): 

160 if isinstance(v, dict) and isinstance(out.get(k), dict): 

161 out[k] = _deep_merge(out[k], v) 

162 else: 

163 out[k] = v 

164 return out 

165 

166 

167def _env_overrides() -> dict: 

168 """收集 GENESIS_ 前缀的条目为嵌套 dict,__ 为嵌套分隔(键统一小写以匹配 yaml)""" 

169 result: dict[str, Any] = {} 

170 for key, value in os.environ.items(): 

171 if key.startswith(ENV_PREFIX): 

172 parts = key[len(ENV_PREFIX):].split("__") 

173 node = result 

174 for part in parts[:-1]: 

175 node = node.setdefault(part.lower(), {}) 

176 node[parts[-1].lower()] = value 

177 return result 

178 

179 

180def _load_yaml(config_dir: Path, name: str) -> dict: 

181 path = config_dir / f"{name}.yaml" 

182 if not path.exists(): 

183 return {} 

184 with path.open("r", encoding="utf-8") as f: 

185 return yaml.safe_load(f) or {} 

186 

187 

188def _redact(data: dict) -> dict: 

189 out = {} 

190 for k, v in data.items(): 

191 if any(kw in str(k).lower() for kw in SECRET_KEYWORDS): 

192 out[k] = "***" 

193 elif isinstance(v, dict): 

194 out[k] = _redact(v) 

195 else: 

196 out[k] = v 

197 return out 

198 

199 

200# ---------- 根 Settings ---------- 

201 

202class Settings(BaseSettings): 

203 model_config = SettingsConfigDict(env_prefix=ENV_PREFIX, env_file=".env", extra="ignore") 

204 

205 app: AppConfig = Field(default_factory=AppConfig) 

206 inference: InferenceConfig = Field(default_factory=InferenceConfig) 

207 rag: RagConfig = Field(default_factory=RagConfig) 

208 writer: WriterConfig = Field(default_factory=WriterConfig) 

209 

210 @classmethod 

211 def from_dir(cls, config_dir: Path | str) -> "Settings": 

212 config_dir = Path(config_dir) 

213 raw = { 

214 "app": _load_yaml(config_dir, "app"), 

215 "inference": _load_yaml(config_dir, "inference"), 

216 "rag": _load_yaml(config_dir, "rag"), 

217 } 

218 env = _env_overrides() 

219 merged = {k: _deep_merge(raw[k], env.get(k, {})) for k in raw} 

220 return cls(**{k: _expand_env(v) for k, v in merged.items()}) 

221 

222 def get_redacted(self) -> dict: 

223 return _redact(self.model_dump(mode="json"))