test_llm_usage_hook.py 2.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374
  1. """LLM usage recording hook stays free of supply_infra at import time."""
  2. from __future__ import annotations
  3. import importlib
  4. from pathlib import Path
  5. from supply_agent.logging.logger import AgentLogger
  6. from supply_agent.logging.usage import (
  7. get_llm_usage_recorder,
  8. record_llm_usage,
  9. set_llm_usage_recorder,
  10. )
  11. from supply_agent.types import Message, Role
  12. def test_usage_module_does_not_import_supply_infra() -> None:
  13. source = importlib.util.find_spec("supply_agent.logging.usage")
  14. assert source is not None and source.origin is not None
  15. text = Path(source.origin).read_text(encoding="utf-8")
  16. assert "import supply_infra" not in text
  17. assert "from supply_infra" not in text
  18. def test_record_llm_usage_is_noop_without_recorder() -> None:
  19. previous = get_llm_usage_recorder()
  20. try:
  21. set_llm_usage_recorder(None)
  22. record_llm_usage({"usage": {"cost": 0.01}})
  23. finally:
  24. set_llm_usage_recorder(previous)
  25. def test_log_llm_output_invokes_usage_recorder(tmp_path: Path) -> None:
  26. previous = get_llm_usage_recorder()
  27. seen: list[dict] = []
  28. def _recorder(payload: dict) -> None:
  29. seen.append(payload)
  30. try:
  31. set_llm_usage_recorder(_recorder)
  32. logger = AgentLogger(tmp_path, enabled=True)
  33. logger.start_run("hello", model="google/gemini-2.5-flash", agent_name="demo_agent")
  34. class _Usage:
  35. def model_dump(self) -> dict:
  36. return {
  37. "prompt_tokens": 10,
  38. "completion_tokens": 5,
  39. "total_tokens": 15,
  40. "cost": 0.0012,
  41. }
  42. class _Response:
  43. model = "google/gemini-2.5-flash"
  44. usage = _Usage()
  45. def model_dump(self) -> dict:
  46. return {"model": self.model, "usage": self.usage.model_dump()}
  47. logger.log_llm_output(
  48. 1,
  49. Message(role=Role.ASSISTANT, content="ok"),
  50. raw_response=_Response(),
  51. model="google/gemini-2.5-flash",
  52. )
  53. assert len(seen) == 1
  54. assert seen[0]["agent_name"] == "demo_agent"
  55. assert seen[0]["model"] == "google/gemini-2.5-flash"
  56. assert seen[0]["usage"]["cost"] == 0.0012
  57. assert seen[0]["usage"]["total_tokens"] == 15
  58. finally:
  59. set_llm_usage_recorder(previous)