Coverage for agentos/core/prompt_manager.py: 0%
195 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 08:01 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-06 08:01 +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 copy
23import difflib
24import json
25import re
26import time
27from collections import defaultdict
28from dataclasses import dataclass, field
29from datetime import datetime, timezone
30from enum import Enum
31from typing import Any, Callable, Dict, List, Optional, Set, Tuple
34# ---------------------------------------------------------------------------
35# Prompt Template
36# ---------------------------------------------------------------------------
39class PromptRole(str, Enum):
40 SYSTEM = "system"
41 USER = "user"
42 ASSISTANT = "assistant"
45@dataclass
46class PromptTemplate:
47 """A versioned prompt template."""
48 name: str
49 version: int = 1
50 content: str = ""
51 role: PromptRole = PromptRole.SYSTEM
52 variables: Set[str] = field(default_factory=set)
53 description: str = ""
54 tags: List[str] = field(default_factory=list)
55 author: str = ""
56 created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
57 updated_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
58 is_active: bool = True
59 metadata: Dict[str, Any] = field(default_factory=dict)
61 # Variable pattern: {{variable_name}}
62 VAR_PATTERN = re.compile(r"\{\{(\w+)\}\}")
64 def extract_variables(self) -> Set[str]:
65 """Extract variable names from template content."""
66 return set(self.VAR_PATTERN.findall(self.content))
68 def render(self, values: Dict[str, str], strict: bool = True) -> str:
69 """
70 Render the template by substituting variables.
72 Args:
73 values: Dict of variable_name → value
74 strict: If True, raise on missing variables. If False, leave placeholders.
76 Raises:
77 ValueError: If strict=True and a variable is missing.
78 """
79 required = self.extract_variables()
81 if strict:
82 missing = required - set(values.keys())
83 if missing:
84 raise ValueError(
85 f"Template '{self.name}' v{self.version} missing variables: {missing}"
86 )
88 result = self.content
89 for var_name in required:
90 if var_name in values:
91 result = result.replace(f"{{{{{var_name}}}}}", values[var_name])
92 elif not strict:
93 # Leave placeholder intact
94 pass
96 return result
98 def validate(self) -> Tuple[bool, List[str]]:
99 """
100 Validate template structure.
102 Returns (is_valid, list_of_issues).
103 """
104 issues = []
105 if not self.name:
106 issues.append("Name is required")
107 if not self.content:
108 issues.append("Content is empty")
109 if self.VAR_PATTERN.findall(self.content):
110 # Variables are OK, just note them
111 pass
112 return len(issues) == 0, issues
114 def diff(self, other: "PromptTemplate") -> str:
115 """Generate a unified diff between this and another template."""
116 a_lines = self.content.splitlines(keepends=True)
117 b_lines = other.content.splitlines(keepends=True)
118 diff = difflib.unified_diff(
119 a_lines, b_lines,
120 fromfile=f"{self.name} v{self.version}",
121 tofile=f"{other.name} v{other.version}",
122 )
123 return "".join(diff)
125 def to_dict(self) -> Dict[str, Any]:
126 return {
127 "name": self.name,
128 "version": self.version,
129 "content": self.content,
130 "role": self.role.value,
131 "variables": sorted(self.variables),
132 "description": self.description,
133 "tags": self.tags,
134 "author": self.author,
135 "created_at": self.created_at,
136 "updated_at": self.updated_at,
137 "is_active": self.is_active,
138 "metadata": self.metadata,
139 }
141 @classmethod
142 def from_dict(cls, data: Dict[str, Any]) -> "PromptTemplate":
143 return cls(
144 name=data["name"],
145 version=data.get("version", 1),
146 content=data["content"],
147 role=PromptRole(data.get("role", "system")),
148 variables=set(data.get("variables", [])),
149 description=data.get("description", ""),
150 tags=data.get("tags", []),
151 author=data.get("author", ""),
152 created_at=data.get("created_at", ""),
153 updated_at=data.get("updated_at", ""),
154 is_active=data.get("is_active", True),
155 metadata=data.get("metadata", {}),
156 )
159# ---------------------------------------------------------------------------
160# Prompt Store
161# ---------------------------------------------------------------------------
164class PromptStore:
165 """
166 Registry of prompt templates with versioning.
168 Supports:
169 - Semantic versioning per template
170 - Latest/active version resolution
171 - Template lineage tracking
172 - Import/export (JSON)
173 """
175 def __init__(self):
176 self._templates: Dict[str, Dict[int, PromptTemplate]] = defaultdict(dict)
177 self._latest: Dict[str, int] = {}
178 self._active: Dict[str, int] = {}
179 self._lineage: Dict[str, List[int]] = defaultdict(list)
181 def add(self, template: PromptTemplate) -> PromptTemplate:
182 """
183 Add a new template or create a new version.
185 Auto-increments version if template name already exists.
186 """
187 name = template.name
189 if name in self._latest:
190 # New version
191 latest = self._latest[name]
192 template.version = latest + 1
193 else:
194 template.version = template.version or 1
196 if not template.variables:
197 template.variables = template.extract_variables()
199 template.updated_at = datetime.now(timezone.utc).isoformat()
200 self._templates[name][template.version] = template
201 self._latest[name] = template.version
202 self._active[name] = template.version
203 self._lineage[name].append(template.version)
205 return template
207 def get(self, name: str, version: Optional[int] = None) -> Optional[PromptTemplate]:
208 """Get a template by name and optional version. Defaults to active version."""
209 if name not in self._templates:
210 return None
212 if version is not None:
213 return self._templates[name].get(version)
215 active_ver = self._active.get(name) or self._latest[name]
216 return self._templates[name].get(active_ver)
218 def get_latest(self, name: str) -> Optional[PromptTemplate]:
219 """Get the latest version of a template."""
220 if name not in self._latest:
221 return None
222 return self._templates[name].get(self._latest[name])
224 def set_active(self, name: str, version: int) -> bool:
225 """Set which version is the active one."""
226 if name not in self._templates or version not in self._templates[name]:
227 return False
228 self._active[name] = version
229 return True
231 def deactivate(self, name: str) -> None:
232 """Deactivate a template (no active version)."""
233 self._active.pop(name, None)
235 def list_templates(self) -> List[Dict[str, Any]]:
236 """List all templates with their versions."""
237 result = []
238 for name, versions in self._templates.items():
239 latest = self._latest[name]
240 active = self._active.get(name, latest)
241 result.append({
242 "name": name,
243 "versions": sorted(versions.keys()),
244 "latest": latest,
245 "active": active,
246 "total_versions": len(versions),
247 })
248 return sorted(result, key=lambda x: x["name"])
250 def get_history(self, name: str) -> List[PromptTemplate]:
251 """Get all versions of a template in chronological order."""
252 if name not in self._templates:
253 return []
254 return [self._templates[name][v] for v in sorted(self._templates[name].keys())]
256 def diff_versions(self, name: str, v1: int, v2: int) -> Optional[str]:
257 """Get diff between two versions of a template."""
258 t1 = self.get(name, v1)
259 t2 = self.get(name, v2)
260 if t1 is None or t2 is None:
261 return None
262 return t1.diff(t2)
264 def remove(self, name: str, version: Optional[int] = None) -> int:
265 """
266 Remove template(s). If version is None, remove all versions.
267 Returns number of versions removed.
268 """
269 if name not in self._templates:
270 return 0
272 if version is not None:
273 if version in self._templates[name]:
274 del self._templates[name][version]
275 if version == self._latest.get(name):
276 remaining = sorted(self._templates[name].keys())
277 self._latest[name] = remaining[-1] if remaining else 0
278 return 1
279 return 0
281 count = len(self._templates[name])
282 del self._templates[name]
283 self._latest.pop(name, None)
284 self._active.pop(name, None)
285 self._lineage.pop(name, None)
286 return count
288 def export_json(self, names: Optional[List[str]] = None) -> str:
289 """Export templates as JSON."""
290 templates = []
291 target = names or list(self._templates.keys())
292 for name in target:
293 for t in self._templates.get(name, {}).values():
294 templates.append(t.to_dict())
295 return json.dumps({"templates": templates}, indent=2, ensure_ascii=False)
297 def import_json(self, json_str: str) -> int:
298 """Import templates from JSON. Returns count of imported templates."""
299 data = json.loads(json_str)
300 count = 0
301 for td in data.get("templates", []):
302 template = PromptTemplate.from_dict(td)
303 self.add(template)
304 count += 1
305 return count
308# ---------------------------------------------------------------------------
309# A/B Test Manager
310# ---------------------------------------------------------------------------
313@dataclass
314class ABTest:
315 """An A/B test comparing two prompt template versions."""
316 name: str
317 template_name: str
318 variant_a_version: int
319 variant_b_version: int
320 split_ratio: float = 0.5 # 0.5 = 50% each
321 is_active: bool = True
322 created_at: str = field(default_factory=lambda: datetime.now(timezone.utc).isoformat())
324 def route(self, session_id: str) -> int:
325 """Route a session to variant A (0) or B (1)."""
326 if not self.is_active:
327 return 0 # Default to A if test is inactive
329 # Deterministic routing based on session hash
330 hash_val = abs(hash(f"{self.name}:{session_id}"))
331 bucket = (hash_val % 100) / 100.0
332 return 0 if bucket < self.split_ratio else 1
335class ABTestManager:
336 """Manage A/B tests for prompt templates."""
338 def __init__(self, store: PromptStore):
339 self._store = store
340 self._tests: Dict[str, ABTest] = {}
341 self._results: Dict[str, Dict[str, int]] = defaultdict(
342 lambda: {"a_served": 0, "b_served": 0}
343 )
345 def create_test(
346 self,
347 name: str,
348 template_name: str,
349 variant_a_version: int,
350 variant_b_version: int,
351 split_ratio: float = 0.5,
352 ) -> ABTest:
353 """Create a new A/B test."""
354 test = ABTest(
355 name=name,
356 template_name=template_name,
357 variant_a_version=variant_a_version,
358 variant_b_version=variant_b_version,
359 split_ratio=split_ratio,
360 )
361 self._tests[name] = test
362 return test
364 def get_template(self, test_name: str, session_id: str) -> Optional[PromptTemplate]:
365 """
366 Get the prompt template for a session in an A/B test.
368 Returns None if the test doesn't exist or template not found.
369 """
370 test = self._tests.get(test_name)
371 if test is None:
372 return None
374 variant = test.route(session_id)
375 version = test.variant_a_version if variant == 0 else test.variant_b_version
377 self._results[test_name][f"{'a' if variant == 0 else 'b'}_served"] += 1
379 return self._store.get(test.template_name, version)
381 def get_results(self, name: str) -> Dict[str, Any]:
382 """Get results for an A/B test."""
383 test = self._tests.get(name)
384 if test is None:
385 return {}
387 results = dict(self._results[name])
388 total = results.get("a_served", 0) + results.get("b_served", 0)
389 return {
390 "test_name": name,
391 "template_name": test.template_name,
392 "variant_a_version": test.variant_a_version,
393 "variant_b_version": test.variant_b_version,
394 "split_ratio": test.split_ratio,
395 "is_active": test.is_active,
396 **results,
397 "total_served": total,
398 "a_pct": round(results.get("a_served", 0) / max(1, total) * 100, 1),
399 "b_pct": round(results.get("b_served", 0) / max(1, total) * 100, 1),
400 }
402 def stop_test(self, name: str) -> None:
403 """Stop an active A/B test."""
404 if name in self._tests:
405 self._tests[name].is_active = False
407 def list_tests(self) -> List[Dict[str, Any]]:
408 """List all A/B tests."""
409 return [self.get_results(name) for name in self._tests]