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
« 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
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)
15prompt_initializer_registry = {}
17DEFAULT_PROMPT_ENVIRONMENT: Final = "development"
18PROMPT_ENVIRONMENT_SERVE_PRECEDENCE: Final = ("production", "staging", "development")
21def get_base_prompt_id(prompt_id: str) -> str:
22 """
23 Extract the base prompt ID by stripping the version suffix if present.
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
40def get_version_number(prompt_id: str) -> int:
41 """
42 Extract the version number from a versioned prompt ID (defaults to 1).
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
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
66 return 1
69def prompt_environment_or_default(environment: str | None) -> str:
70 return environment or DEFAULT_PROMPT_ENVIRONMENT
73def registry_key_for_prompt(prompt: PromptSpec) -> str:
74 return f"{prompt.prompt_id}::{prompt_environment_or_default(prompt.environment)}"
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
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)
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
99def get_prompt_initializer_from_integrations():
100 """
101 Get prompt initializers by discovering them from the prompt_integrations directory structure.
103 Scans the integrations directory for subdirectories containing __init__.py files
104 with either prompt_initializer_registry or initialize_prompt functions.
106 Returns:
107 Dict[str, Callable]: A dictionary mapping guardrail types to their initializer functions
108 """
109 discovered_initializers: Final[dict[str, Callable]] = {}
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")
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
120 # Scan each subdirectory in prompt_integrations
121 for item in os.listdir(integrations_dir):
122 item_path = os.path.join(integrations_dir, item)
124 # Skip files and __pycache__ directories
125 if not os.path.isdir(item_path) or item.startswith("__"):
126 continue
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
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)
138 module = importlib.import_module(module_path)
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 )
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
156 verbose_proxy_logger.debug(
157 "Discovered %s prompt initializers: %s", len(discovered_initializers), list(discovered_initializers.keys())
158 )
160 except Exception as e:
161 verbose_proxy_logger.error("Error discovering prompt initializers: %s", e)
163 return discovered_initializers
166prompt_initializer_registry = get_prompt_initializer_from_integrations()
169class InMemoryPromptRegistry:
170 """
171 Class that handles adding prompt callbacks to the CallbacksManager.
172 """
174 def __init__(self):
175 self.IN_MEMORY_PROMPTS: dict[str, PromptSpec] = {}
176 """
177 Prompt id to Prompt object mapping
178 """
180 self.prompt_id_to_custom_prompt: dict[str, CustomPromptManagement | None] = {}
181 """
182 Guardrail id to CustomGuardrail object mapping
183 """
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
193 Returns a Guardrail object if the guardrail is initialized successfully
194 """
195 import litellm
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]
202 parsed_prompt, custom_prompt_callback = self._build_prompt_callback(prompt=prompt)
203 litellm.logging_callback_manager.add_litellm_callback(custom_prompt_callback)
205 self.IN_MEMORY_PROMPTS[registry_key] = parsed_prompt
206 self.prompt_id_to_custom_prompt[registry_key] = custom_prompt_callback
208 return parsed_prompt
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)
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
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")
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}")
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 )
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
245 def reload_prompt(self, prompt: PromptSpec) -> PromptSpec | None:
246 import litellm
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
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)
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.
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)
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))
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 )
310 def remove_prompt(self, registry_key: str) -> None:
311 import litellm
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)
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.
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 ]
332 for key in keys_to_delete:
333 self.remove_prompt(registry_key=key)
335 return keys_to_delete
338IN_MEMORY_PROMPT_REGISTRY: Final = InMemoryPromptRegistry()