| import sys |
| import unittest |
| from pathlib import Path |
|
|
| from langchain_core.messages import AIMessage, HumanMessage |
|
|
| ROOT = Path(__file__).resolve().parents[2] |
| if str(ROOT) not in sys.path: |
| sys.path.insert(0, str(ROOT)) |
|
|
| from src.agents.agents import DietarySpecialist, OutputMerger |
|
|
|
|
| class StubLLM: |
| def __init__(self, content): |
| self.content = content |
| self.calls = [] |
|
|
| async def ainvoke(self, messages, config=None): |
| self.calls.append(messages) |
| return AIMessage(content=self.content) |
|
|
| async def astream(self, messages, config=None): |
| self.calls.append(messages) |
| yield AIMessage(content=self.content) |
|
|
|
|
| class TestChapter6(unittest.IsolatedAsyncioTestCase): |
| async def test_output_merger_uses_latest_specialist_output(self): |
| llm = StubLLM("merged response") |
| merger = OutputMerger() |
| merger.llm = llm |
|
|
| state = { |
| "messages": [ |
| HumanMessage(content="What should I do next?"), |
| AIMessage(content="Older specialist response"), |
| AIMessage(content="Older user-facing answer"), |
| ], |
| "clinician_outputs": [ |
| "Diagnosis says continue monitoring.", |
| "Treatment says adjust medication dose.", |
| ], |
| } |
|
|
| result = await merger.run(state) |
|
|
| self.assertEqual(result["messages"][-1].content, "merged response") |
| self.assertEqual(result["metrics"][0]["agent"], "OutputMerger") |
|
|
| merged_prompt = "\n".join( |
| message.content for message in llm.calls[-1] |
| ) |
| self.assertIn("Treatment says adjust medication dose.", merged_prompt) |
| self.assertNotIn("Older specialist response", merged_prompt) |
|
|
| async def test_dietary_specialist_returns_metrics(self): |
| llm = StubLLM("dietary guidance") |
| specialist = DietarySpecialist() |
| specialist.llm = llm |
|
|
| state = { |
| "messages": [HumanMessage(content="Suggest a low-carb meal plan for today.")] |
| } |
|
|
| result = await specialist.run(state) |
|
|
| self.assertEqual(result["messages"][-1].content, "dietary guidance") |
| self.assertEqual(result["metrics"][0]["agent"], "DietarySpecialist") |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|