Coverage for .venv/lib/python3.13/site-packages/litellm/proxy/prompts/prompt_registry.py: 46%

167 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 12:01 +0000

1import importlib 

2import os 

3from collections.abc import Callable, Sequence 

4from pathlib import Path 

5from typing import Final 

6 

7from litellm._logging import verbose_proxy_logger 

8from litellm.integrations.custom_prompt_management import CustomPromptManagement 

9from litellm.types.prompts.init_prompts import ( 

10 PromptInfo, 

11 PromptLiteLLMParams, 

12 PromptSpec, 

13) 

14 

15prompt_initializer_registry = {} 

16 

17DEFAULT_PROMPT_ENVIRONMENT: Final = "development" 

18PROMPT_ENVIRONMENT_SERVE_PRECEDENCE: Final = ("production", "staging", "development") 

19 

20 

21def get_base_prompt_id(prompt_id: str) -> str: 

22 """ 

23 Extract the base prompt ID by stripping the version suffix if present. 

24 

25 Examples: 

26 >>> get_base_prompt_id("jack_success.v1") 

27 "jack_success" 

28 >>> get_base_prompt_id("jack_success_v1") 

29 "jack_success" 

30 >>> get_base_prompt_id("jack_success") 

31 "jack_success" 

32 """ 

33 if ".v" in prompt_id: 33 ↛ 34line 33 didn't jump to line 34 because the condition on line 33 was never true

34 return prompt_id.split(".v")[0] 

35 if "_v" in prompt_id: 35 ↛ 36line 35 didn't jump to line 36 because the condition on line 35 was never true

36 return prompt_id.split("_v")[0] 

37 return prompt_id 

38 

39 

40def get_version_number(prompt_id: str) -> int: 

41 """ 

42 Extract the version number from a versioned prompt ID (defaults to 1). 

43 

44 Examples: 

45 >>> get_version_number("jack_success.v2") 

46 2 

47 >>> get_version_number("jack_success_v2") 

48 2 

49 >>> get_version_number("jack_success") 

50 1 

51 """ 

52 if ".v" in prompt_id: 

53 version_str = prompt_id.split(".v")[1] 

54 try: 

55 return int(version_str) 

56 except ValueError: 

57 pass 

58 

59 if "_v" in prompt_id: 

60 version_str = prompt_id.split("_v")[1] 

61 try: 

62 return int(version_str) 

63 except ValueError: 

64 pass 

65 

66 return 1 

67 

68 

69def prompt_environment_or_default(environment: str | None) -> str: 

70 return environment or DEFAULT_PROMPT_ENVIRONMENT 

71 

72 

73def registry_key_for_prompt(prompt: PromptSpec) -> str: 

74 return f"{prompt.prompt_id}::{prompt_environment_or_default(prompt.environment)}" 

75 

76 

77def parse_prompt_version(raw_version: object) -> int | None: 

78 if isinstance(raw_version, bool): 

79 return None 

80 if isinstance(raw_version, int): 

81 return raw_version 

82 if isinstance(raw_version, str) and raw_version.isdigit(): 

83 return int(raw_version) 

84 return None 

85 

86 

87def _spec_version(prompt: PromptSpec) -> int: 

88 return prompt.version if prompt.version is not None else get_version_number(prompt_id=prompt.prompt_id) 

89 

90 

91def _default_serve_environment(prompts: Sequence[PromptSpec]) -> str: 

92 present: Final = frozenset(prompt_environment_or_default(prompt.environment) for prompt in prompts) 

93 ladder_pick: Final = next((env for env in PROMPT_ENVIRONMENT_SERVE_PRECEDENCE if env in present), None) 

94 if ladder_pick is not None: 

95 return ladder_pick 

96 return min(present) if present else DEFAULT_PROMPT_ENVIRONMENT 

97 

98 

99def get_prompt_initializer_from_integrations(): 

100 """ 

101 Get prompt initializers by discovering them from the prompt_integrations directory structure. 

102 

103 Scans the integrations directory for subdirectories containing __init__.py files 

104 with either prompt_initializer_registry or initialize_prompt functions. 

105 

106 Returns: 

107 Dict[str, Callable]: A dictionary mapping guardrail types to their initializer functions 

108 """ 

109 discovered_initializers: Final[dict[str, Callable]] = {} 

110 

111 try: 

112 # Get the path to the prompt_integrations directory 

113 current_dir: Final = Path(__file__).parent.parent.parent 

114 integrations_dir: Final = os.path.join(current_dir, "integrations") 

115 

116 if not os.path.exists(integrations_dir): 116 ↛ 117line 116 didn't jump to line 117 because the condition on line 116 was never true

117 verbose_proxy_logger.debug("integrations directory not found") 

118 return discovered_initializers 

119 

120 # Scan each subdirectory in prompt_integrations 

121 for item in os.listdir(integrations_dir): 

122 item_path = os.path.join(integrations_dir, item) 

123 

124 # Skip files and __pycache__ directories 

125 if not os.path.isdir(item_path) or item.startswith("__"): 

126 continue 

127 

128 # Check if the directory has an __init__.py file 

129 init_file = os.path.join(item_path, "__init__.py") 

130 if not os.path.exists(init_file): 

131 continue 

132 

133 module_path = f"litellm.integrations.{item}" 

134 try: 

135 # Import the module 

136 verbose_proxy_logger.debug("Discovering prompt integrations in: %s", module_path) 

137 

138 module = importlib.import_module(module_path) 

139 

140 # Check for prompt_initializer_registry dictionary 

141 if hasattr(module, "prompt_initializer_registry"): 

142 registry = getattr(module, "prompt_initializer_registry") 

143 if isinstance(registry, dict): 143 ↛ 121line 143 didn't jump to line 121 because the condition on line 143 was always true

144 discovered_initializers.update(registry) 

145 verbose_proxy_logger.debug( 

146 "Found prompt_initializer_registry in %s: %s", module_path, list(registry.keys()) 

147 ) 

148 

149 except ImportError as e: 

150 verbose_proxy_logger.error("Could not import %s: %s", module_path, e) 

151 continue 

152 except Exception as e: 

153 verbose_proxy_logger.error("Error processing %s: %s", module_path, e) 

154 continue 

155 

156 verbose_proxy_logger.debug( 

157 "Discovered %s prompt initializers: %s", len(discovered_initializers), list(discovered_initializers.keys()) 

158 ) 

159 

160 except Exception as e: 

161 verbose_proxy_logger.error("Error discovering prompt initializers: %s", e) 

162 

163 return discovered_initializers 

164 

165 

166prompt_initializer_registry = get_prompt_initializer_from_integrations() 

167 

168 

169class InMemoryPromptRegistry: 

170 """ 

171 Class that handles adding prompt callbacks to the CallbacksManager. 

172 """ 

173 

174 def __init__(self): 

175 self.IN_MEMORY_PROMPTS: dict[str, PromptSpec] = {} 

176 """ 

177 Prompt id to Prompt object mapping 

178 """ 

179 

180 self.prompt_id_to_custom_prompt: dict[str, CustomPromptManagement | None] = {} 

181 """ 

182 Guardrail id to CustomGuardrail object mapping 

183 """ 

184 

185 def initialize_prompt( 

186 self, 

187 prompt: PromptSpec, 

188 config_file_path: str | None = None, 

189 ) -> PromptSpec | None: 

190 """ 

191 Initialize a guardrail from a dictionary and add it to the litellm callback manager 

192 

193 Returns a Guardrail object if the guardrail is initialized successfully 

194 """ 

195 import litellm 

196 

197 registry_key: Final = registry_key_for_prompt(prompt) 

198 if registry_key in self.IN_MEMORY_PROMPTS: 198 ↛ 199line 198 didn't jump to line 199 because the condition on line 198 was never true

199 verbose_proxy_logger.debug("prompt already exists in IN_MEMORY_PROMPTS") 

200 return self.IN_MEMORY_PROMPTS[registry_key] 

201 

202 parsed_prompt, custom_prompt_callback = self._build_prompt_callback(prompt=prompt) 

203 litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback) 

204 

205 self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt 

206 self.prompt_id_to_custom_prompt[registry_key] = custom_prompt_callback 

207 

208 return parsed_prompt 

209 

210 def _build_prompt_callback(self, prompt: PromptSpec) -> tuple[PromptSpec, CustomPromptManagement]: 

211 litellm_params_data: Final = prompt.litellm_params 

212 verbose_proxy_logger.debug("litellm_params= %s", litellm_params_data) 

213 

214 if isinstance(litellm_params_data, dict): 214 ↛ 215line 214 didn't jump to line 215 because the condition on line 214 was never true

215 litellm_params = PromptLiteLLMParams(**litellm_params_data) 

216 else: 

217 litellm_params = litellm_params_data 

218 

219 prompt_integration: Final = litellm_params.prompt_integration 

220 if prompt_integration is None: 220 ↛ 221line 220 didn't jump to line 221 because the condition on line 220 was never true

221 raise ValueError("prompt_integration is required") 

222 

223 initializer: Final = prompt_initializer_registry.get(prompt_integration) 

224 if initializer is None: 224 ↛ 227line 224 didn't jump to line 227 because the condition on line 224 was always true

225 raise ValueError(f"Unsupported prompt: {prompt_integration}") 

226 

227 custom_prompt_callback: Final = initializer(litellm_params, prompt) 

228 if not isinstance(custom_prompt_callback, CustomPromptManagement): 

229 raise ValueError( # noqa: TRY004 # prompt endpoints map ValueError to HTTP 400; keep the existing contract 

230 f"CustomPromptManagement is required, got {type(custom_prompt_callback)}" 

231 ) 

232 

233 parsed_prompt: Final = PromptSpec( 

234 prompt_id=prompt.prompt_id, 

235 litellm_params=litellm_params, 

236 prompt_info=prompt.prompt_info or PromptInfo(prompt_type="config"), 

237 created_at=prompt.created_at, 

238 updated_at=prompt.updated_at, 

239 version=prompt.version, 

240 environment=prompt.environment, 

241 created_by=prompt.created_by, 

242 ) 

243 return parsed_prompt, custom_prompt_callback 

244 

245 def reload_prompt(self, prompt: PromptSpec) -> PromptSpec | None: 

246 import litellm 

247 

248 parsed_prompt, new_callback = self._build_prompt_callback(prompt=prompt) 

249 registry_key: Final = registry_key_for_prompt(parsed_prompt) 

250 stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None) 

251 self.IN_MEMORY_PROMPTS.pop(registry_key, None) 

252 if stale_callback is not None: 

253 litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback) 

254 litellm.logging_callback_manager.add_litellm_callback(new_callback) 

255 self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt 

256 self.prompt_id_to_custom_prompt[registry_key] = new_callback 

257 return parsed_prompt 

258 

259 def sync_prompt_from_db(self, prompt: PromptSpec) -> PromptSpec | None: 

260 existing: Final = self.IN_MEMORY_PROMPTS.get(registry_key_for_prompt(prompt)) 

261 if existing is None: 261 ↛ 263line 261 didn't jump to line 263 because the condition on line 261 was always true

262 return self.initialize_prompt(prompt=prompt) 

263 if existing.litellm_params == prompt.litellm_params and existing.prompt_info == prompt.prompt_info: 

264 return existing 

265 return self.reload_prompt(prompt=prompt) 

266 

267 def resolve_prompt_spec( 

268 self, 

269 prompt_id: str, 

270 version: int | None = None, 

271 environment: str | None = None, 

272 ) -> PromptSpec | None: 

273 """ 

274 Resolve a prompt spec by base prompt id, optional version, and optional environment. 

275 

276 With no environment, resolves within the default serve environment 

277 (production > staging > development > alphabetical first present). 

278 With no version, resolves to the highest version in the chosen environment. 

279 """ 

280 base_prompt_id: Final = get_base_prompt_id(prompt_id=prompt_id) 

281 base_matches: Final = tuple( 

282 spec 

283 for spec in self.IN_MEMORY_PROMPTS.values() 

284 if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id 

285 ) 

286 if not base_matches: 286 ↛ 288line 286 didn't jump to line 288 because the condition on line 286 was always true

287 return None 

288 resolved_environment: Final = ( 

289 environment if environment is not None else _default_serve_environment(base_matches) 

290 ) 

291 env_matches: Final = tuple( 

292 spec for spec in base_matches if prompt_environment_or_default(spec.environment) == resolved_environment 

293 ) 

294 if not env_matches: 

295 return None 

296 if version is not None: 

297 return next((spec for spec in env_matches if _spec_version(spec) == version), None) 

298 return max(env_matches, key=_spec_version) 

299 

300 def get_prompt_callback_for_prompt(self, prompt: PromptSpec) -> CustomPromptManagement | None: 

301 return self.prompt_id_to_custom_prompt.get(registry_key_for_prompt(prompt)) 

302 

303 def has_config_prompt(self, base_prompt_id: str) -> bool: 

304 return any( 

305 spec.prompt_info.prompt_type == "config" 

306 for spec in self.IN_MEMORY_PROMPTS.values() 

307 if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id 

308 ) 

309 

310 def remove_prompt(self, registry_key: str) -> None: 

311 import litellm 

312 

313 self.IN_MEMORY_PROMPTS.pop(registry_key, None) 

314 stale_callback: Final = self.prompt_id_to_custom_prompt.pop(registry_key, None) 

315 if stale_callback is not None: 

316 litellm.logging_callback_manager.remove_callback_from_all_lists(stale_callback) 

317 

318 def delete_prompts_by_base_id(self, base_prompt_id: str, environment: str | None = None) -> list[str]: 

319 """ 

320 Delete matching prompts from memory, along with their registered callbacks, 

321 scoped to one environment when given. 

322 

323 Returns the registry keys that were deleted. 

324 """ 

325 keys_to_delete: Final = [ 

326 key 

327 for key, spec in self.IN_MEMORY_PROMPTS.items() 

328 if get_base_prompt_id(prompt_id=spec.prompt_id) == base_prompt_id 

329 and (environment is None or prompt_environment_or_default(spec.environment) == environment) 

330 ] 

331 

332 for key in keys_to_delete: 

333 self.remove_prompt(registry_key=key) 

334 

335 return keys_to_delete 

336 

337 

338IN_MEMORY_PROMPT_REGISTRY: Final = InMemoryPromptRegistry()