Coverage for agentos/core/prompt_manager.py: 0%

193 statements  

« prev     ^ index     » next       coverage.py v7.14.3, created at 2026-07-06 12:29 +0800

1""" 

2AgentOS Prompt Manager — Versioned Prompt Templates with A/B Testing 

3━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━ 

4 

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 

12 

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""" 

19 

20from __future__ import annotations 

21 

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 

30 

31# --------------------------------------------------------------------------- 

32# Prompt Template 

33# --------------------------------------------------------------------------- 

34 

35 

36class PromptRole(StrEnum): 

37 SYSTEM = "system" 

38 USER = "user" 

39 ASSISTANT = "assistant" 

40 

41 

42@dataclass 

43class PromptTemplate: 

44 """A versioned prompt template.""" 

45 

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) 

58 

59 # Variable pattern: {{variable_name}} 

60 VAR_PATTERN = re.compile(r"\{\{(\w+)\}\}") 

61 

62 def extract_variables(self) -> set[str]: 

63 """Extract variable names from template content.""" 

64 return set(self.VAR_PATTERN.findall(self.content)) 

65 

66 def render(self, values: dict[str, str], strict: bool = True) -> str: 

67 """ 

68 Render the template by substituting variables. 

69 

70 Args: 

71 values: Dict of variable_name → value 

72 strict: If True, raise on missing variables. If False, leave placeholders. 

73 

74 Raises: 

75 ValueError: If strict=True and a variable is missing. 

76 """ 

77 required = self.extract_variables() 

78 

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 ) 

85 

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 

93 

94 return result 

95 

96 def validate(self) -> tuple[bool, list[str]]: 

97 """ 

98 Validate template structure. 

99 

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 

111 

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) 

123 

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 } 

139 

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 ) 

156 

157 

158# --------------------------------------------------------------------------- 

159# Prompt Store 

160# --------------------------------------------------------------------------- 

161 

162 

163class PromptStore: 

164 """ 

165 Registry of prompt templates with versioning. 

166 

167 Supports: 

168 - Semantic versioning per template 

169 - Latest/active version resolution 

170 - Template lineage tracking 

171 - Import/export (JSON) 

172 """ 

173 

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) 

179 

180 def add(self, template: PromptTemplate) -> PromptTemplate: 

181 """ 

182 Add a new template or create a new version. 

183 

184 Auto-increments version if template name already exists. 

185 """ 

186 name = template.name 

187 

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 

194 

195 if not template.variables: 

196 template.variables = template.extract_variables() 

197 

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) 

203 

204 return template 

205 

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 

210 

211 if version is not None: 

212 return self._templates[name].get(version) 

213 

214 active_ver = self._active.get(name) or self._latest[name] 

215 return self._templates[name].get(active_ver) 

216 

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]) 

222 

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 

229 

230 def deactivate(self, name: str) -> None: 

231 """Deactivate a template (no active version).""" 

232 self._active.pop(name, None) 

233 

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"]) 

250 

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())] 

256 

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) 

264 

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 

272 

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 

281 

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 

288 

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) 

297 

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 

307 

308 

309# --------------------------------------------------------------------------- 

310# A/B Test Manager 

311# --------------------------------------------------------------------------- 

312 

313 

314@dataclass 

315class ABTest: 

316 """An A/B test comparing two prompt template versions.""" 

317 

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()) 

325 

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 

330 

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 

335 

336 

337class ABTestManager: 

338 """Manage A/B tests for prompt templates.""" 

339 

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 ) 

346 

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 

365 

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. 

369 

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 

375 

376 variant = test.route(session_id) 

377 version = test.variant_a_version if variant == 0 else test.variant_b_version 

378 

379 self._results[test_name][f"{'a' if variant == 0 else 'b'}_served"] += 1 

380 

381 return self._store.get(test.template_name, version) 

382 

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 {} 

388 

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 } 

403 

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 

408 

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]