Coverage for agentos/core/prompt_manager.py: 0%
193 statements
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 20:40 +0800
« prev ^ index » next coverage.py v7.14.3, created at 2026-07-08 20:40 +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 UTC, datetime
28from enum import StrEnum
29from typing import Any
31# ---------------------------------------------------------------------------
32# Prompt Template
33# ---------------------------------------------------------------------------
36class PromptRole(StrEnum):
37 SYSTEM = "system"
38 USER = "user"
39 ASSISTANT = "assistant"
42@dataclass
43class PromptTemplate:
44 """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(UTC).isoformat())
55 updated_at: str = field(default_factory=lambda: datetime.now(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,
118 b_lines,
119 fromfile=f"{self.name} v{self.version}",
120 tofile=f"{other.name} v{other.version}",
121 )
122 return "".join(diff)
124 def to_dict(self) -> dict[str, Any]:
125 return {
126 "name": self.name,
127 "version": self.version,
128 "content": self.content,
129 "role": self.role.value,
130 "variables": sorted(self.variables),
131 "description": self.description,
132 "tags": self.tags,
133 "author": self.author,
134 "created_at": self.created_at,
135 "updated_at": self.updated_at,
136 "is_active": self.is_active,
137 "metadata": self.metadata,
138 }
140 @classmethod
141 def from_dict(cls, data: dict[str, Any]) -> PromptTemplate:
142 return cls(
143 name=data["name"],
144 version=data.get("version", 1),
145 content=data["content"],
146 role=PromptRole(data.get("role", "system")),
147 variables=set(data.get("variables", [])),
148 description=data.get("description", ""),
149 tags=data.get("tags", []),
150 author=data.get("author", ""),
151 created_at=data.get("created_at", ""),
152 updated_at=data.get("updated_at", ""),
153 is_active=data.get("is_active", True),
154 metadata=data.get("metadata", {}),
155 )
158# ---------------------------------------------------------------------------
159# Prompt Store
160# ---------------------------------------------------------------------------
163class PromptStore:
164 """
165 Registry of prompt templates with versioning.
167 Supports:
168 - Semantic versioning per template
169 - Latest/active version resolution
170 - Template lineage tracking
171 - Import/export (JSON)
172 """
174 def __init__(self):
175 self._templates: dict[str, dict[int, PromptTemplate]] = defaultdict(dict)
176 self._latest: dict[str, int] = {}
177 self._active: dict[str, int] = {}
178 self._lineage: dict[str, list[int]] = defaultdict(list)
180 def add(self, template: PromptTemplate) -> PromptTemplate:
181 """
182 Add a new template or create a new version.
184 Auto-increments version if template name already exists.
185 """
186 name = template.name
188 if name in self._latest:
189 # New version
190 latest = self._latest[name]
191 template.version = latest + 1
192 else:
193 template.version = template.version or 1
195 if not template.variables:
196 template.variables = template.extract_variables()
198 template.updated_at = datetime.now(UTC).isoformat()
199 self._templates[name][template.version] = template
200 self._latest[name] = template.version
201 self._active[name] = template.version
202 self._lineage[name].append(template.version)
204 return template
206 def get(self, name: str, version: int | None = None) -> PromptTemplate | None:
207 """Get a template by name and optional version. Defaults to active version."""
208 if name not in self._templates:
209 return None
211 if version is not None:
212 return self._templates[name].get(version)
214 active_ver = self._active.get(name) or self._latest[name]
215 return self._templates[name].get(active_ver)
217 def get_latest(self, name: str) -> PromptTemplate | None:
218 """Get the latest version of a template."""
219 if name not in self._latest:
220 return None
221 return self._templates[name].get(self._latest[name])
223 def set_active(self, name: str, version: int) -> bool:
224 """Set which version is the active one."""
225 if name not in self._templates or version not in self._templates[name]:
226 return False
227 self._active[name] = version
228 return True
230 def deactivate(self, name: str) -> None:
231 """Deactivate a template (no active version)."""
232 self._active.pop(name, None)
234 def list_templates(self) -> list[dict[str, Any]]:
235 """List all templates with their versions."""
236 result = []
237 for name, versions in self._templates.items():
238 latest = self._latest[name]
239 active = self._active.get(name, latest)
240 result.append(
241 {
242 "name": name,
243 "versions": sorted(versions.keys()),
244 "latest": latest,
245 "active": active,
246 "total_versions": len(versions),
247 }
248 )
249 return sorted(result, key=lambda x: x["name"])
251 def get_history(self, name: str) -> list[PromptTemplate]:
252 """Get all versions of a template in chronological order."""
253 if name not in self._templates:
254 return []
255 return [self._templates[name][v] for v in sorted(self._templates[name].keys())]
257 def diff_versions(self, name: str, v1: int, v2: int) -> str | None:
258 """Get diff between two versions of a template."""
259 t1 = self.get(name, v1)
260 t2 = self.get(name, v2)
261 if t1 is None or t2 is None:
262 return None
263 return t1.diff(t2)
265 def remove(self, name: str, version: int | None = None) -> int:
266 """
267 Remove template(s). If version is None, remove all versions.
268 Returns number of versions removed.
269 """
270 if name not in self._templates:
271 return 0
273 if version is not None:
274 if version in self._templates[name]:
275 del self._templates[name][version]
276 if version == self._latest.get(name):
277 remaining = sorted(self._templates[name].keys())
278 self._latest[name] = remaining[-1] if remaining else 0
279 return 1
280 return 0
282 count = len(self._templates[name])
283 del self._templates[name]
284 self._latest.pop(name, None)
285 self._active.pop(name, None)
286 self._lineage.pop(name, None)
287 return count
289 def export_json(self, names: list[str] | None = None) -> str:
290 """Export templates as JSON."""
291 templates = []
292 target = names or list(self._templates.keys())
293 for name in target:
294 for t in self._templates.get(name, {}).values():
295 templates.append(t.to_dict())
296 return json.dumps({"templates": templates}, indent=2, ensure_ascii=False)
298 def import_json(self, json_str: str) -> int:
299 """Import templates from JSON. Returns count of imported templates."""
300 data = json.loads(json_str)
301 count = 0
302 for td in data.get("templates", []):
303 template = PromptTemplate.from_dict(td)
304 self.add(template)
305 count += 1
306 return count
309# ---------------------------------------------------------------------------
310# A/B Test Manager
311# ---------------------------------------------------------------------------
314@dataclass
315class ABTest:
316 """An A/B test comparing two prompt template versions."""
318 name: str
319 template_name: str
320 variant_a_version: int
321 variant_b_version: int
322 split_ratio: float = 0.5 # 0.5 = 50% each
323 is_active: bool = True
324 created_at: str = field(default_factory=lambda: datetime.now(UTC).isoformat())
326 def route(self, session_id: str) -> int:
327 """Route a session to variant A (0) or B (1)."""
328 if not self.is_active:
329 return 0 # Default to A if test is inactive
331 # Deterministic routing based on session hash
332 hash_val = abs(hash(f"{self.name}:{session_id}"))
333 bucket = (hash_val % 100) / 100.0
334 return 0 if bucket < self.split_ratio else 1
337class ABTestManager:
338 """Manage A/B tests for prompt templates."""
340 def __init__(self, store: PromptStore):
341 self._store = store
342 self._tests: dict[str, ABTest] = {}
343 self._results: dict[str, dict[str, int]] = defaultdict(
344 lambda: {"a_served": 0, "b_served": 0}
345 )
347 def create_test(
348 self,
349 name: str,
350 template_name: str,
351 variant_a_version: int,
352 variant_b_version: int,
353 split_ratio: float = 0.5,
354 ) -> ABTest:
355 """Create a new A/B test."""
356 test = ABTest(
357 name=name,
358 template_name=template_name,
359 variant_a_version=variant_a_version,
360 variant_b_version=variant_b_version,
361 split_ratio=split_ratio,
362 )
363 self._tests[name] = test
364 return test
366 def get_template(self, test_name: str, session_id: str) -> PromptTemplate | None:
367 """
368 Get the prompt template for a session in an A/B test.
370 Returns None if the test doesn't exist or template not found.
371 """
372 test = self._tests.get(test_name)
373 if test is None:
374 return None
376 variant = test.route(session_id)
377 version = test.variant_a_version if variant == 0 else test.variant_b_version
379 self._results[test_name][f"{'a' if variant == 0 else 'b'}_served"] += 1
381 return self._store.get(test.template_name, version)
383 def get_results(self, name: str) -> dict[str, Any]:
384 """Get results for an A/B test."""
385 test = self._tests.get(name)
386 if test is None:
387 return {}
389 results = dict(self._results[name])
390 total = results.get("a_served", 0) + results.get("b_served", 0)
391 return {
392 "test_name": name,
393 "template_name": test.template_name,
394 "variant_a_version": test.variant_a_version,
395 "variant_b_version": test.variant_b_version,
396 "split_ratio": test.split_ratio,
397 "is_active": test.is_active,
398 **results,
399 "total_served": total,
400 "a_pct": round(results.get("a_served", 0) / max(1, total) * 100, 1),
401 "b_pct": round(results.get("b_served", 0) / max(1, total) * 100, 1),
402 }
404 def stop_test(self, name: str) -> None:
405 """Stop an active A/B test."""
406 if name in self._tests:
407 self._tests[name].is_active = False
409 def list_tests(self) -> list[dict[str, Any]]:
410 """List all A/B tests."""
411 return [self.get_results(name) for name in self._tests]