test_agentruntime.py 3.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134
  1. """Tests for lambdagent.agentruntime components."""
  2. import pytest
  3. from lambdagent.agentruntime.action_parser import ActionParser, Action, ParseError
  4. from lambdagent.agentruntime.termination import TerminationOracle
  5. from lambdagent.agentruntime.memory_backend import LocalMemory, SQLiteMemory, MemoryBackend
  6. from lambdagent.agentruntime.trace_store import TraceStore, TraceRecord
  7. from lambdagent.agentruntime.config import RuntimeConfig, LLMConfig, MemoryConfig
  8. # ActionParser tests
  9. def test_parser_json_block():
  10. p = ActionParser(["search", "terminate"])
  11. action = p.parse('Let me search. ```json\n{"action": "search", "input": {"q": "test"}}\n```')
  12. assert action.tool == "search"
  13. def test_parser_inline_json():
  14. p = ActionParser(["calc", "terminate"])
  15. action = p.parse('I will calculate. {"action": "calc", "input": {"expr": "1+1"}}')
  16. assert action.tool == "calc"
  17. def test_parser_implicit_terminate():
  18. p = ActionParser(["search", "terminate"])
  19. action = p.parse("Final Answer: The result is 42.")
  20. assert action.tool == "terminate"
  21. assert "42" in str(action.input)
  22. def test_parser_unknown_tool():
  23. p = ActionParser(["search"])
  24. with pytest.raises(ParseError):
  25. p.parse('{"action": "unknown_tool"}')
  26. # TerminationOracle tests
  27. def test_termination_implicit():
  28. oracle = TerminationOracle()
  29. assert oracle.should_stop("Task complete. Here is the result.", None, 0) == True
  30. assert oracle.should_stop("I need to search more.", None, 0) == False
  31. def test_termination_disabled():
  32. oracle = TerminationOracle(implicit_detection=False)
  33. assert oracle.should_stop("Final Answer: 42", None, 0) == False
  34. # LocalMemory tests
  35. def test_local_memory_write_read():
  36. mem = LocalMemory(size=5, ttl=3600)
  37. mem.write("k1", "v1")
  38. assert mem.read("k1") == "v1"
  39. def test_local_memory_lru():
  40. mem = LocalMemory(size=3, ttl=3600)
  41. for i in range(5):
  42. mem.write(f"k{i}", f"v{i}")
  43. assert mem.read("k0") is None # evicted
  44. assert mem.read("k4") == "v4"
  45. def test_local_memory_recent():
  46. mem = LocalMemory(size=10, ttl=3600)
  47. mem.write("a", "1")
  48. mem.write("b", "2")
  49. recent = mem.read_recent(5)
  50. assert len(recent) == 2
  51. def test_local_memory_clear():
  52. mem = LocalMemory()
  53. mem.write("k", "v")
  54. mem.clear()
  55. assert mem.read("k") is None
  56. # SQLiteMemory tests
  57. def test_sqlite_memory():
  58. mem = SQLiteMemory(size=5, ttl=3600, db_path=":memory:")
  59. mem.write("k1", "v1")
  60. assert mem.read("k1") == "v1"
  61. mem.clear()
  62. assert mem.read("k1") is None
  63. # MemoryBackend factory
  64. def test_memory_factory():
  65. mem = MemoryBackend.create({"strategy": "local", "size": 10, "ttl": 3600})
  66. assert isinstance(mem, LocalMemory)
  67. mem2 = MemoryBackend.create({"strategy": "sqlite", "size": 10, "ttl": 3600})
  68. assert isinstance(mem2, SQLiteMemory)
  69. # TraceStore tests
  70. def test_trace_store():
  71. ts = TraceStore()
  72. ts.append(TraceRecord(step=0, term_name="think", term_type="Lam", duration_ms=100))
  73. ts.append(TraceRecord(step=1, term_name="tool", term_type="Tool", duration_ms=50))
  74. assert len(ts.get_all()) == 2
  75. assert ts.get_step(0).term_name == "think"
  76. stats = ts.stats()
  77. assert stats.total_steps == 2
  78. assert stats.total_time_ms == 150
  79. def test_trace_timeline():
  80. ts = TraceStore()
  81. ts.append(TraceRecord(step=0, term_name="think", term_type="Lam", duration_ms=100))
  82. timeline = ts.to_timeline()
  83. assert "think" in timeline
  84. def test_trace_json():
  85. ts = TraceStore()
  86. ts.append(TraceRecord(step=0, term_name="t", term_type="Lam"))
  87. j = ts.to_json()
  88. assert '"term_name": "t"' in j
  89. # RuntimeConfig tests
  90. def test_config_defaults():
  91. cfg = RuntimeConfig()
  92. assert cfg.llm.model == "claude-sonnet-4-20250514"
  93. assert cfg.memory.strategy == "local"
  94. assert cfg.react.max_steps == 10