autogen.py 3.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111
  1. """
  2. lambdagent_guard.autogen — AutoGen guard wrapper (I13).
  3. Usage:
  4. from lambdagent_guard import guard_autogen
  5. manager = GroupChatManager(groupchat=chat)
  6. guarded = guard_autogen(manager, cost_budget=5.0)
  7. proxy.initiate_chat(guarded, message="Write quicksort")
  8. Catches: AutoGen #108 (blank msg loop), #391 (TERMINATE bypass), #2702 (cost tracking)
  9. """
  10. from __future__ import annotations
  11. from typing import Any, Callable, Optional
  12. from .core import GuardConfig, RuntimeMonitor, run_compile_checks
  13. import warnings
  14. def guard_autogen(
  15. manager,
  16. cost_budget: float = float("inf"),
  17. empty_message_detection: bool = True,
  18. terminate_robustness: bool = True,
  19. loop_detection: bool = True,
  20. cost_alert: Optional[Callable[[float], None]] = None,
  21. ):
  22. """
  23. Wrap an AutoGen GroupChatManager with lambdagent safety guards.
  24. Args:
  25. manager: AutoGen GroupChatManager
  26. cost_budget: Max cost in USD
  27. empty_message_detection: Catch AutoGen #108 blank-message loops
  28. terminate_robustness: Don't rely on exact "TERMINATE" string matching
  29. """
  30. guard_cfg = GuardConfig(
  31. cost_budget=cost_budget,
  32. empty_message_detection=empty_message_detection,
  33. terminate_robustness=terminate_robustness,
  34. loop_detection=loop_detection,
  35. cost_alert=cost_alert,
  36. )
  37. # Phase 1: Compile-time checks
  38. try:
  39. from lambdagent.extractors import extract_config
  40. config = extract_config(manager, framework="autogen")
  41. # Check for fragile termination (AutoGen #108, #391)
  42. term_warning = config.get("_termination_warning")
  43. if term_warning:
  44. warnings.warn(f"[lambdagent-guard] {term_warning}")
  45. issues = run_compile_checks(config, guard_cfg)
  46. for issue in issues:
  47. warnings.warn(f"[lambdagent-guard] {issue}")
  48. except Exception as e:
  49. warnings.warn(f"[lambdagent-guard] Compile check skipped: {e}")
  50. # Phase 2: Runtime hooks
  51. monitor = RuntimeMonitor(guard_cfg)
  52. # Hook into GroupChatManager's run_chat or equivalent
  53. if hasattr(manager, 'run_chat'):
  54. original_run = manager.run_chat
  55. def guarded_run(messages, sender, config_override=None):
  56. for msg in messages if isinstance(messages, list) else [messages]:
  57. content = ""
  58. if isinstance(msg, dict):
  59. content = msg.get("content", "")
  60. elif isinstance(msg, str):
  61. content = msg
  62. monitor.on_step({
  63. "output": str(content)[:500],
  64. "tokens": 0,
  65. })
  66. return original_run(messages, sender, config_override)
  67. manager.run_chat = guarded_run
  68. # Hook terminate robustness: patch is_termination_msg if needed
  69. if terminate_robustness and hasattr(manager, 'groupchat'):
  70. chat = manager.groupchat
  71. original_term = getattr(chat, 'is_termination_msg', None)
  72. if original_term is not None:
  73. def robust_termination(msg):
  74. """Case-insensitive, substring-based termination check."""
  75. content = ""
  76. if isinstance(msg, dict):
  77. content = msg.get("content", "")
  78. elif isinstance(msg, str):
  79. content = msg
  80. if content and "terminate" in content.lower():
  81. return True
  82. # Fall back to original check
  83. try:
  84. return original_term(msg)
  85. except Exception:
  86. return False
  87. chat.is_termination_msg = robust_termination
  88. manager._lambdagent_monitor = monitor
  89. return manager