You can not select more than 25 topics Topics must start with a chinese character,a letter or number, can include dashes ('-') and can be up to 35 characters long.

test_state.py 1.9 kB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465
  1. from typing import Any, Mapping, Sequence
  2. import pytest
  3. from agnext.application import SingleThreadedAgentRuntime
  4. from agnext.core import BaseAgent, MessageContext
  5. class StatefulAgent(BaseAgent):
  6. def __init__(self) -> None:
  7. super().__init__("A stateful agent", [])
  8. self.state = 0
  9. @property
  10. def subscriptions(self) -> Sequence[type]:
  11. return []
  12. async def on_message(self, message: Any, ctx: MessageContext) -> None:
  13. raise NotImplementedError
  14. def save_state(self) -> Mapping[str, Any]:
  15. return {"state": self.state}
  16. def load_state(self, state: Mapping[str, Any]) -> None:
  17. self.state = state["state"]
  18. @pytest.mark.asyncio
  19. async def test_agent_can_save_state() -> None:
  20. runtime = SingleThreadedAgentRuntime()
  21. agent1_id = await runtime.register_and_get("name1", StatefulAgent)
  22. agent1: StatefulAgent = await runtime.try_get_underlying_agent_instance(agent1_id, type=StatefulAgent)
  23. assert agent1.state == 0
  24. agent1.state = 1
  25. assert agent1.state == 1
  26. agent1_state = agent1.save_state()
  27. agent1.state = 2
  28. assert agent1.state == 2
  29. agent1.load_state(agent1_state)
  30. assert agent1.state == 1
  31. @pytest.mark.asyncio
  32. async def test_runtime_can_save_state() -> None:
  33. runtime = SingleThreadedAgentRuntime()
  34. agent1_id = await runtime.register_and_get("name1", StatefulAgent)
  35. agent1: StatefulAgent = await runtime.try_get_underlying_agent_instance(agent1_id, type=StatefulAgent)
  36. assert agent1.state == 0
  37. agent1.state = 1
  38. assert agent1.state == 1
  39. runtime_state = await runtime.save_state()
  40. runtime2 = SingleThreadedAgentRuntime()
  41. agent2_id = await runtime2.register_and_get("name1", StatefulAgent)
  42. agent2: StatefulAgent = await runtime2.try_get_underlying_agent_instance(agent2_id, type=StatefulAgent)
  43. await runtime2.load_state(runtime_state)
  44. assert agent2.state == 1