Coverage for agentos/core/prompt_manager.py: 0%
193 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 10:59 +0800
1"""
2AgentOS Prompt Manager — Versioned Prompt Templates with A/B Testing
3━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
5Production-grade prompt management:
6 - Versioned prompt templates with semantic versioning
7 - Variable interpolation with type validation
8 - A/B testing with traffic splitting
9 - Prompt lineage and diff history
10 - Import/export (JSON, YAML)
11 - System/user/assistant role support
13Architecture:
14 PromptTemplate → single versioned template
15 PromptStore → registry of templates
16 PromptRenderer → interpolate variables into final prompt
17 ABTestManager → split traffic between template variants
18"""
20from __future__ import annotations
22import difflib
23import json
24import re
25from collections import defaultdict
26from dataclasses import dataclass, field
27from datetime import datetime, timezone
28from enum import Enum
29from typing import Any, Dict, List, Optional, Set, Tuple
32# ---------------------------------------------------------------------------
33# Prompt Template
34# ---------------------------------------------------------------------------
37class PromptRole(str, Enum):
38 SYSTEM = "system"
39 USER = "user"
40 ASSISTANT = "assistant"
43@dataclass
44class PromptTemplate:
45 """A versioned prompt template."""
46 name: str
47 version: int = 1
48 content: str = ""
49 role: PromptRole = PromptRole.SYSTEM
50 variables: Set[str] = field(default_factory=set)
51 description: str = ""
52 tags: List[str] = field(default_factory=list)
53 author: str = ""
54 created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
55 updated_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
56 is_active: bool = True
57 metadata: Dict[str, Any] = field(default_factory=dict)
59 # Variable pattern: {{variable_name}}
60 VAR_PATTERN = re.compile(r"\{\{(\w+)\}\}")
62 def extract_variables(self) -> Set[str]:
63 """Extract variable names from template content."""
64 return set(self.VAR_PATTERN.findall(self.content))
66 def render(self, values: Dict[str, str], strict: bool = True) -> str:
67 """
68 Render the template by substituting variables.
70 Args:
71 values: Dict of variable_name → value
72 strict: If True, raise on missing variables. If False, leave placeholders.
74 Raises:
75 ValueError: If strict=True and a variable is missing.
76 """
77 required = self.extract_variables()
79 if strict:
80 missing = required - set(values.keys())
81 if missing:
82 raise ValueError(
83 f"Template '{self.name}' v{self.version} missing variables: {missing}"
84 )
86 result = self.content
87 for var_name in required:
88 if var_name in values:
89 result = result.replace(f"{{{{{var_name}}}}}", values[var_name])
90 elif not strict:
91 # Leave placeholder intact
92 pass
94 return result
96 def validate(self) -> Tuple[bool, List[str]]:
97 """
98 Validate template structure.
100 Returns (is_valid, list_of_issues).
101 """
102 issues = []
103 if not self.name:
104 issues.append("Name is required")
105 if not self.content:
106 issues.append("Content is empty")
107 if self.VAR_PATTERN.findall(self.content):
108 # Variables are OK, just note them
109 pass
110 return len(issues) == 0, issues
112 def diff(self, other: "PromptTemplate") -> str:
113 """Generate a unified diff between this and another template."""
114 a_lines = self.content.splitlines(keepends=True)
115 b_lines = other.content.splitlines(keepends=True)
116 diff = difflib.unified_diff(
117 a_lines, b_lines,
118 fromfile=f"{self.name} v{self.version}",
119 tofile=f"{other.name} v{other.version}",
120 )
121 return "".join(diff)
123 def to_dict(self) -> Dict[str, Any]:
124 return {
125 "name": self.name,
126 "version": self.version,
127 "content": self.content,
128 "role": self.role.value,
129 "variables": sorted(self.variables),
130 "description": self.description,
131 "tags": self.tags,
132 "author": self.author,
133 "created_at": self.created_at,
134 "updated_at": self.updated_at,
135 "is_active": self.is_active,
136 "metadata": self.metadata,
137 }
139 @classmethod
140 def from_dict(cls, data: Dict[str, Any]) -> "PromptTemplate":
141 return cls(
142 name=data["name"],
143 version=data.get("version", 1),
144 content=data["content"],
145 role=PromptRole(data.get("role", "system")),
146 variables=set(data.get("variables", [])),
147 description=data.get("description", ""),
148 tags=data.get("tags", []),
149 author=data.get("author", ""),
150 created_at=data.get("created_at", ""),
151 updated_at=data.get("updated_at", ""),
152 is_active=data.get("is_active", True),
153 metadata=data.get("metadata", {}),
154 )
157# ---------------------------------------------------------------------------
158# Prompt Store
159# ---------------------------------------------------------------------------
162class PromptStore:
163 """
164 Registry of prompt templates with versioning.
166 Supports:
167 - Semantic versioning per template
168 - Latest/active version resolution
169 - Template lineage tracking
170 - Import/export (JSON)
171 """
173 def __init__(self):
174 self._templates: Dict[str, Dict[int, PromptTemplate]] = defaultdict(dict)
175 self._latest: Dict[str, int] = {}
176 self._active: Dict[str, int] = {}
177 self._lineage: Dict[str, List[int]] = defaultdict(list)
179 def add(self, template: PromptTemplate) -> PromptTemplate:
180 """
181 Add a new template or create a new version.
183 Auto-increments version if template name already exists.
184 """
185 name = template.name
187 if name in self._latest:
188 # New version
189 latest = self._latest[name]
190 template.version = latest + 1
191 else:
192 template.version = template.version or 1
194 if not template.variables:
195 template.variables = template.extract_variables()
197 template.updated_at = datetime.now(timezone.utc).isoformat()
198 self._templates[name][template.version] = template
199 self._latest[name] = template.version
200 self._active[name] = template.version
201 self._lineage[name].append(template.version)
203 return template
205 def get(self, name: str, version: Optional[int] = None) -> Optional[PromptTemplate]:
206 """Get a template by name and optional version. Defaults to active version."""
207 if name not in self._templates:
208 return None
210 if version is not None:
211 return self._templates[name].get(version)
213 active_ver = self._active.get(name) or self._latest[name]
214 return self._templates[name].get(active_ver)
216 def get_latest(self, name: str) -> Optional[PromptTemplate]:
217 """Get the latest version of a template."""
218 if name not in self._latest:
219 return None
220 return self._templates[name].get(self._latest[name])
222 def set_active(self, name: str, version: int) -> bool:
223 """Set which version is the active one."""
224 if name not in self._templates or version not in self._templates[name]:
225 return False
226 self._active[name] = version
227 return True
229 def deactivate(self, name: str) -> None:
230 """Deactivate a template (no active version)."""
231 self._active.pop(name, None)
233 def list_templates(self) -> List[Dict[str, Any]]:
234 """List all templates with their versions."""
235 result = []
236 for name, versions in self._templates.items():
237 latest = self._latest[name]
238 active = self._active.get(name, latest)
239 result.append({
240 "name": name,
241 "versions": sorted(versions.keys()),
242 "latest": latest,
243 "active": active,
244 "total_versions": len(versions),
245 })
246 return sorted(result, key=lambda x: x["name"])
248 def get_history(self, name: str) -> List[PromptTemplate]:
249 """Get all versions of a template in chronological order."""
250 if name not in self._templates:
251 return []
252 return [self._templates[name][v] for v in sorted(self._templates[name].keys())]
254 def diff_versions(self, name: str, v1: int, v2: int) -> Optional[str]:
255 """Get diff between two versions of a template."""
256 t1 = self.get(name, v1)
257 t2 = self.get(name, v2)
258 if t1 is None or t2 is None:
259 return None
260 return t1.diff(t2)
262 def remove(self, name: str, version: Optional[int] = None) -> int:
263 """
264 Remove template(s). If version is None, remove all versions.
265 Returns number of versions removed.
266 """
267 if name not in self._templates:
268 return 0
270 if version is not None:
271 if version in self._templates[name]:
272 del self._templates[name][version]
273 if version == self._latest.get(name):
274 remaining = sorted(self._templates[name].keys())
275 self._latest[name] = remaining[-1] if remaining else 0
276 return 1
277 return 0
279 count = len(self._templates[name])
280 del self._templates[name]
281 self._latest.pop(name, None)
282 self._active.pop(name, None)
283 self._lineage.pop(name, None)
284 return count
286 def export_json(self, names: Optional[List[str]] = None) -> str:
287 """Export templates as JSON."""
288 templates = []
289 target = names or list(self._templates.keys())
290 for name in target:
291 for t in self._templates.get(name, {}).values():
292 templates.append(t.to_dict())
293 return json.dumps({"templates": templates}, indent=2, ensure_ascii=False)
295 def import_json(self, json_str: str) -> int:
296 """Import templates from JSON. Returns count of imported templates."""
297 data = json.loads(json_str)
298 count = 0
299 for td in data.get("templates", []):
300 template = PromptTemplate.from_dict(td)
301 self.add(template)
302 count += 1
303 return count
306# ---------------------------------------------------------------------------
307# A/B Test Manager
308# ---------------------------------------------------------------------------
311@dataclass
312class ABTest:
313 """An A/B test comparing two prompt template versions."""
314 name: str
315 template_name: str
316 variant_a_version: int
317 variant_b_version: int
318 split_ratio: float = 0.5 # 0.5 = 50% each
319 is_active: bool = True
320 created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
322 def route(self, session_id: str) -> int:
323 """Route a session to variant A (0) or B (1)."""
324 if not self.is_active:
325 return 0 # Default to A if test is inactive
327 # Deterministic routing based on session hash
328 hash_val = abs(hash(f"{self.name}:{session_id}"))
329 bucket = (hash_val % 100) / 100.0
330 return 0 if bucket < self.split_ratio else 1
333class ABTestManager:
334 """Manage A/B tests for prompt templates."""
336 def __init__(self, store: PromptStore):
337 self._store = store
338 self._tests: Dict[str, ABTest] = {}
339 self._results: Dict[str, Dict[str, int]] = defaultdict(
340 lambda: {"a_served": 0, "b_served": 0}
341 )
343 def create_test(
344 self,
345 name: str,
346 template_name: str,
347 variant_a_version: int,
348 variant_b_version: int,
349 split_ratio: float = 0.5,
350 ) -> ABTest:
351 """Create a new A/B test."""
352 test = ABTest(
353 name=name,
354 template_name=template_name,
355 variant_a_version=variant_a_version,
356 variant_b_version=variant_b_version,
357 split_ratio=split_ratio,
358 )
359 self._tests[name] = test
360 return test
362 def get_template(self, test_name: str, session_id: str) -> Optional[PromptTemplate]:
363 """
364 Get the prompt template for a session in an A/B test.
366 Returns None if the test doesn't exist or template not found.
367 """
368 test = self._tests.get(test_name)
369 if test is None:
370 return None
372 variant = test.route(session_id)
373 version = test.variant_a_version if variant == 0 else test.variant_b_version
375 self._results[test_name][f"{'a' if variant == 0 else 'b'}_served"] += 1
377 return self._store.get(test.template_name, version)
379 def get_results(self, name: str) -> Dict[str, Any]:
380 """Get results for an A/B test."""
381 test = self._tests.get(name)
382 if test is None:
383 return {}
385 results = dict(self._results[name])
386 total = results.get("a_served", 0) + results.get("b_served", 0)
387 return {
388 "test_name": name,
389 "template_name": test.template_name,
390 "variant_a_version": test.variant_a_version,
391 "variant_b_version": test.variant_b_version,
392 "split_ratio": test.split_ratio,
393 "is_active": test.is_active,
394 **results,
395 "total_served": total,
396 "a_pct": round(results.get("a_served", 0) / max(1, total) * 100, 1),
397 "b_pct": round(results.get("b_served", 0) / max(1, total) * 100, 1),
398 }
400 def stop_test(self, name: str) -> None:
401 """Stop an active A/B test."""
402 if name in self._tests:
403 self._tests[name].is_active = False
405 def list_tests(self) -> List[Dict[str, Any]]:
406 """List all A/B tests."""
407 return [self.get_results(name) for name in self._tests]