diff --git a/ravendb/documents/operations/ai/agents/__init__.py b/ravendb/documents/operations/ai/agents/__init__.py index a99c3209..8357366d 100644 --- a/ravendb/documents/operations/ai/agents/__init__.py +++ b/ravendb/documents/operations/ai/agents/__init__.py @@ -3,6 +3,7 @@ AiAgentParameter, AiAgentToolAction, AiAgentToolQuery, + AiAgentToolQueryOptions, AiAgentPersistenceConfiguration, AiAgentChatTrimmingConfiguration, AiAgentSummarizationByTokens, @@ -38,6 +39,7 @@ "AiAgentParameter", "AiAgentToolAction", "AiAgentToolQuery", + "AiAgentToolQueryOptions", "AiAgentPersistenceConfiguration", "AiAgentChatTrimmingConfiguration", "AiAgentSummarizationByTokens", diff --git a/ravendb/tests/ai_agent_tests/test_ai_agents.py b/ravendb/tests/ai_agent_tests/test_ai_agents.py new file mode 100644 index 00000000..ee18c1cd --- /dev/null +++ b/ravendb/tests/ai_agent_tests/test_ai_agents.py @@ -0,0 +1,343 @@ +"""Integration tests for AI agent CRUD operations.""" + +import os +import unittest + +from ravendb import ( + AiAgentConfiguration, + AiAgentParameter, + AiAgentSummarizationByTokens, + AiAgentChatTrimmingConfiguration, + AiAgentToolQuery, + AiAgentToolAction, +) +from ravendb.documents.operations.ai import AiConnectionString, AiModelType +from ravendb.documents.operations.ai.agents import ( + AddOrUpdateAiAgentOperation, + GetAiAgentOperation, + DeleteAiAgentOperation, + AiAgentToolQueryOptions, +) +from ravendb.documents.operations.ai.open_ai_settings import OpenAiSettings +from ravendb.documents.operations.connection_string.put_connection_string_operation import PutConnectionStringOperation +from ravendb.tests.test_base import TestBase + + +@unittest.skipIf(os.environ.get("RAVENDB_LICENSE") is None, "Insufficient license permissions. Skipping on CI/CD.") +class TestAiAgentCrudOperations(TestBase): + """Integration tests for AI agent CRUD operations (require server connection).""" + + CONNECTION_STRING_NAME = "test-ai-agent-cs" + + def setUp(self): + super().setUp() + ai_connection_string = AiConnectionString( + name=self.CONNECTION_STRING_NAME, + identifier="test-ai-agent-cs", + model_type=AiModelType.CHAT, + openai_settings=OpenAiSettings( + api_key="dummy-api-key", + endpoint="https://api.openai.com/v1", + model="gpt-4", + ), + ) + self.store.maintenance.send(PutConnectionStringOperation(ai_connection_string)) + self._created_agent_ids = [] + + def tearDown(self): + for agent_id in self._created_agent_ids: + try: + self.store.maintenance.send(DeleteAiAgentOperation(agent_id)) + except Exception: + pass + super().tearDown() + + def _create_basic_agent(self, name: str, identifier: str) -> AiAgentConfiguration: + """Helper to create a minimal valid AiAgentConfiguration.""" + agent = AiAgentConfiguration( + name=name, + identifier=identifier, + connection_string_name=self.CONNECTION_STRING_NAME, + system_prompt="You are a helpful assistant.", + sample_object='{"answer": "embed your answer here"}', + ) + return agent + + def _create_full_agent(self, name: str, identifier: str) -> AiAgentConfiguration: + """Helper to create a fully configured AiAgentConfiguration.""" + agent = AiAgentConfiguration( + name=name, + identifier=identifier, + connection_string_name=self.CONNECTION_STRING_NAME, + system_prompt="You are a helpful assistant that queries the database.", + sample_object='{"answer": "embed your answer here"}', + parameters=[ + AiAgentParameter("country", "The country to filter by."), + ], + queries=[ + AiAgentToolQuery( + name="get-orders", + description="Retrieve all orders.", + query="from Orders", + parameters_sample_object="{}", + ), + ], + actions=[ + AiAgentToolAction( + name="store-result", + description="Store the result in the database.", + parameters_sample_object='{"result": "embed result here"}', + ), + ], + chat_trimming=AiAgentChatTrimmingConfiguration( + tokens_config=AiAgentSummarizationByTokens( + max_tokens_before_summarization=32768, + max_tokens_after_summarization=1024, + ) + ), + max_model_iterations_per_call=3, + ) + return agent + + # ---- Create ---- + + def test_create_agent_returns_identifier(self): + agent = self._create_basic_agent("TestCreate", "test-create") + result = self.store.maintenance.send(AddOrUpdateAiAgentOperation(agent)) + self._created_agent_ids.append(result.identifier) + + self.assertIsNotNone(result) + self.assertEqual("test-create", result.identifier) + + def test_create_agent_via_store_ai(self): + agent = self._create_basic_agent("TestCreateViaAi", "test-create-via-ai") + result = self.store.ai.add_or_update_agent(agent) + self._created_agent_ids.append(result.identifier) + + self.assertIsNotNone(result) + self.assertEqual("test-create-via-ai", result.identifier) + + def test_create_full_agent(self): + agent = self._create_full_agent("TestCreateFull", "test-create-full") + result = self.store.ai.add_or_update_agent(agent) + self._created_agent_ids.append(result.identifier) + + self.assertIsNotNone(result) + self.assertEqual("test-create-full", result.identifier) + + # ---- Get ---- + + def test_get_all_agents_returns_list(self): + agent1 = self._create_basic_agent("TestGetAll1", "test-get-all-1") + agent2 = self._create_basic_agent("TestGetAll2", "test-get-all-2") + r1 = self.store.ai.add_or_update_agent(agent1) + r2 = self.store.ai.add_or_update_agent(agent2) + self._created_agent_ids.extend([r1.identifier, r2.identifier]) + + response = self.store.ai.get_agents() + + self.assertIsNotNone(response) + ids = [a.identifier for a in response.ai_agents] + self.assertIn("test-get-all-1", ids) + self.assertIn("test-get-all-2", ids) + + def test_get_agent_by_id(self): + agent = self._create_basic_agent("TestGetById", "test-get-by-id") + result = self.store.ai.add_or_update_agent(agent) + self._created_agent_ids.append(result.identifier) + + response = self.store.ai.get_agents("test-get-by-id") + + self.assertIsNotNone(response) + self.assertEqual(1, len(response.ai_agents)) + self.assertEqual("test-get-by-id", response.ai_agents[0].identifier) + + def test_get_agent_by_id_returns_correct_config(self): + agent = self._create_full_agent("TestGetConfig", "test-get-config") + result = self.store.ai.add_or_update_agent(agent) + self._created_agent_ids.append(result.identifier) + + response = self.store.ai.get_agents("test-get-config") + fetched = response.ai_agents[0] + + self.assertEqual("TestGetConfig", fetched.name) + self.assertEqual(self.CONNECTION_STRING_NAME, fetched.connection_string_name) + self.assertEqual(1, len(fetched.queries)) + self.assertEqual("get-orders", fetched.queries[0].name) + self.assertEqual(1, len(fetched.actions)) + self.assertEqual("store-result", fetched.actions[0].name) + self.assertEqual(1, len(fetched.parameters)) + self.assertEqual("country", fetched.parameters[0].name) + self.assertEqual(3, fetched.max_model_iterations_per_call) + + def test_get_agent_via_operation(self): + agent = self._create_basic_agent("TestGetOp", "test-get-op") + result = self.store.maintenance.send(AddOrUpdateAiAgentOperation(agent)) + self._created_agent_ids.append(result.identifier) + + response = self.store.maintenance.send(GetAiAgentOperation("test-get-op")) + + self.assertIsNotNone(response) + self.assertEqual(1, len(response.ai_agents)) + self.assertEqual("test-get-op", response.ai_agents[0].identifier) + + # ---- Update ---- + + def test_update_agent_system_prompt(self): + agent = self._create_basic_agent("TestUpdate", "test-update") + result = self.store.ai.add_or_update_agent(agent) + self._created_agent_ids.append(result.identifier) + + agent.system_prompt = "Updated system prompt." + self.store.ai.add_or_update_agent(agent) + + response = self.store.ai.get_agents("test-update") + self.assertEqual("Updated system prompt.", response.ai_agents[0].system_prompt) + + def test_update_agent_adds_query_tool(self): + agent = self._create_basic_agent("TestUpdateQuery", "test-update-query") + result = self.store.ai.add_or_update_agent(agent) + self._created_agent_ids.append(result.identifier) + + agent.queries.append( + AiAgentToolQuery( + name="new-query", + description="A newly added query tool.", + query="from Employees", + parameters_sample_object="{}", + ) + ) + self.store.ai.add_or_update_agent(agent) + + response = self.store.ai.get_agents("test-update-query") + query_names = [q.name for q in response.ai_agents[0].queries] + self.assertIn("new-query", query_names) + + def test_update_agent_max_iterations(self): + agent = self._create_basic_agent("TestUpdateIter", "test-update-iter") + result = self.store.ai.add_or_update_agent(agent) + self._created_agent_ids.append(result.identifier) + + agent.max_model_iterations_per_call = 10 + self.store.ai.add_or_update_agent(agent) + + response = self.store.ai.get_agents("test-update-iter") + self.assertEqual(10, response.ai_agents[0].max_model_iterations_per_call) + + # ---- Delete ---- + + def test_delete_agent(self): + agent = self._create_basic_agent("TestDelete", "test-delete") + result = self.store.ai.add_or_update_agent(agent) + + self.store.ai.delete_agent(result.identifier) + + with self.assertRaises(RuntimeError): + self.store.ai.get_agents("test-delete") + + def test_delete_agent_via_operation(self): + agent = self._create_basic_agent("TestDeleteOp", "test-delete-op") + result = self.store.maintenance.send(AddOrUpdateAiAgentOperation(agent)) + + self.store.maintenance.send(DeleteAiAgentOperation(result.identifier)) + + with self.assertRaises(RuntimeError): + self.store.maintenance.send(GetAiAgentOperation("test-delete-op")) + + # ---- Full lifecycle ---- + + def test_full_lifecycle(self): + # 1. Create + agent = self._create_full_agent("TestLifecycle", "test-lifecycle") + create_result = self.store.ai.add_or_update_agent(agent) + self.assertEqual("test-lifecycle", create_result.identifier) + + # 2. Get + response = self.store.ai.get_agents("test-lifecycle") + self.assertEqual(1, len(response.ai_agents)) + fetched = response.ai_agents[0] + self.assertEqual("TestLifecycle", fetched.name) + + # 3. Update + agent.system_prompt = "Updated lifecycle prompt." + agent.max_model_iterations_per_call = 5 + self.store.ai.add_or_update_agent(agent) + + # 4. Verify update + updated_response = self.store.ai.get_agents("test-lifecycle") + updated = updated_response.ai_agents[0] + self.assertEqual("Updated lifecycle prompt.", updated.system_prompt) + self.assertEqual(5, updated.max_model_iterations_per_call) + + # 5. Delete + self.store.ai.delete_agent("test-lifecycle") + + # 6. Verify deletion - server raises error when agent doesn't exist + with self.assertRaises(RuntimeError): + self.store.ai.get_agents("test-lifecycle") + + # ---- Query tool options ---- + + def test_query_tool_options_persisted(self): + agent = AiAgentConfiguration( + identifier="test-query-options", + name="TestQueryOptions", + connection_string_name=self.CONNECTION_STRING_NAME, + system_prompt="You help customers with their orders.", + sample_object='{"answer": "embed your answer here"}', + queries=[ + AiAgentToolQuery( + name="GetRecentOrders", + description="Retrieves recent orders for a customer", + query="from Orders where CustomerId = $customerId order by OrderDate desc limit 5", + parameters_sample_object='{"customerId": "embed customer id here"}', + options=AiAgentToolQueryOptions( + add_to_initial_context=True, + allow_model_queries=False, + ), + ), + ], + ) + self._created_agent_ids.append("test-query-options") + + self.store.ai.add_or_update_agent(agent) + + response = self.store.ai.get_agents("test-query-options") + self.assertEqual(1, len(response.ai_agents)) + + fetched = response.ai_agents[0] + self.assertEqual(1, len(fetched.queries)) + + opts = fetched.queries[0].options + self.assertIsNotNone(opts) + self.assertTrue(opts.add_to_initial_context) + self.assertFalse(opts.allow_model_queries) + + def test_query_tool_options_defaults_when_not_set(self): + agent = AiAgentConfiguration( + identifier="test-query-no-options", + name="TestQueryNoOptions", + connection_string_name=self.CONNECTION_STRING_NAME, + system_prompt="Agent without query options.", + sample_object='{"answer": "embed your answer here"}', + queries=[ + AiAgentToolQuery( + name="GetOrders", + description="Retrieves orders", + query="from Orders", + parameters_sample_object="{}", + ), + ], + ) + self._created_agent_ids.append("test-query-no-options") + + self.store.ai.add_or_update_agent(agent) + + response = self.store.ai.get_agents("test-query-no-options") + fetched = response.ai_agents[0] + self.assertEqual(1, len(fetched.queries)) + # No options set — options should be None or have default values + opts = fetched.queries[0].options + if opts is not None: + self.assertIsNone(opts.add_to_initial_context) + self.assertIsNone(opts.allow_model_queries) diff --git a/ravendb/tests/ai_agent_tests/test_ai_agents_conversation_mock.py b/ravendb/tests/ai_agent_tests/test_ai_agents_conversation_mock.py new file mode 100644 index 00000000..13466c8e --- /dev/null +++ b/ravendb/tests/ai_agent_tests/test_ai_agents_conversation_mock.py @@ -0,0 +1,443 @@ +""" +Hybrid tests for AI agent conversation flow. + +Agent CRUD (create/delete) runs against the real embedded server. +The actual conversation call (RunConversationOperation / maintenance.send) is +mocked so no LLM API key is required. +""" + +import json +import os +import unittest +from unittest.mock import patch + +from ravendb import ( + AiAgentConfiguration, + AiAgentToolAction, + AiAgentToolQuery, + AiConversationCreationOptions, +) +from ravendb.documents.ai.ai_answer import AiConversationStatus +from ravendb.documents.ai.ai_conversation import AiHandleErrorStrategy, UnhandledActionEventArgs +from ravendb.documents.operations.ai import AiConnectionString, AiModelType +from ravendb.documents.operations.ai.agents import ( + AiAgentActionRequest, + AiUsage, + ConversationResult, + DeleteAiAgentOperation, + RunConversationOperation, +) +from ravendb.documents.operations.ai.open_ai_settings import OpenAiSettings +from ravendb.documents.operations.connection_string.put_connection_string_operation import PutConnectionStringOperation +from ravendb.tests.test_base import TestBase + +CONNECTION_STRING_NAME = "conv-mock-cs" +AGENT_ID = "conv-mock-agent" + + +def _make_usage(prompt=10, completion=20): + return AiUsage(prompt_tokens=prompt, completion_tokens=completion, total_tokens=prompt + completion) + + +def _done_result(response=None, conversation_id="conversations/1", change_vector="A:1"): + return ConversationResult( + conversation_id=conversation_id, + change_vector=change_vector, + response=response or {"answer": "42"}, + usage=_make_usage(), + action_requests=[], + ) + + +def _action_result(action_name, tool_id, arguments, conversation_id="conversations/1", change_vector="A:1"): + return ConversationResult( + conversation_id=conversation_id, + change_vector=change_vector, + response=None, + usage=_make_usage(), + action_requests=[AiAgentActionRequest(name=action_name, tool_id=tool_id, arguments=json.dumps(arguments))], + ) + + +@unittest.skipIf(os.environ.get("RAVENDB_LICENSE") is None, "Insufficient license permissions. Skipping on CI/CD.") +class TestAiAgentConversationMock(TestBase): + """ + Hybrid tests: agent CRUD uses the real server; conversation calls are mocked. + The mock target is store.maintenance.send - only RunConversationOperation + calls are intercepted; everything else is forwarded to the real executor. + """ + + def setUp(self): + super().setUp() + cs = AiConnectionString( + name=CONNECTION_STRING_NAME, + identifier=CONNECTION_STRING_NAME, + model_type=AiModelType.CHAT, + openai_settings=OpenAiSettings( + api_key="dummy-key", + endpoint="https://api.openai.com/v1", + model="gpt-4", + ), + ) + self.store.maintenance.send(PutConnectionStringOperation(cs)) + agent = AiAgentConfiguration( + identifier=AGENT_ID, + name="ConvMockAgent", + connection_string_name=CONNECTION_STRING_NAME, + system_prompt="You are a helpful assistant.", + sample_object='{"answer": "embed your answer here"}', + queries=[ + AiAgentToolQuery( + name="get-orders", + description="Retrieve orders.", + query="from Orders", + parameters_sample_object="{}", + ), + ], + actions=[ + AiAgentToolAction( + name="store-result", + description="Store the result.", + parameters_sample_object='{"result": "embed result here"}', + ), + ], + ) + self.store.ai.add_or_update_agent(agent) + self._real_send = self.store.maintenance.send + + def tearDown(self): + try: + self.store.maintenance.send(DeleteAiAgentOperation(AGENT_ID)) + except Exception: + pass + super().tearDown() + + def _patched_send(self, mock_conv_result_fn): + real_send = self._real_send + + def _send(operation): + if isinstance(operation, RunConversationOperation): + return mock_conv_result_fn() + return real_send(operation) + + return _send + + def test_basic_conversation_returns_answer(self): + with patch.object( + self.store.maintenance, + "send", + side_effect=self._patched_send(lambda: _done_result(response={"answer": "Hello!"})), + ): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Hi there") + result = chat.run() + self.assertEqual(AiConversationStatus.DONE, result.status) + self.assertEqual("Hello!", result.answer["answer"]) + + def test_conversation_id_and_change_vector_are_stored(self): + with patch.object( + self.store.maintenance, + "send", + side_effect=self._patched_send( + lambda: _done_result(conversation_id="conversations/99", change_vector="A:99") + ), + ): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Hello") + chat.run() + self.assertEqual("conversations/99", chat._conversation_id) + self.assertEqual("A:99", chat._change_vector) + + def test_usage_is_populated(self): + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: _done_result())): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Hello") + result = chat.run() + self.assertIsNotNone(result.usage) + self.assertEqual(10, result.usage.prompt_tokens) + self.assertEqual(20, result.usage.completion_tokens) + self.assertEqual(30, result.usage.total_tokens) + + def test_elapsed_is_populated(self): + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: _done_result())): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Hello") + result = chat.run() + self.assertIsNotNone(result.elapsed) + + def test_context_manager_usage(self): + with patch.object( + self.store.maintenance, + "send", + side_effect=self._patched_send(lambda: _done_result(response={"answer": "ctx"})), + ): + with self.store.ai.conversation(AGENT_ID, "conversations/") as chat: + chat.set_user_prompt("Hello from context manager") + result = chat.run() + self.assertEqual(AiConversationStatus.DONE, result.status) + self.assertEqual("ctx", result.answer["answer"]) + + # ------------------------------------------------------------------ + # Action handler flow (handle / receive) + # ------------------------------------------------------------------ + + def test_handle_invokes_handler_and_sends_response(self): + calls = [] + responses = iter( + [ + _action_result("store-result", "tool-1", {"result": "data"}), + _done_result(response={"answer": "stored"}), + ] + ) + + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: next(responses))): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Store something") + chat.handle( + "store-result", lambda args: calls.append(args) or "ok", AiHandleErrorStrategy.SEND_ERRORS_TO_MODEL + ) + result = chat.run() + + self.assertEqual(AiConversationStatus.DONE, result.status) + self.assertEqual(1, len(calls)) + self.assertEqual({"result": "data"}, calls[0]) + + def test_receive_invokes_handler_with_request_and_args(self): + received = [] + responses = iter( + [ + _action_result("store-result", "tool-2", {"result": "payload"}), + _done_result(), + ] + ) + + def my_receiver(request, args): + received.append((request, args)) + chat.add_action_response(request.tool_id, "done") + + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: next(responses))): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Do something") + chat.receive("store-result", my_receiver, AiHandleErrorStrategy.SEND_ERRORS_TO_MODEL) + result = chat.run() + + self.assertEqual(AiConversationStatus.DONE, result.status) + self.assertEqual(1, len(received)) + req, args = received[0] + self.assertIsInstance(req, AiAgentActionRequest) + self.assertEqual("tool-2", req.tool_id) + self.assertEqual({"result": "payload"}, args) + + def test_multi_turn_action_loop(self): + responses = iter( + [ + _action_result("store-result", "tool-a", {"result": "first"}), + _action_result("store-result", "tool-b", {"result": "second"}), + _done_result(response={"answer": "all done"}), + ] + ) + + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: next(responses))): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Do two things") + chat.handle("store-result", lambda args: "handled", AiHandleErrorStrategy.SEND_ERRORS_TO_MODEL) + result = chat.run() + + self.assertEqual(AiConversationStatus.DONE, result.status) + + # ------------------------------------------------------------------ + # Error handling strategies + # ------------------------------------------------------------------ + + def test_handler_error_send_to_model(self): + responses = iter( + [ + _action_result("store-result", "tool-err", {"result": "x"}), + _done_result(response={"answer": "recovered"}), + ] + ) + + def bad_handler(args): + raise ValueError("something went wrong") + + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: next(responses))): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Trigger error") + chat.handle("store-result", bad_handler, AiHandleErrorStrategy.SEND_ERRORS_TO_MODEL) + result = chat.run() + + # Error was sent to model as a response; conversation continued and finished + self.assertEqual(AiConversationStatus.DONE, result.status) + + def test_handler_error_raise_immediately(self): + def mock_send(): + return _action_result("store-result", "tool-err2", {"result": "x"}) + + def bad_handler(args): + raise RuntimeError("fatal error") + + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(mock_send)): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Trigger fatal error") + chat.handle("store-result", bad_handler, AiHandleErrorStrategy.RAISE_IMMEDIATELY) + with self.assertRaises(RuntimeError): + chat.run() + + # ------------------------------------------------------------------ + # Unhandled action event + # ------------------------------------------------------------------ + + def test_on_unhandled_action_is_called(self): + unhandled = [] + responses = iter( + [ + _action_result("unknown-action", "t-99", {"x": 1}), + _done_result(), + ] + ) + + def on_unhandled(event_args): + unhandled.append(event_args) + event_args.sender.add_action_response(event_args.action.tool_id, "fallback") + + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: next(responses))): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.on_unhandled_action = on_unhandled + chat.set_user_prompt("Do something unhandled") + result = chat.run() + + self.assertEqual(AiConversationStatus.DONE, result.status) + self.assertEqual(1, len(unhandled)) + self.assertIsInstance(unhandled[0], UnhandledActionEventArgs) + self.assertEqual("unknown-action", unhandled[0].action.name) + + def test_no_handler_raises_runtime_error(self): + responses = iter( + [ + _action_result("missing-action", "t-0", {}), + ] + ) + + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: next(responses))): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Trigger missing handler") + with self.assertRaises(RuntimeError): + chat.run() + + # ------------------------------------------------------------------ + # Artificial action injection + # ------------------------------------------------------------------ + + def test_add_artificial_action_with_response(self): + # AiConversation clears _artificial_actions in the finally block after sending, + # so we snapshot the list contents at call time before the clear happens. + captured_artificial_actions = [] + + def capturing_send(operation): + if isinstance(operation, RunConversationOperation): + # snapshot before the finally-block clear + captured_artificial_actions.extend(list(operation._artificial_actions)) + return _done_result(response={"answer": "injected"}) + return self._real_send(operation) + + with patch.object(self.store.maintenance, "send", side_effect=capturing_send): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Use injected context") + chat.add_artificial_action_with_response("get-orders", {"orders": []}) + result = chat.run() + + self.assertEqual(AiConversationStatus.DONE, result.status) + self.assertEqual(1, len(captured_artificial_actions)) + self.assertEqual("get-orders", captured_artificial_actions[0].tool_id) + + def test_add_artificial_action_validates_tool_id(self): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + with self.assertRaises(ValueError): + chat.add_artificial_action_with_response("", "some response") + with self.assertRaises(ValueError): + chat.add_artificial_action_with_response(" ", "some response") + + def test_add_artificial_action_validates_response_not_none(self): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + with self.assertRaises(ValueError): + chat.add_artificial_action_with_response("get-orders", None) + + # ------------------------------------------------------------------ + # Conversation with creation options (parameters / expiration) + # ------------------------------------------------------------------ + + def test_conversation_with_creation_options(self): + options = AiConversationCreationOptions( + parameters={"customerId": "C-42"}, + expiration_in_sec=3600, + ) + sent_operations = [] + + def capturing_send(operation): + sent_operations.append(operation) + return self._patched_send(lambda: _done_result(response={"answer": "ok"}))(operation) + + with patch.object(self.store.maintenance, "send", side_effect=capturing_send): + chat = self.store.ai.conversation(AGENT_ID, "conversations/", creation_options=options) + chat.set_user_prompt("Hello with options") + result = chat.run() + + self.assertEqual(AiConversationStatus.DONE, result.status) + conv_op = next(o for o in sent_operations if isinstance(o, RunConversationOperation)) + self.assertEqual({"customerId": "C-42"}, conv_op._options.parameters) + self.assertEqual(3600, conv_op._options.expiration_in_sec) + + # ------------------------------------------------------------------ + # Streaming + # ------------------------------------------------------------------ + + def test_stream_collects_chunks_and_returns_answer(self): + chunks = [] + + def mock_send(): + return _done_result(response={"answer": "streamed answer"}) + + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(mock_send)): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Stream this") + result = chat.stream(stream_property_path="answer", on_chunk=chunks.append) + + self.assertEqual(AiConversationStatus.DONE, result.status) + + # ------------------------------------------------------------------ + # set_user_prompt validation + # ------------------------------------------------------------------ + + def test_set_user_prompt_empty_raises(self): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + with self.assertRaises(ValueError): + chat.set_user_prompt("") + with self.assertRaises(ValueError): + chat.set_user_prompt(" ") + + def test_required_actions_before_run_raises(self): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + with self.assertRaises(RuntimeError): + _ = chat.required_actions + + def test_required_actions_after_run_returns_list(self): + with patch.object(self.store.maintenance, "send", side_effect=self._patched_send(lambda: _done_result())): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.set_user_prompt("Hello") + chat.run() + self.assertIsInstance(chat.required_actions, list) + self.assertEqual(0, len(chat.required_actions)) + + # ------------------------------------------------------------------ + # Duplicate action handler registration + # ------------------------------------------------------------------ + + def test_duplicate_action_handler_raises(self): + chat = self.store.ai.conversation(AGENT_ID, "conversations/") + chat.handle("store-result", lambda args: None, AiHandleErrorStrategy.SEND_ERRORS_TO_MODEL) + with self.assertRaises(ValueError): + chat.handle("store-result", lambda args: None, AiHandleErrorStrategy.SEND_ERRORS_TO_MODEL) + + +if __name__ == "__main__": + unittest.main()