From d5c0322e2ac109665ae99c65637e97476fc0c815 Mon Sep 17 00:00:00 2001 From: Alex Date: Mon, 30 Mar 2026 16:13:08 +0100 Subject: [PATCH] chore: more tests --- tests/agents/test_node_agent.py | 92 + tests/agents/test_research_agent.py | 473 +++- tests/agents/tools/__init__.py | 0 .../agents/tools/test_api_body_serializer.py | 424 +++ tests/agents/tools/test_api_tool.py | 516 ++++ tests/agents/tools/test_internal_search.py | 596 ++++ tests/agents/tools/test_mcp_tool.py | 1010 +++++++ tests/agents/tools/test_memory.py | 449 +++ .../compression/test_threshold_checker.py | 46 + tests/api/answer/test_base_routes.py | 363 +++ tests/api/answer/test_conversation_service.py | 418 +++ tests/api/answer/test_stream_processor.py | 185 +- tests/api/test_internal_routes.py | 416 +++ tests/api/user/__init__.py | 0 tests/api/user/sources/test_chunks.py | 879 ++++++ tests/api/user/sources/test_source_routes.py | 965 +++++++ tests/api/user/sources/test_upload.py | 1338 +++++++++ tests/api/user/test_agents_routes.py | 2516 +++++++++++++++++ tests/api/user/test_agents_sharing.py | 768 +++++ tests/api/user/test_analytics.py | 388 +++ tests/api/user/test_conversations.py | 360 +++ tests/api/user/test_folders.py | 509 ++++ tests/api/user/test_models.py | 70 + tests/api/user/test_prompts.py | 288 ++ tests/api/user/test_sharing.py | 690 +++++ tests/api/user/test_tools_mcp.py | 1308 +++++++++ tests/api/user/test_tools_routes.py | 1948 +++++++++++++ tests/api/user/test_utils.py | 411 +++ tests/api/user/test_webhooks.py | 225 ++ tests/api/user/test_workflows.py | 406 +++ tests/core/__init__.py | 0 tests/core/test_model_settings.py | 95 + tests/core/test_url_validation.py | 64 + tests/llm/__init__.py | 0 tests/llm/test_anthropic.py | 323 +++ tests/llm/test_base.py | 269 ++ tests/llm/test_google_ai.py | 755 +++++ tests/llm/test_llama_cpp.py | 193 ++ tests/llm/test_openai.py | 717 +++++ tests/llm/test_premai.py | 190 ++ tests/parser/file/__init__.py | 0 tests/parser/file/test_bulk.py | 367 +++ tests/parser/file/test_docling_parser.py | 382 +++ tests/parser/file/test_docs_parser.py | 254 +- tests/parser/remote/__init__.py | 0 tests/parser/remote/test_github_loader.py | 153 + tests/parser/remote/test_sitemap_loader.py | 306 ++ tests/parser/test_chunking.py | 279 ++ tests/parser/test_schema.py | 58 + tests/security/__init__.py | 0 tests/security/test_encryption.py | 80 + tests/storage/__init__.py | 0 tests/storage/test_s3_storage.py | 97 + tests/stt/__init__.py | 0 tests/stt/test_faster_whisper.py | 234 ++ tests/stt/test_live_session.py | 252 ++ tests/stt/test_openai_stt.py | 275 ++ tests/test_auth.py | 83 + tests/test_cache.py | 300 +- tests/test_error.py | 42 +- tests/test_logging.py | 202 ++ tests/test_usage.py | 263 ++ tests/test_utils.py | 131 + tests/vectorstore/test_faiss.py | 155 + 64 files changed, 24436 insertions(+), 140 deletions(-) create mode 100644 tests/agents/test_node_agent.py create mode 100644 tests/agents/tools/__init__.py create mode 100644 tests/agents/tools/test_api_body_serializer.py create mode 100644 tests/agents/tools/test_api_tool.py create mode 100644 tests/agents/tools/test_internal_search.py create mode 100644 tests/agents/tools/test_mcp_tool.py create mode 100644 tests/agents/tools/test_memory.py create mode 100644 tests/api/answer/services/compression/test_threshold_checker.py create mode 100644 tests/api/answer/test_base_routes.py create mode 100644 tests/api/answer/test_conversation_service.py create mode 100644 tests/api/test_internal_routes.py create mode 100644 tests/api/user/__init__.py create mode 100644 tests/api/user/sources/test_chunks.py create mode 100644 tests/api/user/sources/test_source_routes.py create mode 100644 tests/api/user/sources/test_upload.py create mode 100644 tests/api/user/test_agents_routes.py create mode 100644 tests/api/user/test_agents_sharing.py create mode 100644 tests/api/user/test_analytics.py create mode 100644 tests/api/user/test_conversations.py create mode 100644 tests/api/user/test_folders.py create mode 100644 tests/api/user/test_models.py create mode 100644 tests/api/user/test_prompts.py create mode 100644 tests/api/user/test_sharing.py create mode 100644 tests/api/user/test_tools_mcp.py create mode 100644 tests/api/user/test_tools_routes.py create mode 100644 tests/api/user/test_utils.py create mode 100644 tests/api/user/test_webhooks.py create mode 100644 tests/api/user/test_workflows.py create mode 100644 tests/core/__init__.py create mode 100644 tests/llm/__init__.py create mode 100644 tests/llm/test_anthropic.py create mode 100644 tests/llm/test_base.py create mode 100644 tests/llm/test_google_ai.py create mode 100644 tests/llm/test_llama_cpp.py create mode 100644 tests/llm/test_openai.py create mode 100644 tests/llm/test_premai.py create mode 100644 tests/parser/file/__init__.py create mode 100644 tests/parser/file/test_bulk.py create mode 100644 tests/parser/file/test_docling_parser.py create mode 100644 tests/parser/remote/__init__.py create mode 100644 tests/parser/remote/test_sitemap_loader.py create mode 100644 tests/parser/test_chunking.py create mode 100644 tests/parser/test_schema.py create mode 100644 tests/security/__init__.py create mode 100644 tests/storage/__init__.py create mode 100644 tests/stt/__init__.py create mode 100644 tests/stt/test_faster_whisper.py create mode 100644 tests/stt/test_openai_stt.py create mode 100644 tests/test_auth.py create mode 100644 tests/test_logging.py diff --git a/tests/agents/test_node_agent.py b/tests/agents/test_node_agent.py new file mode 100644 index 00000000..d2baef70 --- /dev/null +++ b/tests/agents/test_node_agent.py @@ -0,0 +1,92 @@ + +import pytest + + +@pytest.mark.unit +class TestToolFilterMixin: + + def test_get_user_tools_filters_by_allowed_ids(self): + from application.agents.workflows.node_agent import ToolFilterMixin + + class FakeBase: + def _get_user_tools(self, user="local"): + return { + "t1": {"_id": "id1", "name": "tool1"}, + "t2": {"_id": "id2", "name": "tool2"}, + "t3": {"_id": "id3", "name": "tool3"}, + } + + class TestClass(ToolFilterMixin, FakeBase): + pass + + obj = TestClass() + obj._allowed_tool_ids = ["id1", "id3"] + result = obj._get_user_tools("user1") + assert "t1" in result + assert "t3" in result + assert "t2" not in result + + def test_get_user_tools_returns_empty_when_no_allowed(self): + from application.agents.workflows.node_agent import ToolFilterMixin + + class FakeBase: + def _get_user_tools(self, user="local"): + return {"t1": {"_id": "id1"}} + + class TestClass(ToolFilterMixin, FakeBase): + pass + + obj = TestClass() + obj._allowed_tool_ids = [] + result = obj._get_user_tools() + assert result == {} + + def test_get_tools_filters_by_allowed_ids(self): + from application.agents.workflows.node_agent import ToolFilterMixin + + class FakeBase: + def _get_tools(self, api_key=None): + return { + "t1": {"_id": "id1"}, + "t2": {"_id": "id2"}, + } + + class TestClass(ToolFilterMixin, FakeBase): + pass + + obj = TestClass() + obj._allowed_tool_ids = ["id2"] + result = obj._get_tools("key") + assert "t2" in result + assert "t1" not in result + + def test_get_tools_returns_empty_when_no_allowed(self): + from application.agents.workflows.node_agent import ToolFilterMixin + + class FakeBase: + def _get_tools(self, api_key=None): + return {"t1": {"_id": "id1"}} + + class TestClass(ToolFilterMixin, FakeBase): + pass + + obj = TestClass() + obj._allowed_tool_ids = [] + result = obj._get_tools() + assert result == {} + + +@pytest.mark.unit +class TestWorkflowNodeAgentFactory: + + def test_raises_on_unsupported_type(self): + from application.agents.workflows.node_agent import WorkflowNodeAgentFactory + + with pytest.raises(ValueError, match="Unsupported agent type"): + WorkflowNodeAgentFactory.create( + agent_type="nonexistent", + endpoint="http://example.com", + llm_name="openai", + model_id="gpt-4", + api_key="key", + ) diff --git a/tests/agents/test_research_agent.py b/tests/agents/test_research_agent.py index 30b29ebc..3d2f1c89 100644 --- a/tests/agents/test_research_agent.py +++ b/tests/agents/test_research_agent.py @@ -1,22 +1,31 @@ -"""Tests for ResearchAgent — multi-step research with budget controls.""" +"""Comprehensive tests for application/agents/research_agent.py + +Covers: CitationManager, ResearchAgent (init, budget, timeout, phases: +clarification, planning, research step, synthesis, _extract_text, +JSON parsing, tool setup, is_follow_up). +""" import json import time -from unittest.mock import Mock +from unittest.mock import Mock, patch import pytest + from application.agents.research_agent import ( + COMPLEXITY_CAPS, CitationManager, ResearchAgent, DEFAULT_MAX_STEPS, + DEFAULT_MAX_SUB_ITERATIONS, DEFAULT_TIMEOUT_SECONDS, DEFAULT_TOKEN_BUDGET, + DEFAULT_PARALLEL_WORKERS, ) -# --------------------------------------------------------------------------- +# ===================================================================== # CitationManager -# --------------------------------------------------------------------------- +# ===================================================================== @pytest.mark.unit @@ -41,6 +50,12 @@ class TestCitationManager: assert n1 != n2 assert len(cm.citations) == 2 + def test_add_same_source_different_title(self): + cm = CitationManager() + n1 = cm.add({"source": "s1", "title": "T1"}) + n2 = cm.add({"source": "s1", "title": "T2"}) + assert n1 != n2 + def test_add_docs_returns_mapping(self): cm = CitationManager() docs = [ @@ -51,12 +66,32 @@ class TestCitationManager: assert "[1] Doc A" in text assert "[2] Doc B" in text + def test_add_docs_deduplication(self): + cm = CitationManager() + docs = [ + {"source": "s1", "title": "Doc A"}, + {"source": "s1", "title": "Doc A"}, + ] + text = cm.add_docs(docs) + assert text.count("[1]") == 2 + def test_format_references(self): cm = CitationManager() - cm.add({"source": "http://example.com", "title": "Example", "filename": "ex.md"}) + cm.add({ + "source": "http://example.com", + "title": "Example", + "filename": "ex.md", + }) refs = cm.format_references() assert "[1]" in refs assert "ex.md" in refs + assert "http://example.com" in refs + + def test_format_references_uses_title_when_no_filename(self): + cm = CitationManager() + cm.add({"source": "http://example.com", "title": "My Title"}) + refs = cm.format_references() + assert "My Title" in refs def test_format_references_empty(self): cm = CitationManager() @@ -69,10 +104,21 @@ class TestCitationManager: docs = cm.get_all_docs() assert len(docs) == 2 + def test_format_references_sorted(self): + cm = CitationManager() + cm.add({"source": "s1", "title": "A"}) + cm.add({"source": "s2", "title": "B"}) + cm.add({"source": "s3", "title": "C"}) + refs = cm.format_references() + lines = refs.strip().split("\n") + assert lines[0].startswith("[1]") + assert lines[1].startswith("[2]") + assert lines[2].startswith("[3]") -# --------------------------------------------------------------------------- -# ResearchAgent Init & Budget -# --------------------------------------------------------------------------- + +# ===================================================================== +# ResearchAgent Init & Constants +# ===================================================================== @pytest.mark.unit @@ -86,6 +132,8 @@ class TestResearchAgentInit: assert agent.max_steps == DEFAULT_MAX_STEPS assert agent.timeout_seconds == DEFAULT_TIMEOUT_SECONDS assert agent.token_budget == DEFAULT_TOKEN_BUDGET + assert agent.max_sub_iterations == DEFAULT_MAX_SUB_ITERATIONS + assert agent.parallel_workers == DEFAULT_PARALLEL_WORKERS assert agent.retriever_config == {} def test_custom_budget( @@ -95,11 +143,15 @@ class TestResearchAgentInit: max_steps=3, timeout_seconds=60, token_budget=50_000, + max_sub_iterations=2, + parallel_workers=1, **agent_base_params, ) assert agent.max_steps == 3 assert agent.timeout_seconds == 60 assert agent.token_budget == 50_000 + assert agent.max_sub_iterations == 2 + assert agent.parallel_workers == 1 def test_with_retriever_config( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator @@ -108,18 +160,39 @@ class TestResearchAgentInit: agent = ResearchAgent(retriever_config=rc, **agent_base_params) assert agent.retriever_config == rc + def test_constants(self): + assert DEFAULT_MAX_STEPS == 6 + assert DEFAULT_MAX_SUB_ITERATIONS == 5 + assert DEFAULT_TIMEOUT_SECONDS == 300 + assert DEFAULT_TOKEN_BUDGET == 100_000 + assert DEFAULT_PARALLEL_WORKERS == 3 + + def test_complexity_caps(self): + assert COMPLEXITY_CAPS["simple"] == 2 + assert COMPLEXITY_CAPS["moderate"] == 4 + assert COMPLEXITY_CAPS["complex"] == 6 + + +# ===================================================================== +# Budget & Timeout +# ===================================================================== + @pytest.mark.unit class TestResearchAgentBudget: - def _make_agent(self, agent_base_params, mock_llm_creator, mock_llm_handler_creator, **kwargs): + def _make_agent( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator, **kwargs + ): return ResearchAgent(**kwargs, **agent_base_params) def test_timeout_detection( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): agent = self._make_agent( - agent_base_params, mock_llm_creator, mock_llm_handler_creator, + agent_base_params, + mock_llm_creator, + mock_llm_handler_creator, timeout_seconds=0, ) agent._start_time = time.monotonic() - 1 @@ -129,7 +202,9 @@ class TestResearchAgentBudget: self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): agent = self._make_agent( - agent_base_params, mock_llm_creator, mock_llm_handler_creator, + agent_base_params, + mock_llm_creator, + mock_llm_handler_creator, timeout_seconds=300, ) agent._start_time = time.monotonic() @@ -139,7 +214,9 @@ class TestResearchAgentBudget: self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): agent = self._make_agent( - agent_base_params, mock_llm_creator, mock_llm_handler_creator, + agent_base_params, + mock_llm_creator, + mock_llm_handler_creator, token_budget=1000, ) agent._track_tokens(500) @@ -150,36 +227,55 @@ class TestResearchAgentBudget: assert agent._budget_remaining() == 0 assert agent._is_over_budget() is True - def test_snapshot_llm_tokens_returns_delta( - self, agent_base_params, mock_llm, mock_llm_creator, mock_llm_handler_creator + def test_over_budget_returns_zero_remaining( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): agent = self._make_agent( - agent_base_params, mock_llm_creator, mock_llm_handler_creator, + agent_base_params, + mock_llm_creator, + mock_llm_handler_creator, + token_budget=100, + ) + agent._track_tokens(200) + assert agent._budget_remaining() == 0 + + def test_snapshot_llm_tokens_returns_delta( + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, + ): + agent = self._make_agent( + agent_base_params, + mock_llm_creator, + mock_llm_handler_creator, ) mock_llm.token_usage = {"prompt_tokens": 100, "generated_tokens": 50} delta1 = agent._snapshot_llm_tokens() assert delta1 == 150 - # Simulate more tokens used mock_llm.token_usage = {"prompt_tokens": 200, "generated_tokens": 100} delta2 = agent._snapshot_llm_tokens() - assert delta2 == 150 # 300 - 150 + assert delta2 == 150 def test_elapsed( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): agent = self._make_agent( - agent_base_params, mock_llm_creator, mock_llm_handler_creator, + agent_base_params, + mock_llm_creator, + mock_llm_handler_creator, ) agent._start_time = time.monotonic() - 1.5 elapsed = agent._elapsed() assert elapsed >= 1.0 -# --------------------------------------------------------------------------- -# ResearchAgent Phases -# --------------------------------------------------------------------------- +# ===================================================================== +# Clarification Phase +# ===================================================================== @pytest.mark.unit @@ -195,7 +291,11 @@ class TestResearchAgentClarification: self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): agent_base_params["chat_history"] = [ - {"prompt": "What?", "response": "Clarify", "metadata": {"is_clarification": True}}, + { + "prompt": "What?", + "response": "Clarify", + "metadata": {"is_clarification": True}, + }, ] agent = ResearchAgent(**agent_base_params) assert agent._is_follow_up() is True @@ -209,8 +309,21 @@ class TestResearchAgentClarification: agent = ResearchAgent(**agent_base_params) assert agent._is_follow_up() is False + def test_is_follow_up_empty_metadata( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent_base_params["chat_history"] = [ + {"prompt": "What?", "response": "X", "metadata": {}}, + ] + agent = ResearchAgent(**agent_base_params) + assert agent._is_follow_up() is False + def test_clarification_returns_none_on_no_clarification_needed( - self, agent_base_params, mock_llm, mock_llm_creator, mock_llm_handler_creator + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, ): response = Mock() response.choices = [Mock()] @@ -226,13 +339,16 @@ class TestResearchAgentClarification: assert result is None def test_clarification_returns_questions( - self, agent_base_params, mock_llm, mock_llm_creator, mock_llm_handler_creator + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, ): clarification_json = json.dumps({ "needs_clarification": True, "questions": ["Which version?", "What context?"], }) - # Return a plain string so _extract_text handles it directly mock_llm.gen = Mock(return_value=clarification_json) mock_llm.token_usage = {"prompt_tokens": 10, "generated_tokens": 5} @@ -241,13 +357,76 @@ class TestResearchAgentClarification: assert result is not None assert "Which version?" in result assert "What context?" in result + assert "1." in result + assert "2." in result + + def test_clarification_limits_questions_to_three( + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, + ): + clarification_json = json.dumps({ + "needs_clarification": True, + "questions": ["q1", "q2", "q3", "q4", "q5"], + }) + mock_llm.gen = Mock(return_value=clarification_json) + mock_llm.token_usage = {"prompt_tokens": 10, "generated_tokens": 5} + + agent = ResearchAgent(**agent_base_params) + result = agent._clarification_phase("complex question") + # Should only show 3 questions + assert "3." in result + assert "4." not in result + + def test_clarification_returns_none_on_empty_questions( + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, + ): + clarification_json = json.dumps({ + "needs_clarification": True, + "questions": [], + }) + mock_llm.gen = Mock(return_value=clarification_json) + mock_llm.token_usage = {"prompt_tokens": 10, "generated_tokens": 5} + + agent = ResearchAgent(**agent_base_params) + result = agent._clarification_phase("question") + assert result is None + + def test_clarification_returns_none_on_llm_error( + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, + ): + mock_llm.gen = Mock(side_effect=Exception("LLM error")) + mock_llm.token_usage = {"prompt_tokens": 0, "generated_tokens": 0} + + agent = ResearchAgent(**agent_base_params) + result = agent._clarification_phase("question") + assert result is None + + +# ===================================================================== +# Planning Phase +# ===================================================================== @pytest.mark.unit class TestResearchAgentPlanning: def test_planning_returns_steps_and_complexity( - self, agent_base_params, mock_llm, mock_llm_creator, mock_llm_handler_creator + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, ): plan_json = json.dumps({ "complexity": "moderate", @@ -256,7 +435,6 @@ class TestResearchAgentPlanning: {"query": "sub-question 2", "rationale": "reason 2"}, ], }) - # Return plain string so _extract_text handles it directly mock_llm.gen = Mock(return_value=plan_json) mock_llm.token_usage = {"prompt_tokens": 10, "generated_tokens": 5} @@ -268,13 +446,15 @@ class TestResearchAgentPlanning: assert steps[0]["query"] == "sub-question 1" def test_planning_caps_steps_by_complexity( - self, agent_base_params, mock_llm, mock_llm_creator, mock_llm_handler_creator + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, ): plan_json = json.dumps({ "complexity": "simple", - "steps": [ - {"query": f"q{i}", "rationale": f"r{i}"} for i in range(10) - ], + "steps": [{"query": f"q{i}", "rationale": f"r{i}"} for i in range(10)], }) response = Mock() response.choices = [Mock()] @@ -287,10 +467,34 @@ class TestResearchAgentPlanning: steps, complexity = agent._planning_phase("Simple question") assert complexity == "simple" - assert len(steps) <= 2 # COMPLEXITY_CAPS["simple"] == 2 + assert len(steps) <= 2 + + def test_planning_caps_steps_for_complex( + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, + ): + plan_json = json.dumps({ + "complexity": "complex", + "steps": [{"query": f"q{i}", "rationale": f"r{i}"} for i in range(10)], + }) + mock_llm.gen = Mock(return_value=plan_json) + mock_llm.token_usage = {"prompt_tokens": 10, "generated_tokens": 5} + + agent = ResearchAgent(**agent_base_params) + steps, complexity = agent._planning_phase("Complex analysis") + + assert complexity == "complex" + assert len(steps) <= 6 def test_planning_fallback_on_error( - self, agent_base_params, mock_llm, mock_llm_creator, mock_llm_handler_creator + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, ): mock_llm.gen = Mock(side_effect=Exception("LLM down")) mock_llm.token_usage = {"prompt_tokens": 0, "generated_tokens": 0} @@ -302,23 +506,54 @@ class TestResearchAgentPlanning: assert len(steps) == 1 assert steps[0]["query"] == "Anything" + def test_planning_list_response( + self, + agent_base_params, + mock_llm, + mock_llm_creator, + mock_llm_handler_creator, + ): + plan_json = json.dumps([ + {"query": "q1", "rationale": "r1"}, + {"query": "q2", "rationale": "r2"}, + ]) + mock_llm.gen = Mock(return_value=plan_json) + mock_llm.token_usage = {"prompt_tokens": 10, "generated_tokens": 5} + + agent = ResearchAgent(**agent_base_params) + steps, complexity = agent._planning_phase("question") + + assert complexity == "moderate" + assert len(steps) == 2 + + +# ===================================================================== +# Extract Text +# ===================================================================== + @pytest.mark.unit class TestResearchAgentExtractText: - def _make_agent(self, agent_base_params, mock_llm_creator, mock_llm_handler_creator): + def _make_agent( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): return ResearchAgent(**agent_base_params) def test_extract_from_string( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) assert agent._extract_text("hello") == "hello" def test_extract_from_openai_response( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) response = Mock() response.choices = [Mock()] response.choices[0].message = Mock() @@ -330,7 +565,9 @@ class TestResearchAgentExtractText: def test_extract_from_anthropic_response( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) text_block = Mock() text_block.text = "Anthropic content" response = Mock() @@ -339,47 +576,106 @@ class TestResearchAgentExtractText: response.choices = None assert agent._extract_text(response) == "Anthropic content" + def test_extract_from_message_content( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) + response = Mock() + response.message = Mock() + response.message.content = "From message" + assert agent._extract_text(response) == "From message" + def test_extract_from_none( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) assert agent._extract_text(None) == "" +# ===================================================================== +# Parse JSON +# ===================================================================== + + @pytest.mark.unit class TestResearchAgentParseJson: - def _make_agent(self, agent_base_params, mock_llm_creator, mock_llm_handler_creator): + def _make_agent( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): return ResearchAgent(**agent_base_params) def test_parse_plan_direct_json( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) text = '{"steps": [{"query": "q1"}], "complexity": "simple"}' result = agent._parse_plan_json(text) assert isinstance(result, dict) assert len(result["steps"]) == 1 + def test_parse_plan_list( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) + text = '[{"query": "q1"}]' + result = agent._parse_plan_json(text) + assert isinstance(result, list) + assert len(result) == 1 + def test_parse_plan_from_code_fence( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) text = 'Here is the plan:\n```json\n{"steps": [{"query": "q1"}]}\n```' result = agent._parse_plan_json(text) assert isinstance(result, dict) + def test_parse_plan_from_plain_code_fence( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) + text = 'Result:\n```\n{"steps": [{"query": "q1"}]}\n```' + result = agent._parse_plan_json(text) + assert isinstance(result, dict) + + def test_parse_plan_embedded_json_object( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) + text = 'Here is the plan: {"steps": [{"query": "q1"}]} end.' + result = agent._parse_plan_json(text) + assert isinstance(result, dict) + def test_parse_plan_invalid_returns_empty( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) result = agent._parse_plan_json("not json at all") assert result == [] def test_parse_clarification_json( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) text = '{"needs_clarification": false, "reason": "clear"}' result = agent._parse_clarification_json(text) assert result["needs_clarification"] is False @@ -387,14 +683,101 @@ class TestResearchAgentParseJson: def test_parse_clarification_json_from_code_fence( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) text = '```json\n{"needs_clarification": true, "questions": ["q1"]}\n```' result = agent._parse_clarification_json(text) assert result["needs_clarification"] is True + def test_parse_clarification_embedded_json( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) + text = 'Here: {"needs_clarification": true, "questions": ["q1"]} done.' + result = agent._parse_clarification_json(text) + assert result["needs_clarification"] is True + def test_parse_clarification_json_invalid( self, agent_base_params, mock_llm_creator, mock_llm_handler_creator ): - agent = self._make_agent(agent_base_params, mock_llm_creator, mock_llm_handler_creator) + agent = self._make_agent( + agent_base_params, mock_llm_creator, mock_llm_handler_creator + ) result = agent._parse_clarification_json("not json") assert result is None + + +# ===================================================================== +# Tool Setup +# ===================================================================== + + +@pytest.mark.unit +class TestResearchAgentToolSetup: + + def test_setup_tools_includes_think_and_internal( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = ResearchAgent( + retriever_config={ + "source": {"active_docs": ["abc"]}, + "retriever_name": "classic", + }, + **agent_base_params, + ) + + with patch( + "application.agents.research_agent.add_internal_search_tool" + ) as mock_add: + tools = agent._setup_tools() + mock_add.assert_called_once() + assert "think" in tools + + def test_setup_tools_no_retriever_config( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = ResearchAgent(**agent_base_params) + + with patch( + "application.agents.research_agent.add_internal_search_tool" + ) as mock_add: + tools = agent._setup_tools() + mock_add.assert_called_once() + assert "think" in tools + + +# ===================================================================== +# Collect Step Sources +# ===================================================================== + + +@pytest.mark.unit +class TestCollectStepSources: + + def test_collects_from_internal_search_tool( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = ResearchAgent(**agent_base_params) + + mock_tool = Mock() + mock_tool.retrieved_docs = [ + {"source": "s1", "title": "T1"}, + {"source": "s2", "title": "T2"}, + ] + + cache_key = f"internal_search:internal:{agent.user or ''}" + agent.tool_executor._loaded_tools[cache_key] = mock_tool + + agent._collect_step_sources() + + assert len(agent.citations.citations) == 2 + + def test_no_tool_no_error( + self, agent_base_params, mock_llm_creator, mock_llm_handler_creator + ): + agent = ResearchAgent(**agent_base_params) + agent._collect_step_sources() + assert len(agent.citations.citations) == 0 diff --git a/tests/agents/tools/__init__.py b/tests/agents/tools/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/agents/tools/test_api_body_serializer.py b/tests/agents/tools/test_api_body_serializer.py new file mode 100644 index 00000000..62904211 --- /dev/null +++ b/tests/agents/tools/test_api_body_serializer.py @@ -0,0 +1,424 @@ +"""Comprehensive tests for application/agents/tools/api_body_serializer.py + +Covers: ContentType enum, RequestBodySerializer (JSON, form-urlencoded, +multipart, text/plain, XML, octet-stream, unknown types), encoding rules, +helper methods (_percent_encode, _escape_xml, _dict_to_xml). +""" + +import json + +import pytest + +from application.agents.tools.api_body_serializer import ( + ContentType, + RequestBodySerializer, +) + + +# ===================================================================== +# ContentType Enum +# ===================================================================== + + +@pytest.mark.unit +class TestContentTypeEnum: + + def test_json_value(self): + assert ContentType.JSON == "application/json" + + def test_form_urlencoded_value(self): + assert ContentType.FORM_URLENCODED == "application/x-www-form-urlencoded" + + def test_multipart_value(self): + assert ContentType.MULTIPART_FORM_DATA == "multipart/form-data" + + def test_text_plain_value(self): + assert ContentType.TEXT_PLAIN == "text/plain" + + def test_xml_value(self): + assert ContentType.XML == "application/xml" + + def test_octet_stream_value(self): + assert ContentType.OCTET_STREAM == "application/octet-stream" + + def test_str_enum(self): + assert isinstance(ContentType.JSON, str) + + +# ===================================================================== +# JSON Serialization +# ===================================================================== + + +@pytest.mark.unit +class TestSerializeJson: + + def test_basic_json(self): + body, headers = RequestBodySerializer.serialize( + {"key": "value"}, ContentType.JSON + ) + assert json.loads(body) == {"key": "value"} + assert headers["Content-Type"] == "application/json" + + def test_nested_json(self): + data = {"user": {"name": "Alice", "age": 30}} + body, headers = RequestBodySerializer.serialize(data, ContentType.JSON) + assert json.loads(body) == data + + def test_empty_body_returns_none(self): + body, headers = RequestBodySerializer.serialize({}, ContentType.JSON) + assert body is None + assert headers == {} + + def test_none_body(self): + body, headers = RequestBodySerializer.serialize(None, ContentType.JSON) + assert body is None + + def test_unknown_content_type_falls_back_to_json(self): + body, headers = RequestBodySerializer.serialize( + {"k": "v"}, "application/vnd.custom+json" + ) + assert json.loads(body) == {"k": "v"} + + def test_content_type_with_charset_suffix(self): + body, headers = RequestBodySerializer.serialize( + {"k": "v"}, "application/json; charset=utf-8" + ) + assert json.loads(body) == {"k": "v"} + + def test_compact_json_format(self): + body, _ = RequestBodySerializer.serialize( + {"a": 1, "b": 2}, ContentType.JSON + ) + # Should use compact separators + assert " " not in body + + def test_unicode_json(self): + body, _ = RequestBodySerializer.serialize( + {"name": "Heisenberg"}, ContentType.JSON + ) + parsed = json.loads(body) + assert parsed["name"] == "Heisenberg" + + +# ===================================================================== +# Form URL-Encoded Serialization +# ===================================================================== + + +@pytest.mark.unit +class TestSerializeFormUrlencoded: + + def test_basic_form(self): + body, headers = RequestBodySerializer.serialize( + {"name": "Alice", "age": "30"}, ContentType.FORM_URLENCODED + ) + assert "name=Alice" in body + assert "age=30" in body + assert headers["Content-Type"] == "application/x-www-form-urlencoded" + + def test_none_values_skipped(self): + body, headers = RequestBodySerializer.serialize( + {"name": "Alice", "skip": None}, ContentType.FORM_URLENCODED + ) + assert "name=Alice" in body + assert "skip" not in body + + def test_list_explode_true(self): + body, headers = RequestBodySerializer.serialize( + {"tags": ["a", "b"]}, + ContentType.FORM_URLENCODED, + encoding_rules={"tags": {"style": "form", "explode": True}}, + ) + assert "tags=a" in body + assert "tags=b" in body + + def test_list_explode_false(self): + body, headers = RequestBodySerializer.serialize( + {"tags": ["a", "b"]}, + ContentType.FORM_URLENCODED, + encoding_rules={"tags": {"style": "form", "explode": False}}, + ) + assert "tags=" in body + assert "a" in body and "b" in body + + def test_dict_value_json_content_type(self): + body, headers = RequestBodySerializer.serialize( + {"metadata": {"key": "val"}}, + ContentType.FORM_URLENCODED, + encoding_rules={"metadata": {"contentType": "application/json"}}, + ) + assert "metadata" in body + + def test_dict_value_xml_content_type(self): + body, headers = RequestBodySerializer.serialize( + {"data": {"name": "test"}}, + ContentType.FORM_URLENCODED, + encoding_rules={"data": {"contentType": "application/xml"}}, + ) + assert "data" in body + + def test_dict_value_deep_object_explode(self): + body, headers = RequestBodySerializer.serialize( + {"filter": {"status": "active", "type": "doc"}}, + ContentType.FORM_URLENCODED, + encoding_rules={ + "filter": {"style": "deepObject", "explode": True} + }, + ) + assert "filter" in body + + def test_dict_value_non_exploded(self): + body, headers = RequestBodySerializer.serialize( + {"obj": {"a": "1", "b": "2"}}, + ContentType.FORM_URLENCODED, + encoding_rules={"obj": {"style": "form", "explode": False}}, + ) + assert "obj" in body + + def test_default_explode_for_form_style(self): + """Default explode should be True when style is 'form'.""" + body, headers = RequestBodySerializer.serialize( + {"items": ["x", "y"]}, + ContentType.FORM_URLENCODED, + encoding_rules={"items": {"style": "form"}}, + ) + # explode defaults to True for form style => separate params + assert "items=x" in body + assert "items=y" in body + + def test_special_characters_encoded(self): + body, _ = RequestBodySerializer.serialize( + {"q": "hello world&more"}, ContentType.FORM_URLENCODED + ) + assert "hello" in body + assert "q=" in body + + +# ===================================================================== +# Text Plain Serialization +# ===================================================================== + + +@pytest.mark.unit +class TestSerializeTextPlain: + + def test_single_value(self): + body, headers = RequestBodySerializer.serialize( + {"message": "hello"}, ContentType.TEXT_PLAIN + ) + assert body == "hello" + assert headers["Content-Type"] == "text/plain" + + def test_multiple_values(self): + body, headers = RequestBodySerializer.serialize( + {"name": "Alice", "age": 30}, ContentType.TEXT_PLAIN + ) + assert "name: Alice" in body + assert "age: 30" in body + + +# ===================================================================== +# XML Serialization +# ===================================================================== + + +@pytest.mark.unit +class TestSerializeXml: + + def test_basic_xml(self): + body, headers = RequestBodySerializer.serialize( + {"name": "Alice"}, ContentType.XML + ) + assert 'Alice" in body + assert headers["Content-Type"] == "application/xml" + + def test_nested_xml(self): + body, headers = RequestBodySerializer.serialize( + {"user": {"name": "Alice"}}, ContentType.XML + ) + assert "" in body + assert "Alice" in body + + def test_xml_escapes_special_chars(self): + body, headers = RequestBodySerializer.serialize( + {"data": ""}, ContentType.XML + ) + assert "<script>" in body + + def test_xml_with_list(self): + body, _ = RequestBodySerializer.serialize( + {"items": [1, 2, 3]}, ContentType.XML + ) + assert "1" in body + assert "2" in body + assert "3" in body + + +# ===================================================================== +# Octet Stream Serialization +# ===================================================================== + + +@pytest.mark.unit +class TestSerializeOctetStream: + + def test_dict_body(self): + body, headers = RequestBodySerializer.serialize( + {"key": "val"}, ContentType.OCTET_STREAM + ) + assert isinstance(body, bytes) + assert headers["Content-Type"] == "application/octet-stream" + + def test_bytes_body(self): + body, headers = RequestBodySerializer._serialize_octet_stream(b"\x00\x01") + assert body == b"\x00\x01" + assert headers["Content-Type"] == "application/octet-stream" + + def test_string_body(self): + body, headers = RequestBodySerializer._serialize_octet_stream("hello") + assert body == b"hello" + + +# ===================================================================== +# Multipart Form Data Serialization +# ===================================================================== + + +@pytest.mark.unit +class TestSerializeMultipartFormData: + + def test_basic_multipart(self): + body, headers = RequestBodySerializer.serialize( + {"field": "value"}, ContentType.MULTIPART_FORM_DATA + ) + assert isinstance(body, bytes) + assert "multipart/form-data" in headers["Content-Type"] + assert "boundary=" in headers["Content-Type"] + + def test_none_values_skipped(self): + body, headers = RequestBodySerializer.serialize( + {"field": "value", "empty": None}, ContentType.MULTIPART_FORM_DATA + ) + body_str = body.decode("utf-8", errors="replace") + assert "field" in body_str + assert "empty" not in body_str + + def test_multipart_with_bytes(self): + body, headers = RequestBodySerializer.serialize( + {"file": b"\x00\x01\x02"}, ContentType.MULTIPART_FORM_DATA + ) + assert isinstance(body, bytes) + + def test_multipart_with_dict_json(self): + body, headers = RequestBodySerializer.serialize( + {"meta": {"key": "val"}}, + ContentType.MULTIPART_FORM_DATA, + encoding_rules={"meta": {"contentType": "application/json"}}, + ) + body_str = body.decode("utf-8", errors="replace") + assert "meta" in body_str + assert "application/json" in body_str + + def test_multipart_with_dict_xml(self): + body, headers = RequestBodySerializer.serialize( + {"data": {"name": "test"}}, + ContentType.MULTIPART_FORM_DATA, + encoding_rules={"data": {"contentType": "application/xml"}}, + ) + body_str = body.decode("utf-8", errors="replace") + assert "data" in body_str + + def test_multipart_octet_stream_bytes(self): + body, headers = RequestBodySerializer.serialize( + {"bin": b"\xff\xfe"}, + ContentType.MULTIPART_FORM_DATA, + encoding_rules={"bin": {"contentType": "application/octet-stream"}}, + ) + body_str = body.decode("utf-8", errors="replace") + assert "bin" in body_str + assert "Content-Transfer-Encoding: base64" in body_str + + def test_multipart_string_with_json_content_type(self): + body, headers = RequestBodySerializer.serialize( + {"json_str": '{"a": 1}'}, + ContentType.MULTIPART_FORM_DATA, + encoding_rules={"json_str": {"contentType": "application/json"}}, + ) + body_str = body.decode("utf-8", errors="replace") + assert "json_str" in body_str + + def test_multipart_string_with_non_text_content_type(self): + body, headers = RequestBodySerializer.serialize( + {"custom": "data"}, + ContentType.MULTIPART_FORM_DATA, + encoding_rules={"custom": {"contentType": "application/custom"}}, + ) + body_str = body.decode("utf-8", errors="replace") + assert "custom" in body_str + + +# ===================================================================== +# Helper Methods +# ===================================================================== + + +@pytest.mark.unit +class TestHelpers: + + def test_percent_encode_space(self): + assert RequestBodySerializer._percent_encode("hello world") == "hello%20world" + + def test_percent_encode_slash(self): + assert RequestBodySerializer._percent_encode("a/b") == "a%2Fb" + + def test_percent_encode_safe_chars(self): + assert RequestBodySerializer._percent_encode("a/b", safe_chars="/") == "a/b" + + def test_escape_xml_ampersand(self): + assert "&" in RequestBodySerializer._escape_xml("&") + + def test_escape_xml_lt(self): + assert "<" in RequestBodySerializer._escape_xml("<") + + def test_escape_xml_gt(self): + assert ">" in RequestBodySerializer._escape_xml(">") + + def test_escape_xml_quote(self): + assert """ in RequestBodySerializer._escape_xml('"') + + def test_escape_xml_apos(self): + assert "'" in RequestBodySerializer._escape_xml("'") + + def test_dict_to_xml_list(self): + xml = RequestBodySerializer._dict_to_xml({"items": [1, 2, 3]}) + assert "1" in xml + assert "2" in xml + + def test_dict_to_xml_custom_root(self): + xml = RequestBodySerializer._dict_to_xml({"key": "val"}, root_name="data") + assert "" in xml + assert "val" in xml + + def test_dict_to_xml_deeply_nested(self): + xml = RequestBodySerializer._dict_to_xml({"a": {"b": {"c": "deep"}}}) + assert "deep" in xml + + +# ===================================================================== +# Error Handling +# ===================================================================== + + +@pytest.mark.unit +class TestSerializationErrors: + + def test_serialize_raises_on_internal_error(self): + """Test that serialization errors are wrapped in ValueError.""" + # Patch _serialize_json to raise + with pytest.raises(ValueError, match="Failed to serialize"): + RequestBodySerializer.serialize( + {"key": object()}, # object() is not JSON-serializable + ContentType.JSON, + ) diff --git a/tests/agents/tools/test_api_tool.py b/tests/agents/tools/test_api_tool.py new file mode 100644 index 00000000..31e46d6e --- /dev/null +++ b/tests/agents/tools/test_api_tool.py @@ -0,0 +1,516 @@ +"""Comprehensive tests for application/agents/tools/api_tool.py + +Covers: APITool initialization, all HTTP methods, path param substitution, +SSRF validation, error handling, response parsing, body serialization. +""" + +import json +from unittest.mock import MagicMock, patch + +import pytest +import requests + +from application.agents.tools.api_tool import APITool, DEFAULT_TIMEOUT + + +@pytest.fixture +def get_tool(): + return APITool( + config={ + "url": "https://api.example.com/data", + "method": "GET", + "headers": {"Accept": "application/json"}, + "query_params": {}, + } + ) + + +@pytest.fixture +def post_tool(): + return APITool( + config={ + "url": "https://api.example.com/items", + "method": "POST", + "headers": {}, + "query_params": {}, + } + ) + + +# ===================================================================== +# Initialization +# ===================================================================== + + +@pytest.mark.unit +class TestAPIToolInit: + + def test_default_values(self): + tool = APITool(config={}) + assert tool.url == "" + assert tool.method == "GET" + assert tool.headers == {} + assert tool.query_params == {} + assert tool.body_content_type == "application/json" + assert tool.body_encoding_rules == {} + + def test_custom_config(self): + tool = APITool(config={ + "url": "https://api.test.com", + "method": "POST", + "headers": {"X-Key": "val"}, + "query_params": {"page": "1"}, + "body_content_type": "application/xml", + "body_encoding_rules": {"field": {"style": "form"}}, + }) + assert tool.url == "https://api.test.com" + assert tool.method == "POST" + assert tool.headers == {"X-Key": "val"} + assert tool.query_params == {"page": "1"} + assert tool.body_content_type == "application/xml" + + def test_default_timeout_constant(self): + assert DEFAULT_TIMEOUT == 90 + + +# ===================================================================== +# HTTP Methods +# ===================================================================== + + +@pytest.mark.unit +class TestMakeApiCall: + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_successful_get(self, mock_get, mock_validate, get_tool): + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = {"result": "ok"} + mock_resp.content = b'{"result":"ok"}' + mock_get.return_value = mock_resp + + result = get_tool.execute_action("any_action") + + assert result["status_code"] == 200 + assert result["data"] == {"result": "ok"} + assert result["message"] == "API call successful." + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.post") + def test_successful_post(self, mock_post, mock_validate, post_tool): + mock_resp = MagicMock() + mock_resp.status_code = 201 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = {"id": 1} + mock_resp.content = b'{"id":1}' + mock_post.return_value = mock_resp + + result = post_tool.execute_action("create", name="test") + assert result["status_code"] == 201 + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.put") + def test_put_method(self, mock_put, mock_validate): + tool = APITool(config={"url": "https://example.com/item/1", "method": "PUT"}) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = {} + mock_resp.content = b'{}' + mock_put.return_value = mock_resp + + result = tool.execute_action("update", name="new") + assert result["status_code"] == 200 + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.delete") + def test_delete_method(self, mock_delete, mock_validate): + tool = APITool(config={"url": "https://example.com/item/1", "method": "DELETE"}) + mock_resp = MagicMock() + mock_resp.status_code = 204 + mock_resp.headers = {"Content-Type": "text/plain"} + mock_resp.content = b'' + mock_delete.return_value = mock_resp + + result = tool.execute_action("delete") + assert result["status_code"] == 204 + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.patch") + def test_patch_method(self, mock_patch, mock_validate): + tool = APITool(config={"url": "https://example.com/item/1", "method": "PATCH"}) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = {"patched": True} + mock_resp.content = b'{"patched":true}' + mock_patch.return_value = mock_resp + + result = tool.execute_action("patch", field="val") + assert result["status_code"] == 200 + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.head") + def test_head_method(self, mock_head, mock_validate): + tool = APITool(config={"url": "https://example.com", "method": "HEAD"}) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "text/html"} + mock_resp.content = b'' + mock_head.return_value = mock_resp + + result = tool.execute_action("check") + assert result["status_code"] == 200 + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.options") + def test_options_method(self, mock_options, mock_validate): + tool = APITool(config={"url": "https://example.com", "method": "OPTIONS"}) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "text/plain"} + mock_resp.content = b'' + mock_options.return_value = mock_resp + + result = tool.execute_action("options") + assert result["status_code"] == 200 + + @patch("application.agents.tools.api_tool.validate_url") + def test_unsupported_method(self, mock_validate): + tool = APITool(config={"url": "https://example.com", "method": "CUSTOM"}) + result = tool.execute_action("any") + assert result["status_code"] is None + assert "Unsupported" in result["message"] + + +# ===================================================================== +# SSRF Validation +# ===================================================================== + + +@pytest.mark.unit +class TestSSRFValidation: + + @patch("application.agents.tools.api_tool.validate_url") + def test_ssrf_blocked_initial_url(self, mock_validate, get_tool): + from application.core.url_validation import SSRFError + + mock_validate.side_effect = SSRFError("blocked") + result = get_tool.execute_action("any") + assert result["status_code"] is None + assert "URL validation error" in result["message"] + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_ssrf_blocked_after_param_substitution(self, mock_get, mock_validate): + from application.core.url_validation import SSRFError + + tool = APITool(config={ + "url": "https://api.example.com/{host}/data", + "method": "GET", + "query_params": {"host": "169.254.169.254"}, + }) + + call_count = [0] + + def side_effect(url): + call_count[0] += 1 + if call_count[0] == 2: + raise SSRFError("blocked after substitution") + + mock_validate.side_effect = side_effect + result = tool.execute_action("any") + assert result["status_code"] is None + assert "URL validation error" in result["message"] + + +# ===================================================================== +# Error Handling +# ===================================================================== + + +@pytest.mark.unit +class TestErrorHandling: + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_timeout_error(self, mock_get, mock_validate, get_tool): + mock_get.side_effect = requests.exceptions.Timeout() + result = get_tool.execute_action("any") + assert result["status_code"] is None + assert "timeout" in result["message"].lower() + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_connection_error(self, mock_get, mock_validate, get_tool): + mock_get.side_effect = requests.exceptions.ConnectionError("refused") + result = get_tool.execute_action("any") + assert result["status_code"] is None + assert "Connection error" in result["message"] + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_http_error_with_json(self, mock_get, mock_validate, get_tool): + mock_resp = MagicMock() + mock_resp.status_code = 422 + mock_resp.json.return_value = {"error": "invalid_field"} + mock_resp.raise_for_status.side_effect = requests.exceptions.HTTPError( + response=mock_resp + ) + mock_get.return_value = mock_resp + + result = get_tool.execute_action("any") + assert result["status_code"] == 422 + assert result["data"] == {"error": "invalid_field"} + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_http_error_non_json_body(self, mock_get, mock_validate, get_tool): + mock_resp = MagicMock() + mock_resp.status_code = 404 + mock_resp.text = "Not Found" + mock_resp.json.side_effect = json.JSONDecodeError("", "", 0) + mock_resp.raise_for_status.side_effect = requests.exceptions.HTTPError( + response=mock_resp + ) + mock_get.return_value = mock_resp + + result = get_tool.execute_action("any") + assert result["status_code"] == 404 + assert result["data"] == "Not Found" + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_request_exception(self, mock_get, mock_validate, get_tool): + mock_get.side_effect = requests.exceptions.RequestException("something") + result = get_tool.execute_action("any") + assert "API call failed" in result["message"] + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_unexpected_exception(self, mock_get, mock_validate, get_tool): + mock_get.side_effect = RuntimeError("unexpected") + result = get_tool.execute_action("any") + assert "Unexpected error" in result["message"] + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.post") + def test_body_serialization_error(self, mock_post, mock_validate): + tool = APITool(config={ + "url": "https://example.com", + "method": "POST", + "body_content_type": "application/json", + }) + + with patch( + "application.agents.tools.api_tool.RequestBodySerializer.serialize", + side_effect=ValueError("serialize fail"), + ): + result = tool.execute_action("any", key="val") + assert "serialization error" in result["message"].lower() + + +# ===================================================================== +# Path Param Substitution +# ===================================================================== + + +@pytest.mark.unit +class TestPathParamSubstitution: + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_path_params_substituted(self, mock_get, mock_validate): + tool = APITool(config={ + "url": "https://api.example.com/users/{user_id}/posts/{post_id}", + "method": "GET", + "query_params": {"user_id": "42", "post_id": "7"}, + }) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = [] + mock_resp.content = b'[]' + mock_get.return_value = mock_resp + + tool.execute_action("get") + + called_url = mock_get.call_args[0][0] + assert "/users/42/posts/7" in called_url + assert "{user_id}" not in called_url + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_remaining_query_params_appended(self, mock_get, mock_validate): + tool = APITool(config={ + "url": "https://api.example.com/items", + "method": "GET", + "query_params": {"page": "2", "limit": "10"}, + }) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = [] + mock_resp.content = b'[]' + mock_get.return_value = mock_resp + + tool.execute_action("get") + + called_url = mock_get.call_args[0][0] + assert "page=2" in called_url + assert "limit=10" in called_url + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.get") + def test_query_params_append_with_existing_query_string( + self, mock_get, mock_validate + ): + tool = APITool(config={ + "url": "https://api.example.com/items?existing=true", + "method": "GET", + "query_params": {"page": "1"}, + }) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = [] + mock_resp.content = b'[]' + mock_get.return_value = mock_resp + + tool.execute_action("get") + + called_url = mock_get.call_args[0][0] + assert "&page=1" in called_url + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.post") + def test_empty_body_no_serialization(self, mock_post, mock_validate): + tool = APITool(config={"url": "https://example.com", "method": "POST"}) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = {} + mock_resp.content = b'{}' + mock_post.return_value = mock_resp + + result = tool.execute_action("create") + assert result["status_code"] == 200 + + +# ===================================================================== +# Parse Response +# ===================================================================== + + +@pytest.mark.unit +class TestParseResponse: + + def test_json_response(self, get_tool): + mock_resp = MagicMock() + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = {"key": "val"} + mock_resp.content = b'{"key":"val"}' + + result = get_tool._parse_response(mock_resp) + assert result == {"key": "val"} + + def test_json_decode_error_falls_back_to_text(self, get_tool): + mock_resp = MagicMock() + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.side_effect = json.JSONDecodeError("", "", 0) + mock_resp.text = "not valid json" + mock_resp.content = b"not valid json" + + result = get_tool._parse_response(mock_resp) + assert result == "not valid json" + + def test_text_response(self, get_tool): + mock_resp = MagicMock() + mock_resp.headers = {"Content-Type": "text/plain"} + mock_resp.text = "plain text" + mock_resp.content = b"plain text" + + result = get_tool._parse_response(mock_resp) + assert result == "plain text" + + def test_xml_response(self, get_tool): + mock_resp = MagicMock() + mock_resp.headers = {"Content-Type": "application/xml"} + mock_resp.text = "1" + mock_resp.content = b"1" + + result = get_tool._parse_response(mock_resp) + assert "" in result + + def test_html_response(self, get_tool): + mock_resp = MagicMock() + mock_resp.headers = {"Content-Type": "text/html"} + mock_resp.text = "Hi" + mock_resp.content = b"Hi" + + result = get_tool._parse_response(mock_resp) + assert "" in result + + def test_empty_content(self, get_tool): + mock_resp = MagicMock() + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.content = b"" + + result = get_tool._parse_response(mock_resp) + assert result is None + + def test_binary_response(self, get_tool): + mock_resp = MagicMock() + mock_resp.headers = {"Content-Type": "application/octet-stream"} + mock_resp.text = "binary_text" + mock_resp.content = b"\x00\x01\x02" + + result = get_tool._parse_response(mock_resp) + assert result is not None + + def test_text_xml_content_type(self, get_tool): + mock_resp = MagicMock() + mock_resp.headers = {"Content-Type": "text/xml"} + mock_resp.text = "" + mock_resp.content = b"" + + result = get_tool._parse_response(mock_resp) + assert result == "" + + +# ===================================================================== +# Metadata +# ===================================================================== + + +@pytest.mark.unit +class TestAPIToolMetadata: + + def test_actions_metadata_empty(self, get_tool): + assert get_tool.get_actions_metadata() == [] + + def test_config_requirements_empty(self, get_tool): + assert get_tool.get_config_requirements() == {} + + @patch("application.agents.tools.api_tool.validate_url") + @patch("application.agents.tools.api_tool.requests.post") + def test_content_type_set_for_post_with_no_headers( + self, mock_post, mock_validate + ): + tool = APITool(config={ + "url": "https://example.com", + "method": "POST", + "headers": {}, + }) + mock_resp = MagicMock() + mock_resp.status_code = 200 + mock_resp.headers = {"Content-Type": "application/json"} + mock_resp.json.return_value = {} + mock_resp.content = b'{}' + mock_post.return_value = mock_resp + + tool.execute_action("create") + call_headers = mock_post.call_args[1]["headers"] + assert "Content-Type" in call_headers diff --git a/tests/agents/tools/test_internal_search.py b/tests/agents/tools/test_internal_search.py new file mode 100644 index 00000000..d48698ef --- /dev/null +++ b/tests/agents/tools/test_internal_search.py @@ -0,0 +1,596 @@ +"""Comprehensive tests for application/agents/tools/internal_search.py + +Covers: InternalSearchTool (search, list_files, path_filter, error handling, +directory structure loading), build helpers, add_internal_search_tool, +sources_have_directory_structure. +""" + +import json +from unittest.mock import MagicMock, Mock, patch + +import pytest + +from application.agents.tools.internal_search import ( + INTERNAL_TOOL_ENTRY, + INTERNAL_TOOL_ID, + InternalSearchTool, + add_internal_search_tool, + build_internal_tool_config, + build_internal_tool_entry, + sources_have_directory_structure, +) + + +# ===================================================================== +# InternalSearchTool - Search +# ===================================================================== + + +def _make_tool(**config_overrides): + config = {"source": {}, "retriever_name": "classic", "chunks": 2} + config.update(config_overrides) + return InternalSearchTool(config) + + +@pytest.mark.unit +class TestInternalSearchToolSearch: + + def test_search_no_query_returns_error(self): + tool = _make_tool() + result = tool.execute_action("search", query="") + assert "required" in result.lower() + + def test_search_returns_formatted_docs(self): + tool = _make_tool() + mock_retriever = Mock() + mock_retriever.search.return_value = [ + { + "text": "Hello world", + "title": "Doc1", + "source": "test", + "filename": "doc1.md", + }, + ] + tool._retriever = mock_retriever + + result = tool.execute_action("search", query="hello") + assert "doc1.md" in result + assert "Hello world" in result + assert len(tool.retrieved_docs) == 1 + + def test_search_no_results(self): + tool = _make_tool() + mock_retriever = Mock() + mock_retriever.search.return_value = [] + tool._retriever = mock_retriever + + result = tool.execute_action("search", query="nonexistent") + assert "No documents found" in result + + def test_search_accumulates_docs(self): + tool = _make_tool() + mock_retriever = Mock() + tool._retriever = mock_retriever + + mock_retriever.search.return_value = [ + {"text": "A", "title": "D1", "source": "s1"}, + ] + tool.execute_action("search", query="first") + + mock_retriever.search.return_value = [ + {"text": "B", "title": "D2", "source": "s2"}, + ] + tool.execute_action("search", query="second") + + assert len(tool.retrieved_docs) == 2 + + def test_search_deduplicates_docs(self): + tool = _make_tool() + doc = {"text": "Same", "title": "Same", "source": "same"} + mock_retriever = Mock() + mock_retriever.search.return_value = [doc] + tool._retriever = mock_retriever + + tool.execute_action("search", query="q1") + tool.execute_action("search", query="q2") + + assert len(tool.retrieved_docs) == 1 + + def test_search_with_path_filter(self): + tool = _make_tool() + mock_retriever = Mock() + mock_retriever.search.return_value = [ + {"text": "A", "title": "T", "source": "src/main.py", "filename": "main.py"}, + {"text": "B", "title": "T", "source": "docs/readme.md", "filename": "readme.md"}, + ] + tool._retriever = mock_retriever + + result = tool.execute_action("search", query="code", path_filter="src/") + assert "main.py" in result + assert "readme.md" not in result + + def test_search_path_filter_matches_title(self): + tool = _make_tool() + mock_retriever = Mock() + mock_retriever.search.return_value = [ + {"text": "A", "title": "src/main.py", "source": "other", "filename": ""}, + ] + tool._retriever = mock_retriever + + result = tool.execute_action("search", query="code", path_filter="src/main") + assert "src/main.py" in result + + def test_search_path_filter_no_match(self): + tool = _make_tool() + mock_retriever = Mock() + mock_retriever.search.return_value = [ + {"text": "A", "title": "T", "source": "other/file.txt"}, + ] + tool._retriever = mock_retriever + + result = tool.execute_action("search", query="code", path_filter="src/") + assert "No documents found" in result + + def test_search_retriever_error(self): + tool = _make_tool() + mock_retriever = Mock() + mock_retriever.search.side_effect = Exception("Connection error") + tool._retriever = mock_retriever + + result = tool.execute_action("search", query="test") + assert "failed" in result.lower() or "error" in result.lower() + + def test_unknown_action(self): + tool = _make_tool() + result = tool.execute_action("nonexistent") + assert "Unknown action" in result + + def test_search_formats_with_separator(self): + tool = _make_tool() + mock_retriever = Mock() + mock_retriever.search.return_value = [ + {"text": "A", "title": "D1", "source": "s1", "filename": "f1.md"}, + {"text": "B", "title": "D2", "source": "s2", "filename": "f2.md"}, + ] + tool._retriever = mock_retriever + + result = tool.execute_action("search", query="test") + assert "---" in result + assert "[1]" in result + assert "[2]" in result + + def test_search_uses_title_when_no_filename(self): + tool = _make_tool() + mock_retriever = Mock() + mock_retriever.search.return_value = [ + {"text": "Content", "title": "My Title", "source": "src", "filename": ""}, + ] + tool._retriever = mock_retriever + + result = tool.execute_action("search", query="q") + assert "My Title" in result + + +# ===================================================================== +# InternalSearchTool - List Files +# ===================================================================== + + +@pytest.mark.unit +class TestInternalSearchToolListFiles: + + def test_list_files_no_structure(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = None + + result = tool.execute_action("list_files") + assert "No file structure" in result + + def test_list_files_root(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = { + "src": {"main.py": {}}, + "README.md": {"type": "md", "token_count": 100}, + } + + result = tool.execute_action("list_files") + assert "src/" in result + assert "README.md" in result + + def test_list_files_nested_path(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = { + "src": { + "utils": {"helper.py": {}}, + }, + } + + result = tool.execute_action("list_files", path="src") + assert "utils/" in result + + def test_list_files_invalid_path(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = {"src": {}} + + result = tool.execute_action("list_files", path="nonexistent") + assert "not found" in result + + def test_list_files_empty_directory(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = {"empty_dir": {}} + + result = tool.execute_action("list_files", path="empty_dir") + assert "(empty)" in result + + def test_list_files_file_with_metadata(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = { + "data.csv": { + "type": "text/csv", + "size_bytes": 1024, + "token_count": 500, + }, + } + + result = tool.execute_action("list_files") + assert "data.csv" in result + assert "500 tokens" in result + assert "text/csv" in result + + def test_list_files_file_is_not_directory(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = { + "src": { + "main.py": "plain_file_value", + }, + } + + result = tool.execute_action("list_files", path="src/main.py") + assert "is a file" in result + + def test_list_files_deep_nested_path_with_slashes(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = { + "a": {"b": {"c": {"file.txt": {"type": "text"}}}}, + } + + result = tool.execute_action("list_files", path="a/b/c") + assert "file.txt" in result + + +# ===================================================================== +# Count Files Helper +# ===================================================================== + + +@pytest.mark.unit +class TestCountFiles: + + def test_count_files_nested(self): + tool = InternalSearchTool({"source": {}}) + node = { + "file1.txt": {"type": "text"}, + "dir": { + "file2.txt": {"type": "text"}, + "file3.txt": "plain_value", + }, + } + assert tool._count_files(node) == 3 + + def test_count_files_empty(self): + tool = InternalSearchTool({"source": {}}) + assert tool._count_files({}) == 0 + + +# ===================================================================== +# Directory Structure Loading +# ===================================================================== + + +@pytest.mark.unit +class TestGetDirectoryStructure: + + def test_loads_from_mongo(self): + tool = InternalSearchTool({ + "source": {"active_docs": ["507f1f77bcf86cd799439011"]}, + }) + + mock_collection = MagicMock() + mock_collection.find_one.return_value = { + "_id": "507f1f77bcf86cd799439011", + "name": "test_source", + "directory_structure": {"src": {"main.py": {}}}, + } + + with patch("application.core.mongo_db.MongoDB") as mock_mongo: + mock_db = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + mock_client = MagicMock() + mock_client.__getitem__ = MagicMock(return_value=mock_db) + mock_mongo.get_client.return_value = mock_client + + result = tool._get_directory_structure() + assert result is not None + assert "src" in result + + def test_returns_none_without_active_docs(self): + tool = InternalSearchTool({"source": {}}) + result = tool._get_directory_structure() + assert result is None + assert tool._dir_structure_loaded is True + + def test_caches_after_first_load(self): + tool = InternalSearchTool({"source": {}}) + tool._dir_structure_loaded = True + tool._directory_structure = {"cached": True} + + result = tool._get_directory_structure() + assert result == {"cached": True} + + def test_handles_json_string_structure(self): + tool = InternalSearchTool({ + "source": {"active_docs": ["507f1f77bcf86cd799439011"]}, + }) + + mock_collection = MagicMock() + mock_collection.find_one.return_value = { + "_id": "507f1f77bcf86cd799439011", + "name": "test_source", + "directory_structure": json.dumps({"src": {"app.py": {}}}), + } + + with patch("application.core.mongo_db.MongoDB") as mock_mongo: + mock_db = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + mock_client = MagicMock() + mock_client.__getitem__ = MagicMock(return_value=mock_db) + mock_mongo.get_client.return_value = mock_client + + result = tool._get_directory_structure() + assert result is not None + assert "src" in result + + def test_handles_string_active_docs(self): + tool = InternalSearchTool({ + "source": {"active_docs": "507f1f77bcf86cd799439011"}, + }) + + mock_collection = MagicMock() + mock_collection.find_one.return_value = { + "_id": "507f1f77bcf86cd799439011", + "directory_structure": {"dir": {}}, + } + + with patch("application.core.mongo_db.MongoDB") as mock_mongo: + mock_db = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + mock_client = MagicMock() + mock_client.__getitem__ = MagicMock(return_value=mock_db) + mock_mongo.get_client.return_value = mock_client + + result = tool._get_directory_structure() + assert result is not None + + def test_merges_multiple_sources(self): + tool = InternalSearchTool({ + "source": { + "active_docs": [ + "507f1f77bcf86cd799439011", + "507f1f77bcf86cd799439012", + ], + }, + }) + + mock_collection = MagicMock() + mock_collection.find_one.side_effect = [ + {"name": "src1", "directory_structure": {"a": {}}}, + {"name": "src2", "directory_structure": {"b": {}}}, + ] + + with patch("application.core.mongo_db.MongoDB") as mock_mongo: + mock_db = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + mock_client = MagicMock() + mock_client.__getitem__ = MagicMock(return_value=mock_db) + mock_mongo.get_client.return_value = mock_client + + result = tool._get_directory_structure() + assert "src1" in result + assert "src2" in result + + +# ===================================================================== +# Metadata +# ===================================================================== + + +@pytest.mark.unit +class TestInternalSearchToolMetadata: + + def test_actions_without_directory_structure(self): + tool = InternalSearchTool({"has_directory_structure": False}) + meta = tool.get_actions_metadata() + + action_names = [a["name"] for a in meta] + assert "search" in action_names + assert "list_files" not in action_names + + search = meta[0] + assert "path_filter" not in search["parameters"]["properties"] + + def test_actions_with_directory_structure(self): + tool = InternalSearchTool({"has_directory_structure": True}) + meta = tool.get_actions_metadata() + + action_names = [a["name"] for a in meta] + assert "search" in action_names + assert "list_files" in action_names + + search = next(a for a in meta if a["name"] == "search") + assert "path_filter" in search["parameters"]["properties"] + + def test_config_requirements_empty(self): + tool = InternalSearchTool({}) + assert tool.get_config_requirements() == {} + + +# ===================================================================== +# Build Helpers +# ===================================================================== + + +@pytest.mark.unit +class TestBuildHelpers: + + def test_build_entry_without_directory_structure(self): + entry = build_internal_tool_entry(has_directory_structure=False) + assert entry["name"] == "internal_search" + action_names = [a["name"] for a in entry["actions"]] + assert "search" in action_names + assert "list_files" not in action_names + assert entry["actions"][0].get("active") is True + + def test_build_entry_with_directory_structure(self): + entry = build_internal_tool_entry(has_directory_structure=True) + action_names = [a["name"] for a in entry["actions"]] + assert "list_files" in action_names + # path_filter should be in search params + search_action = next(a for a in entry["actions"] if a["name"] == "search") + assert "path_filter" in search_action["parameters"]["properties"] + + def test_build_config(self): + config = build_internal_tool_config( + source={"active_docs": ["abc"]}, + retriever_name="semantic", + chunks=4, + ) + assert config["source"] == {"active_docs": ["abc"]} + assert config["retriever_name"] == "semantic" + assert config["chunks"] == 4 + + def test_build_config_defaults(self): + config = build_internal_tool_config(source={"active_docs": ["abc"]}) + assert config["retriever_name"] == "classic" + assert config["chunks"] == 2 + assert config["doc_token_limit"] == 50000 + + def test_internal_tool_id(self): + assert INTERNAL_TOOL_ID == "internal" + + def test_internal_tool_entry_constant(self): + assert INTERNAL_TOOL_ENTRY["name"] == "internal_search" + + def test_add_internal_search_tool_with_sources(self): + tools_dict = {} + retriever_config = { + "source": {"active_docs": ["abc"]}, + "retriever_name": "classic", + "chunks": 2, + "model_id": "gpt-4", + "llm_name": "openai", + "api_key": "key", + } + + with patch( + "application.agents.tools.internal_search.sources_have_directory_structure", + return_value=False, + ): + add_internal_search_tool(tools_dict, retriever_config) + + assert INTERNAL_TOOL_ID in tools_dict + assert tools_dict[INTERNAL_TOOL_ID]["name"] == "internal_search" + assert "config" in tools_dict[INTERNAL_TOOL_ID] + + def test_add_internal_search_tool_no_sources(self): + tools_dict = {} + retriever_config = {"source": {}} + + add_internal_search_tool(tools_dict, retriever_config) + assert INTERNAL_TOOL_ID not in tools_dict + + def test_add_internal_search_tool_empty_config(self): + tools_dict = {} + add_internal_search_tool(tools_dict, {}) + assert INTERNAL_TOOL_ID not in tools_dict + + +# ===================================================================== +# sources_have_directory_structure +# ===================================================================== + + +@pytest.mark.unit +class TestSourcesHaveDirectoryStructure: + + def test_no_active_docs(self): + assert sources_have_directory_structure({}) is False + + def test_with_directory_structure(self): + mock_collection = MagicMock() + mock_collection.find_one.return_value = { + "directory_structure": {"src": {}}, + } + + with patch("application.core.mongo_db.MongoDB") as mock_mongo: + mock_db = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + mock_client = MagicMock() + mock_client.__getitem__ = MagicMock(return_value=mock_db) + mock_mongo.get_client.return_value = mock_client + + result = sources_have_directory_structure( + {"active_docs": ["507f1f77bcf86cd799439011"]} + ) + assert result is True + + def test_without_directory_structure(self): + mock_collection = MagicMock() + mock_collection.find_one.return_value = {"directory_structure": None} + + with patch("application.core.mongo_db.MongoDB") as mock_mongo: + mock_db = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + mock_client = MagicMock() + mock_client.__getitem__ = MagicMock(return_value=mock_db) + mock_mongo.get_client.return_value = mock_client + + result = sources_have_directory_structure( + {"active_docs": ["507f1f77bcf86cd799439011"]} + ) + assert result is False + + def test_handles_exception_gracefully(self): + with patch( + "application.core.mongo_db.MongoDB.get_client", + side_effect=Exception("DB down"), + ): + result = sources_have_directory_structure( + {"active_docs": ["507f1f77bcf86cd799439011"]} + ) + assert result is False + + def test_string_active_docs(self): + mock_collection = MagicMock() + mock_collection.find_one.return_value = { + "directory_structure": {"a": {}}, + } + + with patch("application.core.mongo_db.MongoDB") as mock_mongo: + mock_db = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + mock_client = MagicMock() + mock_client.__getitem__ = MagicMock(return_value=mock_db) + mock_mongo.get_client.return_value = mock_client + + result = sources_have_directory_structure( + {"active_docs": "507f1f77bcf86cd799439011"} + ) + assert result is True diff --git a/tests/agents/tools/test_mcp_tool.py b/tests/agents/tools/test_mcp_tool.py new file mode 100644 index 00000000..981486eb --- /dev/null +++ b/tests/agents/tools/test_mcp_tool.py @@ -0,0 +1,1010 @@ +"""Comprehensive tests for application/agents/tools/mcp_tool.py + +Covers: MCPTool init, cache key generation, transport creation, tool formatting, +result formatting, execute_action, discover_tools, test_connection, +get_actions_metadata, DocsGPTOAuth, NonInteractiveOAuth, DBTokenStorage, +MCPOAuthManager. +""" + +import asyncio +import concurrent.futures +from unittest.mock import MagicMock, patch + +import pytest + + +# ---- Fixtures to isolate module-level side effects ---- + +@pytest.fixture(autouse=True) +def _patch_mcp_globals(monkeypatch): + """Patch module-level MongoDB and cache to avoid real connections.""" + import sys + + if "application.agents.tools.mcp_tool" in sys.modules: + mcp_mod = sys.modules["application.agents.tools.mcp_tool"] + else: + mock_tasks = MagicMock() + monkeypatch.setitem(sys.modules, "application.api.user.tasks", mock_tasks) + import application.agents.tools.mcp_tool as mcp_mod + + mock_mongo = MagicMock() + mock_db = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=MagicMock()) + monkeypatch.setattr(mcp_mod, "mongo", mock_mongo) + monkeypatch.setattr(mcp_mod, "db", mock_db) + monkeypatch.setattr(mcp_mod, "_mcp_clients_cache", {}) + + +@pytest.fixture +def mcp_config(): + return { + "server_url": "https://mcp.example.com/api", + "transport_type": "http", + "auth_type": "none", + "timeout": 10, + } + + +@pytest.fixture +def bearer_config(): + return { + "server_url": "https://mcp.example.com/api", + "transport_type": "http", + "auth_type": "bearer", + "auth_credentials": {"bearer_token": "tok_123"}, + "timeout": 10, + } + + +def _make_tool(config, **kwargs): + from application.agents.tools.mcp_tool import MCPTool + + with patch.object(MCPTool, "_setup_client"): + return MCPTool(config, **kwargs) + + +# ===================================================================== +# MCPTool Initialization +# ===================================================================== + + +@pytest.mark.unit +class TestMCPToolInit: + + def test_basic_init(self, mcp_config): + tool = _make_tool(mcp_config) + assert tool.server_url == "https://mcp.example.com/api" + assert tool.transport_type == "http" + assert tool.auth_type == "none" + assert tool.timeout == 10 + assert tool.available_tools == [] + + def test_bearer_auth_credentials(self, bearer_config): + tool = _make_tool(bearer_config) + assert tool.auth_credentials["bearer_token"] == "tok_123" + + def test_no_server_url_skips_setup(self): + from application.agents.tools.mcp_tool import MCPTool + + with patch.object(MCPTool, "_setup_client") as mock_setup: + MCPTool({"server_url": "", "auth_type": "none"}) + mock_setup.assert_not_called() + + def test_oauth_skips_setup(self): + from application.agents.tools.mcp_tool import MCPTool + + with patch.object(MCPTool, "_setup_client") as mock_setup: + MCPTool({ + "server_url": "https://mcp.example.com", + "auth_type": "oauth", + }) + mock_setup.assert_not_called() + + def test_encrypted_credentials_decryption(self): + from application.agents.tools.mcp_tool import MCPTool + + with patch.object(MCPTool, "_setup_client"), \ + patch("application.agents.tools.mcp_tool.decrypt_credentials", + return_value={"bearer_token": "decrypted_tok"}): + tool = MCPTool( + { + "server_url": "https://mcp.example.com", + "auth_type": "bearer", + "encrypted_credentials": "enc_data", + }, + user_id="user1", + ) + assert tool.auth_credentials == {"bearer_token": "decrypted_tok"} + + def test_query_mode_default_false(self, mcp_config): + tool = _make_tool(mcp_config) + assert tool.query_mode is False + + def test_query_mode_explicit_true(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "none", + "query_mode": True, + }) + assert tool.query_mode is True + + def test_custom_headers_stored(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "none", + "headers": {"X-Custom": "val"}, + }) + assert tool.custom_headers == {"X-Custom": "val"} + + +# ===================================================================== +# Redirect URI Resolution +# ===================================================================== + + +@pytest.mark.unit +class TestResolveRedirectUri: + + def test_configured_redirect_uri_used(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "none", + "redirect_uri": "https://my.app/callback/", + }) + assert tool.redirect_uri == "https://my.app/callback" + + def test_fallback_to_settings(self, monkeypatch): + from application.core.settings import settings + + monkeypatch.setattr(settings, "API_URL", "https://api.docsgpt.co") + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "none", + }) + assert "/api/mcp_server/callback" in tool.redirect_uri + + +# ===================================================================== +# Cache Key Generation +# ===================================================================== + + +@pytest.mark.unit +class TestGenerateCacheKey: + + def test_none_auth(self, mcp_config): + tool = _make_tool(mcp_config) + assert "none" in tool._cache_key + assert "mcp.example.com" in tool._cache_key + + def test_bearer_auth(self, bearer_config): + tool = _make_tool(bearer_config) + assert "bearer:" in tool._cache_key + + def test_api_key_auth(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "api_key", + "auth_credentials": {"api_key": "sk-test12345678"}, + }) + assert "apikey:" in tool._cache_key + + def test_basic_auth(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "basic", + "auth_credentials": {"username": "user1", "password": "pass"}, + }) + assert "basic:user1" in tool._cache_key + + def test_oauth_auth_includes_scopes(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "oauth", + "oauth_scopes": ["read", "write"], + }) + assert "oauth:" in tool._cache_key + assert "read,write" in tool._cache_key + + def test_bearer_empty_token(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "bearer", + "auth_credentials": {}, + }) + assert "bearer:none" in tool._cache_key + + def test_api_key_empty(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "api_key", + "auth_credentials": {}, + }) + assert "apikey:none" in tool._cache_key + + +# ===================================================================== +# Transport Creation +# ===================================================================== + + +@pytest.mark.unit +class TestCreateTransport: + + def test_http_transport(self, mcp_config): + tool = _make_tool(mcp_config) + transport = tool._create_transport() + assert transport is not None + + def test_sse_transport(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com/sse", + "transport_type": "sse", + "auth_type": "none", + }) + transport = tool._create_transport() + assert transport is not None + + def test_auto_detects_sse(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com/sse", + "transport_type": "auto", + "auth_type": "none", + }) + transport = tool._create_transport() + assert transport is not None + + def test_auto_defaults_to_http(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com/api", + "transport_type": "auto", + "auth_type": "none", + }) + transport = tool._create_transport() + assert transport is not None + + def test_stdio_transport_disabled(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "transport_type": "stdio", + "auth_type": "none", + }) + with pytest.raises(ValueError, match="STDIO transport is disabled"): + tool._create_transport() + + def test_api_key_header_injected(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "transport_type": "http", + "auth_type": "api_key", + "auth_credentials": { + "api_key": "sk-test", + "api_key_header": "X-Custom-Key", + }, + }) + transport = tool._create_transport() + assert transport is not None + + def test_basic_auth_header_injected(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "transport_type": "http", + "auth_type": "basic", + "auth_credentials": {"username": "user", "password": "pass"}, + }) + transport = tool._create_transport() + assert transport is not None + + def test_unknown_transport_type_defaults_to_http(self): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "transport_type": "grpc", + "auth_type": "none", + }) + transport = tool._create_transport() + assert transport is not None + + +# ===================================================================== +# Format Tools +# ===================================================================== + + +@pytest.mark.unit +class TestFormatTools: + + def test_format_list_of_dicts(self, mcp_config): + tool = _make_tool(mcp_config) + result = tool._format_tools([{"name": "tool1", "description": "desc"}]) + assert len(result) == 1 + assert result[0]["name"] == "tool1" + + def test_format_tools_with_name_attribute(self, mcp_config): + tool = _make_tool(mcp_config) + mock_tool = MagicMock() + mock_tool.name = "my_tool" + mock_tool.description = "A tool" + mock_tool.inputSchema = {"type": "object", "properties": {}} + + result = tool._format_tools([mock_tool]) + assert len(result) == 1 + assert result[0]["name"] == "my_tool" + assert result[0]["inputSchema"] == {"type": "object", "properties": {}} + + def test_format_tools_with_model_dump(self, mcp_config): + tool = _make_tool(mcp_config) + mock_tool = MagicMock(spec=[]) + mock_tool.model_dump = MagicMock( + return_value={"name": "dumped", "description": "from dump"} + ) + # Ensure no "name" attribute + result = tool._format_tools([mock_tool]) + assert len(result) == 1 + assert result[0]["name"] == "dumped" + + def test_format_tools_fallback_str(self, mcp_config): + tool = _make_tool(mcp_config) + + class BareTool: + def __str__(self): + return "bare_tool" + + result = tool._format_tools([BareTool()]) + assert len(result) == 1 + assert result[0]["name"] == "bare_tool" + assert result[0]["description"] == "" + + def test_format_tools_response_object(self, mcp_config): + tool = _make_tool(mcp_config) + resp = MagicMock() + resp.tools = [{"name": "t1", "description": "d1"}] + + result = tool._format_tools(resp) + assert len(result) == 1 + + def test_format_tools_without_input_schema(self, mcp_config): + tool = _make_tool(mcp_config) + mock_tool = MagicMock() + mock_tool.name = "simple" + mock_tool.description = "no schema" + del mock_tool.inputSchema + + result = tool._format_tools([mock_tool]) + assert "inputSchema" not in result[0] + + def test_format_empty(self, mcp_config): + tool = _make_tool(mcp_config) + assert tool._format_tools([]) == [] + assert tool._format_tools("unexpected") == [] + + +# ===================================================================== +# Format Result +# ===================================================================== + + +@pytest.mark.unit +class TestFormatResult: + + def test_format_result_with_text_content(self, mcp_config): + tool = _make_tool(mcp_config) + mock_result = MagicMock() + text_item = MagicMock() + text_item.text = "Hello" + del text_item.data + mock_result.content = [text_item] + mock_result.isError = False + + result = tool._format_result(mock_result) + assert result["content"][0]["type"] == "text" + assert result["content"][0]["text"] == "Hello" + assert result["isError"] is False + + def test_format_result_with_data_content(self, mcp_config): + tool = _make_tool(mcp_config) + mock_result = MagicMock() + data_item = MagicMock() + del data_item.text + data_item.data = {"key": "value"} + mock_result.content = [data_item] + mock_result.isError = False + + result = tool._format_result(mock_result) + assert result["content"][0]["type"] == "data" + assert result["content"][0]["data"] == {"key": "value"} + + def test_format_result_unknown_content_type(self, mcp_config): + tool = _make_tool(mcp_config) + mock_result = MagicMock() + unknown_item = MagicMock() + del unknown_item.text + del unknown_item.data + mock_result.content = [unknown_item] + mock_result.isError = False + + result = tool._format_result(mock_result) + assert result["content"][0]["type"] == "unknown" + + def test_format_raw_result(self, mcp_config): + tool = _make_tool(mcp_config) + raw = {"key": "value"} + assert tool._format_result(raw) == raw + + +# ===================================================================== +# Execute Action +# ===================================================================== + + +@pytest.mark.unit +class TestExecuteAction: + + def test_no_server_raises(self): + tool = _make_tool({"server_url": "", "auth_type": "none"}) + with pytest.raises(Exception, match="No MCP server configured"): + tool.execute_action("test_action") + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_successful_execute(self, mock_run, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + mock_run.return_value = {"key": "value"} + + result = tool.execute_action("test_action", param1="val1") + + mock_run.assert_called_once_with("call_tool", "test_action", param1="val1") + assert result == {"key": "value"} + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_empty_kwargs_cleaned(self, mock_run, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + mock_run.return_value = {} + + tool.execute_action("test", param1="", param2=None, param3="real") + + call_kwargs = mock_run.call_args[1] + assert "param1" not in call_kwargs + assert "param2" not in call_kwargs + assert call_kwargs["param3"] == "real" + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_auth_error_retries_for_non_oauth(self, mock_run, mcp_config): + from application.agents.tools.mcp_tool import MCPTool + + tool = _make_tool(mcp_config) + tool._client = MagicMock() + + mock_run.side_effect = [ + Exception("401 Unauthorized"), + {"key": "retry_ok"}, + ] + + with patch.object(MCPTool, "_setup_client"): + result = tool.execute_action("act") + assert result == {"key": "retry_ok"} + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_auth_error_raises_for_oauth(self, mock_run): + tool = _make_tool({ + "server_url": "https://mcp.example.com", + "auth_type": "oauth", + }) + tool._client = MagicMock() + mock_run.side_effect = Exception("401 Unauthorized") + + with pytest.raises(Exception, match="OAuth session expired"): + tool.execute_action("act") + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_non_auth_error_raises(self, mock_run, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + mock_run.side_effect = Exception("Something weird happened") + + with pytest.raises(Exception, match="Failed to execute action"): + tool.execute_action("act") + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_no_client_calls_setup(self, mock_run, mcp_config): + from application.agents.tools.mcp_tool import MCPTool + + tool = _make_tool(mcp_config) + tool._client = None + mock_run.return_value = {"ok": True} + + with patch.object(MCPTool, "_setup_client") as mock_setup: + tool.execute_action("act") + mock_setup.assert_called() + + +# ===================================================================== +# Discover Tools +# ===================================================================== + + +@pytest.mark.unit +class TestDiscoverTools: + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_discover_tools_success(self, mock_run, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + mock_run.return_value = [{"name": "t1", "description": "d1"}] + + result = tool.discover_tools() + assert len(result) == 1 + assert result[0]["name"] == "t1" + + def test_discover_tools_no_server_url(self): + tool = _make_tool({"server_url": "", "auth_type": "none"}) + assert tool.discover_tools() == [] + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_discover_tools_error(self, mock_run, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + mock_run.side_effect = Exception("connection lost") + + with pytest.raises(Exception, match="Failed to discover tools"): + tool.discover_tools() + + +# ===================================================================== +# Test Connection +# ===================================================================== + + +@pytest.mark.unit +class TestTestConnection: + + def test_no_server_url(self): + tool = _make_tool({"server_url": "", "auth_type": "none"}) + result = tool.test_connection() + assert result["success"] is False + assert "No server URL" in result["message"] + + def test_invalid_scheme(self): + tool = _make_tool({"server_url": "ftp://bad.com", "auth_type": "none"}) + result = tool.test_connection() + assert result["success"] is False + assert "Invalid URL scheme" in result["message"] + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_regular_connection_success(self, mock_run, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + mock_run.side_effect = [ + None, # ping + [{"name": "t1", "description": "d1"}], # list_tools + ] + + result = tool.test_connection() + assert result["success"] is True + assert result["tools_count"] == 1 + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_regular_connection_ping_fails_tools_work(self, mock_run, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + mock_run.side_effect = [ + Exception("ping failed"), # ping + [{"name": "t1", "description": "d1"}], # list_tools + ] + + result = tool.test_connection() + assert result["success"] is True + + @patch("application.agents.tools.mcp_tool.MCPTool._run_async_operation") + def test_regular_connection_both_fail(self, mock_run, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + mock_run.side_effect = [ + Exception("ping failed"), + Exception("tools failed"), + ] + + result = tool.test_connection() + assert result["success"] is False + + def test_client_init_failure(self): + from application.agents.tools.mcp_tool import MCPTool + + tool = _make_tool({ + "server_url": "https://good.example.com", + "auth_type": "none", + }) + tool._client = None + + with patch.object(MCPTool, "_setup_client", side_effect=Exception("init fail")): + result = tool.test_connection() + assert result["success"] is False + assert "Client init failed" in result["message"] + + +# ===================================================================== +# Map Error +# ===================================================================== + + +@pytest.mark.unit +class TestMapError: + + def test_timeout_error(self, mcp_config): + tool = _make_tool(mcp_config) + err = tool._map_error("test", concurrent.futures.TimeoutError()) + assert "Timed out" in str(err) + + def test_connection_refused(self, mcp_config): + tool = _make_tool(mcp_config) + err = tool._map_error("test", ConnectionRefusedError()) + assert "Connection refused" in str(err) + + def test_403_forbidden(self, mcp_config): + tool = _make_tool(mcp_config) + err = tool._map_error("test", Exception("403 Forbidden")) + assert "Access denied" in str(err) + + def test_401_unauthorized(self, mcp_config): + tool = _make_tool(mcp_config) + err = tool._map_error("test", Exception("401 Unauthorized")) + assert "Authentication failed" in str(err) + + def test_econnrefused_pattern(self, mcp_config): + tool = _make_tool(mcp_config) + err = tool._map_error("test", Exception("ECONNREFUSED error")) + assert "Connection refused" in str(err) + + def test_ssl_error(self, mcp_config): + tool = _make_tool(mcp_config) + err = tool._map_error("test", Exception("SSL certificate verify failed")) + assert "SSL" in str(err) + + def test_unknown_error_passthrough(self, mcp_config): + tool = _make_tool(mcp_config) + original = RuntimeError("something weird") + err = tool._map_error("test", original) + assert err is original + + +# ===================================================================== +# Get Actions Metadata +# ===================================================================== + + +@pytest.mark.unit +class TestGetActionsMetadata: + + def test_empty_tools(self, mcp_config): + tool = _make_tool(mcp_config) + tool.available_tools = [] + assert tool.get_actions_metadata() == [] + + def test_tools_with_input_schema(self, mcp_config): + tool = _make_tool(mcp_config) + tool.available_tools = [ + { + "name": "search", + "description": "Search things", + "inputSchema": { + "type": "object", + "properties": {"query": {"type": "string"}}, + "required": ["query"], + "additionalProperties": False, + "description": "Search params", + }, + } + ] + meta = tool.get_actions_metadata() + assert len(meta) == 1 + assert meta[0]["name"] == "search" + assert "query" in meta[0]["parameters"]["properties"] + assert meta[0]["parameters"]["additionalProperties"] is False + + def test_tools_without_schema(self, mcp_config): + tool = _make_tool(mcp_config) + tool.available_tools = [{"name": "ping", "description": "Ping"}] + meta = tool.get_actions_metadata() + assert meta[0]["parameters"]["properties"] == {} + + def test_tools_with_flat_schema(self, mcp_config): + tool = _make_tool(mcp_config) + tool.available_tools = [ + { + "name": "flat", + "description": "Flat schema", + "inputSchema": {"query": {"type": "string"}}, + } + ] + meta = tool.get_actions_metadata() + assert "query" in meta[0]["parameters"]["properties"] + + def test_tools_with_alternate_schema_keys(self, mcp_config): + tool = _make_tool(mcp_config) + tool.available_tools = [ + { + "name": "alt", + "description": "Alt", + "parameters": { + "type": "object", + "properties": {"x": {"type": "number"}}, + }, + } + ] + meta = tool.get_actions_metadata() + assert "x" in meta[0]["parameters"]["properties"] + + def test_config_requirements(self, mcp_config): + tool = _make_tool(mcp_config) + reqs = tool.get_config_requirements() + assert "server_url" in reqs + assert "auth_type" in reqs + assert reqs["server_url"]["required"] is True + assert "timeout" in reqs + + +# ===================================================================== +# Setup Client (caching) +# ===================================================================== + + +@pytest.mark.unit +class TestSetupClient: + + def test_setup_client_caches_client(self, mcp_config): + from application.agents.tools.mcp_tool import MCPTool + + tool = MCPTool.__new__(MCPTool) + tool.config = mcp_config + tool.server_url = mcp_config["server_url"] + tool.transport_type = "http" + tool.auth_type = "none" + tool.timeout = 10 + tool.custom_headers = {} + tool.auth_credentials = {} + tool.oauth_scopes = [] + tool.oauth_task_id = None + tool.oauth_client_name = "DocsGPT-MCP" + tool.redirect_uri = "https://example.com/callback" + tool.query_mode = False + tool._cache_key = "test_cache_key" + tool._client = None + tool.available_tools = [] + + mock_client = MagicMock() + with patch.object(MCPTool, "_create_transport", return_value=MagicMock()), \ + patch("application.agents.tools.mcp_tool.Client", return_value=mock_client): + tool._setup_client() + assert tool._client is mock_client + + +# ===================================================================== +# MCPOAuthManager +# ===================================================================== + + +@pytest.mark.unit +class TestMCPOAuthManager: + + def test_handle_callback_success(self): + from application.agents.tools.mcp_tool import MCPOAuthManager + + mock_redis = MagicMock() + manager = MCPOAuthManager(mock_redis) + + result = manager.handle_oauth_callback(state="abc123", code="auth_code") + assert result is True + mock_redis.setex.assert_called() + + def test_handle_callback_no_redis(self): + from application.agents.tools.mcp_tool import MCPOAuthManager + + manager = MCPOAuthManager(None) + result = manager.handle_oauth_callback(state="abc", code="code") + assert result is False + + def test_handle_callback_no_state(self): + from application.agents.tools.mcp_tool import MCPOAuthManager + + mock_redis = MagicMock() + manager = MCPOAuthManager(mock_redis) + result = manager.handle_oauth_callback(state="", code="code") + assert result is False + + def test_handle_callback_error(self): + from application.agents.tools.mcp_tool import MCPOAuthManager + + mock_redis = MagicMock() + manager = MCPOAuthManager(mock_redis) + + result = manager.handle_oauth_callback( + state="abc", code="", error="access_denied" + ) + assert result is False + + def test_get_oauth_status_no_task(self): + from application.agents.tools.mcp_tool import MCPOAuthManager + + manager = MCPOAuthManager(MagicMock()) + result = manager.get_oauth_status("") + assert result["status"] == "not_started" + + def test_get_oauth_status_with_task(self): + from application.agents.tools.mcp_tool import MCPOAuthManager + + with patch("application.agents.tools.mcp_tool.mcp_oauth_status_task", + return_value={"status": "complete"}): + manager = MCPOAuthManager(MagicMock()) + result = manager.get_oauth_status("task123") + assert result["status"] == "complete" + + +# ===================================================================== +# DBTokenStorage +# ===================================================================== + + +@pytest.mark.unit +class TestDBTokenStorage: + + def test_get_base_url(self): + from application.agents.tools.mcp_tool import DBTokenStorage + + assert ( + DBTokenStorage.get_base_url("https://mcp.example.com/api/v1") + == "https://mcp.example.com" + ) + + def test_get_base_url_with_port(self): + from application.agents.tools.mcp_tool import DBTokenStorage + + assert ( + DBTokenStorage.get_base_url("http://localhost:8080/path") + == "http://localhost:8080" + ) + + def test_get_db_key(self): + from application.agents.tools.mcp_tool import DBTokenStorage + + mock_db = MagicMock() + storage = DBTokenStorage( + server_url="https://mcp.example.com/api", + user_id="user1", + db_client=mock_db, + ) + key = storage.get_db_key() + assert key["server_url"] == "https://mcp.example.com" + assert key["user_id"] == "user1" + + def test_get_tokens_none(self): + from application.agents.tools.mcp_tool import DBTokenStorage + + mock_db = MagicMock() + mock_collection = MagicMock() + mock_collection.find_one.return_value = None + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + + storage = DBTokenStorage( + server_url="https://mcp.example.com", + user_id="user1", + db_client=mock_db, + ) + + loop = asyncio.new_event_loop() + try: + result = loop.run_until_complete(storage.get_tokens()) + assert result is None + finally: + loop.close() + + def test_serialize_client_info(self): + from application.agents.tools.mcp_tool import DBTokenStorage + + mock_db = MagicMock() + storage = DBTokenStorage( + server_url="https://mcp.example.com", + user_id="user1", + db_client=mock_db, + ) + info = {"redirect_uris": ["https://example.com/cb"]} + result = storage._serialize_client_info(info) + assert result["redirect_uris"] == ["https://example.com/cb"] + + def test_clear(self): + from application.agents.tools.mcp_tool import DBTokenStorage + + mock_db = MagicMock() + mock_collection = MagicMock() + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + + storage = DBTokenStorage( + server_url="https://mcp.example.com", + user_id="user1", + db_client=mock_db, + ) + + loop = asyncio.new_event_loop() + try: + loop.run_until_complete(storage.clear()) + mock_collection.delete_one.assert_called_once() + finally: + loop.close() + + +# ===================================================================== +# NonInteractiveOAuth +# ===================================================================== + + +@pytest.mark.unit +class TestNonInteractiveOAuth: + + def test_redirect_handler_raises(self): + from application.agents.tools.mcp_tool import NonInteractiveOAuth + + mock_db = MagicMock() + mock_collection = MagicMock() + mock_collection.find_one.return_value = None + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + + oauth = NonInteractiveOAuth( + mcp_url="https://mcp.example.com/api", + scopes=["read"], + redis_client=None, + redirect_uri="https://example.com/callback", + db=mock_db, + user_id="user1", + ) + + loop = asyncio.new_event_loop() + try: + with pytest.raises(Exception, match="OAuth session expired"): + loop.run_until_complete( + oauth.redirect_handler("https://auth.example.com/authorize?state=x") + ) + finally: + loop.close() + + def test_callback_handler_raises(self): + from application.agents.tools.mcp_tool import NonInteractiveOAuth + + mock_db = MagicMock() + mock_collection = MagicMock() + mock_collection.find_one.return_value = None + mock_db.__getitem__ = MagicMock(return_value=mock_collection) + + oauth = NonInteractiveOAuth( + mcp_url="https://mcp.example.com/api", + scopes=["read"], + redis_client=None, + redirect_uri="https://example.com/callback", + db=mock_db, + user_id="user1", + ) + + loop = asyncio.new_event_loop() + try: + with pytest.raises(Exception, match="OAuth session expired"): + loop.run_until_complete(oauth.callback_handler()) + finally: + loop.close() + + +# ===================================================================== +# Run Async Operation +# ===================================================================== + + +@pytest.mark.unit +class TestRunAsyncOperation: + + def test_run_in_new_loop(self, mcp_config): + tool = _make_tool(mcp_config) + tool._client = MagicMock() + + async def fake_execute(*args, **kwargs): + return "ok" + + with patch.object(tool, "_execute_with_client", side_effect=fake_execute): + result = tool._run_in_new_loop("ping") + assert result == "ok" diff --git a/tests/agents/tools/test_memory.py b/tests/agents/tools/test_memory.py new file mode 100644 index 00000000..3757f883 --- /dev/null +++ b/tests/agents/tools/test_memory.py @@ -0,0 +1,449 @@ +"""Comprehensive tests for application/agents/tools/memory.py + +Covers: MemoryTool initialization, path validation, all actions +(view, create, str_replace, insert, delete, rename), directory operations, +error handling, and metadata. +""" + +import mongomock +import pytest + + +def _get_settings(): + from application.core.settings import settings + return settings + + +@pytest.fixture +def mock_memory_db(monkeypatch): + """Set up a mongomock-based memory collection.""" + settings = _get_settings() + mock_client = mongomock.MongoClient() + mock_db = mock_client[settings.MONGO_DB_NAME] + + def get_mock_client(): + return {settings.MONGO_DB_NAME: mock_db} + + monkeypatch.setattr( + "application.core.mongo_db.MongoDB.get_client", get_mock_client + ) + return mock_db + + +@pytest.fixture +def memory_tool(mock_memory_db): + from application.agents.tools.memory import MemoryTool + + return MemoryTool( + tool_config={"tool_id": "test_tool_001"}, + user_id="test_user", + ) + + +# ===================================================================== +# Initialization +# ===================================================================== + + +@pytest.mark.unit +class TestMemoryToolInit: + + def test_init_with_config(self, mock_memory_db): + from application.agents.tools.memory import MemoryTool + + tool = MemoryTool( + tool_config={"tool_id": "custom_id"}, user_id="user1" + ) + assert tool.tool_id == "custom_id" + assert tool.user_id == "user1" + + def test_init_fallback_to_user_id(self, mock_memory_db): + from application.agents.tools.memory import MemoryTool + + tool = MemoryTool(tool_config={}, user_id="user1") + assert tool.tool_id == "default_user1" + + def test_init_no_user_no_config(self, mock_memory_db): + from application.agents.tools.memory import MemoryTool + + tool = MemoryTool() + assert tool.tool_id is not None # UUID fallback + assert tool.user_id is None + + +# ===================================================================== +# Path Validation +# ===================================================================== + + +@pytest.mark.unit +class TestPathValidation: + + def test_valid_path(self, memory_tool): + assert memory_tool._validate_path("/notes.txt") == "/notes.txt" + + def test_adds_leading_slash(self, memory_tool): + assert memory_tool._validate_path("notes.txt") == "/notes.txt" + + def test_empty_path_returns_none(self, memory_tool): + assert memory_tool._validate_path("") is None + + def test_double_dots_rejected(self, memory_tool): + assert memory_tool._validate_path("/../../etc/passwd") is None + + def test_double_slash_rejected(self, memory_tool): + assert memory_tool._validate_path("//path") is None + + def test_preserves_trailing_slash(self, memory_tool): + result = memory_tool._validate_path("/project/") + assert result.endswith("/") + + def test_root_path(self, memory_tool): + assert memory_tool._validate_path("/") == "/" + + def test_whitespace_stripped(self, memory_tool): + result = memory_tool._validate_path(" /notes.txt ") + assert result == "/notes.txt" + + +# ===================================================================== +# Execute Action - No User +# ===================================================================== + + +@pytest.mark.unit +class TestNoUser: + + def test_requires_user_id(self, mock_memory_db): + from application.agents.tools.memory import MemoryTool + + tool = MemoryTool(tool_config={"tool_id": "t"}, user_id=None) + result = tool.execute_action("view", path="/") + assert "Error" in result + assert "user_id" in result + + def test_unknown_action(self, memory_tool): + result = memory_tool.execute_action("fly") + assert "Unknown action" in result + + +# ===================================================================== +# View Action +# ===================================================================== + + +@pytest.mark.unit +class TestViewAction: + + def test_view_empty_directory(self, memory_tool): + result = memory_tool.execute_action("view", path="/") + assert "Directory: /" in result + assert "(empty)" in result + + def test_view_directory_with_files(self, memory_tool): + memory_tool.execute_action("create", path="/notes.txt", file_text="content") + memory_tool.execute_action("create", path="/todo.txt", file_text="tasks") + + result = memory_tool.execute_action("view", path="/") + assert "notes.txt" in result + assert "todo.txt" in result + + def test_view_file_content(self, memory_tool): + memory_tool.execute_action("create", path="/hello.txt", file_text="Hello World") + result = memory_tool.execute_action("view", path="/hello.txt") + assert "Hello World" in result + + def test_view_nonexistent_file(self, memory_tool): + result = memory_tool.execute_action("view", path="/missing.txt") + assert "Error" in result + assert "not found" in result.lower() + + def test_view_file_with_range(self, memory_tool): + memory_tool.execute_action( + "create", path="/lines.txt", file_text="line1\nline2\nline3\nline4" + ) + result = memory_tool.execute_action( + "view", path="/lines.txt", view_range=[2, 3] + ) + assert "line2" in result + assert "line3" in result + + def test_view_file_range_out_of_bounds(self, memory_tool): + memory_tool.execute_action("create", path="/short.txt", file_text="only") + result = memory_tool.execute_action( + "view", path="/short.txt", view_range=[100, 200] + ) + assert "out of bounds" in result.lower() + + def test_view_invalid_path(self, memory_tool): + result = memory_tool.execute_action("view", path="") + assert "Error" in result + + def test_view_subdirectory(self, memory_tool): + memory_tool.execute_action( + "create", path="/project/src/main.py", file_text="code" + ) + result = memory_tool.execute_action("view", path="/project/") + assert "src/main.py" in result + + +# ===================================================================== +# Create Action +# ===================================================================== + + +@pytest.mark.unit +class TestCreateAction: + + def test_create_file(self, memory_tool): + result = memory_tool.execute_action( + "create", path="/test.txt", file_text="content" + ) + assert "File created" in result + + content = memory_tool.execute_action("view", path="/test.txt") + assert "content" in content + + def test_overwrite_file(self, memory_tool): + memory_tool.execute_action("create", path="/test.txt", file_text="old") + memory_tool.execute_action("create", path="/test.txt", file_text="new") + + content = memory_tool.execute_action("view", path="/test.txt") + assert "new" in content + + def test_create_at_directory_path(self, memory_tool): + result = memory_tool.execute_action("create", path="/dir/", file_text="text") + assert "Error" in result + assert "directory path" in result.lower() + + def test_create_invalid_path(self, memory_tool): + result = memory_tool.execute_action("create", path="", file_text="text") + assert "Error" in result + + def test_create_nested_path(self, memory_tool): + result = memory_tool.execute_action( + "create", path="/a/b/c/file.txt", file_text="deep" + ) + assert "File created" in result + + +# ===================================================================== +# String Replace Action +# ===================================================================== + + +@pytest.mark.unit +class TestStrReplaceAction: + + def test_replace_text(self, memory_tool): + memory_tool.execute_action( + "create", path="/doc.txt", file_text="Hello World" + ) + result = memory_tool.execute_action( + "str_replace", path="/doc.txt", old_str="Hello", new_str="Hi" + ) + assert "File updated" in result + + content = memory_tool.execute_action("view", path="/doc.txt") + assert "Hi World" in content + + def test_replace_not_found(self, memory_tool): + memory_tool.execute_action("create", path="/doc.txt", file_text="Hello") + result = memory_tool.execute_action( + "str_replace", path="/doc.txt", old_str="Missing", new_str="X" + ) + assert "not found" in result.lower() + + def test_replace_empty_old_str(self, memory_tool): + memory_tool.execute_action("create", path="/doc.txt", file_text="Hello") + result = memory_tool.execute_action( + "str_replace", path="/doc.txt", old_str="", new_str="X" + ) + assert "Error" in result + + def test_replace_file_not_found(self, memory_tool): + result = memory_tool.execute_action( + "str_replace", path="/missing.txt", old_str="a", new_str="b" + ) + assert "not found" in result.lower() + + def test_replace_case_insensitive(self, memory_tool): + memory_tool.execute_action( + "create", path="/doc.txt", file_text="Hello World" + ) + result = memory_tool.execute_action( + "str_replace", path="/doc.txt", old_str="hello", new_str="Hi" + ) + assert "File updated" in result + + +# ===================================================================== +# Insert Action +# ===================================================================== + + +@pytest.mark.unit +class TestInsertAction: + + def test_insert_text(self, memory_tool): + memory_tool.execute_action( + "create", path="/doc.txt", file_text="line1\nline2" + ) + result = memory_tool.execute_action( + "insert", path="/doc.txt", insert_line=2, insert_text="inserted" + ) + assert "inserted" in result.lower() + + content = memory_tool.execute_action("view", path="/doc.txt") + assert "inserted" in content + + def test_insert_empty_text(self, memory_tool): + memory_tool.execute_action("create", path="/doc.txt", file_text="line1") + result = memory_tool.execute_action( + "insert", path="/doc.txt", insert_line=1, insert_text="" + ) + assert "Error" in result + + def test_insert_file_not_found(self, memory_tool): + result = memory_tool.execute_action( + "insert", path="/missing.txt", insert_line=1, insert_text="text" + ) + assert "not found" in result.lower() + + def test_insert_invalid_line_number(self, memory_tool): + memory_tool.execute_action("create", path="/doc.txt", file_text="line1") + result = memory_tool.execute_action( + "insert", path="/doc.txt", insert_line=-5, insert_text="text" + ) + assert "Error" in result + + +# ===================================================================== +# Delete Action +# ===================================================================== + + +@pytest.mark.unit +class TestDeleteAction: + + def test_delete_file(self, memory_tool): + memory_tool.execute_action("create", path="/test.txt", file_text="data") + result = memory_tool.execute_action("delete", path="/test.txt") + assert "Deleted" in result + + content = memory_tool.execute_action("view", path="/test.txt") + assert "not found" in content.lower() + + def test_delete_nonexistent_file(self, memory_tool): + result = memory_tool.execute_action("delete", path="/missing.txt") + assert "not found" in result.lower() + + def test_delete_root_clears_all(self, memory_tool): + memory_tool.execute_action("create", path="/a.txt", file_text="a") + memory_tool.execute_action("create", path="/b.txt", file_text="b") + + result = memory_tool.execute_action("delete", path="/") + assert "Deleted" in result + assert "2" in result + + def test_delete_directory(self, memory_tool): + memory_tool.execute_action("create", path="/dir/f1.txt", file_text="1") + memory_tool.execute_action("create", path="/dir/f2.txt", file_text="2") + + result = memory_tool.execute_action("delete", path="/dir/") + assert "Deleted" in result + + def test_delete_directory_without_trailing_slash(self, memory_tool): + memory_tool.execute_action("create", path="/dir/f1.txt", file_text="1") + + result = memory_tool.execute_action("delete", path="/dir") + assert "Deleted" in result + + def test_delete_invalid_path(self, memory_tool): + result = memory_tool.execute_action("delete", path="") + assert "Error" in result + + +# ===================================================================== +# Rename Action +# ===================================================================== + + +@pytest.mark.unit +class TestRenameAction: + + def test_rename_file(self, memory_tool): + memory_tool.execute_action("create", path="/old.txt", file_text="data") + result = memory_tool.execute_action( + "rename", old_path="/old.txt", new_path="/new.txt" + ) + assert "Renamed" in result + + content = memory_tool.execute_action("view", path="/new.txt") + assert "data" in content + + def test_rename_file_not_found(self, memory_tool): + result = memory_tool.execute_action( + "rename", old_path="/missing.txt", new_path="/new.txt" + ) + assert "not found" in result.lower() + + def test_rename_target_exists(self, memory_tool): + memory_tool.execute_action("create", path="/a.txt", file_text="a") + memory_tool.execute_action("create", path="/b.txt", file_text="b") + + result = memory_tool.execute_action( + "rename", old_path="/a.txt", new_path="/b.txt" + ) + assert "already exists" in result.lower() + + def test_rename_root_rejected(self, memory_tool): + result = memory_tool.execute_action( + "rename", old_path="/", new_path="/new/" + ) + assert "Cannot rename root" in result + + def test_rename_directory(self, memory_tool): + memory_tool.execute_action("create", path="/old/f.txt", file_text="data") + result = memory_tool.execute_action( + "rename", old_path="/old/", new_path="/new/" + ) + assert "Renamed" in result + + content = memory_tool.execute_action("view", path="/new/f.txt") + assert "data" in content + + def test_rename_directory_not_found(self, memory_tool): + result = memory_tool.execute_action( + "rename", old_path="/missing/", new_path="/new/" + ) + assert "not found" in result.lower() + + def test_rename_invalid_path(self, memory_tool): + result = memory_tool.execute_action( + "rename", old_path="", new_path="/new.txt" + ) + assert "Error" in result + + +# ===================================================================== +# Metadata +# ===================================================================== + + +@pytest.mark.unit +class TestMemoryToolMetadata: + + def test_actions_metadata(self, memory_tool): + meta = memory_tool.get_actions_metadata() + action_names = [a["name"] for a in meta] + assert "view" in action_names + assert "create" in action_names + assert "str_replace" in action_names + assert "insert" in action_names + assert "delete" in action_names + assert "rename" in action_names + assert len(meta) == 6 + + def test_config_requirements(self, memory_tool): + assert memory_tool.get_config_requirements() == {} diff --git a/tests/api/answer/services/compression/test_threshold_checker.py b/tests/api/answer/services/compression/test_threshold_checker.py new file mode 100644 index 00000000..9938df06 --- /dev/null +++ b/tests/api/answer/services/compression/test_threshold_checker.py @@ -0,0 +1,46 @@ +from unittest.mock import patch + +import pytest + + +@pytest.mark.unit +class TestCompressionThresholdChecker: + + def _make_checker(self, pct=0.7): + from application.api.answer.services.compression.threshold_checker import ( + CompressionThresholdChecker, + ) + + return CompressionThresholdChecker(threshold_percentage=pct) + + @patch( + "application.api.answer.services.compression.threshold_checker.get_token_limit", + return_value=8000, + ) + @patch( + "application.api.answer.services.compression.threshold_checker.TokenCounter.count_message_tokens", + return_value=6000, + ) + def test_check_message_tokens_above_threshold(self, mock_count, mock_limit): + checker = self._make_checker(0.7) + assert checker.check_message_tokens([{"role": "user"}], "gpt-4") is True + + @patch( + "application.api.answer.services.compression.threshold_checker.get_token_limit", + return_value=8000, + ) + @patch( + "application.api.answer.services.compression.threshold_checker.TokenCounter.count_message_tokens", + return_value=1000, + ) + def test_check_message_tokens_below_threshold(self, mock_count, mock_limit): + checker = self._make_checker(0.7) + assert checker.check_message_tokens([{"role": "user"}], "gpt-4") is False + + @patch( + "application.api.answer.services.compression.threshold_checker.TokenCounter.count_message_tokens", + side_effect=Exception("Token error"), + ) + def test_check_message_tokens_exception_returns_false(self, mock_count): + checker = self._make_checker(0.7) + assert checker.check_message_tokens([], "gpt-4") is False diff --git a/tests/api/answer/test_base_routes.py b/tests/api/answer/test_base_routes.py new file mode 100644 index 00000000..6e208841 --- /dev/null +++ b/tests/api/answer/test_base_routes.py @@ -0,0 +1,363 @@ +"""Unit tests for application/api/answer/routes/base.py — BaseAnswerResource. + +Additional coverage beyond tests/api/answer/routes/test_base.py: + - _prepare_tool_calls_for_logging: truncation, non-dict items + - complete_stream: tool_calls, thoughts, structured output, metadata, + isNoneDoc, GeneratorExit handling, compression metadata + - process_response_stream: structured answer, incomplete stream + - error_stream_generate: format + - check_usage: string boolean parsing ("True" strings) +""" + +import json +from unittest.mock import MagicMock + +import pytest +from bson import ObjectId + + +@pytest.mark.unit +class TestPrepareToolCallsForLogging: + + def test_empty_list(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + assert resource._prepare_tool_calls_for_logging([]) == [] + + def test_none_returns_empty(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + assert resource._prepare_tool_calls_for_logging(None) == [] + + def test_truncates_long_result(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + tool_calls = [{"result": "x" * 20000}] + prepared = resource._prepare_tool_calls_for_logging(tool_calls, max_chars=100) + assert len(prepared[0]["result"]) == 100 + + def test_truncates_result_full(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + tool_calls = [{"result_full": "y" * 20000}] + prepared = resource._prepare_tool_calls_for_logging(tool_calls, max_chars=50) + assert len(prepared[0]["result_full"]) == 50 + + def test_non_dict_items_wrapped(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + tool_calls = ["string_item", 42] + prepared = resource._prepare_tool_calls_for_logging(tool_calls) + assert prepared[0] == {"result": "string_item"} + assert prepared[1] == {"result": "42"} + + def test_preserves_short_results(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + tool_calls = [{"tool_name": "search", "result": "short text"}] + prepared = resource._prepare_tool_calls_for_logging(tool_calls) + assert prepared[0]["result"] == "short text" + assert prepared[0]["tool_name"] == "search" + + +@pytest.mark.unit +class TestCompleteStreamToolCalls: + + def test_streams_tool_calls(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + mock_agent = MagicMock() + mock_agent.gen.return_value = iter( + [ + {"answer": "Using tool..."}, + {"tool_calls": [{"name": "search", "result": "found"}]}, + ] + ) + + stream = list( + resource.complete_stream( + question="Search for X", + agent=mock_agent, + conversation_id=None, + user_api_key=None, + decoded_token={"sub": "u"}, + should_save_conversation=False, + ) + ) + tool_chunks = [s for s in stream if '"type": "tool_calls"' in s] + assert len(tool_chunks) == 1 + + def test_streams_thought_events(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + mock_agent = MagicMock() + mock_agent.gen.return_value = iter( + [ + {"thought": "Let me think..."}, + {"answer": "Here is the answer"}, + ] + ) + + stream = list( + resource.complete_stream( + question="Q", + agent=mock_agent, + conversation_id=None, + user_api_key=None, + decoded_token={"sub": "u"}, + should_save_conversation=False, + ) + ) + thought_chunks = [s for s in stream if '"type": "thought"' in s] + assert len(thought_chunks) == 1 + assert "Let me think" in thought_chunks[0] + + +@pytest.mark.unit +class TestCompleteStreamStructuredOutput: + + def test_streams_structured_answer(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + mock_agent = MagicMock() + mock_agent.gen.return_value = iter( + [ + { + "answer": '{"key": "value"}', + "structured": True, + "schema": {"type": "object"}, + }, + ] + ) + + stream = list( + resource.complete_stream( + question="Q", + agent=mock_agent, + conversation_id=None, + user_api_key=None, + decoded_token={"sub": "u"}, + should_save_conversation=False, + ) + ) + structured_chunks = [ + s for s in stream if '"type": "structured_answer"' in s + ] + assert len(structured_chunks) == 1 + + +@pytest.mark.unit +class TestCompleteStreamMetadata: + + def test_metadata_collected(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + mock_agent = MagicMock() + mock_agent.gen.return_value = iter( + [ + {"metadata": {"search_query": "test"}}, + {"answer": "result"}, + ] + ) + + stream = list( + resource.complete_stream( + question="Q", + agent=mock_agent, + conversation_id=None, + user_api_key=None, + decoded_token={"sub": "u"}, + should_save_conversation=False, + ) + ) + # Should not crash, metadata handled silently + answer_chunks = [s for s in stream if '"type": "answer"' in s] + assert len(answer_chunks) == 1 + + +@pytest.mark.unit +class TestCompleteStreamIsNoneDoc: + + def test_isNoneDoc_sets_source_to_none(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + mock_agent = MagicMock() + mock_agent.gen.return_value = iter( + [ + {"answer": "answer"}, + {"sources": [{"text": "doc", "source": "real_source"}]}, + ] + ) + + stream = list( + resource.complete_stream( + question="Q", + agent=mock_agent, + conversation_id=None, + user_api_key=None, + decoded_token={"sub": "u"}, + isNoneDoc=True, + should_save_conversation=False, + ) + ) + # Verify stream completes without error + end_chunks = [s for s in stream if '"type": "end"' in s] + assert len(end_chunks) == 1 + + +@pytest.mark.unit +class TestCompleteStreamErrorType: + + def test_error_type_event_sanitized(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + mock_agent = MagicMock() + mock_agent.gen.return_value = iter( + [ + {"type": "error", "error": "API key invalid: sk-xxx"}, + ] + ) + + stream = list( + resource.complete_stream( + question="Q", + agent=mock_agent, + conversation_id=None, + user_api_key=None, + decoded_token={"sub": "u"}, + should_save_conversation=False, + ) + ) + error_chunks = [s for s in stream if '"type": "error"' in s] + assert len(error_chunks) == 1 + + def test_non_error_type_event_passed_through(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + mock_agent = MagicMock() + mock_agent.gen.return_value = iter( + [ + {"type": "custom_event", "data": "value"}, + ] + ) + + stream = list( + resource.complete_stream( + question="Q", + agent=mock_agent, + conversation_id=None, + user_api_key=None, + decoded_token={"sub": "u"}, + should_save_conversation=False, + ) + ) + custom_chunks = [s for s in stream if '"type": "custom_event"' in s] + assert len(custom_chunks) == 1 + + +@pytest.mark.unit +class TestProcessResponseStreamExtended: + + def test_handles_structured_answer(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + stream = [ + f'data: {json.dumps({"type": "structured_answer", "answer": "{}", "structured": True, "schema": None})}\n\n', + f'data: {json.dumps({"type": "id", "id": str(ObjectId())})}\n\n', + f'data: {json.dumps({"type": "end"})}\n\n', + ] + result = resource.process_response_stream(iter(stream)) + assert result[1] == "{}" + # Structured output adds extra tuple element + assert len(result) == 7 + assert result[6]["structured"] is True + + def test_handles_tool_calls_event(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + stream = [ + f'data: {json.dumps({"type": "answer", "answer": "result"})}\n\n', + f'data: {json.dumps({"type": "tool_calls", "tool_calls": [{"name": "t1"}]})}\n\n', + f'data: {json.dumps({"type": "id", "id": "conv1"})}\n\n', + f'data: {json.dumps({"type": "end"})}\n\n', + ] + result = resource.process_response_stream(iter(stream)) + assert result[3] == [{"name": "t1"}] + + def test_incomplete_stream(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + stream = [ + f'data: {json.dumps({"type": "answer", "answer": "partial"})}\n\n', + ] + result = resource.process_response_stream(iter(stream)) + assert result[4] == "Stream ended unexpectedly" + + def test_handles_thought_event(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + + with flask_app.app_context(): + resource = BaseAnswerResource() + stream = [ + f'data: {json.dumps({"type": "thought", "thought": "thinking..."})}\n\n', + f'data: {json.dumps({"type": "end"})}\n\n', + ] + result = resource.process_response_stream(iter(stream)) + assert result[4] == "thinking..." + + +@pytest.mark.unit +class TestCheckUsageStringBooleans: + + def test_string_true_parsed_correctly(self, mock_mongo_db, flask_app): + from application.api.answer.routes.base import BaseAnswerResource + from application.core.settings import settings + + with flask_app.app_context(): + agents_collection = mock_mongo_db[settings.MONGO_DB_NAME]["agents"] + agents_collection.insert_one( + { + "_id": ObjectId(), + "key": "str_bool_key", + "limited_token_mode": "True", + "token_limit": 1000000, + "limited_request_mode": "True", + "request_limit": 1000000, + } + ) + resource = BaseAnswerResource() + result = resource.check_usage({"user_api_key": "str_bool_key"}) + # Should not exceed limits, so returns None + assert result is None diff --git a/tests/api/answer/test_conversation_service.py b/tests/api/answer/test_conversation_service.py new file mode 100644 index 00000000..c032f1b9 --- /dev/null +++ b/tests/api/answer/test_conversation_service.py @@ -0,0 +1,418 @@ +"""Unit tests for application/api/answer/services/conversation_service.py. + +Additional coverage beyond tests/api/answer/services/test_conversation_service.py: + - save_conversation: index-based update, metadata persistence, agent key tracking + - update_compression_metadata + - append_compression_message + - get_compression_metadata + - Edge cases: None token, empty summary, shared_with access +""" + +from datetime import datetime, timezone +from unittest.mock import Mock + +import pytest +from bson import ObjectId + + +@pytest.mark.unit +class TestConversationServiceGetExtended: + + def test_returns_conversation_for_shared_user(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + conv_id = ObjectId() + collection.insert_one( + { + "_id": conv_id, + "user": "owner_123", + "shared_with": ["shared_user"], + "name": "Shared Conv", + "queries": [], + } + ) + + result = service.get_conversation(str(conv_id), "shared_user") + assert result is not None + assert result["name"] == "Shared Conv" + + def test_handles_exception_gracefully(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + + service = ConversationService() + # Pass an invalid ObjectId + result = service.get_conversation("not-an-objectid", "user_123") + assert result is None + + +@pytest.mark.unit +class TestSaveConversationExtended: + + def test_raises_for_none_token(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + + service = ConversationService() + with pytest.raises(ValueError, match="Invalid or missing authentication"): + service.save_conversation( + conversation_id=None, + question="Q", + response="A", + thought="", + sources=[], + tool_calls=[], + llm=Mock(), + model_id="m", + decoded_token=None, + ) + + def test_update_existing_at_index(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + conv_id = ObjectId() + collection.insert_one( + { + "_id": conv_id, + "user": "user_123", + "name": "Conv", + "queries": [ + { + "prompt": "Q1", + "response": "A1", + "thought": "", + "sources": [], + "tool_calls": [], + }, + { + "prompt": "Q2", + "response": "A2", + "thought": "", + "sources": [], + "tool_calls": [], + }, + ], + } + ) + + result = service.save_conversation( + conversation_id=str(conv_id), + question="Q1_updated", + response="A1_updated", + thought="thinking", + sources=[], + tool_calls=[], + llm=Mock(), + model_id="gpt-4", + decoded_token={"sub": "user_123"}, + index=0, + ) + assert result == str(conv_id) + + saved = collection.find_one({"_id": conv_id}) + assert saved["queries"][0]["prompt"] == "Q1_updated" + assert saved["queries"][0]["response"] == "A1_updated" + + def test_update_at_index_unauthorized(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + conv_id = ObjectId() + collection.insert_one( + { + "_id": conv_id, + "user": "owner", + "queries": [{"prompt": "Q", "response": "A"}], + } + ) + + with pytest.raises(ValueError, match="not found or unauthorized"): + service.save_conversation( + conversation_id=str(conv_id), + question="Hack", + response="Attempt", + thought="", + sources=[], + tool_calls=[], + llm=Mock(), + model_id="m", + decoded_token={"sub": "hacker"}, + index=0, + ) + + def test_saves_metadata(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + mock_llm = Mock() + mock_llm.gen.return_value = "Title" + + conv_id = service.save_conversation( + conversation_id=None, + question="Q", + response="A", + thought="", + sources=[], + tool_calls=[], + llm=mock_llm, + model_id="m", + decoded_token={"sub": "user_123"}, + metadata={"search_query": "rewritten query"}, + ) + + saved = collection.find_one({"_id": ObjectId(conv_id)}) + assert saved["queries"][0]["metadata"] == {"search_query": "rewritten query"} + + def test_no_metadata_when_none(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + mock_llm = Mock() + mock_llm.gen.return_value = "Title" + + conv_id = service.save_conversation( + conversation_id=None, + question="Q", + response="A", + thought="", + sources=[], + tool_calls=[], + llm=mock_llm, + model_id="m", + decoded_token={"sub": "user_123"}, + metadata=None, + ) + + saved = collection.find_one({"_id": ObjectId(conv_id)}) + assert "metadata" not in saved["queries"][0] + + def test_saves_with_api_key_and_agent(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + agents_collection = mock_mongo_db[settings.MONGO_DB_NAME]["agents"] + + agent_id = ObjectId() + agents_collection.insert_one( + {"_id": agent_id, "key": "agent_key_123", "user": "user_123"} + ) + + mock_llm = Mock() + mock_llm.gen.return_value = "Title" + + conv_id = service.save_conversation( + conversation_id=None, + question="Q", + response="A", + thought="", + sources=[], + tool_calls=[], + llm=mock_llm, + model_id="m", + decoded_token={"sub": "user_123"}, + api_key="agent_key_123", + agent_id=str(agent_id), + ) + + saved = collection.find_one({"_id": ObjectId(conv_id)}) + assert saved["api_key"] == "agent_key_123" + assert saved["agent_id"] == str(agent_id) + + def test_empty_completion_uses_question_prefix(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + mock_llm = Mock() + mock_llm.gen.return_value = " " # whitespace only + + conv_id = service.save_conversation( + conversation_id=None, + question="What is the meaning of life in programming?", + response="42", + thought="", + sources=[], + tool_calls=[], + llm=mock_llm, + model_id="m", + decoded_token={"sub": "user_123"}, + ) + + saved = collection.find_one({"_id": ObjectId(conv_id)}) + assert saved["name"] == "What is the meaning of life in programming?"[:50] + + +@pytest.mark.unit +class TestUpdateCompressionMetadata: + + def test_updates_compression_fields(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + conv_id = ObjectId() + collection.insert_one( + {"_id": conv_id, "user": "u", "queries": []} + ) + + meta = { + "timestamp": datetime.now(timezone.utc), + "compressed_summary": "Summary of conversation", + "model_used": "gpt-4", + } + + service.update_compression_metadata(str(conv_id), meta) + + saved = collection.find_one({"_id": conv_id}) + assert saved["compression_metadata"]["is_compressed"] is True + assert len(saved["compression_metadata"]["compression_points"]) == 1 + + +@pytest.mark.unit +class TestAppendCompressionMessage: + + def test_appends_summary_query(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + conv_id = ObjectId() + collection.insert_one( + {"_id": conv_id, "user": "u", "queries": []} + ) + + meta = { + "compressed_summary": "This is the summary", + "timestamp": datetime.now(timezone.utc), + "model_used": "gpt-4", + } + + service.append_compression_message(str(conv_id), meta) + + saved = collection.find_one({"_id": conv_id}) + assert len(saved["queries"]) == 1 + assert saved["queries"][0]["prompt"] == "[Context Compression Summary]" + assert saved["queries"][0]["response"] == "This is the summary" + + def test_empty_summary_does_nothing(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + conv_id = ObjectId() + collection.insert_one( + {"_id": conv_id, "user": "u", "queries": []} + ) + + service.append_compression_message(str(conv_id), {"compressed_summary": ""}) + + saved = collection.find_one({"_id": conv_id}) + assert len(saved["queries"]) == 0 + + +@pytest.mark.unit +class TestGetCompressionMetadata: + + def test_returns_metadata(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + conv_id = ObjectId() + collection.insert_one( + { + "_id": conv_id, + "user": "u", + "compression_metadata": {"is_compressed": True}, + } + ) + + result = service.get_compression_metadata(str(conv_id)) + assert result is not None + assert result["is_compressed"] is True + + def test_returns_none_for_no_metadata(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + from application.core.settings import settings + + service = ConversationService() + collection = mock_mongo_db[settings.MONGO_DB_NAME]["conversations"] + + conv_id = ObjectId() + collection.insert_one({"_id": conv_id, "user": "u"}) + + result = service.get_compression_metadata(str(conv_id)) + assert result is None + + def test_returns_none_for_missing_conversation(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + + service = ConversationService() + result = service.get_compression_metadata(str(ObjectId())) + assert result is None + + def test_handles_invalid_id(self, mock_mongo_db): + from application.api.answer.services.conversation_service import ( + ConversationService, + ) + + service = ConversationService() + result = service.get_compression_metadata("invalid-id") + assert result is None diff --git a/tests/api/answer/test_stream_processor.py b/tests/api/answer/test_stream_processor.py index 1ba4a130..ae7fdc8e 100644 --- a/tests/api/answer/test_stream_processor.py +++ b/tests/api/answer/test_stream_processor.py @@ -1,4 +1,17 @@ -"""Tests for application/api/answer/services/stream_processor.py — get_prompt and helpers.""" +"""Tests for application/api/answer/services/stream_processor.py — get_prompt and helpers. + +Extended coverage for StreamProcessor including: + - get_prompt: all presets and DB fallback + - StreamProcessor init, _resolve_agent_id, _get_prompt_content + - _get_required_tool_actions + - _get_attachments_content: valid, invalid, empty + - _configure_retriever + - _validate_and_set_model + - _get_agent_key + - _get_data_from_api_key + - _configure_source + - pre_fetch_docs +""" from unittest.mock import MagicMock, patch @@ -119,6 +132,22 @@ class TestStreamProcessorInit: sp = StreamProcessor(request_data={}, decoded_token=None) assert sp.initial_user_id is None + @pytest.mark.unit + def test_init_default_model_and_config(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor(request_data={}, decoded_token={"sub": "u"}) + assert sp.model_id is None + assert sp.is_shared_usage is False + assert sp.shared_token is None + assert sp.compressed_summary is None + assert sp.compressed_summary_tokens == 0 + class TestGetAttachmentsContent: @@ -169,6 +198,19 @@ class TestGetAttachmentsContent: result = sp._get_attachments_content(["bad"], "u") assert result == [] + @pytest.mark.unit + def test_none_ids_returns_empty(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor(request_data={}, decoded_token={"sub": "u"}) + result = sp._get_attachments_content(None, "u") + assert result == [] + class TestResolveAgentId: @@ -250,6 +292,23 @@ class TestResolveAgentId: sp.conversation_service.get_conversation.side_effect = Exception("db error") assert sp._resolve_agent_id() is None + @pytest.mark.unit + def test_conversation_without_agent_id(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor( + request_data={"conversation_id": "conv1"}, + decoded_token={"sub": "u"}, + ) + sp.conversation_service = MagicMock() + sp.conversation_service.get_conversation.return_value = {"name": "test conv"} + assert sp._resolve_agent_id() is None + class TestGetPromptContent: @@ -299,6 +358,19 @@ class TestGetPromptContent: sp.agent_config = {"prompt_id": "bad_id"} assert sp._get_prompt_content() is None + @pytest.mark.unit + def test_agent_config_not_dict(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor(request_data={}, decoded_token={"sub": "u"}) + sp.agent_config = "not_a_dict" + assert sp._get_prompt_content() is None + class TestGetRequiredToolActions: @@ -329,3 +401,114 @@ class TestGetRequiredToolActions: sp._prompt_content = "No template syntax here" result = sp._get_required_tool_actions() assert result == {} + + @pytest.mark.unit + def test_caches_result(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor(request_data={}, decoded_token={"sub": "u"}) + sp._required_tool_actions = {"tool1": {"action1"}} + result = sp._get_required_tool_actions() + assert result == {"tool1": {"action1"}} + + +class TestConfigureRetriever: + + @pytest.mark.unit + def test_default_values(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor( + request_data={"question": "Q"}, + decoded_token={"sub": "u"}, + ) + sp.model_id = "test-model" + sp.agent_key = None + sp._configure_retriever() + assert sp.retriever_config["retriever_name"] == "classic" + assert sp.retriever_config["chunks"] == 2 + + @pytest.mark.unit + def test_isNoneDoc_sets_zero_chunks(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor( + request_data={"question": "Q", "isNoneDoc": True}, + decoded_token={"sub": "u"}, + ) + sp.model_id = "test-model" + sp.agent_key = None + sp._configure_retriever() + assert sp.retriever_config["chunks"] == 0 + + @pytest.mark.unit + def test_custom_retriever_and_chunks(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor( + request_data={"question": "Q", "retriever": "hybrid", "chunks": "5"}, + decoded_token={"sub": "u"}, + ) + sp.model_id = "test-model" + sp.agent_key = None + sp._configure_retriever() + assert sp.retriever_config["retriever_name"] == "hybrid" + assert sp.retriever_config["chunks"] == 5 + + +class TestConfigureSource: + + @pytest.mark.unit + def test_active_docs_from_request(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor( + request_data={"question": "Q", "active_docs": "source_123"}, + decoded_token={"sub": "u"}, + ) + sp.agent_key = None + sp._configure_source() + assert sp.source == {"active_docs": "source_123"} + + @pytest.mark.unit + def test_no_source_config(self): + mock_db = MagicMock() + with patch("application.api.answer.services.stream_processor.MongoDB") as MockMongo, \ + patch("application.api.answer.services.stream_processor.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + MockMongo.get_client.return_value = {"docsgpt": mock_db} + + from application.api.answer.services.stream_processor import StreamProcessor + sp = StreamProcessor( + request_data={"question": "Q"}, + decoded_token={"sub": "u"}, + ) + sp.agent_key = None + sp._configure_source() + assert sp.source == {} + assert sp.all_sources == [] diff --git a/tests/api/test_internal_routes.py b/tests/api/test_internal_routes.py new file mode 100644 index 00000000..dd704e23 --- /dev/null +++ b/tests/api/test_internal_routes.py @@ -0,0 +1,416 @@ +"""Unit tests for application/api/internal/routes.py. + +Covers: + - verify_internal_key: key validation + - /api/download: file download + - /api/upload_index: index file upload (existing & new entries) +""" + +import io +import json +from unittest.mock import MagicMock + +import pytest +from bson.objectid import ObjectId + + +@pytest.fixture +def internal_app(monkeypatch, mock_mongo_db): + """Create a Flask app with the internal blueprint registered.""" + from flask import Flask + + # Patch module-level MongoDB references before importing routes + from application.core.settings import settings + + db = mock_mongo_db[settings.MONGO_DB_NAME] + monkeypatch.setattr( + "application.api.internal.routes.conversations_collection", + db["conversations"], + ) + monkeypatch.setattr( + "application.api.internal.routes.sources_collection", + db["sources"], + ) + + from application.api.internal.routes import internal + + app = Flask(__name__) + app.register_blueprint(internal) + app.config["TESTING"] = True + return app, db + + +# --------------------------------------------------------------------------- +# verify_internal_key +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestVerifyInternalKey: + + def test_no_internal_key_configured_allows_access( + self, internal_app, monkeypatch + ): + app, db = internal_app + monkeypatch.setattr( + "application.api.internal.routes.settings", + MagicMock( + INTERNAL_KEY=None, + UPLOAD_FOLDER="uploads", + VECTOR_STORE="faiss", + EMBEDDINGS_NAME="test", + MONGO_DB_NAME="docsgpt", + ), + ) + with app.test_client() as client: + # download will fail for missing file but should not be 401 + resp = client.get("/api/download?user=u&name=n&file=f") + assert resp.status_code != 401 + + def test_missing_key_returns_401(self, internal_app, monkeypatch): + app, db = internal_app + monkeypatch.setattr( + "application.api.internal.routes.settings", + MagicMock( + INTERNAL_KEY="secret123", + UPLOAD_FOLDER="uploads", + VECTOR_STORE="faiss", + EMBEDDINGS_NAME="test", + MONGO_DB_NAME="docsgpt", + ), + ) + with app.test_client() as client: + resp = client.get("/api/download?user=u&name=n&file=f") + assert resp.status_code == 401 + + def test_wrong_key_returns_401(self, internal_app, monkeypatch): + app, db = internal_app + monkeypatch.setattr( + "application.api.internal.routes.settings", + MagicMock( + INTERNAL_KEY="secret123", + UPLOAD_FOLDER="uploads", + VECTOR_STORE="faiss", + EMBEDDINGS_NAME="test", + MONGO_DB_NAME="docsgpt", + ), + ) + with app.test_client() as client: + resp = client.get( + "/api/download?user=u&name=n&file=f", + headers={"X-Internal-Key": "wrong"}, + ) + assert resp.status_code == 401 + + def test_correct_key_allows_access(self, internal_app, monkeypatch): + app, db = internal_app + monkeypatch.setattr( + "application.api.internal.routes.settings", + MagicMock( + INTERNAL_KEY="secret123", + UPLOAD_FOLDER="uploads", + VECTOR_STORE="faiss", + EMBEDDINGS_NAME="test", + MONGO_DB_NAME="docsgpt", + ), + ) + with app.test_client() as client: + # Will 404 for missing file, but should pass auth check + resp = client.get( + "/api/download?user=u&name=n&file=f", + headers={"X-Internal-Key": "secret123"}, + ) + assert resp.status_code != 401 + + +# --------------------------------------------------------------------------- +# /api/upload_index +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestUploadIndex: + + def _make_settings(self, vector_store="faiss"): + return MagicMock( + INTERNAL_KEY=None, + UPLOAD_FOLDER="uploads", + VECTOR_STORE=vector_store, + EMBEDDINGS_NAME="test_embeddings", + MONGO_DB_NAME="docsgpt", + ) + + def test_missing_user_returns_no_user(self, internal_app, monkeypatch): + app, db = internal_app + monkeypatch.setattr( + "application.api.internal.routes.settings", self._make_settings() + ) + with app.test_client() as client: + resp = client.post("/api/upload_index", data={}) + assert resp.json["status"] == "no user" + + def test_missing_name_returns_no_name(self, internal_app, monkeypatch): + app, db = internal_app + monkeypatch.setattr( + "application.api.internal.routes.settings", self._make_settings() + ) + with app.test_client() as client: + resp = client.post("/api/upload_index", data={"user": "testuser"}) + assert resp.json["status"] == "no name" + + def test_creates_new_source_entry(self, internal_app, monkeypatch): + app, db = internal_app + doc_id = str(ObjectId()) + settings_mock = self._make_settings(vector_store="other") + monkeypatch.setattr( + "application.api.internal.routes.settings", settings_mock + ) + mock_storage = MagicMock() + monkeypatch.setattr( + "application.api.internal.routes.StorageCreator", + MagicMock(get_storage=MagicMock(return_value=mock_storage)), + ) + + with app.test_client() as client: + resp = client.post( + "/api/upload_index", + data={ + "user": "testuser", + "name": "testjob", + "tokens": "100", + "retriever": "classic", + "id": doc_id, + "type": "local", + }, + ) + assert resp.json["status"] == "ok" + + entry = db["sources"].find_one({"_id": ObjectId(doc_id)}) + assert entry is not None + assert entry["user"] == "testuser" + assert entry["name"] == "testjob" + + def test_updates_existing_source_entry(self, internal_app, monkeypatch): + app, db = internal_app + doc_id = ObjectId() + settings_mock = self._make_settings(vector_store="other") + monkeypatch.setattr( + "application.api.internal.routes.settings", settings_mock + ) + mock_storage = MagicMock() + monkeypatch.setattr( + "application.api.internal.routes.StorageCreator", + MagicMock(get_storage=MagicMock(return_value=mock_storage)), + ) + + # Insert existing entry + db["sources"].insert_one( + {"_id": doc_id, "user": "old_user", "name": "old_name"} + ) + + with app.test_client() as client: + resp = client.post( + "/api/upload_index", + data={ + "user": "new_user", + "name": "new_name", + "tokens": "200", + "retriever": "hybrid", + "id": str(doc_id), + "type": "remote", + }, + ) + assert resp.json["status"] == "ok" + + entry = db["sources"].find_one({"_id": doc_id}) + assert entry["user"] == "new_user" + assert entry["name"] == "new_name" + + def test_parses_directory_structure_json(self, internal_app, monkeypatch): + app, db = internal_app + doc_id = str(ObjectId()) + settings_mock = self._make_settings(vector_store="other") + monkeypatch.setattr( + "application.api.internal.routes.settings", settings_mock + ) + mock_storage = MagicMock() + monkeypatch.setattr( + "application.api.internal.routes.StorageCreator", + MagicMock(get_storage=MagicMock(return_value=mock_storage)), + ) + + dir_struct = {"root": {"files": ["a.txt", "b.txt"]}} + with app.test_client() as client: + resp = client.post( + "/api/upload_index", + data={ + "user": "u", + "name": "n", + "tokens": "0", + "retriever": "classic", + "id": doc_id, + "type": "local", + "directory_structure": json.dumps(dir_struct), + }, + ) + assert resp.json["status"] == "ok" + + entry = db["sources"].find_one({"_id": ObjectId(doc_id)}) + assert entry["directory_structure"] == dir_struct + + def test_invalid_directory_structure_defaults_empty( + self, internal_app, monkeypatch + ): + app, db = internal_app + doc_id = str(ObjectId()) + settings_mock = self._make_settings(vector_store="other") + monkeypatch.setattr( + "application.api.internal.routes.settings", settings_mock + ) + mock_storage = MagicMock() + monkeypatch.setattr( + "application.api.internal.routes.StorageCreator", + MagicMock(get_storage=MagicMock(return_value=mock_storage)), + ) + + with app.test_client() as client: + resp = client.post( + "/api/upload_index", + data={ + "user": "u", + "name": "n", + "tokens": "0", + "retriever": "classic", + "id": doc_id, + "type": "local", + "directory_structure": "not valid json", + }, + ) + assert resp.json["status"] == "ok" + + entry = db["sources"].find_one({"_id": ObjectId(doc_id)}) + assert entry["directory_structure"] == {} + + def test_file_name_map_parsed(self, internal_app, monkeypatch): + app, db = internal_app + doc_id = str(ObjectId()) + settings_mock = self._make_settings(vector_store="other") + monkeypatch.setattr( + "application.api.internal.routes.settings", settings_mock + ) + mock_storage = MagicMock() + monkeypatch.setattr( + "application.api.internal.routes.StorageCreator", + MagicMock(get_storage=MagicMock(return_value=mock_storage)), + ) + + fmap = {"hash1": "file1.txt"} + with app.test_client() as client: + resp = client.post( + "/api/upload_index", + data={ + "user": "u", + "name": "n", + "tokens": "0", + "retriever": "classic", + "id": doc_id, + "type": "local", + "file_name_map": json.dumps(fmap), + }, + ) + assert resp.json["status"] == "ok" + + entry = db["sources"].find_one({"_id": ObjectId(doc_id)}) + assert entry["file_name_map"] == fmap + + def test_faiss_missing_files_returns_no_file( + self, internal_app, monkeypatch + ): + app, db = internal_app + doc_id = str(ObjectId()) + settings_mock = self._make_settings(vector_store="faiss") + monkeypatch.setattr( + "application.api.internal.routes.settings", settings_mock + ) + mock_storage = MagicMock() + monkeypatch.setattr( + "application.api.internal.routes.StorageCreator", + MagicMock(get_storage=MagicMock(return_value=mock_storage)), + ) + + with app.test_client() as client: + resp = client.post( + "/api/upload_index", + data={ + "user": "u", + "name": "n", + "tokens": "0", + "retriever": "classic", + "id": doc_id, + "type": "local", + }, + ) + assert resp.json["status"] == "no file" + + def test_faiss_empty_filename_returns_no_file_name( + self, internal_app, monkeypatch + ): + app, db = internal_app + doc_id = str(ObjectId()) + settings_mock = self._make_settings(vector_store="faiss") + monkeypatch.setattr( + "application.api.internal.routes.settings", settings_mock + ) + mock_storage = MagicMock() + monkeypatch.setattr( + "application.api.internal.routes.StorageCreator", + MagicMock(get_storage=MagicMock(return_value=mock_storage)), + ) + + with app.test_client() as client: + resp = client.post( + "/api/upload_index", + data={ + "user": "u", + "name": "n", + "tokens": "0", + "retriever": "classic", + "id": doc_id, + "type": "local", + "file_faiss": (io.BytesIO(b""), ""), + }, + ) + assert resp.json["status"] == "no file name" + + def test_remote_data_and_sync_frequency(self, internal_app, monkeypatch): + app, db = internal_app + doc_id = str(ObjectId()) + settings_mock = self._make_settings(vector_store="other") + monkeypatch.setattr( + "application.api.internal.routes.settings", settings_mock + ) + mock_storage = MagicMock() + monkeypatch.setattr( + "application.api.internal.routes.StorageCreator", + MagicMock(get_storage=MagicMock(return_value=mock_storage)), + ) + + with app.test_client() as client: + resp = client.post( + "/api/upload_index", + data={ + "user": "u", + "name": "n", + "tokens": "0", + "retriever": "classic", + "id": doc_id, + "type": "remote", + "remote_data": '{"url":"http://example.com"}', + "sync_frequency": "daily", + }, + ) + assert resp.json["status"] == "ok" + + entry = db["sources"].find_one({"_id": ObjectId(doc_id)}) + assert entry["sync_frequency"] == "daily" + assert entry["remote_data"] == '{"url":"http://example.com"}' diff --git a/tests/api/user/__init__.py b/tests/api/user/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/api/user/sources/test_chunks.py b/tests/api/user/sources/test_chunks.py new file mode 100644 index 00000000..7702210e --- /dev/null +++ b/tests/api/user/sources/test_chunks.py @@ -0,0 +1,879 @@ +"""Tests for source chunk management routes.""" + +import pytest +from unittest.mock import Mock, patch +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +def _status(response): + if isinstance(response, tuple): + return response[1] + return response.status_code + + +def _json(response): + if isinstance(response, tuple): + return response[0].json + return response.json + + +# --------------------------------------------------------------------------- +# GetChunks +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestGetChunks: + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.chunks import GetChunks + + with app.test_request_context("/api/get_chunks?id=abc"): + from flask import request + + request.decoded_token = None + response = GetChunks().get() + + assert _status(response) == 401 + + def test_returns_400_for_invalid_doc_id(self, app): + from application.api.user.sources.chunks import GetChunks + + with app.test_request_context("/api/get_chunks?id=invalid"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetChunks().get() + + assert _status(response) == 400 + assert "Invalid doc_id" in _json(response)["error"] + + def test_returns_404_when_doc_not_found(self, app): + from application.api.user.sources.chunks import GetChunks + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ): + with app.test_request_context(f"/api/get_chunks?id={doc_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetChunks().get() + + assert _status(response) == 404 + assert "not found" in _json(response)["error"] + + def test_returns_paginated_chunks(self, app): + from application.api.user.sources.chunks import GetChunks + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + + chunks = [ + {"text": f"chunk {i}", "metadata": {}, "doc_id": f"c{i}"} + for i in range(25) + ] + mock_store = Mock() + mock_store.get_chunks.return_value = chunks + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + f"/api/get_chunks?id={doc_id}&page=2&per_page=10" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = GetChunks().get() + + assert _status(response) == 200 + data = _json(response) + assert data["total"] == 25 + assert data["page"] == 2 + assert data["per_page"] == 10 + assert len(data["chunks"]) == 10 + + def test_filters_chunks_by_path(self, app): + from application.api.user.sources.chunks import GetChunks + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + + chunks = [ + {"text": "a", "metadata": {"source": "inputs/dir/file.pdf"}, "doc_id": "c1"}, + {"text": "b", "metadata": {"source": "inputs/other.txt"}, "doc_id": "c2"}, + {"text": "c", "metadata": {"file_path": "guides/setup.md"}, "doc_id": "c3"}, + ] + mock_store = Mock() + mock_store.get_chunks.return_value = chunks + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + f"/api/get_chunks?id={doc_id}&path=file.pdf" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = GetChunks().get() + + data = _json(response) + assert data["total"] == 1 + assert data["chunks"][0]["text"] == "a" + assert data["path"] == "file.pdf" + + def test_filters_chunks_by_file_path_metadata(self, app): + from application.api.user.sources.chunks import GetChunks + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + + chunks = [ + {"text": "a", "metadata": {"source": "inputs/dir/file.pdf"}, "doc_id": "c1"}, + {"text": "c", "metadata": {"file_path": "guides/setup.md"}, "doc_id": "c3"}, + ] + mock_store = Mock() + mock_store.get_chunks.return_value = chunks + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + f"/api/get_chunks?id={doc_id}&path=setup.md" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = GetChunks().get() + + data = _json(response) + assert data["total"] == 1 + assert data["chunks"][0]["text"] == "c" + + def test_filters_chunks_by_search_term(self, app): + from application.api.user.sources.chunks import GetChunks + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + + chunks = [ + {"text": "Python is great", "metadata": {"title": "intro"}, "doc_id": "c1"}, + {"text": "Java tutorial", "metadata": {"title": "java guide"}, "doc_id": "c2"}, + {"text": "Hello world", "metadata": {"title": "Python Basics"}, "doc_id": "c3"}, + ] + mock_store = Mock() + mock_store.get_chunks.return_value = chunks + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + f"/api/get_chunks?id={doc_id}&search=python" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = GetChunks().get() + + data = _json(response) + assert data["total"] == 2 + assert data["search"] == "python" + + def test_combines_path_and_search_filters(self, app): + from application.api.user.sources.chunks import GetChunks + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + + chunks = [ + {"text": "Python intro", "metadata": {"source": "dir/intro.md", "title": ""}, "doc_id": "c1"}, + {"text": "Python deep", "metadata": {"source": "dir/deep.md", "title": ""}, "doc_id": "c2"}, + {"text": "Java intro", "metadata": {"source": "dir/intro.md", "title": ""}, "doc_id": "c3"}, + ] + mock_store = Mock() + mock_store.get_chunks.return_value = chunks + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + f"/api/get_chunks?id={doc_id}&path=intro.md&search=python" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = GetChunks().get() + + data = _json(response) + assert data["total"] == 1 + assert data["chunks"][0]["doc_id"] == "c1" + + def test_returns_500_on_store_error(self, app): + from application.api.user.sources.chunks import GetChunks + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.get_chunks.side_effect = Exception("Store error") + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context(f"/api/get_chunks?id={doc_id}"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = GetChunks().get() + + assert _status(response) == 500 + + def test_no_path_or_search_returns_null_fields(self, app): + from application.api.user.sources.chunks import GetChunks + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.get_chunks.return_value = [] + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context(f"/api/get_chunks?id={doc_id}"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = GetChunks().get() + + data = _json(response) + assert data["path"] is None + assert data["search"] is None + + +# --------------------------------------------------------------------------- +# AddChunk +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestAddChunk: + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.chunks import AddChunk + + with app.test_request_context( + "/api/add_chunk", method="POST", json={"id": "abc", "text": "hi"} + ): + from flask import request + + request.decoded_token = None + response = AddChunk().post() + + assert _status(response) == 401 + + def test_returns_400_missing_required_fields(self, app): + from application.api.user.sources.chunks import AddChunk + + with app.test_request_context( + "/api/add_chunk", method="POST", json={"id": str(ObjectId())} + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = AddChunk().post() + + # check_required_fields returns a tuple (response, status) + assert response is not None + + def test_returns_400_for_invalid_doc_id(self, app): + from application.api.user.sources.chunks import AddChunk + + with app.test_request_context( + "/api/add_chunk", method="POST", json={"id": "bad", "text": "hi"} + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = AddChunk().post() + + assert _status(response) == 400 + assert "Invalid doc_id" in _json(response)["error"] + + def test_returns_404_when_doc_not_found(self, app): + from application.api.user.sources.chunks import AddChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/add_chunk", method="POST", + json={"id": doc_id, "text": "hello"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = AddChunk().post() + + assert _status(response) == 404 + + def test_adds_chunk_successfully(self, app): + from application.api.user.sources.chunks import AddChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.add_chunk.return_value = "new-chunk-id" + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ), patch( + "application.api.user.sources.chunks.num_tokens_from_string", + return_value=5, + ): + with app.test_request_context( + "/api/add_chunk", method="POST", + json={"id": doc_id, "text": "hello world"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = AddChunk().post() + + assert _status(response) == 201 + data = _json(response) + assert data["chunk_id"] == "new-chunk-id" + assert "successfully" in data["message"] + call_args = mock_store.add_chunk.call_args + assert call_args[0][0] == "hello world" + assert call_args[0][1]["token_count"] == 5 + + def test_adds_chunk_with_custom_metadata(self, app): + from application.api.user.sources.chunks import AddChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.add_chunk.return_value = "cid" + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ), patch( + "application.api.user.sources.chunks.num_tokens_from_string", + return_value=3, + ): + with app.test_request_context( + "/api/add_chunk", method="POST", + json={ + "id": doc_id, + "text": "hi", + "metadata": {"source": "test.pdf"}, + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = AddChunk().post() + + assert _status(response) == 201 + meta = mock_store.add_chunk.call_args[0][1] + assert meta["source"] == "test.pdf" + assert meta["token_count"] == 3 + + def test_returns_500_on_store_error(self, app): + from application.api.user.sources.chunks import AddChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.add_chunk.side_effect = Exception("fail") + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ), patch( + "application.api.user.sources.chunks.num_tokens_from_string", + return_value=1, + ): + with app.test_request_context( + "/api/add_chunk", method="POST", + json={"id": doc_id, "text": "hello"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = AddChunk().post() + + assert _status(response) == 500 + + +# --------------------------------------------------------------------------- +# DeleteChunk +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestDeleteChunk: + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.chunks import DeleteChunk + + with app.test_request_context("/api/delete_chunk?id=abc&chunk_id=xyz"): + from flask import request + + request.decoded_token = None + response = DeleteChunk().delete() + + assert _status(response) == 401 + + def test_returns_400_for_invalid_doc_id(self, app): + from application.api.user.sources.chunks import DeleteChunk + + with app.test_request_context("/api/delete_chunk?id=bad&chunk_id=xyz"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteChunk().delete() + + assert _status(response) == 400 + assert "Invalid doc_id" in _json(response)["error"] + + def test_returns_404_when_doc_not_found(self, app): + from application.api.user.sources.chunks import DeleteChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/delete_chunk?id={doc_id}&chunk_id=cid" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteChunk().delete() + + assert _status(response) == 404 + + def test_deletes_chunk_successfully(self, app): + from application.api.user.sources.chunks import DeleteChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.delete_chunk.return_value = True + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + f"/api/delete_chunk?id={doc_id}&chunk_id=cid" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteChunk().delete() + + assert _status(response) == 200 + assert "successfully" in _json(response)["message"] + mock_store.delete_chunk.assert_called_once_with("cid") + + def test_returns_404_when_chunk_not_deleted(self, app): + from application.api.user.sources.chunks import DeleteChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.delete_chunk.return_value = False + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + f"/api/delete_chunk?id={doc_id}&chunk_id=missing" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteChunk().delete() + + assert _status(response) == 404 + assert "not found" in _json(response)["message"] + + def test_returns_500_on_store_error(self, app): + from application.api.user.sources.chunks import DeleteChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.delete_chunk.side_effect = Exception("boom") + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + f"/api/delete_chunk?id={doc_id}&chunk_id=cid" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteChunk().delete() + + assert _status(response) == 500 + + +# --------------------------------------------------------------------------- +# UpdateChunk +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestUpdateChunk: + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.chunks import UpdateChunk + + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={"id": "abc", "chunk_id": "cid"}, + ): + from flask import request + + request.decoded_token = None + response = UpdateChunk().put() + + assert _status(response) == 401 + + def test_returns_400_missing_required_fields(self, app): + from application.api.user.sources.chunks import UpdateChunk + + with app.test_request_context( + "/api/update_chunk", method="PUT", json={"id": str(ObjectId())} + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + assert response is not None + + def test_returns_400_for_invalid_doc_id(self, app): + from application.api.user.sources.chunks import UpdateChunk + + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={"id": "bad", "chunk_id": "cid"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + assert _status(response) == 400 + + def test_returns_404_when_doc_not_found(self, app): + from application.api.user.sources.chunks import UpdateChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={"id": doc_id, "chunk_id": "cid"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + assert _status(response) == 404 + + def test_returns_404_when_chunk_not_found(self, app): + from application.api.user.sources.chunks import UpdateChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.get_chunks.return_value = [ + {"doc_id": "other", "text": "x", "metadata": {}}, + ] + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={"id": doc_id, "chunk_id": "missing"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + assert _status(response) == 404 + assert "Chunk not found" in _json(response)["error"] + + def test_updates_chunk_text_successfully(self, app): + from application.api.user.sources.chunks import UpdateChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.get_chunks.return_value = [ + {"doc_id": "cid", "text": "old text", "metadata": {"source": "f.pdf"}}, + ] + mock_store.add_chunk.return_value = "new-cid" + mock_store.delete_chunk.return_value = True + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ), patch( + "application.api.user.sources.chunks.num_tokens_from_string", + return_value=7, + ): + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={"id": doc_id, "chunk_id": "cid", "text": "new text"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + assert _status(response) == 200 + data = _json(response) + assert data["chunk_id"] == "new-cid" + assert data["original_chunk_id"] == "cid" + # Verify add was called with new text and merged metadata + add_call = mock_store.add_chunk.call_args + assert add_call[0][0] == "new text" + assert add_call[0][1]["source"] == "f.pdf" + assert add_call[0][1]["token_count"] == 7 + mock_store.delete_chunk.assert_called_once_with("cid") + + def test_updates_chunk_metadata_only(self, app): + from application.api.user.sources.chunks import UpdateChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.get_chunks.return_value = [ + {"doc_id": "cid", "text": "keep me", "metadata": {"source": "f.pdf"}}, + ] + mock_store.add_chunk.return_value = "new-cid" + mock_store.delete_chunk.return_value = True + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={ + "id": doc_id, + "chunk_id": "cid", + "metadata": {"title": "new title"}, + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + assert _status(response) == 200 + add_call = mock_store.add_chunk.call_args + # text should be preserved + assert add_call[0][0] == "keep me" + # metadata should be merged + assert add_call[0][1]["source"] == "f.pdf" + assert add_call[0][1]["title"] == "new title" + + def test_update_warns_when_old_chunk_delete_fails(self, app): + from application.api.user.sources.chunks import UpdateChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.get_chunks.return_value = [ + {"doc_id": "cid", "text": "text", "metadata": {}}, + ] + mock_store.add_chunk.return_value = "new-cid" + mock_store.delete_chunk.return_value = False # delete fails + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={"id": doc_id, "chunk_id": "cid"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + # Still returns 200 with a warning logged + assert _status(response) == 200 + + def test_returns_500_when_add_chunk_fails(self, app): + from application.api.user.sources.chunks import UpdateChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.get_chunks.return_value = [ + {"doc_id": "cid", "text": "text", "metadata": {}}, + ] + mock_store.add_chunk.side_effect = Exception("add failed") + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={"id": doc_id, "chunk_id": "cid", "text": "new"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + assert _status(response) == 500 + assert "addition failed" in _json(response)["error"] + + def test_returns_500_on_general_store_error(self, app): + from application.api.user.sources.chunks import UpdateChunk + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = {"_id": ObjectId(doc_id), "user": "u1"} + mock_store = Mock() + mock_store.get_chunks.side_effect = Exception("connection lost") + + with patch( + "application.api.user.sources.chunks.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.chunks.get_vector_store", + return_value=mock_store, + ): + with app.test_request_context( + "/api/update_chunk", method="PUT", + json={"id": doc_id, "chunk_id": "cid"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UpdateChunk().put() + + assert _status(response) == 500 diff --git a/tests/api/user/sources/test_source_routes.py b/tests/api/user/sources/test_source_routes.py new file mode 100644 index 00000000..eed326aa --- /dev/null +++ b/tests/api/user/sources/test_source_routes.py @@ -0,0 +1,965 @@ +"""Tests for source management routes (CombinedJson, PaginatedSources, +DeleteByIds, DeleteOldIndexes, ManageSync, DirectoryStructure). + +Note: SyncSource and _get_provider_from_remote_data are already covered in +test_routes.py and are NOT duplicated here. +""" + +import json + +import pytest +from unittest.mock import MagicMock, Mock, patch +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +def _status(response): + if isinstance(response, tuple): + return response[1] + return response.status_code + + +def _json(response): + if isinstance(response, tuple): + return response[0].json + return response.json + + +# --------------------------------------------------------------------------- +# CombinedJson (/api/sources) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestCombinedJson: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.routes import CombinedJson + + with app.test_request_context("/api/sources"): + from flask import request + + request.decoded_token = None + response = CombinedJson().get() + + assert _status(response) == 401 + + def test_returns_default_source_plus_user_sources(self, app): + from application.api.user.sources.routes import CombinedJson + + src_id = ObjectId() + mock_cursor = MagicMock() + mock_cursor.sort.return_value = [ + { + "_id": src_id, + "name": "My Doc", + "date": "2024-01-01", + "tokens": "100", + "retriever": "classic", + "sync_frequency": "daily", + "remote_data": json.dumps({"provider": "github"}), + "directory_structure": None, + "type": "file", + } + ] + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/sources"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CombinedJson().get() + + assert _status(response) == 200 + data = _json(response) + # First entry is always the Default + assert data[0]["name"] == "Default" + assert data[0]["date"] == "default" + # Second entry is user source + assert data[1]["id"] == str(src_id) + assert data[1]["name"] == "My Doc" + assert data[1]["provider"] == "github" + assert data[1]["syncFrequency"] == "daily" + assert data[1]["is_nested"] is False + + def test_is_nested_true_when_directory_structure_present(self, app): + from application.api.user.sources.routes import CombinedJson + + mock_cursor = MagicMock() + mock_cursor.sort.return_value = [ + { + "_id": ObjectId(), + "name": "Nested", + "date": "2024-01-01", + "directory_structure": {"files": ["a.txt"]}, + } + ] + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/sources"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = CombinedJson().get() + + data = _json(response) + assert data[1]["is_nested"] is True + + def test_returns_400_on_db_error(self, app): + from application.api.user.sources.routes import CombinedJson + + mock_collection = Mock() + mock_collection.find.side_effect = Exception("db err") + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/sources"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = CombinedJson().get() + + assert _status(response) == 400 + + def test_type_defaults_to_file(self, app): + from application.api.user.sources.routes import CombinedJson + + mock_cursor = MagicMock() + mock_cursor.sort.return_value = [ + {"_id": ObjectId(), "name": "X", "date": "d"} + ] + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/sources"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = CombinedJson().get() + + data = _json(response) + assert data[1]["type"] == "file" + + +# --------------------------------------------------------------------------- +# PaginatedSources (/api/sources/paginated) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestPaginatedSources: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.routes import PaginatedSources + + with app.test_request_context("/api/sources/paginated"): + from flask import request + + request.decoded_token = None + response = PaginatedSources().get() + + assert _status(response) == 401 + + def test_returns_paginated_results(self, app): + from application.api.user.sources.routes import PaginatedSources + + ids = [ObjectId() for _ in range(3)] + docs = [ + {"_id": ids[i], "name": f"Doc{i}", "date": f"2024-0{i + 1}-01"} + for i in range(3) + ] + + mock_cursor = MagicMock() + mock_cursor.sort.return_value = mock_cursor + mock_cursor.skip.return_value = mock_cursor + mock_cursor.limit.return_value = docs + + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + mock_collection.count_documents.return_value = 3 + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/sources/paginated?page=1&rows=10&sort=date&order=desc" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = PaginatedSources().get() + + assert _status(response) == 200 + data = _json(response) + assert data["total"] == 3 + assert data["totalPages"] == 1 + assert data["currentPage"] == 1 + assert len(data["paginated"]) == 3 + + def test_search_filter_applies_regex(self, app): + from application.api.user.sources.routes import PaginatedSources + + mock_cursor = MagicMock() + mock_cursor.sort.return_value = mock_cursor + mock_cursor.skip.return_value = mock_cursor + mock_cursor.limit.return_value = [] + + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + mock_collection.count_documents.return_value = 0 + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/sources/paginated?search=test%20doc" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = PaginatedSources().get() + + assert _status(response) == 200 + # Verify search query was passed + query_arg = mock_collection.count_documents.call_args[0][0] + assert query_arg["name"]["$regex"] == "test doc" + assert query_arg["name"]["$options"] == "i" + + def test_ascending_sort_order(self, app): + from application.api.user.sources.routes import PaginatedSources + + mock_cursor = MagicMock() + mock_cursor.sort.return_value = mock_cursor + mock_cursor.skip.return_value = mock_cursor + mock_cursor.limit.return_value = [] + + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + mock_collection.count_documents.return_value = 0 + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/sources/paginated?order=asc&sort=name" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = PaginatedSources().get() + + assert _status(response) == 200 + mock_cursor.sort.assert_called_once_with("name", 1) + + def test_page_clamped_to_valid_range(self, app): + from application.api.user.sources.routes import PaginatedSources + + mock_cursor = MagicMock() + mock_cursor.sort.return_value = mock_cursor + mock_cursor.skip.return_value = mock_cursor + mock_cursor.limit.return_value = [] + + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + mock_collection.count_documents.return_value = 5 # 1 page with default 10 rows + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/sources/paginated?page=999" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = PaginatedSources().get() + + data = _json(response) + assert data["currentPage"] == 1 # clamped + + def test_returns_400_on_db_error(self, app): + from application.api.user.sources.routes import PaginatedSources + + mock_collection = Mock() + mock_collection.count_documents.return_value = 0 + mock_collection.find.side_effect = Exception("db error") + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/sources/paginated"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = PaginatedSources().get() + + assert _status(response) == 400 + + def test_paginated_includes_provider_and_is_nested(self, app): + from application.api.user.sources.routes import PaginatedSources + + doc = { + "_id": ObjectId(), + "name": "S3 Src", + "date": "2024-01-01", + "remote_data": {"provider": "s3"}, + "directory_structure": {"dirs": ["a"]}, + "type": "s3", + } + + mock_cursor = MagicMock() + mock_cursor.sort.return_value = mock_cursor + mock_cursor.skip.return_value = mock_cursor + mock_cursor.limit.return_value = [doc] + + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + mock_collection.count_documents.return_value = 1 + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/sources/paginated"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = PaginatedSources().get() + + data = _json(response) + entry = data["paginated"][0] + assert entry["provider"] == "s3" + assert entry["isNested"] is True + assert entry["type"] == "s3" + + +# --------------------------------------------------------------------------- +# DeleteByIds (/api/delete_by_ids) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestDeleteByIds: + + def test_returns_400_when_path_missing(self, app): + from application.api.user.sources.routes import DeleteByIds + + with app.test_request_context("/api/delete_by_ids"): + response = DeleteByIds().get() + + assert _status(response) == 400 + assert "Missing" in _json(response)["message"] + + def test_returns_200_on_successful_delete(self, app): + from application.api.user.sources.routes import DeleteByIds + + mock_collection = Mock() + mock_collection.delete_index.return_value = True + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/delete_by_ids?path=id1,id2"): + response = DeleteByIds().get() + + assert _status(response) == 200 + assert _json(response)["success"] is True + mock_collection.delete_index.assert_called_once_with(ids="id1,id2") + + def test_returns_400_when_delete_returns_false(self, app): + from application.api.user.sources.routes import DeleteByIds + + mock_collection = Mock() + mock_collection.delete_index.return_value = False + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/delete_by_ids?path=id1"): + response = DeleteByIds().get() + + assert _status(response) == 400 + + def test_returns_400_on_exception(self, app): + from application.api.user.sources.routes import DeleteByIds + + mock_collection = Mock() + mock_collection.delete_index.side_effect = Exception("fail") + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context("/api/delete_by_ids?path=id1"): + response = DeleteByIds().get() + + assert _status(response) == 400 + + +# --------------------------------------------------------------------------- +# DeleteOldIndexes (/api/delete_old) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestDeleteOldIndexes: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.routes import DeleteOldIndexes + + with app.test_request_context("/api/delete_old?source_id=abc"): + from flask import request + + request.decoded_token = None + response = DeleteOldIndexes().get() + + assert _status(response) == 401 + + def test_returns_400_when_source_id_missing(self, app): + from application.api.user.sources.routes import DeleteOldIndexes + + with app.test_request_context("/api/delete_old"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteOldIndexes().get() + + assert _status(response) == 400 + + def test_returns_404_when_doc_not_found(self, app): + from application.api.user.sources.routes import DeleteOldIndexes + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context(f"/api/delete_old?source_id={source_id}"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteOldIndexes().get() + + assert _status(response) == 404 + + def test_deletes_faiss_index_and_file(self, app): + from application.api.user.sources.routes import DeleteOldIndexes + + source_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": source_id, + "user": "u1", + "file_path": "uploads/u1/doc.pdf", + } + mock_storage = Mock() + mock_storage.file_exists.return_value = True + mock_storage.is_directory.return_value = False + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.routes.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.sources.routes.settings" + ) as mock_settings: + mock_settings.VECTOR_STORE = "faiss" + with app.test_request_context( + f"/api/delete_old?source_id={source_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteOldIndexes().get() + + assert _status(response) == 200 + assert _json(response)["success"] is True + # Should have checked and deleted faiss files + assert mock_storage.delete_file.call_count >= 1 + mock_collection.delete_one.assert_called_once() + + def test_deletes_non_faiss_vector_index(self, app): + from application.api.user.sources.routes import DeleteOldIndexes + + source_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": source_id, + "user": "u1", + } + mock_storage = Mock() + mock_vectorstore = Mock() + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.routes.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.sources.routes.VectorCreator.create_vectorstore", + return_value=mock_vectorstore, + ), patch( + "application.api.user.sources.routes.settings" + ) as mock_settings: + mock_settings.VECTOR_STORE = "elasticsearch" + with app.test_request_context( + f"/api/delete_old?source_id={source_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteOldIndexes().get() + + assert _status(response) == 200 + mock_vectorstore.delete_index.assert_called_once() + + def test_deletes_directory_of_files(self, app): + from application.api.user.sources.routes import DeleteOldIndexes + + source_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": source_id, + "user": "u1", + "file_path": "uploads/u1/mydir", + } + mock_storage = Mock() + mock_storage.is_directory.return_value = True + mock_storage.list_files.return_value = ["uploads/u1/mydir/a.txt", "uploads/u1/mydir/b.txt"] + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.routes.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.sources.routes.settings" + ) as mock_settings: + mock_settings.VECTOR_STORE = "faiss" + mock_storage.file_exists.return_value = False + with app.test_request_context( + f"/api/delete_old?source_id={source_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteOldIndexes().get() + + assert _status(response) == 200 + # Each file in directory should be deleted + assert mock_storage.delete_file.call_count == 2 + + def test_handles_file_not_found_gracefully(self, app): + from application.api.user.sources.routes import DeleteOldIndexes + + source_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": source_id, + "user": "u1", + "file_path": "uploads/missing.pdf", + } + mock_storage = Mock() + mock_storage.is_directory.side_effect = FileNotFoundError() + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.routes.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.sources.routes.settings" + ) as mock_settings: + mock_settings.VECTOR_STORE = "faiss" + mock_storage.file_exists.return_value = False + with app.test_request_context( + f"/api/delete_old?source_id={source_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteOldIndexes().get() + + assert _status(response) == 200 + mock_collection.delete_one.assert_called_once() + + def test_returns_400_on_general_error(self, app): + from application.api.user.sources.routes import DeleteOldIndexes + + source_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": source_id, + "user": "u1", + } + mock_storage = Mock() + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.routes.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.sources.routes.settings" + ) as mock_settings: + mock_settings.VECTOR_STORE = "faiss" + mock_storage.file_exists.side_effect = RuntimeError("disk error") + with app.test_request_context( + f"/api/delete_old?source_id={source_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DeleteOldIndexes().get() + + assert _status(response) == 400 + + +# --------------------------------------------------------------------------- +# ManageSync (/api/manage_sync) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestManageSync: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.routes import ManageSync + + with app.test_request_context( + "/api/manage_sync", method="POST", + json={"source_id": "x", "sync_frequency": "daily"}, + ): + from flask import request + + request.decoded_token = None + response = ManageSync().post() + + assert _status(response) == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.sources.routes import ManageSync + + with app.test_request_context( + "/api/manage_sync", method="POST", + json={"source_id": "abc"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSync().post() + + assert response is not None + + def test_returns_400_for_invalid_frequency(self, app): + from application.api.user.sources.routes import ManageSync + + with app.test_request_context( + "/api/manage_sync", method="POST", + json={"source_id": str(ObjectId()), "sync_frequency": "hourly"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSync().post() + + assert _status(response) == 400 + assert "Invalid frequency" in _json(response)["message"] + + def test_updates_sync_frequency_successfully(self, app): + from application.api.user.sources.routes import ManageSync + + source_id = str(ObjectId()) + mock_collection = Mock() + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/manage_sync", method="POST", + json={"source_id": source_id, "sync_frequency": "weekly"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSync().post() + + assert _status(response) == 200 + assert _json(response)["success"] is True + call_args = mock_collection.update_one.call_args + assert call_args[0][0]["_id"] == ObjectId(source_id) + assert call_args[0][0]["user"] == "u1" + assert call_args[0][1]["$set"]["sync_frequency"] == "weekly" + + def test_accepts_all_valid_frequencies(self, app): + from application.api.user.sources.routes import ManageSync + + mock_collection = Mock() + + for freq in ["never", "daily", "weekly", "monthly"]: + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/manage_sync", method="POST", + json={"source_id": str(ObjectId()), "sync_frequency": freq}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSync().post() + + assert _status(response) == 200 + + def test_returns_400_on_db_error(self, app): + from application.api.user.sources.routes import ManageSync + + mock_collection = Mock() + mock_collection.update_one.side_effect = Exception("db err") + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/manage_sync", method="POST", + json={"source_id": str(ObjectId()), "sync_frequency": "daily"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSync().post() + + assert _status(response) == 400 + + +# --------------------------------------------------------------------------- +# RedirectToSources (/api/combine) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRedirectToSources: + + def test_redirects_to_sources(self, app): + from application.api.user.sources.routes import RedirectToSources + + with app.test_request_context("/api/combine"): + response = RedirectToSources().get() + + assert response.status_code == 301 + assert response.location == "/api/sources" + + +# --------------------------------------------------------------------------- +# DirectoryStructure (/api/directory_structure) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestDirectoryStructure: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.routes import DirectoryStructure + + with app.test_request_context("/api/directory_structure?id=abc"): + from flask import request + + request.decoded_token = None + response = DirectoryStructure().get() + + assert _status(response) == 401 + + def test_returns_400_when_id_missing(self, app): + from application.api.user.sources.routes import DirectoryStructure + + with app.test_request_context("/api/directory_structure"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DirectoryStructure().get() + + assert _status(response) == 400 + assert "required" in _json(response)["error"] + + def test_returns_400_for_invalid_doc_id(self, app): + from application.api.user.sources.routes import DirectoryStructure + + with app.test_request_context("/api/directory_structure?id=invalid"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DirectoryStructure().get() + + assert _status(response) == 400 + assert "Invalid" in _json(response)["error"] + + def test_returns_404_when_doc_not_found(self, app): + from application.api.user.sources.routes import DirectoryStructure + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context(f"/api/directory_structure?id={doc_id}"): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DirectoryStructure().get() + + assert _status(response) == 404 + assert "not found" in _json(response)["error"] + + def test_returns_directory_structure(self, app): + from application.api.user.sources.routes import DirectoryStructure + + doc_id = ObjectId() + dir_struct = {"dirs": ["a", "b"], "files": ["c.txt"]} + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": doc_id, + "user": "u1", + "directory_structure": dir_struct, + "file_path": "uploads/u1/mydir", + "remote_data": json.dumps({"provider": "github"}), + } + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/directory_structure?id={doc_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DirectoryStructure().get() + + assert _status(response) == 200 + data = _json(response) + assert data["success"] is True + assert data["directory_structure"] == dir_struct + assert data["base_path"] == "uploads/u1/mydir" + assert data["provider"] == "github" + + def test_returns_none_provider_when_no_remote_data(self, app): + from application.api.user.sources.routes import DirectoryStructure + + doc_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": doc_id, + "user": "u1", + "directory_structure": {}, + "file_path": "path", + } + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/directory_structure?id={doc_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DirectoryStructure().get() + + data = _json(response) + assert data["provider"] is None + + def test_handles_invalid_remote_data_json(self, app): + from application.api.user.sources.routes import DirectoryStructure + + doc_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": doc_id, + "user": "u1", + "directory_structure": {}, + "file_path": "path", + "remote_data": "not-valid-json{", + } + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/directory_structure?id={doc_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DirectoryStructure().get() + + data = _json(response) + assert data["success"] is True + assert data["provider"] is None + + def test_returns_500_on_general_error(self, app): + from application.api.user.sources.routes import DirectoryStructure + + doc_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.side_effect = Exception("db error") + + with patch( + "application.api.user.sources.routes.sources_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/directory_structure?id={doc_id}" + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = DirectoryStructure().get() + + assert _status(response) == 500 diff --git a/tests/api/user/sources/test_upload.py b/tests/api/user/sources/test_upload.py new file mode 100644 index 00000000..f0bb6a75 --- /dev/null +++ b/tests/api/user/sources/test_upload.py @@ -0,0 +1,1338 @@ +"""Tests for source upload routes (UploadFile, UploadRemote, ManageSourceFiles, +TaskStatus) and the _enforce_audio_path_size_limit helper. + +Note: test_audio_upload.py already covers audio-extension pass-through and +oversized-audio rejection for UploadFile; those are NOT duplicated here. +""" + +import io +import json + +import pytest +from types import SimpleNamespace +from unittest.mock import MagicMock, Mock, patch +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +def _status(response): + if isinstance(response, tuple): + return response[1] + return response.status_code + + +def _json(response): + if isinstance(response, tuple): + return response[0].json + return response.json + + +# --------------------------------------------------------------------------- +# _enforce_audio_path_size_limit helper +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestEnforceAudioPathSizeLimit: + + def test_skips_non_audio_file(self): + from application.api.user.sources.upload import _enforce_audio_path_size_limit + + # Should not raise for non-audio file regardless of size + with patch( + "application.api.user.sources.upload.is_audio_filename", + return_value=False, + ): + _enforce_audio_path_size_limit("/tmp/big.pdf", "big.pdf") + + def test_raises_for_oversized_audio(self): + from application.api.user.sources.upload import _enforce_audio_path_size_limit + from application.stt.upload_limits import AudioFileTooLargeError + + with patch( + "application.api.user.sources.upload.is_audio_filename", + return_value=True, + ), patch( + "application.api.user.sources.upload.os.path.getsize", + return_value=999_999_999, + ), patch( + "application.api.user.sources.upload.enforce_audio_file_size_limit", + side_effect=AudioFileTooLargeError("too big"), + ): + with pytest.raises(AudioFileTooLargeError): + _enforce_audio_path_size_limit("/tmp/big.wav", "big.wav") + + def test_passes_for_small_audio(self): + from application.api.user.sources.upload import _enforce_audio_path_size_limit + + with patch( + "application.api.user.sources.upload.is_audio_filename", + return_value=True, + ), patch( + "application.api.user.sources.upload.os.path.getsize", + return_value=1024, + ), patch( + "application.api.user.sources.upload.enforce_audio_file_size_limit", + ) as mock_enforce: + _enforce_audio_path_size_limit("/tmp/small.wav", "small.wav") + mock_enforce.assert_called_once_with(1024) + + +# --------------------------------------------------------------------------- +# UploadFile (/api/upload) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestUploadFile: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.upload import UploadFile + + with app.test_request_context( + "/api/upload", method="POST", + data={"user": "u1", "name": "test", "file": (io.BytesIO(b"x"), "f.txt")}, + content_type="multipart/form-data", + ): + from flask import request + + request.decoded_token = None + response = UploadFile().post() + + assert _status(response) == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.sources.upload import UploadFile + + with app.test_request_context( + "/api/upload", method="POST", + data={"user": "u1"}, + content_type="multipart/form-data", + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UploadFile().post() + + assert _status(response) == 400 + + def test_returns_400_when_no_files(self, app): + from application.api.user.sources.upload import UploadFile + + with app.test_request_context( + "/api/upload", method="POST", + data={"user": "u1", "name": "test"}, + content_type="multipart/form-data", + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UploadFile().post() + + assert _status(response) == 400 + + def test_returns_400_when_files_have_empty_filenames(self, app): + from application.api.user.sources.upload import UploadFile + + with app.test_request_context( + "/api/upload", method="POST", + data={"user": "u1", "name": "test", "file": (io.BytesIO(b""), "")}, + content_type="multipart/form-data", + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UploadFile().post() + + assert _status(response) == 400 + + def test_successful_upload(self, app): + from application.api.user.sources.upload import UploadFile + + mock_storage = MagicMock() + mock_task = SimpleNamespace(id="task-abc") + + with app.test_request_context( + "/api/upload", method="POST", + data={ + "user": "u1", + "name": "My Doc", + "file": (io.BytesIO(b"hello"), "test.txt"), + }, + content_type="multipart/form-data", + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.sources.upload.ingest" + ) as mock_ingest, patch( + "application.api.user.sources.upload._enforce_audio_path_size_limit", + ): + mock_ingest.delay.return_value = mock_task + response = UploadFile().post() + + assert _status(response) == 200 + data = _json(response) + assert data["success"] is True + assert data["task_id"] == "task-abc" + mock_ingest.delay.assert_called_once() + + def test_returns_400_on_general_error(self, app): + from application.api.user.sources.upload import UploadFile + + with app.test_request_context( + "/api/upload", method="POST", + data={ + "user": "u1", + "name": "Doc", + "file": (io.BytesIO(b"data"), "f.txt"), + }, + content_type="multipart/form-data", + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + side_effect=Exception("storage down"), + ): + response = UploadFile().post() + + assert _status(response) == 400 + + +# --------------------------------------------------------------------------- +# UploadRemote (/api/remote) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestUploadRemote: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.upload import UploadRemote + + with app.test_request_context( + "/api/remote", method="POST", + data={"user": "u1", "source": "github", "name": "repo", "data": "{}"}, + ): + from flask import request + + request.decoded_token = None + response = UploadRemote().post() + + assert _status(response) == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.sources.upload import UploadRemote + + with app.test_request_context( + "/api/remote", method="POST", + data={"user": "u1", "source": "github"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UploadRemote().post() + + assert response is not None + + def test_github_source(self, app): + from application.api.user.sources.upload import UploadRemote + + mock_task = SimpleNamespace(id="task-gh") + config = json.dumps({"repo_url": "https://github.com/test/repo"}) + + with app.test_request_context( + "/api/remote", method="POST", + data={"user": "u1", "source": "github", "name": "repo", "data": config}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ingest_remote" + ) as mock_ingest: + mock_ingest.delay.return_value = mock_task + response = UploadRemote().post() + + assert _status(response) == 200 + data = _json(response) + assert data["task_id"] == "task-gh" + call_kwargs = mock_ingest.delay.call_args[1] + assert call_kwargs["source_data"] == "https://github.com/test/repo" + assert call_kwargs["loader"] == "github" + + def test_crawler_source(self, app): + from application.api.user.sources.upload import UploadRemote + + mock_task = SimpleNamespace(id="task-cr") + config = json.dumps({"url": "https://example.com"}) + + with app.test_request_context( + "/api/remote", method="POST", + data={"user": "u1", "source": "crawler", "name": "site", "data": config}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ingest_remote" + ) as mock_ingest: + mock_ingest.delay.return_value = mock_task + response = UploadRemote().post() + + assert _status(response) == 200 + call_kwargs = mock_ingest.delay.call_args[1] + assert call_kwargs["source_data"] == "https://example.com" + + def test_url_source(self, app): + from application.api.user.sources.upload import UploadRemote + + mock_task = SimpleNamespace(id="task-url") + config = json.dumps({"url": "https://example.com/doc"}) + + with app.test_request_context( + "/api/remote", method="POST", + data={"user": "u1", "source": "url", "name": "url-src", "data": config}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ingest_remote" + ) as mock_ingest: + mock_ingest.delay.return_value = mock_task + response = UploadRemote().post() + + assert _status(response) == 200 + call_kwargs = mock_ingest.delay.call_args[1] + assert call_kwargs["source_data"] == "https://example.com/doc" + + def test_reddit_source(self, app): + from application.api.user.sources.upload import UploadRemote + + mock_task = SimpleNamespace(id="task-reddit") + config_data = {"subreddit": "python", "limit": 10} + config = json.dumps(config_data) + + with app.test_request_context( + "/api/remote", method="POST", + data={"user": "u1", "source": "reddit", "name": "reddit-src", "data": config}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ingest_remote" + ) as mock_ingest: + mock_ingest.delay.return_value = mock_task + response = UploadRemote().post() + + assert _status(response) == 200 + call_kwargs = mock_ingest.delay.call_args[1] + assert call_kwargs["source_data"] == config_data + + def test_s3_source(self, app): + from application.api.user.sources.upload import UploadRemote + + mock_task = SimpleNamespace(id="task-s3") + config_data = {"bucket": "my-bucket", "key": "data/"} + config = json.dumps(config_data) + + with app.test_request_context( + "/api/remote", method="POST", + data={"user": "u1", "source": "s3", "name": "s3-src", "data": config}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ingest_remote" + ) as mock_ingest: + mock_ingest.delay.return_value = mock_task + response = UploadRemote().post() + + assert _status(response) == 200 + call_kwargs = mock_ingest.delay.call_args[1] + assert call_kwargs["source_data"] == config_data + + def test_connector_source_success(self, app): + from application.api.user.sources.upload import UploadRemote + + mock_task = SimpleNamespace(id="task-gd") + config = json.dumps({ + "session_token": "token123", + "file_ids": ["f1", "f2"], + "folder_ids": "fold1, fold2", + "recursive": True, + "retriever": "classic", + }) + + with app.test_request_context( + "/api/remote", method="POST", + data={ + "user": "u1", + "source": "google_drive", + "name": "gdrive", + "data": config, + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ConnectorCreator.get_supported_connectors", + return_value=["google_drive", "share_point"], + ), patch( + "application.api.user.sources.upload.ingest_connector_task" + ) as mock_connector: + mock_connector.delay.return_value = mock_task + response = UploadRemote().post() + + assert _status(response) == 200 + data = _json(response) + assert data["task_id"] == "task-gd" + call_kwargs = mock_connector.delay.call_args[1] + assert call_kwargs["session_token"] == "token123" + assert call_kwargs["file_ids"] == ["f1", "f2"] + assert call_kwargs["folder_ids"] == ["fold1", "fold2"] + assert call_kwargs["recursive"] is True + + def test_connector_source_missing_session_token(self, app): + from application.api.user.sources.upload import UploadRemote + + config = json.dumps({"file_ids": ["f1"]}) + + with app.test_request_context( + "/api/remote", method="POST", + data={ + "user": "u1", + "source": "google_drive", + "name": "gdrive", + "data": config, + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ConnectorCreator.get_supported_connectors", + return_value=["google_drive"], + ): + response = UploadRemote().post() + + assert _status(response) == 400 + assert "session_token" in _json(response)["error"] + + def test_connector_file_ids_as_string(self, app): + from application.api.user.sources.upload import UploadRemote + + mock_task = SimpleNamespace(id="task-sp") + config = json.dumps({ + "session_token": "tok", + "file_ids": "a, b, c", + "folder_ids": [], + }) + + with app.test_request_context( + "/api/remote", method="POST", + data={ + "user": "u1", + "source": "share_point", + "name": "sp", + "data": config, + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ConnectorCreator.get_supported_connectors", + return_value=["share_point"], + ), patch( + "application.api.user.sources.upload.ingest_connector_task" + ) as mock_ct: + mock_ct.delay.return_value = mock_task + response = UploadRemote().post() + + assert _status(response) == 200 + call_kwargs = mock_ct.delay.call_args[1] + assert call_kwargs["file_ids"] == ["a", "b", "c"] + + def test_connector_non_list_file_ids_becomes_empty(self, app): + from application.api.user.sources.upload import UploadRemote + + mock_task = SimpleNamespace(id="task-sp2") + config = json.dumps({ + "session_token": "tok", + "file_ids": 42, + "folder_ids": True, + }) + + with app.test_request_context( + "/api/remote", method="POST", + data={ + "user": "u1", + "source": "share_point", + "name": "sp", + "data": config, + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + + with patch( + "application.api.user.sources.upload.ConnectorCreator.get_supported_connectors", + return_value=["share_point"], + ), patch( + "application.api.user.sources.upload.ingest_connector_task" + ) as mock_ct: + mock_ct.delay.return_value = mock_task + response = UploadRemote().post() + + assert _status(response) == 200 + call_kwargs = mock_ct.delay.call_args[1] + assert call_kwargs["file_ids"] == [] + assert call_kwargs["folder_ids"] == [] + + def test_returns_400_on_error(self, app): + from application.api.user.sources.upload import UploadRemote + + with app.test_request_context( + "/api/remote", method="POST", + data={ + "user": "u1", + "source": "github", + "name": "repo", + "data": "invalid-json", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = UploadRemote().post() + + assert _status(response) == 400 + + +# --------------------------------------------------------------------------- +# ManageSourceFiles (/api/manage_source_files) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestManageSourceFiles: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"source_id": "abc", "operation": "add"}, + ): + from flask import request + + request.decoded_token = None + response = ManageSourceFiles().post() + + assert _status(response) == 401 + + def test_returns_400_missing_source_id(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"operation": "add"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + + def test_returns_400_missing_operation(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"source_id": str(ObjectId())}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + + def test_returns_400_invalid_operation(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"source_id": str(ObjectId()), "operation": "invalid"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + assert "must be" in _json(response)["message"] + + def test_returns_400_invalid_source_id_format(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"source_id": "bad-id", "operation": "add"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + assert "Invalid source ID" in _json(response)["message"] + + def test_returns_404_when_source_not_found(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + mock_collection = Mock() + mock_collection.find_one.return_value = None + source_id = str(ObjectId()) + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"source_id": source_id, "operation": "add"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 404 + + def test_add_operation_no_files_returns_400(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"source_id": source_id, "operation": "add"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + assert "No files" in _json(response)["message"] + + def test_add_operation_success(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + "file_name_map": {}, + } + mock_storage = MagicMock() + mock_task = SimpleNamespace(id="reingest-1") + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.tasks.reingest_source_task" + ) as mock_reingest: + mock_reingest.delay.return_value = mock_task + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "add", + "file": (io.BytesIO(b"content"), "new_file.txt"), + }, + content_type="multipart/form-data", + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 200 + data = _json(response) + assert data["success"] is True + assert "1 files" in data["message"] + assert data["reingest_task_id"] == "reingest-1" + + def test_add_operation_with_parent_dir(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + "file_name_map": {}, + } + mock_storage = MagicMock() + mock_task = SimpleNamespace(id="reingest-2") + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.tasks.reingest_source_task" + ) as mock_reingest: + mock_reingest.delay.return_value = mock_task + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "add", + "parent_dir": "subdir", + "file": (io.BytesIO(b"content"), "f.txt"), + }, + content_type="multipart/form-data", + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 200 + data = _json(response) + assert data["parent_dir"] == "subdir" + # Verify storage.save_file was called with path including parent_dir + save_call = mock_storage.save_file.call_args + assert "subdir" in save_call[0][1] + + def test_add_rejects_invalid_parent_dir(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "add", + "parent_dir": "../escape", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + assert "Invalid parent" in _json(response)["message"] + + def test_add_rejects_absolute_parent_dir(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "add", + "parent_dir": "/etc/passwd", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + + def test_remove_operation_success(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + "file_name_map": {"old.txt": "Original Name.txt"}, + } + mock_storage = MagicMock() + mock_storage.file_exists.return_value = True + mock_task = SimpleNamespace(id="reingest-3") + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.tasks.reingest_source_task" + ) as mock_reingest: + mock_reingest.delay.return_value = mock_task + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove", + "file_paths": json.dumps(["old.txt"]), + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 200 + data = _json(response) + assert "1 files" in data["message"] + mock_storage.delete_file.assert_called_once() + + def test_remove_missing_file_paths_returns_400(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"source_id": source_id, "operation": "remove"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + assert "file_paths required" in _json(response)["message"] + + def test_remove_invalid_file_paths_format(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove", + "file_paths": "not-valid-json{", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + assert "Invalid file_paths" in _json(response)["message"] + + def test_remove_directory_success(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + "file_name_map": {"subdir/a.txt": "A.txt", "subdir/b.txt": "B.txt", "other.txt": "Other.txt"}, + } + mock_storage = MagicMock() + mock_storage.is_directory.return_value = True + mock_storage.remove_directory.return_value = True + mock_task = SimpleNamespace(id="reingest-4") + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.tasks.reingest_source_task" + ) as mock_reingest: + mock_reingest.delay.return_value = mock_task + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove_directory", + "directory_path": "subdir", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 200 + data = _json(response) + assert data["removed_directory"] == "subdir" + # file_name_map should have subdir entries removed + update_call = mock_collection.update_one.call_args + updated_map = update_call[0][1]["$set"]["file_name_map"] + assert "subdir/a.txt" not in updated_map + assert "other.txt" in updated_map + + def test_remove_directory_missing_path_returns_400(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove_directory", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + assert "directory_path required" in _json(response)["message"] + + def test_remove_directory_path_traversal_rejected(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove_directory", + "directory_path": "../../../etc", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 400 + assert "Invalid directory" in _json(response)["message"] + + def test_remove_directory_not_found_returns_404(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + mock_storage.is_directory.return_value = False + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove_directory", + "directory_path": "nonexistent", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 404 + + def test_remove_directory_storage_failure_returns_500(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + mock_storage = MagicMock() + mock_storage.is_directory.return_value = True + mock_storage.remove_directory.return_value = False + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove_directory", + "directory_path": "mydir", + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 500 + + def test_returns_500_on_db_find_error(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.side_effect = Exception("db error") + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={"source_id": source_id, "operation": "add"}, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 500 + + def test_file_name_map_as_json_string(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + "file_name_map": json.dumps({"old.txt": "Old File.txt"}), + } + mock_storage = MagicMock() + mock_storage.file_exists.return_value = True + mock_task = SimpleNamespace(id="reingest-5") + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + return_value=mock_storage, + ), patch( + "application.api.user.tasks.reingest_source_task" + ) as mock_reingest: + mock_reingest.delay.return_value = mock_task + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove", + "file_paths": json.dumps(["old.txt"]), + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 200 + + def test_general_operation_error_returns_500(self, app): + from application.api.user.sources.upload import ManageSourceFiles + + source_id = str(ObjectId()) + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": ObjectId(source_id), + "user": "u1", + "file_path": "uploads/u1/src", + } + + with patch( + "application.api.user.sources.upload.sources_collection", + mock_collection, + ), patch( + "application.api.user.sources.upload.StorageCreator.get_storage", + side_effect=Exception("storage crash"), + ): + with app.test_request_context( + "/api/manage_source_files", method="POST", + data={ + "source_id": source_id, + "operation": "remove", + "file_paths": json.dumps(["x.txt"]), + }, + ): + from flask import request + + request.decoded_token = {"sub": "u1"} + response = ManageSourceFiles().post() + + assert _status(response) == 500 + assert "Operation failed" in _json(response)["message"] + + +# --------------------------------------------------------------------------- +# TaskStatus (/api/task_status) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestTaskStatus: + + def test_returns_400_when_task_id_missing(self, app): + from application.api.user.sources.upload import TaskStatus + + with app.test_request_context("/api/task_status"): + response = TaskStatus().get() + + assert _status(response) == 400 + assert "Task ID is required" in _json(response)["message"] + + def test_returns_task_status_success(self, app): + from application.api.user.sources.upload import TaskStatus + + mock_task = Mock() + mock_task.status = "SUCCESS" + mock_task.info = {"result": "done"} + + mock_celery = Mock() + mock_celery.AsyncResult.return_value = mock_task + + with patch( + "application.celery_init.celery", mock_celery + ): + with app.test_request_context("/api/task_status?task_id=tid-123"): + response = TaskStatus().get() + + assert _status(response) == 200 + data = _json(response) + assert data["status"] == "SUCCESS" + assert data["result"] == {"result": "done"} + + def test_returns_task_status_pending_with_workers(self, app): + from application.api.user.sources.upload import TaskStatus + + mock_task = Mock() + mock_task.status = "PENDING" + mock_task.info = None + + mock_inspect = Mock() + mock_inspect.ping.return_value = {"worker1": {"ok": "pong"}} + mock_celery = Mock() + mock_celery.AsyncResult.return_value = mock_task + mock_celery.control.inspect.return_value = mock_inspect + + with patch( + "application.celery_init.celery", mock_celery + ): + with app.test_request_context("/api/task_status?task_id=tid-pend"): + response = TaskStatus().get() + + assert _status(response) == 200 + + def test_returns_503_when_no_workers(self, app): + from application.api.user.sources.upload import TaskStatus + + mock_task = Mock() + mock_task.status = "PENDING" + mock_task.info = None + + mock_inspect = Mock() + mock_inspect.ping.return_value = None + mock_celery = Mock() + mock_celery.AsyncResult.return_value = mock_task + mock_celery.control.inspect.return_value = mock_inspect + + with patch( + "application.celery_init.celery", mock_celery + ): + with app.test_request_context("/api/task_status?task_id=tid-nw"): + response = TaskStatus().get() + + assert _status(response) == 503 + assert "unavailable" in _json(response)["message"] + + def test_handles_non_serializable_task_meta(self, app): + from application.api.user.sources.upload import TaskStatus + + mock_task = Mock() + mock_task.status = "SUCCESS" + mock_task.info = object() # non-serializable + + mock_celery = Mock() + mock_celery.AsyncResult.return_value = mock_task + + with patch( + "application.celery_init.celery", mock_celery + ): + with app.test_request_context("/api/task_status?task_id=tid-ns"): + response = TaskStatus().get() + + assert _status(response) == 200 + data = _json(response) + # Non-serializable info should be converted to string + assert isinstance(data["result"], str) + + def test_returns_400_on_general_error(self, app): + from application.api.user.sources.upload import TaskStatus + + with patch( + "application.celery_init.celery", + ) as mock_celery: + mock_celery.AsyncResult.side_effect = Exception("broken") + with app.test_request_context("/api/task_status?task_id=tid-err"): + response = TaskStatus().get() + + assert _status(response) == 400 diff --git a/tests/api/user/test_agents_routes.py b/tests/api/user/test_agents_routes.py new file mode 100644 index 00000000..ed996081 --- /dev/null +++ b/tests/api/user/test_agents_routes.py @@ -0,0 +1,2516 @@ +"""Tests for application.api.user.agents.routes module.""" + +import uuid +from unittest.mock import Mock, patch + +import pytest +from bson import DBRef, ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +# --------------------------------------------------------------------------- +# Helper functions +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestNormalizeWorkflowReference: + + def test_returns_none_for_none(self): + from application.api.user.agents.routes import normalize_workflow_reference + + assert normalize_workflow_reference(None) is None + + def test_extracts_id_from_dict(self): + from application.api.user.agents.routes import normalize_workflow_reference + + result = normalize_workflow_reference({"id": "abc123"}) + assert result == "abc123" + + def test_extracts_underscore_id_from_dict(self): + from application.api.user.agents.routes import normalize_workflow_reference + + result = normalize_workflow_reference({"_id": "abc123"}) + assert result == "abc123" + + def test_extracts_workflow_id_from_dict(self): + from application.api.user.agents.routes import normalize_workflow_reference + + result = normalize_workflow_reference({"workflow_id": "abc123"}) + assert result == "abc123" + + def test_returns_empty_string_for_blank_string(self): + from application.api.user.agents.routes import normalize_workflow_reference + + assert normalize_workflow_reference(" ") == "" + + def test_returns_plain_string_value(self): + from application.api.user.agents.routes import normalize_workflow_reference + + oid = str(ObjectId()) + assert normalize_workflow_reference(oid) == oid + + def test_parses_json_string_value(self): + from application.api.user.agents.routes import normalize_workflow_reference + + result = normalize_workflow_reference('"some_id"') + assert result == "some_id" + + def test_parses_json_dict_value(self): + from application.api.user.agents.routes import normalize_workflow_reference + + result = normalize_workflow_reference('{"id": "xyz"}') + assert result == "xyz" + + def test_returns_string_for_non_json_string(self): + from application.api.user.agents.routes import normalize_workflow_reference + + assert normalize_workflow_reference("plain_value") == "plain_value" + + def test_converts_non_string_non_dict_to_str(self): + from application.api.user.agents.routes import normalize_workflow_reference + + assert normalize_workflow_reference(42) == "42" + + def test_dict_priority_id_over_workflow_id(self): + from application.api.user.agents.routes import normalize_workflow_reference + + result = normalize_workflow_reference( + {"id": "first", "_id": "second", "workflow_id": "third"} + ) + assert result == "first" + + +@pytest.mark.unit +class TestValidateWorkflowAccess: + + def test_returns_none_when_not_required_and_empty(self, app): + from application.api.user.agents.routes import validate_workflow_access + + with app.app_context(): + wf_id, err = validate_workflow_access(None, "user1", required=False) + assert wf_id is None + assert err is None + + def test_returns_error_when_required_and_empty(self, app): + from application.api.user.agents.routes import validate_workflow_access + + with app.app_context(): + wf_id, err = validate_workflow_access(None, "user1", required=True) + assert wf_id is None + assert err is not None + assert err.status_code == 400 + + def test_returns_error_for_invalid_oid_format(self, app): + from application.api.user.agents.routes import validate_workflow_access + + with app.app_context(): + wf_id, err = validate_workflow_access("not-a-valid-oid", "user1") + assert wf_id is None + assert err.status_code == 400 + + def test_returns_404_when_workflow_not_found(self, app): + from application.api.user.agents.routes import validate_workflow_access + + mock_wf_col = Mock() + mock_wf_col.find_one.return_value = None + oid = str(ObjectId()) + + with app.app_context(): + with patch( + "application.api.user.agents.routes.workflows_collection", mock_wf_col + ): + wf_id, err = validate_workflow_access(oid, "user1") + assert err.status_code == 404 + + def test_returns_workflow_id_on_success(self, app): + from application.api.user.agents.routes import validate_workflow_access + + mock_wf_col = Mock() + mock_wf_col.find_one.return_value = {"_id": ObjectId(), "user": "user1"} + oid = str(ObjectId()) + + with app.app_context(): + with patch( + "application.api.user.agents.routes.workflows_collection", mock_wf_col + ): + wf_id, err = validate_workflow_access(oid, "user1") + assert wf_id == oid + assert err is None + + +@pytest.mark.unit +class TestBuildAgentDocument: + + def test_builds_classic_agent(self): + from application.api.user.agents.routes import build_agent_document + + data = { + "name": "Test Agent", + "description": "desc", + "status": "draft", + "chunks": "5", + "retriever": "classic", + "prompt_id": "default", + } + doc = build_agent_document( + data, "user1", "key123", "classic", image_url="img.png" + ) + assert doc["user"] == "user1" + assert doc["name"] == "Test Agent" + assert doc["image"] == "img.png" + assert doc["agent_type"] == "classic" + assert "createdAt" in doc + assert "updatedAt" in doc + + def test_builds_workflow_agent(self): + from application.api.user.agents.routes import build_agent_document + + data = { + "name": "WF Agent", + "status": "published", + "workflow": "wf123", + "folder_id": "folder1", + } + doc = build_agent_document(data, "user1", "key123", "workflow") + assert doc["workflow"] == "wf123" + assert doc["folder_id"] == "folder1" + # Workflow agents should not have classic-specific fields + assert "image" not in doc + assert "source" not in doc + + def test_defaults_to_classic_for_unknown_type(self): + from application.api.user.agents.routes import build_agent_document + + data = {"name": "Agent", "status": "draft"} + doc = build_agent_document(data, "user1", "k", "unknown_type") + assert doc["agent_type"] == "classic" + + def test_defaults_to_classic_for_empty_type(self): + from application.api.user.agents.routes import build_agent_document + + data = {"name": "Agent", "status": "draft"} + doc = build_agent_document(data, "user1", "k", "") + assert doc["agent_type"] == "classic" + + def test_limited_token_mode_string_true(self): + from application.api.user.agents.routes import build_agent_document + + data = {"name": "A", "status": "draft", "limited_token_mode": "True"} + doc = build_agent_document(data, "user1", "k", "classic") + assert doc["limited_token_mode"] is True + + def test_limited_token_mode_string_false(self): + from application.api.user.agents.routes import build_agent_document + + data = {"name": "A", "status": "draft", "limited_token_mode": "False"} + doc = build_agent_document(data, "user1", "k", "classic") + assert doc["limited_token_mode"] is False + + def test_limited_request_mode_bool(self): + from application.api.user.agents.routes import build_agent_document + + data = {"name": "A", "status": "draft", "limited_request_mode": True} + doc = build_agent_document(data, "user1", "k", "classic") + assert doc["limited_request_mode"] is True + + def test_filters_to_allowed_fields(self): + from application.api.user.agents.routes import build_agent_document + + data = {"name": "A", "status": "draft"} + doc = build_agent_document(data, "user1", "k", "workflow") + # Workflow doc should not have classic-only fields + for classic_field in ["image", "source", "sources", "chunks", "retriever"]: + assert classic_field not in doc + + def test_source_and_sources_passed_through(self): + from application.api.user.agents.routes import build_agent_document + + source_ref = DBRef("sources", ObjectId()) + doc = build_agent_document( + {"name": "A", "status": "draft"}, + "user1", + "k", + "classic", + source_field=source_ref, + sources_list=[source_ref], + ) + assert doc["source"] == source_ref + assert doc["sources"] == [source_ref] + + +# --------------------------------------------------------------------------- +# Route classes +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestGetAgent: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import GetAgent + + with app.test_request_context("/api/get_agent?id=abc"): + from flask import request + + request.decoded_token = None + result = GetAgent().get() + # Returns tuple (dict, status_code) + assert result[1] == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.agents.routes import GetAgent + + with app.test_request_context("/api/get_agent"): + from flask import request + + request.decoded_token = {"sub": "user1"} + result = GetAgent().get() + assert result[1] == 400 + + def test_returns_404_agent_not_found(self, app): + from application.api.user.agents.routes import GetAgent + + agent_id = ObjectId() + mock_col = Mock() + mock_col.find_one.return_value = None + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/get_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + result = GetAgent().get() + assert result[1] == 404 + + def test_returns_agent_data_on_success(self, app): + from application.api.user.agents.routes import GetAgent + + agent_id = ObjectId() + key = str(uuid.uuid4()) + mock_col = Mock() + mock_col.find_one.return_value = { + "_id": agent_id, + "user": "user1", + "name": "Test Agent", + "description": "desc", + "chunks": "5", + "retriever": "classic", + "prompt_id": "default", + "tools": [], + "agent_type": "classic", + "status": "published", + "key": key, + } + + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.routes.db", mock_db + ): + with app.test_request_context(f"/api/get_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetAgent().get() + assert response.status_code == 200 + data = response.json + assert data["id"] == str(agent_id) + assert data["name"] == "Test Agent" + assert data["key"].startswith(key[:4]) + + def test_returns_400_on_exception(self, app): + from application.api.user.agents.routes import GetAgent + + mock_col = Mock() + mock_col.find_one.side_effect = Exception("DB error") + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/get_agent?id={ObjectId()}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + result = GetAgent().get() + assert result[1] == 400 + + +@pytest.mark.unit +class TestGetAgents: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import GetAgents + + with app.test_request_context("/api/get_agents"): + from flask import request + + request.decoded_token = None + result = GetAgents().get() + assert result[1] == 401 + + def test_returns_agents_list(self, app): + from application.api.user.agents.routes import GetAgents + + agent_id = ObjectId() + key = str(uuid.uuid4()) + mock_agents_col = Mock() + mock_agents_col.find.return_value = [ + { + "_id": agent_id, + "user": "user1", + "name": "Agent1", + "source": "default", + "retriever": "classic", + "key": key, + } + ] + + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": [str(agent_id)]}, + } + ) + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.routes.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.routes.db", mock_db + ): + with app.test_request_context("/api/get_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetAgents().get() + assert response.status_code == 200 + data = response.json + assert len(data) == 1 + assert data[0]["name"] == "Agent1" + assert data[0]["pinned"] is True + + def test_filters_agents_without_source_or_retriever(self, app): + from application.api.user.agents.routes import GetAgents + + mock_agents_col = Mock() + # Agent without source/retriever and not workflow type -> filtered out + mock_agents_col.find.return_value = [ + { + "_id": ObjectId(), + "user": "user1", + "name": "BadAgent", + "agent_type": "classic", + } + ] + + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": []}, + } + ) + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.routes.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.routes.db", mock_db + ): + with app.test_request_context("/api/get_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetAgents().get() + assert response.status_code == 200 + assert len(response.json) == 0 + + def test_includes_workflow_agent_without_source(self, app): + from application.api.user.agents.routes import GetAgents + + agent_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.find.return_value = [ + { + "_id": agent_id, + "user": "user1", + "name": "WFAgent", + "agent_type": "workflow", + "key": "abcd1234efgh", + } + ] + + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": []}, + } + ) + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.routes.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.routes.db", mock_db + ): + with app.test_request_context("/api/get_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetAgents().get() + assert response.status_code == 200 + assert len(response.json) == 1 + assert response.json[0]["name"] == "WFAgent" + + def test_returns_400_on_exception(self, app): + from application.api.user.agents.routes import GetAgents + + mock_ensure = Mock(side_effect=Exception("DB error")) + + with patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ): + with app.test_request_context("/api/get_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetAgents().get() + assert response.status_code == 400 + + +@pytest.mark.unit +class TestCreateAgent: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import CreateAgent + + with app.test_request_context( + "/api/create_agent", + method="POST", + json={"name": "A"}, + ): + from flask import request + + request.decoded_token = None + result = CreateAgent().post() + assert result[1] == 401 + + def test_returns_400_invalid_status(self, app): + from application.api.user.agents.routes import CreateAgent + + with app.test_request_context( + "/api/create_agent", + method="POST", + json={"name": "A", "status": "invalid"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 400 + + def test_creates_draft_agent_success(self, app): + from application.api.user.agents.routes import CreateAgent + + inserted_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.insert_one.return_value = Mock(inserted_id=inserted_id) + + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "Draft Agent", + "status": "draft", + "agent_type": "classic", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 201 + data = response.json + assert data["id"] == str(inserted_id) + # Draft agents get empty key + assert data["key"] == "" + + def test_creates_published_classic_agent(self, app): + from application.api.user.agents.routes import CreateAgent + + inserted_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.insert_one.return_value = Mock(inserted_id=inserted_id) + mock_handle_img = Mock(return_value=("img.png", None)) + source_id = str(ObjectId()) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "Published Agent", + "description": "A test agent", + "status": "published", + "agent_type": "classic", + "source": source_id, + "chunks": "5", + "retriever": "classic", + "prompt_id": "default", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 201 + data = response.json + assert data["id"] == str(inserted_id) + # Published agents get a uuid key + assert data["key"] != "" + + def test_published_classic_requires_source(self, app): + from application.api.user.agents.routes import CreateAgent + + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "No Source", + "description": "desc", + "status": "published", + "agent_type": "classic", + "chunks": "5", + "retriever": "classic", + "prompt_id": "default", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 400 + + def test_creates_workflow_agent(self, app): + from application.api.user.agents.routes import CreateAgent + + inserted_id = ObjectId() + wf_id = str(ObjectId()) + mock_agents_col = Mock() + mock_agents_col.insert_one.return_value = Mock(inserted_id=inserted_id) + mock_handle_img = Mock(return_value=("", None)) + mock_wf_col = Mock() + mock_wf_col.find_one.return_value = {"_id": ObjectId(wf_id), "user": "user1"} + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ), patch( + "application.api.user.agents.routes.workflows_collection", mock_wf_col + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "WF Agent", + "status": "published", + "agent_type": "workflow", + "workflow": wf_id, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 201 + + def test_image_upload_failure(self, app): + from application.api.user.agents.routes import CreateAgent + + mock_handle_img = Mock(return_value=(None, Mock())) + + with patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "Agent", + "status": "draft", + "agent_type": "classic", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 400 + + def test_invalid_json_schema(self, app): + from application.api.user.agents.routes import CreateAgent + from application.core.json_schema_utils import JsonSchemaValidationError + + def raise_exc(val): + raise JsonSchemaValidationError("is invalid") + + with patch( + "application.api.user.agents.routes.normalize_json_schema_payload", + side_effect=raise_exc, + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "Agent", + "status": "draft", + "json_schema": {"bad": True}, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 400 + + def test_folder_id_not_found(self, app): + from application.api.user.agents.routes import CreateAgent + + mock_handle_img = Mock(return_value=("", None)) + mock_folders = Mock() + mock_folders.find_one.return_value = None + folder_id = str(ObjectId()) + + with patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ), patch( + "application.api.user.agents.routes.agent_folders_collection", mock_folders + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "Agent", + "status": "draft", + "agent_type": "classic", + "folder_id": folder_id, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 404 + + def test_invalid_folder_id_format(self, app): + from application.api.user.agents.routes import CreateAgent + + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "Agent", + "status": "draft", + "agent_type": "classic", + "folder_id": "not-valid", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 400 + + def test_form_data_with_json_fields(self, app): + from application.api.user.agents.routes import CreateAgent + + inserted_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.insert_one.return_value = Mock(inserted_id=inserted_id) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + content_type="multipart/form-data", + data={ + "name": "FormAgent", + "status": "draft", + "agent_type": "classic", + "tools": '["tool1", "tool2"]', + "sources": '["src1"]', + "models": '["model1"]', + "json_schema": "null", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 201 + + def test_form_data_invalid_json_tools(self, app): + from application.api.user.agents.routes import CreateAgent + + inserted_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.insert_one.return_value = Mock(inserted_id=inserted_id) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + content_type="multipart/form-data", + data={ + "name": "FormAgent", + "status": "draft", + "agent_type": "classic", + "tools": "not-valid-json", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + # invalid json for tools falls back to [] + assert response.status_code == 201 + + def test_create_with_sources_list(self, app): + from application.api.user.agents.routes import CreateAgent + + inserted_id = ObjectId() + src_id = str(ObjectId()) + mock_agents_col = Mock() + mock_agents_col.insert_one.return_value = Mock(inserted_id=inserted_id) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "MultiSource", + "description": "desc", + "status": "published", + "agent_type": "classic", + "sources": [src_id, "default"], + "chunks": "5", + "retriever": "classic", + "prompt_id": "default", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 201 + + def test_db_insert_failure(self, app): + from application.api.user.agents.routes import CreateAgent + + mock_agents_col = Mock() + mock_agents_col.insert_one.side_effect = Exception("insert error") + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + "/api/create_agent", + method="POST", + json={ + "name": "Agent", + "status": "draft", + "agent_type": "classic", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateAgent().post() + assert response.status_code == 400 + + +@pytest.mark.unit +class TestUpdateAgent: + + def _make_existing_agent(self, agent_id=None): + return { + "_id": agent_id or ObjectId(), + "user": "user1", + "name": "Existing Agent", + "description": "existing desc", + "source": "default", + "chunks": "5", + "retriever": "classic", + "prompt_id": "default", + "status": "published", + "agent_type": "classic", + "key": str(uuid.uuid4()), + } + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import UpdateAgent + + with app.test_request_context( + "/api/update_agent/abc", + method="PUT", + json={"name": "A"}, + ): + from flask import request + + request.decoded_token = None + response = UpdateAgent().put("abc") + assert response.status_code == 401 + + def test_returns_400_invalid_agent_id(self, app): + from application.api.user.agents.routes import UpdateAgent + + with app.test_request_context( + "/api/update_agent/not-valid", + method="PUT", + json={"name": "A"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put("not-valid") + assert response.status_code == 400 + + def test_returns_404_agent_not_found(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = str(ObjectId()) + mock_col = Mock() + mock_col.find_one.return_value = None + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"name": "Updated"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(agent_id) + assert response.status_code == 404 + + def test_updates_agent_name_success(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"name": "Updated Name"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + assert response.json["success"] is True + + def test_returns_400_invalid_status(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"status": "invalid_status"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_returns_400_negative_chunks(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"chunks": -1}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_returns_400_tools_not_list(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"tools": "not-a-list"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_returns_400_no_update_data(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"unknown_field": "value"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_limited_token_mode_without_limit(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"limited_token_mode": True}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_limited_request_mode_without_limit(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"limited_request_mode": True}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_token_limit_without_mode(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"token_limit": 1000}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_request_limit_without_mode(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"request_limit": 100}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_source_with_invalid_oid(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"source": "not-valid-oid"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_sources_list_with_invalid_oid(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"sources": ["bad-id"]}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_update_source_to_default(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"source": "default"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + + def test_update_empty_source_on_draft(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + existing["status"] = "draft" + existing["key"] = "" + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"source": ""}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + + def test_update_empty_source_on_published_fails(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + # Published agent with only "default" source + existing["source"] = "default" + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"source": ""}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_update_chunks_empty_defaults_to_2(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"chunks": ""}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + call_args = mock_col.update_one.call_args[0][1]["$set"] + assert call_args["chunks"] == "2" + + def test_update_invalid_chunks_value(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"chunks": "abc"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_update_matched_but_not_modified(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=0) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"name": "Same Name"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + assert "No changes detected" in response.json["message"] + + def test_update_matched_zero_returns_404(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=0, modified_count=0) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"name": "New Name"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 404 + + def test_publish_draft_generates_key(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + existing["status"] = "draft" + existing["key"] = "" + existing["source"] = "default" + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"status": "published"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + assert "key" in response.json + + def test_publish_missing_required_fields(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = { + "_id": agent_id, + "user": "user1", + "name": "", + "description": "", + "source": "", + "chunks": "", + "retriever": "", + "prompt_id": "", + "status": "draft", + "agent_type": "classic", + "key": "", + } + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"status": "published"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_db_update_exception(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.side_effect = Exception("DB error") + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"name": "Updated"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 500 + + def test_update_json_schema_valid(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + schema_val = {"type": "object", "properties": {"x": {"type": "string"}}} + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ), patch( + "application.api.user.agents.routes.normalize_json_schema_payload", + return_value=schema_val, + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"json_schema": schema_val}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + + def test_update_json_schema_invalid(self, app): + from application.api.user.agents.routes import UpdateAgent + from application.core.json_schema_utils import JsonSchemaValidationError + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + def raise_exc(val): + raise JsonSchemaValidationError("is invalid") + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ), patch( + "application.api.user.agents.routes.normalize_json_schema_payload", + side_effect=raise_exc, + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"json_schema": {"bad": True}}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_update_json_schema_none(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"json_schema": None}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + + def test_update_folder_id_valid(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + folder_id = str(ObjectId()) + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + mock_folders = Mock() + mock_folders.find_one.return_value = {"_id": ObjectId(folder_id), "user": "user1"} + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ), patch( + "application.api.user.agents.routes.agent_folders_collection", mock_folders + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"folder_id": folder_id}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + + def test_update_folder_id_invalid_format(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"folder_id": "invalid"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_update_folder_not_found(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + folder_id = str(ObjectId()) + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + mock_folders = Mock() + mock_folders.find_one.return_value = None + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ), patch( + "application.api.user.agents.routes.agent_folders_collection", mock_folders + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"folder_id": folder_id}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 404 + + def test_update_folder_id_empty_sets_none(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"folder_id": ""}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + call_args = mock_col.update_one.call_args[0][1]["$set"] + assert call_args["folder_id"] is None + + def test_empty_name_field_rejected(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"name": ""}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_update_workflow_field(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + wf_id = str(ObjectId()) + existing = self._make_existing_agent(agent_id) + existing["agent_type"] = "workflow" + existing["status"] = "draft" + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + mock_wf_col = Mock() + mock_wf_col.find_one.return_value = {"_id": ObjectId(wf_id), "user": "user1"} + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ), patch( + "application.api.user.agents.routes.workflows_collection", mock_wf_col + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"workflow": wf_id}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + + def test_publish_workflow_without_workflow_field(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = { + "_id": agent_id, + "user": "user1", + "name": "WF Agent", + "status": "draft", + "agent_type": "workflow", + "key": "", + "workflow": None, + } + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"status": "published"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 400 + + def test_form_data_json_parse_error(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = str(ObjectId()) + existing = self._make_existing_agent(ObjectId(agent_id)) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + content_type="multipart/form-data", + data={"tools": "not-valid-json"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(agent_id) + assert response.status_code == 400 + + def test_limited_token_mode_string_true(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={ + "limited_token_mode": "True", + "token_limit": 5000, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + + def test_update_sources_with_default(self, app): + from application.api.user.agents.routes import UpdateAgent + + agent_id = ObjectId() + src_id = str(ObjectId()) + existing = self._make_existing_agent(agent_id) + mock_col = Mock() + mock_col.find_one.return_value = existing + mock_col.update_one.return_value = Mock(matched_count=1, modified_count=1) + mock_handle_img = Mock(return_value=("", None)) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.handle_image_upload", mock_handle_img + ): + with app.test_request_context( + f"/api/update_agent/{agent_id}", + method="PUT", + json={"sources": ["default", src_id]}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateAgent().put(str(agent_id)) + assert response.status_code == 200 + + +@pytest.mark.unit +class TestDeleteAgent: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import DeleteAgent + + with app.test_request_context("/api/delete_agent"): + from flask import request + + request.decoded_token = None + response = DeleteAgent().delete() + assert response.status_code == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.agents.routes import DeleteAgent + + with app.test_request_context("/api/delete_agent"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteAgent().delete() + assert response.status_code == 400 + + def test_returns_404_agent_not_found(self, app): + from application.api.user.agents.routes import DeleteAgent + + mock_col = Mock() + mock_col.find_one_and_delete.return_value = None + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/delete_agent?id={ObjectId()}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteAgent().delete() + assert response.status_code == 404 + + def test_deletes_classic_agent_success(self, app): + from application.api.user.agents.routes import DeleteAgent + + agent_id = ObjectId() + mock_col = Mock() + mock_col.find_one_and_delete.return_value = { + "_id": agent_id, + "user": "user1", + "agent_type": "classic", + } + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/delete_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteAgent().delete() + assert response.status_code == 200 + assert response.json["id"] == str(agent_id) + + def test_deletes_workflow_agent_cleans_up(self, app): + from application.api.user.agents.routes import DeleteAgent + + agent_id = ObjectId() + wf_id = str(ObjectId()) + mock_agents_col = Mock() + mock_agents_col.find_one_and_delete.return_value = { + "_id": agent_id, + "user": "user1", + "agent_type": "workflow", + "workflow": wf_id, + } + mock_wf_col = Mock() + mock_wf_col.find_one.return_value = {"_id": ObjectId(wf_id), "user": "user1"} + mock_nodes_col = Mock() + mock_edges_col = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.workflows_collection", mock_wf_col + ), patch( + "application.api.user.agents.routes.workflow_nodes_collection", mock_nodes_col + ), patch( + "application.api.user.agents.routes.workflow_edges_collection", mock_edges_col + ): + with app.test_request_context(f"/api/delete_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteAgent().delete() + assert response.status_code == 200 + mock_nodes_col.delete_many.assert_called_once() + mock_edges_col.delete_many.assert_called_once() + mock_wf_col.delete_one.assert_called_once() + + def test_deletes_workflow_agent_non_owned_skips_cleanup(self, app): + from application.api.user.agents.routes import DeleteAgent + + agent_id = ObjectId() + wf_id = str(ObjectId()) + mock_agents_col = Mock() + mock_agents_col.find_one_and_delete.return_value = { + "_id": agent_id, + "user": "user1", + "agent_type": "workflow", + "workflow": wf_id, + } + mock_wf_col = Mock() + mock_wf_col.find_one.return_value = None # Not owned + mock_nodes_col = Mock() + mock_edges_col = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.workflows_collection", mock_wf_col + ), patch( + "application.api.user.agents.routes.workflow_nodes_collection", mock_nodes_col + ), patch( + "application.api.user.agents.routes.workflow_edges_collection", mock_edges_col + ): + with app.test_request_context(f"/api/delete_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteAgent().delete() + assert response.status_code == 200 + mock_nodes_col.delete_many.assert_not_called() + + def test_returns_400_on_exception(self, app): + from application.api.user.agents.routes import DeleteAgent + + mock_col = Mock() + mock_col.find_one_and_delete.side_effect = Exception("DB error") + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/delete_agent?id={ObjectId()}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteAgent().delete() + assert response.status_code == 400 + + +@pytest.mark.unit +class TestPinnedAgents: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import PinnedAgents + + with app.test_request_context("/api/pinned_agents"): + from flask import request + + request.decoded_token = None + response = PinnedAgents().get() + assert response.status_code == 401 + + def test_returns_empty_when_no_pinned(self, app): + from application.api.user.agents.routes import PinnedAgents + + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": []}, + } + ) + + with patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ): + with app.test_request_context("/api/pinned_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinnedAgents().get() + assert response.status_code == 200 + assert response.json == [] + + def test_returns_pinned_agents(self, app): + from application.api.user.agents.routes import PinnedAgents + + agent_id = ObjectId() + key = str(uuid.uuid4()) + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": [str(agent_id)]}, + } + ) + mock_agents_col = Mock() + mock_agents_col.find.return_value = [ + { + "_id": agent_id, + "name": "Pinned Agent", + "source": "default", + "retriever": "classic", + "key": key, + } + ] + mock_resolve = Mock(return_value=[]) + mock_users_col = Mock() + mock_db = Mock() + + with patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.routes.users_collection", mock_users_col + ), patch( + "application.api.user.agents.routes.db", mock_db + ): + with app.test_request_context("/api/pinned_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinnedAgents().get() + assert response.status_code == 200 + data = response.json + assert len(data) == 1 + assert data[0]["pinned"] is True + + def test_cleans_up_stale_pinned_ids(self, app): + from application.api.user.agents.routes import PinnedAgents + + stale_id = str(ObjectId()) + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": [stale_id]}, + } + ) + mock_agents_col = Mock() + mock_agents_col.find.return_value = [] + mock_users_col = Mock() + + with patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.routes.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.routes.users_collection", mock_users_col + ): + with app.test_request_context("/api/pinned_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinnedAgents().get() + assert response.status_code == 200 + mock_users_col.update_one.assert_called_once() + + def test_returns_400_on_exception(self, app): + from application.api.user.agents.routes import PinnedAgents + + mock_ensure = Mock(side_effect=Exception("DB error")) + + with patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ): + with app.test_request_context("/api/pinned_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinnedAgents().get() + assert response.status_code == 400 + + +@pytest.mark.unit +class TestGetTemplateAgents: + + def test_returns_template_agents(self, app): + from application.api.user.agents.routes import GetTemplateAgents + + agent_id = ObjectId() + mock_col = Mock() + mock_col.find.return_value = [ + { + "_id": agent_id, + "name": "Template1", + "description": "A template", + "image": "img.png", + } + ] + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context("/api/template_agents"): + response = GetTemplateAgents().get() + assert response.status_code == 200 + data = response.json + assert len(data) == 1 + assert data[0]["name"] == "Template1" + + def test_returns_400_on_exception(self, app): + from application.api.user.agents.routes import GetTemplateAgents + + mock_col = Mock() + mock_col.find.side_effect = Exception("DB error") + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context("/api/template_agents"): + response = GetTemplateAgents().get() + assert response.status_code == 400 + + +@pytest.mark.unit +class TestAdoptAgent: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import AdoptAgent + + with app.test_request_context("/api/adopt_agent?id=abc"): + from flask import request + + request.decoded_token = None + response = AdoptAgent().post() + assert response.status_code == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.agents.routes import AdoptAgent + + with app.test_request_context("/api/adopt_agent"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AdoptAgent().post() + assert response.status_code == 400 + + def test_returns_404_template_not_found(self, app): + from application.api.user.agents.routes import AdoptAgent + + mock_col = Mock() + mock_col.find_one.return_value = None + agent_id = str(ObjectId()) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/adopt_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AdoptAgent().post() + assert response.status_code == 404 + + def test_adopts_agent_success(self, app): + from application.api.user.agents.routes import AdoptAgent + + agent_id = ObjectId() + new_id = ObjectId() + mock_col = Mock() + mock_col.find_one.return_value = { + "_id": agent_id, + "user": "system", + "name": "Template Agent", + "description": "A template", + "source": "default", + "tools": [], + } + mock_col.insert_one.return_value = Mock(inserted_id=new_id) + mock_resolve = Mock(return_value=[]) + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.resolve_tool_details", mock_resolve + ): + with app.test_request_context(f"/api/adopt_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AdoptAgent().post() + assert response.status_code == 200 + data = response.json + assert data["success"] is True + assert data["agent"]["id"] == str(new_id) + + def test_returns_400_on_exception(self, app): + from application.api.user.agents.routes import AdoptAgent + + mock_col = Mock() + mock_col.find_one.side_effect = Exception("DB error") + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/adopt_agent?id={ObjectId()}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AdoptAgent().post() + assert response.status_code == 400 + + +@pytest.mark.unit +class TestPinAgent: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import PinAgent + + with app.test_request_context("/api/pin_agent"): + from flask import request + + request.decoded_token = None + response = PinAgent().post() + assert response.status_code == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.agents.routes import PinAgent + + with app.test_request_context("/api/pin_agent"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinAgent().post() + assert response.status_code == 400 + + def test_returns_404_agent_not_found(self, app): + from application.api.user.agents.routes import PinAgent + + mock_col = Mock() + mock_col.find_one.return_value = None + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/pin_agent?id={ObjectId()}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinAgent().post() + assert response.status_code == 404 + + def test_pins_agent(self, app): + from application.api.user.agents.routes import PinAgent + + agent_id = str(ObjectId()) + mock_col = Mock() + mock_col.find_one.return_value = {"_id": ObjectId(agent_id)} + + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": []}, + } + ) + mock_users_col = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.routes.users_collection", mock_users_col + ): + with app.test_request_context(f"/api/pin_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinAgent().post() + assert response.status_code == 200 + assert response.json["action"] == "pinned" + + def test_unpins_agent(self, app): + from application.api.user.agents.routes import PinAgent + + agent_id = str(ObjectId()) + mock_col = Mock() + mock_col.find_one.return_value = {"_id": ObjectId(agent_id)} + + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": [agent_id]}, + } + ) + mock_users_col = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.routes.users_collection", mock_users_col + ): + with app.test_request_context(f"/api/pin_agent?id={agent_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinAgent().post() + assert response.status_code == 200 + assert response.json["action"] == "unpinned" + + def test_returns_500_on_exception(self, app): + from application.api.user.agents.routes import PinAgent + + mock_col = Mock() + mock_col.find_one.side_effect = Exception("DB error") + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context(f"/api/pin_agent?id={ObjectId()}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = PinAgent().post() + assert response.status_code == 500 + + +@pytest.mark.unit +class TestRemoveSharedAgent: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.routes import RemoveSharedAgent + + with app.test_request_context("/api/remove_shared_agent"): + from flask import request + + request.decoded_token = None + response = RemoveSharedAgent().delete() + assert response.status_code == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.agents.routes import RemoveSharedAgent + + with app.test_request_context("/api/remove_shared_agent"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = RemoveSharedAgent().delete() + assert response.status_code == 400 + + def test_returns_404_shared_agent_not_found(self, app): + from application.api.user.agents.routes import RemoveSharedAgent + + mock_col = Mock() + mock_col.find_one.return_value = None + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context( + f"/api/remove_shared_agent?id={ObjectId()}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = RemoveSharedAgent().delete() + assert response.status_code == 404 + + def test_removes_shared_agent_success(self, app): + from application.api.user.agents.routes import RemoveSharedAgent + + agent_id = str(ObjectId()) + mock_col = Mock() + mock_col.find_one.return_value = { + "_id": ObjectId(agent_id), + "shared_publicly": True, + } + + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"pinned": [], "shared_with_me": [agent_id]}, + } + ) + mock_users_col = Mock() + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ), patch( + "application.api.user.agents.routes.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.routes.users_collection", mock_users_col + ): + with app.test_request_context( + f"/api/remove_shared_agent?id={agent_id}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = RemoveSharedAgent().delete() + assert response.status_code == 200 + assert response.json["action"] == "removed" + mock_users_col.update_one.assert_called_once() + + def test_returns_500_on_exception(self, app): + from application.api.user.agents.routes import RemoveSharedAgent + + mock_col = Mock() + mock_col.find_one.side_effect = Exception("DB error") + + with patch( + "application.api.user.agents.routes.agents_collection", mock_col + ): + with app.test_request_context( + f"/api/remove_shared_agent?id={ObjectId()}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = RemoveSharedAgent().delete() + assert response.status_code == 500 diff --git a/tests/api/user/test_agents_sharing.py b/tests/api/user/test_agents_sharing.py new file mode 100644 index 00000000..eb626508 --- /dev/null +++ b/tests/api/user/test_agents_sharing.py @@ -0,0 +1,768 @@ +"""Tests for application.api.user.agents.sharing module.""" + +from unittest.mock import Mock, patch + +import pytest +from bson import DBRef, ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +# --------------------------------------------------------------------------- +# SharedAgent (GET /shared_agent) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestSharedAgent: + + def test_returns_400_missing_token(self, app): + from application.api.user.agents.sharing import SharedAgent + + with app.test_request_context("/api/shared_agent"): + response = SharedAgent().get() + assert response.status_code == 400 + + def test_returns_404_agent_not_found(self, app): + from application.api.user.agents.sharing import SharedAgent + + mock_col = Mock() + mock_col.find_one.return_value = None + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_col + ): + with app.test_request_context("/api/shared_agent?token=abc123"): + response = SharedAgent().get() + assert response.status_code == 404 + + def test_returns_shared_agent_data(self, app): + from application.api.user.agents.sharing import SharedAgent + + agent_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.find_one.return_value = { + "_id": agent_id, + "user": "owner1", + "name": "Shared Agent", + "description": "A shared agent", + "chunks": "5", + "retriever": "classic", + "prompt_id": "default", + "tools": [], + "agent_type": "classic", + "status": "published", + "shared_publicly": True, + "shared_token": "abc123", + } + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.db", mock_db + ): + with app.test_request_context("/api/shared_agent?token=abc123"): + from flask import request + + # No decoded_token -> anonymous access + request.decoded_token = None + response = SharedAgent().get() + assert response.status_code == 200 + data = response.json + assert data["id"] == str(agent_id) + assert data["name"] == "Shared Agent" + assert data["shared"] is True + + def test_adds_to_shared_with_me_for_different_user(self, app): + from application.api.user.agents.sharing import SharedAgent + + agent_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.find_one.return_value = { + "_id": agent_id, + "user": "owner1", + "name": "Agent", + "tools": [], + "shared_publicly": True, + "shared_token": "abc123", + } + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + mock_ensure = Mock(return_value={"user_id": "user2"}) + mock_users_col = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.db", mock_db + ), patch( + "application.api.user.agents.sharing.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.sharing.users_collection", mock_users_col + ): + with app.test_request_context("/api/shared_agent?token=abc123"): + from flask import request + + request.decoded_token = {"sub": "user2"} + response = SharedAgent().get() + assert response.status_code == 200 + mock_ensure.assert_called_once_with("user2") + mock_users_col.update_one.assert_called_once() + + def test_does_not_add_to_shared_for_owner(self, app): + from application.api.user.agents.sharing import SharedAgent + + agent_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.find_one.return_value = { + "_id": agent_id, + "user": "owner1", + "name": "Agent", + "tools": [], + "shared_publicly": True, + "shared_token": "abc123", + } + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + mock_ensure = Mock() + mock_users_col = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.db", mock_db + ), patch( + "application.api.user.agents.sharing.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.sharing.users_collection", mock_users_col + ): + with app.test_request_context("/api/shared_agent?token=abc123"): + from flask import request + + request.decoded_token = {"sub": "owner1"} + response = SharedAgent().get() + assert response.status_code == 200 + mock_ensure.assert_not_called() + mock_users_col.update_one.assert_not_called() + + def test_enriches_tool_names(self, app): + from application.api.user.agents.sharing import SharedAgent + + agent_id = ObjectId() + tool_id = str(ObjectId()) + mock_agents_col = Mock() + mock_agents_col.find_one.return_value = { + "_id": agent_id, + "user": "owner1", + "name": "Agent", + "tools": [tool_id], + "shared_publicly": True, + "shared_token": "tok", + } + mock_tools_col = Mock() + mock_tools_col.find_one.return_value = { + "_id": ObjectId(tool_id), + "name": "calculator", + } + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.user_tools_collection", mock_tools_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.db", mock_db + ): + with app.test_request_context("/api/shared_agent?token=tok"): + from flask import request + + request.decoded_token = None + response = SharedAgent().get() + assert response.status_code == 200 + assert response.json["tools"] == ["calculator"] + + def test_handles_source_dbref(self, app): + from application.api.user.agents.sharing import SharedAgent + + agent_id = ObjectId() + source_id = ObjectId() + source_ref = DBRef("sources", source_id) + mock_agents_col = Mock() + mock_agents_col.find_one.return_value = { + "_id": agent_id, + "user": "owner1", + "name": "Agent", + "source": source_ref, + "tools": [], + "shared_publicly": True, + "shared_token": "tok", + } + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + mock_db.dereference.return_value = {"_id": source_id} + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.db", mock_db + ): + with app.test_request_context("/api/shared_agent?token=tok"): + from flask import request + + request.decoded_token = None + response = SharedAgent().get() + assert response.status_code == 200 + assert response.json["source"] == str(source_id) + + def test_returns_400_on_exception(self, app): + from application.api.user.agents.sharing import SharedAgent + + mock_col = Mock() + mock_col.find_one.side_effect = Exception("DB error") + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_col + ): + with app.test_request_context("/api/shared_agent?token=tok"): + response = SharedAgent().get() + assert response.status_code == 400 + + def test_tool_enrichment_handles_missing_tool(self, app): + from application.api.user.agents.sharing import SharedAgent + + agent_id = ObjectId() + tool_id = str(ObjectId()) + mock_agents_col = Mock() + mock_agents_col.find_one.return_value = { + "_id": agent_id, + "user": "owner1", + "name": "Agent", + "tools": [tool_id], + "shared_publicly": True, + "shared_token": "tok", + } + mock_tools_col = Mock() + mock_tools_col.find_one.return_value = None + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.user_tools_collection", mock_tools_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.db", mock_db + ): + with app.test_request_context("/api/shared_agent?token=tok"): + from flask import request + + request.decoded_token = None + response = SharedAgent().get() + assert response.status_code == 200 + # Missing tools are skipped + assert response.json["tools"] == [] + + def test_image_url_generated_when_present(self, app): + from application.api.user.agents.sharing import SharedAgent + + agent_id = ObjectId() + mock_agents_col = Mock() + mock_agents_col.find_one.return_value = { + "_id": agent_id, + "user": "owner1", + "name": "Agent", + "image": "path/to/img.png", + "tools": [], + "shared_publicly": True, + "shared_token": "tok", + } + mock_resolve = Mock(return_value=[]) + mock_db = Mock() + mock_generate = Mock(return_value="http://example.com/img.png") + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.db", mock_db + ), patch( + "application.api.user.agents.sharing.generate_image_url", mock_generate + ): + with app.test_request_context("/api/shared_agent?token=tok"): + from flask import request + + request.decoded_token = None + response = SharedAgent().get() + assert response.status_code == 200 + assert response.json["image"] == "http://example.com/img.png" + mock_generate.assert_called_once_with("path/to/img.png") + + +# --------------------------------------------------------------------------- +# SharedAgents (GET /shared_agents) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestSharedAgents: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.sharing import SharedAgents + + with app.test_request_context("/api/shared_agents"): + from flask import request + + request.decoded_token = None + response = SharedAgents().get() + assert response.status_code == 401 + + def test_returns_shared_agents_list(self, app): + from application.api.user.agents.sharing import SharedAgents + + agent_id = ObjectId() + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": { + "shared_with_me": [str(agent_id)], + "pinned": [str(agent_id)], + }, + } + ) + mock_agents_col = Mock() + mock_agents_col.find.return_value = [ + { + "_id": agent_id, + "name": "Shared Agent", + "description": "desc", + "tools": [], + "agent_type": "classic", + "status": "published", + "shared_publicly": True, + "shared_token": "tok123", + } + ] + mock_resolve = Mock(return_value=[]) + mock_users_col = Mock() + + with patch( + "application.api.user.agents.sharing.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.users_collection", mock_users_col + ): + with app.test_request_context("/api/shared_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = SharedAgents().get() + assert response.status_code == 200 + data = response.json + assert len(data) == 1 + assert data[0]["name"] == "Shared Agent" + assert data[0]["pinned"] is True + + def test_removes_stale_shared_ids(self, app): + from application.api.user.agents.sharing import SharedAgents + + stale_id = str(ObjectId()) + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": { + "shared_with_me": [stale_id], + "pinned": [], + }, + } + ) + mock_agents_col = Mock() + mock_agents_col.find.return_value = [] # None found + mock_users_col = Mock() + + with patch( + "application.api.user.agents.sharing.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.users_collection", mock_users_col + ): + with app.test_request_context("/api/shared_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = SharedAgents().get() + assert response.status_code == 200 + mock_users_col.update_one.assert_called_once() + call_args = mock_users_col.update_one.call_args + assert stale_id in call_args[0][1]["$pullAll"][ + "agent_preferences.shared_with_me" + ] + + def test_returns_empty_when_no_shared_ids(self, app): + from application.api.user.agents.sharing import SharedAgents + + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": {"shared_with_me": [], "pinned": []}, + } + ) + mock_agents_col = Mock() + mock_agents_col.find.return_value = [] + + with patch( + "application.api.user.agents.sharing.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ): + with app.test_request_context("/api/shared_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = SharedAgents().get() + assert response.status_code == 200 + assert response.json == [] + + def test_returns_400_on_exception(self, app): + from application.api.user.agents.sharing import SharedAgents + + mock_ensure = Mock(side_effect=Exception("DB error")) + + with patch( + "application.api.user.agents.sharing.ensure_user_doc", mock_ensure + ): + with app.test_request_context("/api/shared_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = SharedAgents().get() + assert response.status_code == 400 + + def test_image_url_generated(self, app): + from application.api.user.agents.sharing import SharedAgents + + agent_id = ObjectId() + mock_ensure = Mock( + return_value={ + "user_id": "user1", + "agent_preferences": { + "shared_with_me": [str(agent_id)], + "pinned": [], + }, + } + ) + mock_agents_col = Mock() + mock_agents_col.find.return_value = [ + { + "_id": agent_id, + "name": "Agent", + "image": "path.png", + "tools": [], + "shared_publicly": True, + } + ] + mock_resolve = Mock(return_value=[]) + mock_generate = Mock(return_value="http://example.com/path.png") + + with patch( + "application.api.user.agents.sharing.ensure_user_doc", mock_ensure + ), patch( + "application.api.user.agents.sharing.agents_collection", mock_agents_col + ), patch( + "application.api.user.agents.sharing.resolve_tool_details", mock_resolve + ), patch( + "application.api.user.agents.sharing.generate_image_url", mock_generate + ), patch( + "application.api.user.agents.sharing.users_collection", Mock() + ): + with app.test_request_context("/api/shared_agents"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = SharedAgents().get() + assert response.status_code == 200 + assert response.json[0]["image"] == "http://example.com/path.png" + + +# --------------------------------------------------------------------------- +# ShareAgent (PUT /share_agent) +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestShareAgent: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.sharing import ShareAgent + + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={"id": "abc", "shared": True}, + ): + from flask import request + + request.decoded_token = None + response = ShareAgent().put() + assert response.status_code == 401 + + def test_returns_400_missing_json_body(self, app): + from application.api.user.agents.sharing import ShareAgent + + with app.test_request_context( + "/api/share_agent", + method="PUT", + content_type="application/json", + data=b"{}", + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + # Empty JSON object -> no id, no shared -> 400 + response = ShareAgent().put() + assert response.status_code == 400 + assert response.json["success"] is False + + def test_returns_400_missing_id(self, app): + from application.api.user.agents.sharing import ShareAgent + + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={"shared": True}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 400 + + def test_returns_400_missing_shared_param(self, app): + from application.api.user.agents.sharing import ShareAgent + + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={"id": str(ObjectId())}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 400 + + def test_returns_400_invalid_agent_id(self, app): + from application.api.user.agents.sharing import ShareAgent + + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={"id": "invalid-oid", "shared": True}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 400 + + def test_returns_404_agent_not_found(self, app): + from application.api.user.agents.sharing import ShareAgent + + mock_col = Mock() + mock_col.find_one.return_value = None + agent_id = str(ObjectId()) + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_col + ): + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={"id": agent_id, "shared": True}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 404 + + def test_shares_agent_success(self, app): + from application.api.user.agents.sharing import ShareAgent + + agent_id = ObjectId() + mock_col = Mock() + mock_col.find_one.return_value = { + "_id": agent_id, + "user": "user1", + } + mock_col.update_one.return_value = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_col + ): + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={ + "id": str(agent_id), + "shared": True, + "username": "TestUser", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 200 + data = response.json + assert data["success"] is True + assert data["shared_token"] is not None + mock_col.update_one.assert_called_once() + + def test_unshares_agent_success(self, app): + from application.api.user.agents.sharing import ShareAgent + + agent_id = ObjectId() + mock_col = Mock() + mock_col.find_one.return_value = { + "_id": agent_id, + "user": "user1", + } + mock_col.update_one.return_value = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_col + ): + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={ + "id": str(agent_id), + "shared": False, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 200 + data = response.json + assert data["success"] is True + assert data["shared_token"] is None + + def test_returns_400_on_db_exception(self, app): + from application.api.user.agents.sharing import ShareAgent + + agent_id = ObjectId() + mock_col = Mock() + mock_col.find_one.return_value = { + "_id": agent_id, + "user": "user1", + } + mock_col.update_one.side_effect = Exception("DB error") + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_col + ): + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={ + "id": str(agent_id), + "shared": True, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 400 + + def test_share_with_username(self, app): + from application.api.user.agents.sharing import ShareAgent + + agent_id = ObjectId() + mock_col = Mock() + mock_col.find_one.return_value = { + "_id": agent_id, + "user": "user1", + } + mock_col.update_one.return_value = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_col + ): + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={ + "id": str(agent_id), + "shared": True, + "username": "SharedByUser", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 200 + # Verify the update call includes shared_metadata with username + update_call = mock_col.update_one.call_args[0][1]["$set"] + assert update_call["shared_metadata"]["shared_by"] == "SharedByUser" + assert update_call["shared_publicly"] is True + assert "shared_token" in update_call + + def test_shared_false_explicitly(self, app): + from application.api.user.agents.sharing import ShareAgent + + agent_id = ObjectId() + mock_col = Mock() + mock_col.find_one.return_value = { + "_id": agent_id, + "user": "user1", + } + mock_col.update_one.return_value = Mock() + + with patch( + "application.api.user.agents.sharing.agents_collection", mock_col + ): + with app.test_request_context( + "/api/share_agent", + method="PUT", + json={ + "id": str(agent_id), + "shared": False, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareAgent().put() + assert response.status_code == 200 + update_call = mock_col.update_one.call_args[0][1] + assert update_call["$set"]["shared_publicly"] is False + assert update_call["$set"]["shared_token"] is None diff --git a/tests/api/user/test_analytics.py b/tests/api/user/test_analytics.py new file mode 100644 index 00000000..12e98bb0 --- /dev/null +++ b/tests/api/user/test_analytics.py @@ -0,0 +1,388 @@ +import datetime +from unittest.mock import Mock, patch + +import pytest +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +@pytest.mark.unit +class TestGetMessageAnalytics: + + def test_returns_message_analytics_last_30_days(self, app): + from application.api.user.analytics.routes import GetMessageAnalytics + + mock_conversations = Mock() + mock_conversations.aggregate.return_value = [ + {"_id": "2024-06-01", "count": 5}, + {"_id": "2024-06-02", "count": 3}, + ] + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_message_analytics", + method="POST", + json={"filter_option": "last_30_days"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetMessageAnalytics().post() + + assert response.status_code == 200 + assert response.json["success"] is True + assert "messages" in response.json + + def test_returns_401_unauthenticated(self, app): + from application.api.user.analytics.routes import GetMessageAnalytics + + with app.test_request_context( + "/api/get_message_analytics", + method="POST", + json={"filter_option": "last_30_days"}, + ): + from flask import request + + request.decoded_token = None + response = GetMessageAnalytics().post() + + assert response.status_code == 401 + + def test_returns_400_invalid_filter_option(self, app): + from application.api.user.analytics.routes import GetMessageAnalytics + + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_message_analytics", + method="POST", + json={"filter_option": "invalid_option"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetMessageAnalytics().post() + + assert response.status_code == 400 + + def test_filters_by_api_key(self, app): + from application.api.user.analytics.routes import GetMessageAnalytics + + agent_id = ObjectId() + mock_agents = Mock() + mock_agents.find_one.return_value = { + "_id": agent_id, + "key": "api_key_value", + } + mock_conversations = Mock() + mock_conversations.aggregate.return_value = [] + + with patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ), patch( + "application.api.user.analytics.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + "/api/get_message_analytics", + method="POST", + json={ + "filter_option": "last_7_days", + "api_key_id": str(agent_id), + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetMessageAnalytics().post() + + assert response.status_code == 200 + pipeline = mock_conversations.aggregate.call_args[0][0] + assert pipeline[0]["$match"].get("api_key") == "api_key_value" + + def test_last_hour_filter(self, app): + from application.api.user.analytics.routes import GetMessageAnalytics + + mock_conversations = Mock() + mock_conversations.aggregate.return_value = [] + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_message_analytics", + method="POST", + json={"filter_option": "last_hour"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetMessageAnalytics().post() + + assert response.status_code == 200 + + def test_last_24_hour_filter(self, app): + from application.api.user.analytics.routes import GetMessageAnalytics + + mock_conversations = Mock() + mock_conversations.aggregate.return_value = [] + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_message_analytics", + method="POST", + json={"filter_option": "last_24_hour"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetMessageAnalytics().post() + + assert response.status_code == 200 + + +@pytest.mark.unit +class TestGetTokenAnalytics: + + def test_returns_token_analytics(self, app): + from application.api.user.analytics.routes import GetTokenAnalytics + + mock_token_usage = Mock() + mock_token_usage.aggregate.return_value = [ + {"_id": {"day": "2024-06-01"}, "total_tokens": 1000} + ] + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.token_usage_collection", + mock_token_usage, + ), patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_token_analytics", + method="POST", + json={"filter_option": "last_30_days"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetTokenAnalytics().post() + + assert response.status_code == 200 + assert response.json["success"] is True + assert "token_usage" in response.json + + def test_returns_400_invalid_filter(self, app): + from application.api.user.analytics.routes import GetTokenAnalytics + + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_token_analytics", + method="POST", + json={"filter_option": "invalid"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetTokenAnalytics().post() + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestGetFeedbackAnalytics: + + def test_returns_feedback_analytics(self, app): + from application.api.user.analytics.routes import GetFeedbackAnalytics + + mock_conversations = Mock() + mock_conversations.aggregate.return_value = [ + {"_id": "2024-06-01", "positive": 10, "negative": 2} + ] + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_feedback_analytics", + method="POST", + json={"filter_option": "last_30_days"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetFeedbackAnalytics().post() + + assert response.status_code == 200 + assert response.json["success"] is True + assert "feedback" in response.json + + def test_returns_400_invalid_filter(self, app): + from application.api.user.analytics.routes import GetFeedbackAnalytics + + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_feedback_analytics", + method="POST", + json={"filter_option": "bad"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetFeedbackAnalytics().post() + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestGetUserLogs: + + def test_returns_paginated_logs(self, app): + from application.api.user.analytics.routes import GetUserLogs + + log_id = ObjectId() + mock_cursor = Mock() + mock_cursor.sort.return_value.skip.return_value.limit.return_value = [ + { + "_id": log_id, + "action": "query", + "level": "info", + "user": "user1", + "question": "test?", + "sources": [], + "retriever_params": {}, + "timestamp": datetime.datetime(2024, 6, 1), + } + ] + mock_user_logs = Mock() + mock_user_logs.find.return_value = mock_cursor + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.user_logs_collection", + mock_user_logs, + ), patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_user_logs", + method="POST", + json={"page": 1, "page_size": 10}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetUserLogs().post() + + assert response.status_code == 200 + assert response.json["success"] is True + assert response.json["page"] == 1 + assert len(response.json["logs"]) == 1 + assert response.json["has_more"] is False + + def test_detects_has_more(self, app): + from application.api.user.analytics.routes import GetUserLogs + + items = [ + {"_id": ObjectId(), "action": f"q{i}", "level": "info"} + for i in range(3) + ] + mock_cursor = Mock() + mock_cursor.sort.return_value.skip.return_value.limit.return_value = items + mock_user_logs = Mock() + mock_user_logs.find.return_value = mock_cursor + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.analytics.routes.user_logs_collection", + mock_user_logs, + ), patch( + "application.api.user.analytics.routes.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/get_user_logs", + method="POST", + json={"page": 1, "page_size": 2}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetUserLogs().post() + + assert response.status_code == 200 + assert response.json["has_more"] is True + assert len(response.json["logs"]) == 2 + + def test_returns_401_unauthenticated(self, app): + from application.api.user.analytics.routes import GetUserLogs + + with app.test_request_context( + "/api/get_user_logs", + method="POST", + json={"page": 1}, + ): + from flask import request + + request.decoded_token = None + response = GetUserLogs().post() + + assert response.status_code == 401 diff --git a/tests/api/user/test_conversations.py b/tests/api/user/test_conversations.py new file mode 100644 index 00000000..6fc67fa8 --- /dev/null +++ b/tests/api/user/test_conversations.py @@ -0,0 +1,360 @@ +from unittest.mock import Mock, patch + +import pytest +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +@pytest.mark.unit +class TestDeleteConversation: + + def test_deletes_conversation(self, app): + from application.api.user.conversations.routes import DeleteConversation + + conv_id = ObjectId() + mock_collection = Mock() + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_collection, + ): + with app.test_request_context(f"/api/delete_conversation?id={conv_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteConversation().post() + + assert response.status_code == 200 + assert response.json["success"] is True + mock_collection.delete_one.assert_called_once_with( + {"_id": conv_id, "user": "user1"} + ) + + def test_returns_401_unauthenticated(self, app): + from application.api.user.conversations.routes import DeleteConversation + + with app.test_request_context("/api/delete_conversation?id=abc"): + from flask import request + + request.decoded_token = None + response = DeleteConversation().post() + + assert response.status_code == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.conversations.routes import DeleteConversation + + with app.test_request_context("/api/delete_conversation"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteConversation().post() + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestDeleteAllConversations: + + def test_deletes_all_for_user(self, app): + from application.api.user.conversations.routes import DeleteAllConversations + + mock_collection = Mock() + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_collection, + ): + with app.test_request_context("/api/delete_all_conversations"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteAllConversations().get() + + assert response.status_code == 200 + mock_collection.delete_many.assert_called_once_with({"user": "user1"}) + + def test_returns_401_unauthenticated(self, app): + from application.api.user.conversations.routes import DeleteAllConversations + + with app.test_request_context("/api/delete_all_conversations"): + from flask import request + + request.decoded_token = None + response = DeleteAllConversations().get() + + assert response.status_code == 401 + + +@pytest.mark.unit +class TestGetConversations: + + def test_returns_conversations(self, app): + from application.api.user.conversations.routes import GetConversations + + conv_id = ObjectId() + mock_cursor = Mock() + mock_cursor.sort.return_value.limit.return_value = [ + { + "_id": conv_id, + "name": "Test Chat", + "agent_id": "agent1", + "is_shared_usage": False, + } + ] + mock_collection = Mock() + mock_collection.find.return_value = mock_cursor + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_collection, + ): + with app.test_request_context("/api/get_conversations"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetConversations().get() + + assert response.status_code == 200 + data = response.json + assert len(data) == 1 + assert data[0]["id"] == str(conv_id) + assert data[0]["name"] == "Test Chat" + + def test_returns_401_unauthenticated(self, app): + from application.api.user.conversations.routes import GetConversations + + with app.test_request_context("/api/get_conversations"): + from flask import request + + request.decoded_token = None + response = GetConversations().get() + + assert response.status_code == 401 + + +@pytest.mark.unit +class TestGetSingleConversation: + + def test_returns_conversation(self, app): + from application.api.user.conversations.routes import GetSingleConversation + + conv_id = ObjectId() + mock_conv_collection = Mock() + mock_conv_collection.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [{"prompt": "hi", "response": "hello"}], + "agent_id": "agent1", + } + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_conv_collection, + ): + with app.test_request_context( + f"/api/get_single_conversation?id={conv_id}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSingleConversation().get() + + assert response.status_code == 200 + assert response.json["queries"] == [{"prompt": "hi", "response": "hello"}] + assert response.json["agent_id"] == "agent1" + + def test_returns_404_not_found(self, app): + from application.api.user.conversations.routes import GetSingleConversation + + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/get_single_conversation?id={ObjectId()}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSingleConversation().get() + + assert response.status_code == 404 + + def test_returns_400_missing_id(self, app): + from application.api.user.conversations.routes import GetSingleConversation + + with app.test_request_context("/api/get_single_conversation"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSingleConversation().get() + + assert response.status_code == 400 + + def test_resolves_attachments(self, app): + from application.api.user.conversations.routes import GetSingleConversation + + conv_id = ObjectId() + att_id = ObjectId() + mock_conv_collection = Mock() + mock_conv_collection.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [ + {"prompt": "hi", "response": "hello", "attachments": [str(att_id)]} + ], + } + mock_att_collection = Mock() + mock_att_collection.find_one.return_value = { + "_id": att_id, + "filename": "doc.pdf", + } + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_conv_collection, + ), patch( + "application.api.user.conversations.routes.attachments_collection", + mock_att_collection, + ): + with app.test_request_context( + f"/api/get_single_conversation?id={conv_id}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSingleConversation().get() + + assert response.status_code == 200 + attachments = response.json["queries"][0]["attachments"] + assert len(attachments) == 1 + assert attachments[0]["fileName"] == "doc.pdf" + + +@pytest.mark.unit +class TestUpdateConversationName: + + def test_updates_name(self, app): + from application.api.user.conversations.routes import UpdateConversationName + + conv_id = ObjectId() + mock_collection = Mock() + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_conversation_name", + method="POST", + json={"id": str(conv_id), "name": "New Name"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateConversationName().post() + + assert response.status_code == 200 + assert response.json["success"] is True + mock_collection.update_one.assert_called_once() + + def test_returns_400_missing_fields(self, app): + from application.api.user.conversations.routes import UpdateConversationName + + with app.test_request_context( + "/api/update_conversation_name", + method="POST", + json={"id": str(ObjectId())}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateConversationName().post() + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestSubmitFeedback: + + def test_submits_positive_feedback(self, app): + from application.api.user.conversations.routes import SubmitFeedback + + conv_id = ObjectId() + mock_collection = Mock() + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_collection, + ): + with app.test_request_context( + "/api/feedback", + method="POST", + json={ + "feedback": "LIKE", + "conversation_id": str(conv_id), + "question_index": 0, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = SubmitFeedback().post() + + assert response.status_code == 200 + assert response.json["success"] is True + call_args = mock_collection.update_one.call_args + assert "$set" in call_args[0][1] + + def test_removes_feedback_when_null(self, app): + from application.api.user.conversations.routes import SubmitFeedback + + conv_id = ObjectId() + mock_collection = Mock() + + with patch( + "application.api.user.conversations.routes.conversations_collection", + mock_collection, + ): + with app.test_request_context( + "/api/feedback", + method="POST", + json={ + "feedback": None, + "conversation_id": str(conv_id), + "question_index": 0, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = SubmitFeedback().post() + + assert response.status_code == 200 + call_args = mock_collection.update_one.call_args + assert "$unset" in call_args[0][1] + + def test_returns_400_missing_fields(self, app): + from application.api.user.conversations.routes import SubmitFeedback + + with app.test_request_context( + "/api/feedback", + method="POST", + json={"feedback": "LIKE"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = SubmitFeedback().post() + + assert response.status_code == 400 diff --git a/tests/api/user/test_folders.py b/tests/api/user/test_folders.py new file mode 100644 index 00000000..70286eb4 --- /dev/null +++ b/tests/api/user/test_folders.py @@ -0,0 +1,509 @@ +import datetime +from unittest.mock import Mock, patch + +import pytest +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +@pytest.mark.unit +class TestAgentFoldersGet: + + def test_returns_folders(self, app): + from application.api.user.agents.folders import AgentFolders + + now = datetime.datetime(2024, 6, 15, tzinfo=datetime.timezone.utc) + folder_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + { + "_id": folder_id, + "name": "My Folder", + "parent_id": None, + "created_at": now, + "updated_at": now, + } + ] + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_collection, + ): + with app.test_request_context("/api/agents/folders/", method="GET"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolders().get() + + assert response.status_code == 200 + folders = response.json["folders"] + assert len(folders) == 1 + assert folders[0]["id"] == str(folder_id) + assert folders[0]["name"] == "My Folder" + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.folders import AgentFolders + + with app.test_request_context("/api/agents/folders/", method="GET"): + from flask import request + + request.decoded_token = None + response = AgentFolders().get() + + assert response.status_code == 401 + + +@pytest.mark.unit +class TestAgentFoldersCreate: + + def test_creates_folder(self, app): + from application.api.user.agents.folders import AgentFolders + + inserted_id = ObjectId() + mock_collection = Mock() + mock_collection.insert_one.return_value = Mock(inserted_id=inserted_id) + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_collection, + ): + with app.test_request_context( + "/api/agents/folders/", + method="POST", + json={"name": "New Folder"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolders().post() + + assert response.status_code == 201 + assert response.json["id"] == str(inserted_id) + assert response.json["name"] == "New Folder" + + def test_returns_400_missing_name(self, app): + from application.api.user.agents.folders import AgentFolders + + with app.test_request_context( + "/api/agents/folders/", + method="POST", + json={}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolders().post() + + assert response.status_code == 400 + + def test_validates_parent_folder_exists(self, app): + from application.api.user.agents.folders import AgentFolders + + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_collection, + ): + with app.test_request_context( + "/api/agents/folders/", + method="POST", + json={"name": "Sub", "parent_id": str(ObjectId())}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolders().post() + + assert response.status_code == 404 + + +@pytest.mark.unit +class TestAgentFolderGet: + + def test_returns_folder_with_agents_and_subfolders(self, app): + from application.api.user.agents.folders import AgentFolder + + folder_id = ObjectId() + agent_id = ObjectId() + subfolder_id = ObjectId() + mock_folders = Mock() + mock_folders.find_one.return_value = { + "_id": folder_id, + "name": "Folder", + "parent_id": None, + } + mock_folders.find.return_value = [ + {"_id": subfolder_id, "name": "Subfolder"} + ] + mock_agents = Mock() + mock_agents.find.return_value = [ + {"_id": agent_id, "name": "Agent 1", "description": "Desc"} + ] + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_folders, + ), patch( + "application.api.user.agents.folders.agents_collection", + mock_agents, + ): + with app.test_request_context( + f"/api/agents/folders/{folder_id}", method="GET" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolder().get(str(folder_id)) + + assert response.status_code == 200 + assert response.json["name"] == "Folder" + assert len(response.json["agents"]) == 1 + assert len(response.json["subfolders"]) == 1 + + def test_returns_404_not_found(self, app): + from application.api.user.agents.folders import AgentFolder + + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/agents/folders/{ObjectId()}", method="GET" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolder().get(str(ObjectId())) + + assert response.status_code == 404 + + +@pytest.mark.unit +class TestAgentFolderUpdate: + + def test_updates_folder_name(self, app): + from application.api.user.agents.folders import AgentFolder + + folder_id = ObjectId() + mock_collection = Mock() + mock_collection.update_one.return_value = Mock(matched_count=1) + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/agents/folders/{folder_id}", + method="PUT", + json={"name": "Renamed"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolder().put(str(folder_id)) + + assert response.status_code == 200 + assert response.json["success"] is True + + def test_prevents_self_parent(self, app): + from application.api.user.agents.folders import AgentFolder + + folder_id = str(ObjectId()) + + with app.test_request_context( + f"/api/agents/folders/{folder_id}", + method="PUT", + json={"parent_id": folder_id}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolder().put(folder_id) + + assert response.status_code == 400 + assert "own parent" in response.json["message"] + + def test_returns_404_when_not_found(self, app): + from application.api.user.agents.folders import AgentFolder + + mock_collection = Mock() + mock_collection.update_one.return_value = Mock(matched_count=0) + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/agents/folders/{ObjectId()}", + method="PUT", + json={"name": "X"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolder().put(str(ObjectId())) + + assert response.status_code == 404 + + +@pytest.mark.unit +class TestAgentFolderDelete: + + def test_deletes_folder_and_unsets_references(self, app): + from application.api.user.agents.folders import AgentFolder + + folder_id = str(ObjectId()) + mock_folders = Mock() + mock_folders.delete_one.return_value = Mock(deleted_count=1) + mock_agents = Mock() + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_folders, + ), patch( + "application.api.user.agents.folders.agents_collection", + mock_agents, + ): + with app.test_request_context( + f"/api/agents/folders/{folder_id}", method="DELETE" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolder().delete(folder_id) + + assert response.status_code == 200 + mock_agents.update_many.assert_called_once() + mock_folders.update_many.assert_called_once() + mock_folders.delete_one.assert_called_once() + + def test_returns_404_not_found(self, app): + from application.api.user.agents.folders import AgentFolder + + mock_folders = Mock() + mock_folders.delete_one.return_value = Mock(deleted_count=0) + mock_agents = Mock() + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_folders, + ), patch( + "application.api.user.agents.folders.agents_collection", + mock_agents, + ): + with app.test_request_context( + f"/api/agents/folders/{ObjectId()}", method="DELETE" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentFolder().delete(str(ObjectId())) + + assert response.status_code == 404 + + +@pytest.mark.unit +class TestMoveAgentToFolder: + + def test_moves_agent_to_folder(self, app): + from application.api.user.agents.folders import MoveAgentToFolder + + agent_id = ObjectId() + folder_id = ObjectId() + mock_agents = Mock() + mock_agents.find_one.return_value = {"_id": agent_id, "user": "user1"} + mock_folders = Mock() + mock_folders.find_one.return_value = {"_id": folder_id} + + with patch( + "application.api.user.agents.folders.agents_collection", + mock_agents, + ), patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_folders, + ): + with app.test_request_context( + "/api/agents/folders/move_agent", + method="POST", + json={ + "agent_id": str(agent_id), + "folder_id": str(folder_id), + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MoveAgentToFolder().post() + + assert response.status_code == 200 + mock_agents.update_one.assert_called_once() + + def test_removes_agent_from_folder(self, app): + from application.api.user.agents.folders import MoveAgentToFolder + + agent_id = ObjectId() + mock_agents = Mock() + mock_agents.find_one.return_value = {"_id": agent_id, "user": "user1"} + + with patch( + "application.api.user.agents.folders.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/agents/folders/move_agent", + method="POST", + json={"agent_id": str(agent_id), "folder_id": None}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MoveAgentToFolder().post() + + assert response.status_code == 200 + call_args = mock_agents.update_one.call_args + assert "$unset" in call_args[0][1] + + def test_returns_404_agent_not_found(self, app): + from application.api.user.agents.folders import MoveAgentToFolder + + mock_agents = Mock() + mock_agents.find_one.return_value = None + + with patch( + "application.api.user.agents.folders.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/agents/folders/move_agent", + method="POST", + json={"agent_id": str(ObjectId())}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MoveAgentToFolder().post() + + assert response.status_code == 404 + + def test_returns_400_missing_agent_id(self, app): + from application.api.user.agents.folders import MoveAgentToFolder + + with app.test_request_context( + "/api/agents/folders/move_agent", + method="POST", + json={}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MoveAgentToFolder().post() + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestBulkMoveAgents: + + def test_bulk_moves_to_folder(self, app): + from application.api.user.agents.folders import BulkMoveAgents + + folder_id = ObjectId() + agent_ids = [str(ObjectId()), str(ObjectId())] + mock_agents = Mock() + mock_folders = Mock() + mock_folders.find_one.return_value = {"_id": folder_id} + + with patch( + "application.api.user.agents.folders.agents_collection", + mock_agents, + ), patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_folders, + ): + with app.test_request_context( + "/api/agents/folders/bulk_move", + method="POST", + json={"agent_ids": agent_ids, "folder_id": str(folder_id)}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = BulkMoveAgents().post() + + assert response.status_code == 200 + mock_agents.update_many.assert_called_once() + + def test_bulk_removes_from_folders(self, app): + from application.api.user.agents.folders import BulkMoveAgents + + agent_ids = [str(ObjectId())] + mock_agents = Mock() + + with patch( + "application.api.user.agents.folders.agents_collection", + mock_agents, + ): + with app.test_request_context( + "/api/agents/folders/bulk_move", + method="POST", + json={"agent_ids": agent_ids}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = BulkMoveAgents().post() + + assert response.status_code == 200 + call_args = mock_agents.update_many.call_args + assert "$unset" in call_args[0][1] + + def test_returns_400_missing_agent_ids(self, app): + from application.api.user.agents.folders import BulkMoveAgents + + with app.test_request_context( + "/api/agents/folders/bulk_move", + method="POST", + json={}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = BulkMoveAgents().post() + + assert response.status_code == 400 + + def test_returns_404_folder_not_found(self, app): + from application.api.user.agents.folders import BulkMoveAgents + + mock_folders = Mock() + mock_folders.find_one.return_value = None + + with patch( + "application.api.user.agents.folders.agent_folders_collection", + mock_folders, + ): + with app.test_request_context( + "/api/agents/folders/bulk_move", + method="POST", + json={ + "agent_ids": [str(ObjectId())], + "folder_id": str(ObjectId()), + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = BulkMoveAgents().post() + + assert response.status_code == 404 diff --git a/tests/api/user/test_models.py b/tests/api/user/test_models.py new file mode 100644 index 00000000..2b2d9370 --- /dev/null +++ b/tests/api/user/test_models.py @@ -0,0 +1,70 @@ +from unittest.mock import Mock, patch + +import pytest +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +@pytest.mark.unit +class TestModelsListResource: + + def test_returns_models(self, app): + from application.api.user.models.routes import ModelsListResource + + mock_model = Mock() + mock_model.to_dict.return_value = { + "id": "gpt-4", + "name": "GPT-4", + "provider": "openai", + } + + mock_registry = Mock() + mock_registry.get_enabled_models.return_value = [mock_model] + mock_registry.default_model_id = "gpt-4" + + with patch( + "application.api.user.models.routes.ModelRegistry.get_instance", + return_value=mock_registry, + ): + with app.test_request_context("/api/models"): + response = ModelsListResource().get() + + assert response.status_code == 200 + assert response.json["count"] == 1 + assert response.json["default_model_id"] == "gpt-4" + assert response.json["models"][0]["id"] == "gpt-4" + + def test_returns_empty_models(self, app): + from application.api.user.models.routes import ModelsListResource + + mock_registry = Mock() + mock_registry.get_enabled_models.return_value = [] + mock_registry.default_model_id = None + + with patch( + "application.api.user.models.routes.ModelRegistry.get_instance", + return_value=mock_registry, + ): + with app.test_request_context("/api/models"): + response = ModelsListResource().get() + + assert response.status_code == 200 + assert response.json["count"] == 0 + assert response.json["models"] == [] + + def test_returns_500_on_error(self, app): + from application.api.user.models.routes import ModelsListResource + + with patch( + "application.api.user.models.routes.ModelRegistry.get_instance", + side_effect=Exception("Registry error"), + ): + with app.test_request_context("/api/models"): + response = ModelsListResource().get() + + assert response.status_code == 500 diff --git a/tests/api/user/test_prompts.py b/tests/api/user/test_prompts.py new file mode 100644 index 00000000..3cd969aa --- /dev/null +++ b/tests/api/user/test_prompts.py @@ -0,0 +1,288 @@ +from unittest.mock import Mock, mock_open, patch + +import pytest +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +@pytest.mark.unit +class TestCreatePrompt: + + def test_creates_prompt(self, app): + from application.api.user.prompts.routes import CreatePrompt + + mock_collection = Mock() + inserted_id = ObjectId() + mock_collection.insert_one.return_value = Mock(inserted_id=inserted_id) + + with patch( + "application.api.user.prompts.routes.prompts_collection", + mock_collection, + ): + with app.test_request_context( + "/api/create_prompt", + method="POST", + json={"name": "My Prompt", "content": "You are helpful."}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreatePrompt().post() + + assert response.status_code == 200 + assert response.json["id"] == str(inserted_id) + mock_collection.insert_one.assert_called_once() + doc = mock_collection.insert_one.call_args[0][0] + assert doc["name"] == "My Prompt" + assert doc["user"] == "user1" + + def test_returns_401_unauthenticated(self, app): + from application.api.user.prompts.routes import CreatePrompt + + with app.test_request_context( + "/api/create_prompt", + method="POST", + json={"name": "P", "content": "C"}, + ): + from flask import request + + request.decoded_token = None + response = CreatePrompt().post() + + assert response.status_code == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.prompts.routes import CreatePrompt + + with app.test_request_context( + "/api/create_prompt", + method="POST", + json={"name": "P"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreatePrompt().post() + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestGetPrompts: + + def test_returns_prompts_with_defaults(self, app): + from application.api.user.prompts.routes import GetPrompts + + user_prompt_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + {"_id": user_prompt_id, "name": "Custom Prompt"} + ] + + with patch( + "application.api.user.prompts.routes.prompts_collection", + mock_collection, + ): + with app.test_request_context("/api/get_prompts"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetPrompts().get() + + assert response.status_code == 200 + data = response.json + public_names = [p["name"] for p in data if p["type"] == "public"] + assert "default" in public_names + assert "creative" in public_names + assert "strict" in public_names + private = [p for p in data if p["type"] == "private"] + assert len(private) == 1 + assert private[0]["name"] == "Custom Prompt" + + def test_returns_401_unauthenticated(self, app): + from application.api.user.prompts.routes import GetPrompts + + with app.test_request_context("/api/get_prompts"): + from flask import request + + request.decoded_token = None + response = GetPrompts().get() + + assert response.status_code == 401 + + +@pytest.mark.unit +class TestGetSinglePrompt: + + def test_returns_default_prompt(self, app): + from application.api.user.prompts.routes import GetSinglePrompt + + with patch("builtins.open", mock_open(read_data="Default prompt content")): + with app.test_request_context("/api/get_single_prompt?id=default"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSinglePrompt().get() + + assert response.status_code == 200 + assert response.json["content"] == "Default prompt content" + + def test_returns_creative_prompt(self, app): + from application.api.user.prompts.routes import GetSinglePrompt + + with patch("builtins.open", mock_open(read_data="Creative content")): + with app.test_request_context("/api/get_single_prompt?id=creative"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSinglePrompt().get() + + assert response.status_code == 200 + assert response.json["content"] == "Creative content" + + def test_returns_strict_prompt(self, app): + from application.api.user.prompts.routes import GetSinglePrompt + + with patch("builtins.open", mock_open(read_data="Strict content")): + with app.test_request_context("/api/get_single_prompt?id=strict"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSinglePrompt().get() + + assert response.status_code == 200 + assert response.json["content"] == "Strict content" + + def test_returns_custom_prompt(self, app): + from application.api.user.prompts.routes import GetSinglePrompt + + prompt_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": prompt_id, + "content": "Custom content", + } + + with patch( + "application.api.user.prompts.routes.prompts_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/get_single_prompt?id={prompt_id}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSinglePrompt().get() + + assert response.status_code == 200 + assert response.json["content"] == "Custom content" + + def test_returns_400_missing_id(self, app): + from application.api.user.prompts.routes import GetSinglePrompt + + with app.test_request_context("/api/get_single_prompt"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetSinglePrompt().get() + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestDeletePrompt: + + def test_deletes_prompt(self, app): + from application.api.user.prompts.routes import DeletePrompt + + prompt_id = ObjectId() + mock_collection = Mock() + + with patch( + "application.api.user.prompts.routes.prompts_collection", + mock_collection, + ): + with app.test_request_context( + "/api/delete_prompt", + method="POST", + json={"id": str(prompt_id)}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeletePrompt().post() + + assert response.status_code == 200 + assert response.json["success"] is True + mock_collection.delete_one.assert_called_once_with( + {"_id": prompt_id, "user": "user1"} + ) + + def test_returns_400_missing_id(self, app): + from application.api.user.prompts.routes import DeletePrompt + + with app.test_request_context( + "/api/delete_prompt", + method="POST", + json={}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeletePrompt().post() + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestUpdatePrompt: + + def test_updates_prompt(self, app): + from application.api.user.prompts.routes import UpdatePrompt + + prompt_id = ObjectId() + mock_collection = Mock() + + with patch( + "application.api.user.prompts.routes.prompts_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_prompt", + method="POST", + json={ + "id": str(prompt_id), + "name": "Updated", + "content": "New content", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdatePrompt().post() + + assert response.status_code == 200 + assert response.json["success"] is True + mock_collection.update_one.assert_called_once() + + def test_returns_400_missing_fields(self, app): + from application.api.user.prompts.routes import UpdatePrompt + + with app.test_request_context( + "/api/update_prompt", + method="POST", + json={"id": str(ObjectId()), "name": "Updated"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdatePrompt().post() + + assert response.status_code == 400 diff --git a/tests/api/user/test_sharing.py b/tests/api/user/test_sharing.py new file mode 100644 index 00000000..b719f45c --- /dev/null +++ b/tests/api/user/test_sharing.py @@ -0,0 +1,690 @@ +import uuid +from unittest.mock import Mock, patch + +import pytest +from bson import ObjectId +from bson.binary import Binary, UuidRepresentation +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +@pytest.mark.unit +class TestShareConversation: + + def test_shares_non_promptable_conversation(self, app): + from application.api.user.sharing.routes import ShareConversation + + conv_id = ObjectId() + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Test Chat", + "queries": [{"prompt": "hi"}], + } + mock_shared = Mock() + mock_shared.find_one.return_value = None + mock_shared.insert_one.return_value = Mock() + + with patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ): + with app.test_request_context( + "/api/share?isPromptable=false", + method="POST", + json={"conversation_id": str(conv_id)}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareConversation().post() + + assert response.status_code == 201 + assert response.json["success"] is True + assert "identifier" in response.json + mock_shared.insert_one.assert_called_once() + + def test_returns_existing_shared_link(self, app): + from application.api.user.sharing.routes import ShareConversation + + conv_id = ObjectId() + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Test Chat", + "queries": [{"prompt": "hi"}], + } + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": conv_id, + } + + with patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ): + with app.test_request_context( + "/api/share?isPromptable=false", + method="POST", + json={"conversation_id": str(conv_id)}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareConversation().post() + + assert response.status_code == 200 + assert response.json["identifier"] == str(test_uuid) + + def test_returns_401_unauthenticated(self, app): + from application.api.user.sharing.routes import ShareConversation + + with app.test_request_context( + "/api/share?isPromptable=false", + method="POST", + json={"conversation_id": str(ObjectId())}, + ): + from flask import request + + request.decoded_token = None + response = ShareConversation().post() + + assert response.status_code == 401 + + def test_returns_400_missing_conversation_id(self, app): + from application.api.user.sharing.routes import ShareConversation + + with app.test_request_context( + "/api/share?isPromptable=false", + method="POST", + json={}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareConversation().post() + + assert response.status_code == 400 + + def test_returns_400_missing_isPromptable(self, app): + from application.api.user.sharing.routes import ShareConversation + + with app.test_request_context( + "/api/share", + method="POST", + json={"conversation_id": str(ObjectId())}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareConversation().post() + + assert response.status_code == 400 + assert "isPromptable" in response.json["message"] + + def test_returns_404_conversation_not_found(self, app): + from application.api.user.sharing.routes import ShareConversation + + mock_conversations = Mock() + mock_conversations.find_one.return_value = None + + with patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + "/api/share?isPromptable=false", + method="POST", + json={"conversation_id": str(ObjectId())}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareConversation().post() + + assert response.status_code == 404 + + +@pytest.mark.unit +class TestGetPubliclySharedConversations: + + def test_returns_shared_conversation(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": conv_id, + "first_n_queries": 2, + "isPromptable": False, + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Shared Chat", + "queries": [ + {"prompt": "q1", "response": "a1"}, + {"prompt": "q2", "response": "a2"}, + {"prompt": "q3", "response": "a3"}, + ], + } + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 200 + assert response.json["success"] is True + assert response.json["title"] == "Shared Chat" + assert len(response.json["queries"]) == 2 + + def test_returns_404_not_found(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + mock_shared = Mock() + mock_shared.find_one.return_value = None + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 404 + + def test_returns_404_conversation_deleted(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": conv_id, + "first_n_queries": 1, + "isPromptable": False, + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = None + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 404 + + def test_includes_api_key_when_promptable(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": conv_id, + "first_n_queries": 1, + "isPromptable": True, + "api_key": "shared_api_key", + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [{"prompt": "q1", "response": "a1"}], + } + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 200 + assert response.json["api_key"] == "shared_api_key" + + def test_handles_dbref_conversation_id(self, app): + from bson.dbref import DBRef + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": DBRef("conversations", conv_id), + "first_n_queries": 1, + "isPromptable": False, + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [{"prompt": "q1", "response": "a1"}], + } + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 200 + mock_conversations.find_one.assert_called_once_with({"_id": conv_id}) + + def test_handles_dict_oid_conversation_id(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": {"$id": {"$oid": str(conv_id)}}, + "first_n_queries": 1, + "isPromptable": False, + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [{"prompt": "q1", "response": "a1"}], + } + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 200 + + def test_handles_dict_id_string_conversation_id(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": {"$id": str(conv_id)}, + "first_n_queries": 1, + "isPromptable": False, + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [{"prompt": "q1", "response": "a1"}], + } + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 200 + + def test_handles_dict_underscore_id_conversation_id(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": {"_id": str(conv_id)}, + "first_n_queries": 1, + "isPromptable": False, + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [{"prompt": "q1", "response": "a1"}], + } + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 200 + + def test_handles_string_conversation_id(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": str(conv_id), + "first_n_queries": 1, + "isPromptable": False, + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [{"prompt": "q1", "response": "a1"}], + } + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 200 + + def test_resolves_attachments_in_shared(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + conv_id = ObjectId() + att_id = ObjectId() + + mock_shared = Mock() + mock_shared.find_one.return_value = { + "uuid": binary_uuid, + "conversation_id": conv_id, + "first_n_queries": 1, + "isPromptable": False, + } + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Chat", + "queries": [ + {"prompt": "q1", "response": "a1", "attachments": [str(att_id)]} + ], + } + mock_attachments = Mock() + mock_attachments.find_one.return_value = { + "_id": att_id, + "filename": "file.pdf", + } + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ), patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.sharing.routes.attachments_collection", + mock_attachments, + ): + with app.test_request_context( + f"/api/shared_conversation/{test_uuid}" + ): + response = GetPubliclySharedConversations().get(str(test_uuid)) + + assert response.status_code == 200 + assert response.json["queries"][0]["attachments"][0]["fileName"] == "file.pdf" + + def test_handles_general_exception(self, app): + from application.api.user.sharing.routes import ( + GetPubliclySharedConversations, + ) + + mock_shared = Mock() + mock_shared.find_one.side_effect = Exception("DB error") + + with patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ): + with app.test_request_context( + f"/api/shared_conversation/{uuid.uuid4()}" + ): + response = GetPubliclySharedConversations().get(str(uuid.uuid4())) + + assert response.status_code == 400 + + +@pytest.mark.unit +class TestShareConversationPromptable: + + def test_promptable_with_existing_api_key_and_existing_share(self, app): + from application.api.user.sharing.routes import ShareConversation + + conv_id = ObjectId() + test_uuid = uuid.uuid4() + binary_uuid = Binary.from_uuid(test_uuid, UuidRepresentation.STANDARD) + + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Test Chat", + "queries": [{"prompt": "hi"}], + } + mock_agents = Mock() + mock_agents.find_one.return_value = {"key": "existing_api_uuid"} + mock_shared = Mock() + mock_shared.find_one.return_value = {"uuid": binary_uuid} + + with patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.sharing.routes.agents_collection", + mock_agents, + ), patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ): + with app.test_request_context( + "/api/share?isPromptable=true", + method="POST", + json={ + "conversation_id": str(conv_id), + "prompt_id": "default", + "chunks": "3", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareConversation().post() + + assert response.status_code == 200 + assert response.json["identifier"] == str(test_uuid) + + def test_promptable_with_existing_api_key_new_share(self, app): + from application.api.user.sharing.routes import ShareConversation + + conv_id = ObjectId() + + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Test Chat", + "queries": [{"prompt": "hi"}], + } + mock_agents = Mock() + mock_agents.find_one.return_value = {"key": "existing_api_uuid"} + mock_shared = Mock() + mock_shared.find_one.return_value = None + + with patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.sharing.routes.agents_collection", + mock_agents, + ), patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ): + with app.test_request_context( + "/api/share?isPromptable=true", + method="POST", + json={ + "conversation_id": str(conv_id), + "source": str(ObjectId()), + "retriever": "classic", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareConversation().post() + + assert response.status_code == 201 + mock_shared.insert_one.assert_called_once() + + def test_promptable_creates_new_api_key(self, app): + from application.api.user.sharing.routes import ShareConversation + + conv_id = ObjectId() + + mock_conversations = Mock() + mock_conversations.find_one.return_value = { + "_id": conv_id, + "name": "Test Chat", + "queries": [{"prompt": "hi"}], + } + mock_agents = Mock() + mock_agents.find_one.return_value = None + mock_shared = Mock() + + with patch( + "application.api.user.sharing.routes.conversations_collection", + mock_conversations, + ), patch( + "application.api.user.sharing.routes.agents_collection", + mock_agents, + ), patch( + "application.api.user.sharing.routes.shared_conversations_collections", + mock_shared, + ): + with app.test_request_context( + "/api/share?isPromptable=true", + method="POST", + json={ + "conversation_id": str(conv_id), + "source": str(ObjectId()), + "retriever": "classic", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ShareConversation().post() + + assert response.status_code == 201 + mock_agents.insert_one.assert_called_once() + mock_shared.insert_one.assert_called_once() diff --git a/tests/api/user/test_tools_mcp.py b/tests/api/user/test_tools_mcp.py new file mode 100644 index 00000000..2a8b7a61 --- /dev/null +++ b/tests/api/user/test_tools_mcp.py @@ -0,0 +1,1308 @@ +"""Unit tests for application.api.user.tools.mcp.""" + +import json +from unittest.mock import Mock, patch + +import pytest +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +# --------------------------------------------------------------------------- +# Helper: _sanitize_mcp_transport +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestSanitizeMcpTransport: + + def test_defaults_to_auto(self): + from application.api.user.tools.mcp import _sanitize_mcp_transport + + config = {} + result = _sanitize_mcp_transport(config) + assert result == "auto" + assert config["transport_type"] == "auto" + + def test_accepts_sse(self): + from application.api.user.tools.mcp import _sanitize_mcp_transport + + config = {"transport_type": "SSE"} + result = _sanitize_mcp_transport(config) + assert result == "sse" + assert config["transport_type"] == "sse" + + def test_accepts_http(self): + from application.api.user.tools.mcp import _sanitize_mcp_transport + + config = {"transport_type": "HTTP"} + result = _sanitize_mcp_transport(config) + assert result == "http" + + def test_rejects_unsupported_transport(self): + from application.api.user.tools.mcp import _sanitize_mcp_transport + + config = {"transport_type": "stdio"} + with pytest.raises(ValueError, match="Unsupported transport_type"): + _sanitize_mcp_transport(config) + + def test_strips_command_and_args(self): + from application.api.user.tools.mcp import _sanitize_mcp_transport + + config = { + "transport_type": "auto", + "command": "/usr/bin/mcp", + "args": ["--flag"], + } + _sanitize_mcp_transport(config) + assert "command" not in config + assert "args" not in config + + def test_handles_none_transport_type(self): + from application.api.user.tools.mcp import _sanitize_mcp_transport + + config = {"transport_type": None} + result = _sanitize_mcp_transport(config) + assert result == "auto" + + +# --------------------------------------------------------------------------- +# Helper: _extract_auth_credentials +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestExtractAuthCredentials: + + def test_api_key_auth(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = { + "auth_type": "api_key", + "api_key": "my-key", + "api_key_header": "X-API-Key", + } + result = _extract_auth_credentials(config) + assert result == {"api_key": "my-key", "api_key_header": "X-API-Key"} + + def test_api_key_auth_only_key(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = {"auth_type": "api_key", "api_key": "my-key"} + result = _extract_auth_credentials(config) + assert result == {"api_key": "my-key"} + assert "api_key_header" not in result + + def test_bearer_auth(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = {"auth_type": "bearer", "bearer_token": "tok123"} + result = _extract_auth_credentials(config) + assert result == {"bearer_token": "tok123"} + + def test_bearer_auth_empty_token(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = {"auth_type": "bearer"} + result = _extract_auth_credentials(config) + assert result == {} + + def test_basic_auth(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = {"auth_type": "basic", "username": "user", "password": "pass"} + result = _extract_auth_credentials(config) + assert result == {"username": "user", "password": "pass"} + + def test_basic_auth_partial(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = {"auth_type": "basic", "username": "user"} + result = _extract_auth_credentials(config) + assert result == {"username": "user"} + + def test_none_auth(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = {"auth_type": "none"} + result = _extract_auth_credentials(config) + assert result == {} + + def test_default_no_auth_type(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = {} + result = _extract_auth_credentials(config) + assert result == {} + + def test_unknown_auth_type(self): + from application.api.user.tools.mcp import _extract_auth_credentials + + config = {"auth_type": "oauth"} + result = _extract_auth_credentials(config) + assert result == {} + + +# --------------------------------------------------------------------------- +# Route: TestMCPServerConfig +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestTestMCPServerConfig: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.mcp import TestMCPServerConfig + + with app.test_request_context( + "/api/mcp_server/test", + method="POST", + json={"config": {}}, + ): + from flask import request + + request.decoded_token = None + response = TestMCPServerConfig().post() + + assert response.status_code == 401 + + def test_returns_400_missing_config(self, app): + from application.api.user.tools.mcp import TestMCPServerConfig + + with app.test_request_context( + "/api/mcp_server/test", + method="POST", + json={}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = TestMCPServerConfig().post() + + assert response.status_code == 400 + + def test_returns_400_unsupported_transport(self, app): + from application.api.user.tools.mcp import TestMCPServerConfig + + with app.test_request_context( + "/api/mcp_server/test", + method="POST", + json={"config": {"transport_type": "stdio"}}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = TestMCPServerConfig().post() + + assert response.status_code == 400 + assert "Unsupported transport_type" in response.json["error"] + + def test_successful_connection_test(self, app): + from application.api.user.tools.mcp import TestMCPServerConfig + + mock_mcp_tool = Mock() + mock_mcp_tool.test_connection.return_value = { + "success": True, + "tools_count": 3, + } + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ): + with app.test_request_context( + "/api/mcp_server/test", + method="POST", + json={ + "config": { + "server_url": "http://localhost:8080", + "transport_type": "http", + "auth_type": "none", + } + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = TestMCPServerConfig().post() + + assert response.status_code == 200 + assert response.json["success"] is True + + def test_returns_oauth_required(self, app): + from application.api.user.tools.mcp import TestMCPServerConfig + + mock_mcp_tool = Mock() + mock_mcp_tool.test_connection.return_value = { + "requires_oauth": True, + "authorization_url": "https://auth.example.com", + } + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ): + with app.test_request_context( + "/api/mcp_server/test", + method="POST", + json={ + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "none", + } + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = TestMCPServerConfig().post() + + assert response.status_code == 200 + assert response.json["requires_oauth"] is True + + def test_redacts_failure_message(self, app): + from application.api.user.tools.mcp import TestMCPServerConfig + + mock_mcp_tool = Mock() + mock_mcp_tool.test_connection.return_value = { + "success": False, + "message": "SSL certificate verify failed", + } + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ): + with app.test_request_context( + "/api/mcp_server/test", + method="POST", + json={ + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "none", + } + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = TestMCPServerConfig().post() + + assert response.status_code == 200 + assert response.json["message"] == "Connection test failed" + + def test_returns_500_on_exception(self, app): + from application.api.user.tools.mcp import TestMCPServerConfig + + with patch( + "application.api.user.tools.mcp.MCPTool", + side_effect=RuntimeError("boom"), + ): + with app.test_request_context( + "/api/mcp_server/test", + method="POST", + json={ + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "none", + } + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = TestMCPServerConfig().post() + + assert response.status_code == 500 + assert "Connection test failed" in response.json["error"] + + def test_passes_auth_credentials_to_mcp_tool(self, app): + from application.api.user.tools.mcp import TestMCPServerConfig + + mock_mcp_tool = Mock() + mock_mcp_tool.test_connection.return_value = {"success": True} + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ) as mock_cls: + with app.test_request_context( + "/api/mcp_server/test", + method="POST", + json={ + "config": { + "server_url": "http://localhost:8080", + "transport_type": "http", + "auth_type": "bearer", + "bearer_token": "tok123", + } + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = TestMCPServerConfig().post() + + assert response.status_code == 200 + call_kwargs = mock_cls.call_args + config_arg = call_kwargs[1]["config"] if "config" in call_kwargs[1] else call_kwargs[0][0] + assert config_arg["auth_credentials"]["bearer_token"] == "tok123" + + +# --------------------------------------------------------------------------- +# Route: MCPServerSave +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestMCPServerSave: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.mcp import MCPServerSave + + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={"displayName": "My MCP", "config": {}}, + ): + from flask import request + + request.decoded_token = None + response = MCPServerSave().post() + + assert response.status_code == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.tools.mcp import MCPServerSave + + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={"displayName": "My MCP"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 400 + + def test_returns_400_unsupported_transport(self, app): + from application.api.user.tools.mcp import MCPServerSave + + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "My MCP", + "config": {"transport_type": "stdio"}, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 400 + assert "Unsupported transport_type" in response.json["error"] + + def test_creates_new_mcp_server_no_auth(self, app): + from application.api.user.tools.mcp import MCPServerSave + + inserted_id = ObjectId() + mock_mcp_tool = Mock() + mock_mcp_tool.discover_tools.return_value = None + mock_mcp_tool.get_actions_metadata.return_value = [ + {"name": "tool1", "parameters": {"properties": {"q": {"type": "string"}}}}, + ] + mock_collection = Mock() + mock_collection.insert_one.return_value = Mock(inserted_id=inserted_id) + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ), patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "My MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "none", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 200 + assert response.json["success"] is True + assert response.json["id"] == str(inserted_id) + assert response.json["tools_count"] == 1 + mock_collection.insert_one.assert_called_once() + + def test_creates_with_bearer_auth(self, app): + from application.api.user.tools.mcp import MCPServerSave + + inserted_id = ObjectId() + mock_mcp_tool = Mock() + mock_mcp_tool.discover_tools.return_value = None + mock_mcp_tool.get_actions_metadata.return_value = [] + mock_collection = Mock() + mock_collection.insert_one.return_value = Mock(inserted_id=inserted_id) + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ), patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.mcp.encrypt_credentials", + return_value="enc-blob", + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "My MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "bearer", + "bearer_token": "tok123", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 200 + call_arg = mock_collection.insert_one.call_args[0][0] + assert call_arg["config"]["encrypted_credentials"] == "enc-blob" + assert "bearer_token" not in call_arg["config"] + + def test_updates_existing_mcp_server(self, app): + from application.api.user.tools.mcp import MCPServerSave + + tool_id = ObjectId() + mock_mcp_tool = Mock() + mock_mcp_tool.discover_tools.return_value = None + mock_mcp_tool.get_actions_metadata.return_value = [ + {"name": "tool1"}, + {"name": "tool2"}, + ] + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": tool_id, + "config": {}, + } + mock_collection.update_one.return_value = Mock(matched_count=1) + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ), patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "id": str(tool_id), + "displayName": "Updated MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "none", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 200 + assert response.json["id"] == str(tool_id) + assert response.json["tools_count"] == 2 + assert "updated" in response.json["message"].lower() + + def test_returns_404_update_not_found(self, app): + from application.api.user.tools.mcp import MCPServerSave + + tool_id = ObjectId() + mock_mcp_tool = Mock() + mock_mcp_tool.discover_tools.return_value = None + mock_mcp_tool.get_actions_metadata.return_value = [] + mock_collection = Mock() + mock_collection.find_one.return_value = None + mock_collection.update_one.return_value = Mock(matched_count=0) + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ), patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "id": str(tool_id), + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "none", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 404 + + def test_oauth_auth_without_task_id(self, app): + from application.api.user.tools.mcp import MCPServerSave + + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "oauth", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 400 + assert "OAuth authorization" in response.json["error"] + + def test_oauth_auth_not_completed(self, app): + from application.api.user.tools.mcp import MCPServerSave + + mock_manager = Mock() + mock_manager.get_oauth_status.return_value = {"status": "pending"} + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + return_value=Mock(), + ), patch( + "application.api.user.tools.mcp.MCPOAuthManager", + return_value=mock_manager, + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "oauth", + "oauth_task_id": "task123", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 400 + assert "OAuth failed" in response.json["error"] + + def test_oauth_auth_completed_successfully(self, app): + from application.api.user.tools.mcp import MCPServerSave + + inserted_id = ObjectId() + mock_manager = Mock() + mock_manager.get_oauth_status.return_value = { + "status": "completed", + "tools": [{"name": "tool1"}], + } + mock_collection = Mock() + mock_collection.insert_one.return_value = Mock(inserted_id=inserted_id) + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + return_value=Mock(), + ), patch( + "application.api.user.tools.mcp.MCPOAuthManager", + return_value=mock_manager, + ), patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "oauth", + "oauth_task_id": "task123", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 200 + assert response.json["success"] is True + + def test_no_credentials_for_non_none_auth_raises(self, app): + from application.api.user.tools.mcp import MCPServerSave + + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "bearer", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 500 + + def test_returns_500_on_exception(self, app): + from application.api.user.tools.mcp import MCPServerSave + + with patch( + "application.api.user.tools.mcp.MCPTool", + side_effect=RuntimeError("boom"), + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "none", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 500 + assert "Failed to save MCP server" in response.json["error"] + + def test_strips_sensitive_fields_from_storage(self, app): + from application.api.user.tools.mcp import MCPServerSave + + inserted_id = ObjectId() + mock_mcp_tool = Mock() + mock_mcp_tool.discover_tools.return_value = None + mock_mcp_tool.get_actions_metadata.return_value = [] + mock_collection = Mock() + mock_collection.insert_one.return_value = Mock(inserted_id=inserted_id) + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ), patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.mcp.encrypt_credentials", + return_value="enc", + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "api_key", + "api_key": "secret", + "api_key_header": "X-Key", + "username": "u", + "password": "p", + "bearer_token": "bt", + "redirect_uri": "http://cb", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 200 + stored_config = mock_collection.insert_one.call_args[0][0]["config"] + for field in ["api_key", "bearer_token", "username", "password", "api_key_header", "redirect_uri"]: + assert field not in stored_config + + def test_merges_existing_encrypted_credentials_on_update(self, app): + from application.api.user.tools.mcp import MCPServerSave + + tool_id = ObjectId() + mock_mcp_tool = Mock() + mock_mcp_tool.discover_tools.return_value = None + mock_mcp_tool.get_actions_metadata.return_value = [] + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": tool_id, + "config": {"encrypted_credentials": "old-enc"}, + } + mock_collection.update_one.return_value = Mock(matched_count=1) + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ), patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.mcp.decrypt_credentials", + return_value={"api_key": "old-key"}, + ), patch( + "application.api.user.tools.mcp.encrypt_credentials", + return_value="merged-enc", + ) as mock_encrypt: + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "id": str(tool_id), + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "api_key", + "api_key": "new-key", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 200 + merged_call = mock_encrypt.call_args[0][0] + assert merged_call["api_key"] == "new-key" + + def test_preserves_existing_encrypted_when_no_new_credentials(self, app): + from application.api.user.tools.mcp import MCPServerSave + + tool_id = ObjectId() + mock_mcp_tool = Mock() + mock_mcp_tool.discover_tools.return_value = None + mock_mcp_tool.get_actions_metadata.return_value = [] + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": tool_id, + "config": {"encrypted_credentials": "existing-enc"}, + } + mock_collection.update_one.return_value = Mock(matched_count=1) + + with patch( + "application.api.user.tools.mcp.MCPTool", + return_value=mock_mcp_tool, + ), patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/mcp_server/save", + method="POST", + json={ + "id": str(tool_id), + "displayName": "MCP", + "config": { + "server_url": "http://localhost:8080", + "transport_type": "auto", + "auth_type": "none", + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPServerSave().post() + + assert response.status_code == 200 + update_call = mock_collection.update_one.call_args[0][1]["$set"] + assert update_call["config"]["encrypted_credentials"] == "existing-enc" + + +# --------------------------------------------------------------------------- +# Route: MCPOAuthCallback +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestMCPOAuthCallback: + + def test_redirects_on_error_param(self, app): + from application.api.user.tools.mcp import MCPOAuthCallback + + with app.test_request_context( + "/api/mcp_server/callback?error=access_denied&code=abc&state=xyz" + ): + response = MCPOAuthCallback().get() + + assert response.status_code == 302 + assert "error" in response.headers["Location"] + assert "access_denied" in response.headers["Location"] + + def test_redirects_on_missing_code_or_state(self, app): + from application.api.user.tools.mcp import MCPOAuthCallback + + with app.test_request_context("/api/mcp_server/callback"): + response = MCPOAuthCallback().get() + + assert response.status_code == 302 + assert "error" in response.headers["Location"] + + def test_redirects_on_missing_code(self, app): + from application.api.user.tools.mcp import MCPOAuthCallback + + with app.test_request_context("/api/mcp_server/callback?state=xyz"): + response = MCPOAuthCallback().get() + + assert response.status_code == 302 + assert "error" in response.headers["Location"] + + def test_redirects_success_on_valid_callback(self, app): + from application.api.user.tools.mcp import MCPOAuthCallback + + mock_manager = Mock() + mock_manager.handle_oauth_callback.return_value = True + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + return_value=Mock(), + ), patch( + "application.api.user.tools.mcp.MCPOAuthManager", + return_value=mock_manager, + ): + with app.test_request_context( + "/api/mcp_server/callback?code=authcode&state=statetoken" + ): + response = MCPOAuthCallback().get() + + assert response.status_code == 302 + assert "success" in response.headers["Location"] + + def test_redirects_error_on_failed_callback(self, app): + from application.api.user.tools.mcp import MCPOAuthCallback + + mock_manager = Mock() + mock_manager.handle_oauth_callback.return_value = False + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + return_value=Mock(), + ), patch( + "application.api.user.tools.mcp.MCPOAuthManager", + return_value=mock_manager, + ): + with app.test_request_context( + "/api/mcp_server/callback?code=authcode&state=statetoken" + ): + response = MCPOAuthCallback().get() + + assert response.status_code == 302 + assert "error" in response.headers["Location"] + assert "failed" in response.headers["Location"].lower() + + def test_redirects_error_when_redis_unavailable(self, app): + from application.api.user.tools.mcp import MCPOAuthCallback + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + return_value=None, + ): + with app.test_request_context( + "/api/mcp_server/callback?code=authcode&state=statetoken" + ): + response = MCPOAuthCallback().get() + + assert response.status_code == 302 + assert "Redis" in response.headers["Location"] + + def test_redirects_error_on_exception(self, app): + from application.api.user.tools.mcp import MCPOAuthCallback + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + side_effect=RuntimeError("redis down"), + ): + with app.test_request_context( + "/api/mcp_server/callback?code=authcode&state=statetoken" + ): + response = MCPOAuthCallback().get() + + assert response.status_code == 302 + assert "error" in response.headers["Location"] + + +# --------------------------------------------------------------------------- +# Route: MCPOAuthStatus +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestMCPOAuthStatus: + + def test_returns_pending_when_no_status(self, app): + from application.api.user.tools.mcp import MCPOAuthStatus + + mock_redis = Mock() + mock_redis.get.return_value = None + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + return_value=mock_redis, + ): + with app.test_request_context("/api/mcp_server/oauth_status/task123"): + response = MCPOAuthStatus().get("task123") + + assert response.status_code == 200 + assert response.json["status"] == "pending" + assert response.json["task_id"] == "task123" + + def test_returns_status_with_tools(self, app): + from application.api.user.tools.mcp import MCPOAuthStatus + + status_data = { + "status": "completed", + "tools": [ + {"name": "tool1", "description": "desc1", "extra": "should_be_stripped"}, + {"name": "tool2", "description": "desc2"}, + ], + } + mock_redis = Mock() + mock_redis.get.return_value = json.dumps(status_data) + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + return_value=mock_redis, + ): + with app.test_request_context("/api/mcp_server/oauth_status/task123"): + response = MCPOAuthStatus().get("task123") + + assert response.status_code == 200 + assert response.json["status"] == "completed" + tools = response.json["tools"] + assert len(tools) == 2 + assert tools[0]["name"] == "tool1" + assert "extra" not in tools[0] + + def test_returns_status_without_tools(self, app): + from application.api.user.tools.mcp import MCPOAuthStatus + + status_data = {"status": "in_progress", "message": "Authorizing..."} + mock_redis = Mock() + mock_redis.get.return_value = json.dumps(status_data) + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + return_value=mock_redis, + ): + with app.test_request_context("/api/mcp_server/oauth_status/task123"): + response = MCPOAuthStatus().get("task123") + + assert response.status_code == 200 + assert response.json["status"] == "in_progress" + + def test_returns_500_on_exception(self, app): + from application.api.user.tools.mcp import MCPOAuthStatus + + with patch( + "application.api.user.tools.mcp.get_redis_instance", + side_effect=RuntimeError("redis down"), + ): + with app.test_request_context("/api/mcp_server/oauth_status/task123"): + response = MCPOAuthStatus().get("task123") + + assert response.status_code == 500 + assert "Failed to get OAuth status" in response.json["error"] + + +# --------------------------------------------------------------------------- +# Route: MCPAuthStatus +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestMCPAuthStatus: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = None + response = MCPAuthStatus().get() + + assert response.status_code == 401 + + def test_returns_empty_statuses_when_no_mcp_tools(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + mock_collection = Mock() + mock_collection.find.return_value = [] + + with patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPAuthStatus().get() + + assert response.status_code == 200 + assert response.json["statuses"] == {} + + def test_returns_configured_for_non_oauth_tools(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + { + "_id": tool_id, + "config": {"auth_type": "api_key"}, + } + ] + + with patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPAuthStatus().get() + + assert response.status_code == 200 + assert response.json["statuses"][str(tool_id)] == "configured" + + def test_returns_connected_for_oauth_with_tokens(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + { + "_id": tool_id, + "config": { + "auth_type": "oauth", + "server_url": "https://api.example.com/mcp", + }, + } + ] + mock_sessions = Mock() + mock_sessions.find.return_value = [ + { + "server_url": "https://api.example.com", + "tokens": {"access_token": "tok123"}, + } + ] + + with patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.mcp._connector_sessions", + mock_sessions, + ): + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPAuthStatus().get() + + assert response.status_code == 200 + assert response.json["statuses"][str(tool_id)] == "connected" + + def test_returns_needs_auth_for_oauth_without_tokens(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + { + "_id": tool_id, + "config": { + "auth_type": "oauth", + "server_url": "https://api.example.com/mcp", + }, + } + ] + mock_sessions = Mock() + mock_sessions.find.return_value = [] + + with patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.mcp._connector_sessions", + mock_sessions, + ): + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPAuthStatus().get() + + assert response.status_code == 200 + assert response.json["statuses"][str(tool_id)] == "needs_auth" + + def test_returns_needs_auth_for_oauth_without_server_url(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + { + "_id": tool_id, + "config": {"auth_type": "oauth", "server_url": ""}, + } + ] + + with patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPAuthStatus().get() + + assert response.status_code == 200 + assert response.json["statuses"][str(tool_id)] == "needs_auth" + + def test_returns_configured_for_none_auth_type(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + {"_id": tool_id, "config": {}} + ] + + with patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPAuthStatus().get() + + assert response.status_code == 200 + assert response.json["statuses"][str(tool_id)] == "configured" + + def test_returns_500_on_exception(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + mock_collection = Mock() + mock_collection.find.side_effect = RuntimeError("db fail") + + with patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ): + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPAuthStatus().get() + + assert response.status_code == 500 + assert "Failed to check auth status" in response.json["error"] + + def test_multiple_tools_mixed_auth(self, app): + from application.api.user.tools.mcp import MCPAuthStatus + + tool_id_1 = ObjectId() + tool_id_2 = ObjectId() + tool_id_3 = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + {"_id": tool_id_1, "config": {"auth_type": "api_key"}}, + { + "_id": tool_id_2, + "config": { + "auth_type": "oauth", + "server_url": "https://api.example.com/mcp", + }, + }, + { + "_id": tool_id_3, + "config": { + "auth_type": "oauth", + "server_url": "https://other.example.com/mcp", + }, + }, + ] + mock_sessions = Mock() + mock_sessions.find.return_value = [ + { + "server_url": "https://api.example.com", + "tokens": {"access_token": "tok"}, + }, + ] + + with patch( + "application.api.user.tools.mcp.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.mcp._connector_sessions", + mock_sessions, + ): + with app.test_request_context("/api/mcp_server/auth_status"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = MCPAuthStatus().get() + + statuses = response.json["statuses"] + assert statuses[str(tool_id_1)] == "configured" + assert statuses[str(tool_id_2)] == "connected" + assert statuses[str(tool_id_3)] == "needs_auth" diff --git a/tests/api/user/test_tools_routes.py b/tests/api/user/test_tools_routes.py new file mode 100644 index 00000000..38f6b5a8 --- /dev/null +++ b/tests/api/user/test_tools_routes.py @@ -0,0 +1,1948 @@ +"""Unit tests for application.api.user.tools.routes.""" + +from datetime import datetime +from unittest.mock import Mock, patch + +import pytest +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +# --------------------------------------------------------------------------- +# Helper: _encrypt_secret_fields +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestEncryptSecretFields: + + def test_encrypts_secret_keys(self): + from application.api.user.tools.routes import _encrypt_secret_fields + + config = {"api_key": "my-secret", "base_url": "https://example.com"} + config_requirements = { + "api_key": {"secret": True}, + "base_url": {"secret": False}, + } + with patch( + "application.api.user.tools.routes.encrypt_credentials", + return_value="encrypted-blob", + ): + result = _encrypt_secret_fields(config, config_requirements, "user1") + + assert "api_key" not in result + assert result["encrypted_credentials"] == "encrypted-blob" + assert result["base_url"] == "https://example.com" + + def test_returns_config_unchanged_when_no_secrets(self): + from application.api.user.tools.routes import _encrypt_secret_fields + + config = {"base_url": "https://example.com"} + config_requirements = {"base_url": {"secret": False}} + result = _encrypt_secret_fields(config, config_requirements, "user1") + assert result == config + + def test_skips_empty_secret_values(self): + from application.api.user.tools.routes import _encrypt_secret_fields + + config = {"api_key": "", "base_url": "https://example.com"} + config_requirements = {"api_key": {"secret": True}} + result = _encrypt_secret_fields(config, config_requirements, "user1") + assert result == config + + def test_skips_secret_key_not_in_config(self): + from application.api.user.tools.routes import _encrypt_secret_fields + + config = {"base_url": "https://example.com"} + config_requirements = {"api_key": {"secret": True}} + result = _encrypt_secret_fields(config, config_requirements, "user1") + assert result == config + + +# --------------------------------------------------------------------------- +# Helper: _validate_config +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestValidateConfig: + + def test_returns_empty_on_valid_config(self): + from application.api.user.tools.routes import _validate_config + + config = {"api_key": "abc123"} + config_requirements = { + "api_key": {"required": True, "label": "API Key"}, + } + errors = _validate_config(config, config_requirements) + assert errors == {} + + def test_reports_missing_required_field(self): + from application.api.user.tools.routes import _validate_config + + config = {} + config_requirements = { + "api_key": {"required": True, "label": "API Key"}, + } + errors = _validate_config(config, config_requirements) + assert "api_key" in errors + + def test_skips_required_secret_when_existing_secrets(self): + from application.api.user.tools.routes import _validate_config + + config = {} + config_requirements = { + "api_key": {"required": True, "secret": True, "label": "API Key"}, + } + errors = _validate_config(config, config_requirements, has_existing_secrets=True) + assert errors == {} + + def test_validates_number_type(self): + from application.api.user.tools.routes import _validate_config + + config = {"timeout": "abc"} + config_requirements = { + "timeout": {"type": "number", "label": "Timeout"}, + } + errors = _validate_config(config, config_requirements) + assert "timeout" in errors + + def test_validates_timeout_range_too_low(self): + from application.api.user.tools.routes import _validate_config + + config = {"timeout": "0"} + config_requirements = { + "timeout": {"type": "number", "label": "Timeout"}, + } + errors = _validate_config(config, config_requirements) + assert "timeout" in errors + assert "between 1 and 300" in errors["timeout"] + + def test_validates_timeout_range_too_high(self): + from application.api.user.tools.routes import _validate_config + + config = {"timeout": "500"} + config_requirements = { + "timeout": {"type": "number", "label": "Timeout"}, + } + errors = _validate_config(config, config_requirements) + assert "timeout" in errors + + def test_valid_timeout(self): + from application.api.user.tools.routes import _validate_config + + config = {"timeout": "60"} + config_requirements = { + "timeout": {"type": "number", "label": "Timeout"}, + } + errors = _validate_config(config, config_requirements) + assert errors == {} + + def test_validates_enum_value(self): + from application.api.user.tools.routes import _validate_config + + config = {"mode": "invalid"} + config_requirements = { + "mode": {"enum": ["fast", "slow"], "label": "Mode"}, + } + errors = _validate_config(config, config_requirements) + assert "mode" in errors + + def test_valid_enum_value(self): + from application.api.user.tools.routes import _validate_config + + config = {"mode": "fast"} + config_requirements = { + "mode": {"enum": ["fast", "slow"], "label": "Mode"}, + } + errors = _validate_config(config, config_requirements) + assert errors == {} + + def test_depends_on_skips_when_condition_not_met(self): + from application.api.user.tools.routes import _validate_config + + config = {"mode": "simple"} + config_requirements = { + "mode": {"required": True, "label": "Mode"}, + "advanced_key": { + "required": True, + "label": "Advanced Key", + "depends_on": {"mode": "advanced"}, + }, + } + errors = _validate_config(config, config_requirements) + assert errors == {} + + def test_depends_on_validates_when_condition_met(self): + from application.api.user.tools.routes import _validate_config + + config = {"mode": "advanced"} + config_requirements = { + "mode": {"required": True, "label": "Mode"}, + "advanced_key": { + "required": True, + "label": "Advanced Key", + "depends_on": {"mode": "advanced"}, + }, + } + errors = _validate_config(config, config_requirements) + assert "advanced_key" in errors + + def test_empty_string_not_treated_as_value_for_required(self): + from application.api.user.tools.routes import _validate_config + + config = {"api_key": ""} + config_requirements = { + "api_key": {"required": True, "label": "API Key"}, + } + errors = _validate_config(config, config_requirements) + assert "api_key" in errors + + def test_uses_key_name_when_no_label(self): + from application.api.user.tools.routes import _validate_config + + config = {} + config_requirements = { + "api_key": {"required": True}, + } + errors = _validate_config(config, config_requirements) + assert "api_key" in errors + assert "api_key is required" in errors["api_key"] + + +# --------------------------------------------------------------------------- +# Helper: _merge_secrets_on_update +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestMergeSecretsOnUpdate: + + def test_no_secret_keys_returns_new_config(self): + from application.api.user.tools.routes import _merge_secrets_on_update + + new_config = {"base_url": "https://new.example.com"} + existing_config = {"base_url": "https://old.example.com"} + config_requirements = {"base_url": {"secret": False}} + + result = _merge_secrets_on_update( + new_config, existing_config, config_requirements, "user1" + ) + assert result == new_config + + def test_merges_existing_encrypted_with_new_secret(self): + from application.api.user.tools.routes import _merge_secrets_on_update + + new_config = {"api_key": "new-key", "base_url": "https://example.com"} + existing_config = { + "base_url": "https://old.com", + "encrypted_credentials": "old-blob", + } + config_requirements = { + "api_key": {"secret": True}, + "base_url": {"secret": False}, + } + with patch( + "application.api.user.tools.routes.decrypt_credentials", + return_value={"api_key": "old-key"}, + ), patch( + "application.api.user.tools.routes.encrypt_credentials", + return_value="new-blob", + ) as mock_encrypt: + result = _merge_secrets_on_update( + new_config, existing_config, config_requirements, "user1" + ) + + assert result["encrypted_credentials"] == "new-blob" + assert "api_key" not in result + assert result["base_url"] == "https://example.com" + encrypted_call = mock_encrypt.call_args[0][0] + assert encrypted_call["api_key"] == "new-key" + + def test_keeps_existing_secret_when_not_in_new_config(self): + from application.api.user.tools.routes import _merge_secrets_on_update + + new_config = {"base_url": "https://example.com"} + existing_config = { + "base_url": "https://old.com", + "encrypted_credentials": "old-blob", + } + config_requirements = { + "api_key": {"secret": True}, + "base_url": {"secret": False}, + } + with patch( + "application.api.user.tools.routes.decrypt_credentials", + return_value={"api_key": "old-key"}, + ), patch( + "application.api.user.tools.routes.encrypt_credentials", + return_value="new-blob", + ) as mock_encrypt: + _merge_secrets_on_update( + new_config, existing_config, config_requirements, "user1" + ) + + encrypted_call = mock_encrypt.call_args[0][0] + assert encrypted_call["api_key"] == "old-key" + + def test_removes_encrypted_credentials_when_no_secrets(self): + from application.api.user.tools.routes import _merge_secrets_on_update + + new_config = {"base_url": "https://example.com"} + existing_config = {"base_url": "https://old.com"} + config_requirements = { + "api_key": {"secret": True}, + "base_url": {"secret": False}, + } + with patch( + "application.api.user.tools.routes.decrypt_credentials", + return_value={}, + ): + result = _merge_secrets_on_update( + new_config, existing_config, config_requirements, "user1" + ) + + assert "encrypted_credentials" not in result + + def test_strips_has_encrypted_credentials_flag(self): + from application.api.user.tools.routes import _merge_secrets_on_update + + new_config = {"api_key": "k", "has_encrypted_credentials": True} + existing_config = {"encrypted_credentials": "blob"} + config_requirements = {"api_key": {"secret": True}} + + with patch( + "application.api.user.tools.routes.decrypt_credentials", + return_value={}, + ), patch( + "application.api.user.tools.routes.encrypt_credentials", + return_value="blob2", + ): + result = _merge_secrets_on_update( + new_config, existing_config, config_requirements, "user1" + ) + + assert "has_encrypted_credentials" not in result + + +# --------------------------------------------------------------------------- +# Helper: transform_actions +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestTransformActions: + + def test_sets_active_and_param_defaults(self): + from application.api.user.tools.routes import transform_actions + + actions = [ + { + "name": "search", + "parameters": { + "properties": { + "query": {"type": "string"}, + "limit": {"type": "integer"}, + } + }, + } + ] + result = transform_actions(actions) + assert len(result) == 1 + assert result[0]["active"] is True + props = result[0]["parameters"]["properties"] + assert props["query"]["filled_by_llm"] is True + assert props["query"]["value"] == "" + assert props["limit"]["filled_by_llm"] is True + + def test_handles_action_without_parameters(self): + from application.api.user.tools.routes import transform_actions + + actions = [{"name": "ping"}] + result = transform_actions(actions) + assert result[0]["active"] is True + assert "parameters" not in result[0] + + def test_handles_empty_properties(self): + from application.api.user.tools.routes import transform_actions + + actions = [{"name": "noop", "parameters": {"properties": {}}}] + result = transform_actions(actions) + assert result[0]["active"] is True + + def test_handles_empty_list(self): + from application.api.user.tools.routes import transform_actions + + assert transform_actions([]) == [] + + +# --------------------------------------------------------------------------- +# Route: AvailableTools +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestAvailableTools: + + def test_returns_tools_metadata(self, app): + from application.api.user.tools.routes import AvailableTools + + mock_tool = Mock() + mock_tool.__doc__ = "My Tool\nA great tool description" + mock_tool.get_config_requirements.return_value = {"key": {"required": True}} + mock_tool.get_actions_metadata.return_value = [{"name": "do_thing"}] + + mock_manager = Mock() + mock_manager.tools = {"my_tool": mock_tool} + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context("/api/available_tools"): + response = AvailableTools().get() + + assert response.status_code == 200 + data = response.json + assert data["success"] is True + assert len(data["data"]) == 1 + assert data["data"][0]["name"] == "my_tool" + assert data["data"][0]["displayName"] == "My Tool" + assert data["data"][0]["description"] == "A great tool description" + + def test_returns_400_on_error(self, app): + from application.api.user.tools.routes import AvailableTools + + mock_tool = Mock() + mock_tool.__doc__ = "Bad Tool" + mock_tool.get_config_requirements.side_effect = Exception("fail") + + mock_manager = Mock() + mock_manager.tools = {"bad_tool": mock_tool} + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context("/api/available_tools"): + response = AvailableTools().get() + + assert response.status_code == 400 + + def test_single_line_docstring(self, app): + from application.api.user.tools.routes import AvailableTools + + mock_tool = Mock() + mock_tool.__doc__ = "Simple Tool" + mock_tool.get_config_requirements.return_value = {} + mock_tool.get_actions_metadata.return_value = [] + + mock_manager = Mock() + mock_manager.tools = {"simple": mock_tool} + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context("/api/available_tools"): + response = AvailableTools().get() + + assert response.status_code == 200 + assert response.json["data"][0]["displayName"] == "Simple Tool" + assert response.json["data"][0]["description"] == "" + + +# --------------------------------------------------------------------------- +# Route: GetTools +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestGetTools: + + def test_returns_user_tools(self, app): + from application.api.user.tools.routes import GetTools + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + { + "_id": tool_id, + "name": "my_tool", + "user": "user1", + "config": {"base_url": "http://example.com"}, + "configRequirements": {"base_url": {"secret": False}}, + } + ] + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", Mock() + ): + with app.test_request_context("/api/get_tools"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetTools().get() + + assert response.status_code == 200 + assert response.json["success"] is True + assert len(response.json["tools"]) == 1 + assert response.json["tools"][0]["id"] == str(tool_id) + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import GetTools + + with app.test_request_context("/api/get_tools"): + from flask import request + + request.decoded_token = None + response = GetTools().get() + + assert response.status_code == 401 + + def test_masks_encrypted_credentials(self, app): + from application.api.user.tools.routes import GetTools + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + { + "_id": tool_id, + "name": "my_tool", + "user": "user1", + "config": { + "base_url": "http://example.com", + "encrypted_credentials": "blob", + }, + "configRequirements": { + "api_key": {"secret": True}, + }, + } + ] + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", Mock() + ): + with app.test_request_context("/api/get_tools"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetTools().get() + + tool_data = response.json["tools"][0] + assert tool_data["config"].get("has_encrypted_credentials") is True + assert "encrypted_credentials" not in tool_data["config"] + + def test_loads_config_requirements_from_tool_manager(self, app): + from application.api.user.tools.routes import GetTools + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find.return_value = [ + { + "_id": tool_id, + "name": "my_tool", + "user": "user1", + "config": {"base_url": "http://example.com"}, + "configRequirements": {}, + } + ] + + mock_tool_instance = Mock() + mock_tool_instance.get_config_requirements.return_value = { + "base_url": {"secret": False} + } + mock_manager = Mock() + mock_manager.tools = {"my_tool": mock_tool_instance} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context("/api/get_tools"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetTools().get() + + tool_data = response.json["tools"][0] + assert "base_url" in tool_data["configRequirements"] + + def test_returns_400_on_error(self, app): + from application.api.user.tools.routes import GetTools + + mock_collection = Mock() + mock_collection.find.side_effect = Exception("db error") + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context("/api/get_tools"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetTools().get() + + assert response.status_code == 400 + + +# --------------------------------------------------------------------------- +# Route: CreateTool +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestCreateTool: + + def _make_tool_instance(self): + tool_instance = Mock() + tool_instance.get_actions_metadata.return_value = [ + { + "name": "search", + "parameters": { + "properties": {"q": {"type": "string"}} + }, + } + ] + tool_instance.get_config_requirements.return_value = { + "api_key": {"required": True, "secret": True, "label": "API Key"}, + } + return tool_instance + + def test_creates_tool_successfully(self, app): + from application.api.user.tools.routes import CreateTool + + tool_instance = self._make_tool_instance() + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + mock_collection = Mock() + inserted_id = ObjectId() + mock_collection.insert_one.return_value = Mock(inserted_id=inserted_id) + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ), patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.encrypt_credentials", + return_value="blob", + ): + with app.test_request_context( + "/api/create_tool", + method="POST", + json={ + "name": "my_tool", + "displayName": "My Tool", + "description": "Desc", + "config": {"api_key": "secret123"}, + "status": True, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateTool().post() + + assert response.status_code == 200 + assert response.json["id"] == str(inserted_id) + mock_collection.insert_one.assert_called_once() + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import CreateTool + + with app.test_request_context( + "/api/create_tool", method="POST", json={} + ): + from flask import request + + request.decoded_token = None + response = CreateTool().post() + + assert response.status_code == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.tools.routes import CreateTool + + with app.test_request_context( + "/api/create_tool", + method="POST", + json={"name": "my_tool"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateTool().post() + + assert response.status_code == 400 + + def test_returns_404_tool_not_found(self, app): + from application.api.user.tools.routes import CreateTool + + mock_manager = Mock() + mock_manager.tools = {} + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/create_tool", + method="POST", + json={ + "name": "nonexistent", + "displayName": "X", + "description": "D", + "config": {}, + "status": True, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateTool().post() + + assert response.status_code == 404 + + def test_returns_400_on_validation_error(self, app): + from application.api.user.tools.routes import CreateTool + + tool_instance = self._make_tool_instance() + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/create_tool", + method="POST", + json={ + "name": "my_tool", + "displayName": "My Tool", + "description": "Desc", + "config": {}, + "status": True, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateTool().post() + + assert response.status_code == 400 + assert response.json["message"] == "Validation failed" + + def test_returns_400_on_actions_error(self, app): + from application.api.user.tools.routes import CreateTool + + tool_instance = Mock() + tool_instance.get_actions_metadata.side_effect = Exception("boom") + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/create_tool", + method="POST", + json={ + "name": "my_tool", + "displayName": "My Tool", + "description": "Desc", + "config": {}, + "status": True, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateTool().post() + + assert response.status_code == 400 + + def test_returns_400_on_insert_error(self, app): + from application.api.user.tools.routes import CreateTool + + tool_instance = Mock() + tool_instance.get_actions_metadata.return_value = [] + tool_instance.get_config_requirements.return_value = {} + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + mock_collection = Mock() + mock_collection.insert_one.side_effect = Exception("db fail") + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ), patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/create_tool", + method="POST", + json={ + "name": "my_tool", + "displayName": "My Tool", + "description": "Desc", + "config": {}, + "status": True, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateTool().post() + + assert response.status_code == 400 + + def test_includes_custom_name(self, app): + from application.api.user.tools.routes import CreateTool + + tool_instance = Mock() + tool_instance.get_actions_metadata.return_value = [] + tool_instance.get_config_requirements.return_value = {} + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + mock_collection = Mock() + inserted_id = ObjectId() + mock_collection.insert_one.return_value = Mock(inserted_id=inserted_id) + + with patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ), patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/create_tool", + method="POST", + json={ + "name": "my_tool", + "displayName": "My Tool", + "description": "Desc", + "config": {}, + "status": True, + "customName": "Custom", + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = CreateTool().post() + + assert response.status_code == 200 + call_arg = mock_collection.insert_one.call_args[0][0] + assert call_arg["customName"] == "Custom" + + +# --------------------------------------------------------------------------- +# Route: UpdateTool +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestUpdateTool: + + def test_updates_tool_successfully(self, app): + from application.api.user.tools.routes import UpdateTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": tool_id, + "name": "my_tool", + "config": {"base_url": "http://old.com"}, + } + + tool_instance = Mock() + tool_instance.get_config_requirements.return_value = {} + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/update_tool", + method="POST", + json={ + "id": str(tool_id), + "displayName": "Updated Name", + "config": {"base_url": "http://new.com"}, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateTool().post() + + assert response.status_code == 200 + assert response.json["success"] is True + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import UpdateTool + + with app.test_request_context( + "/api/update_tool", method="POST", json={"id": "abc"} + ): + from flask import request + + request.decoded_token = None + response = UpdateTool().post() + + assert response.status_code == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.tools.routes import UpdateTool + + with app.test_request_context( + "/api/update_tool", method="POST", json={} + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateTool().post() + + assert response.status_code == 400 + + def test_returns_404_tool_not_found(self, app): + from application.api.user.tools.routes import UpdateTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = None + + mock_manager = Mock() + mock_manager.tools = {} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/update_tool", + method="POST", + json={ + "id": str(tool_id), + "config": {"base_url": "http://new.com"}, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateTool().post() + + assert response.status_code == 404 + + def test_returns_400_on_invalid_function_name(self, app): + from application.api.user.tools.routes import UpdateTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_manager = Mock() + mock_manager.tools = {} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/update_tool", + method="POST", + json={ + "id": str(tool_id), + "config": { + "actions": {"invalid name!": {}} + }, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateTool().post() + + assert response.status_code == 400 + assert "Invalid function name" in response.json["message"] + + def test_returns_400_on_validation_error(self, app): + from application.api.user.tools.routes import UpdateTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": tool_id, + "name": "my_tool", + "config": {}, + } + + tool_instance = Mock() + tool_instance.get_config_requirements.return_value = { + "api_key": {"required": True, "label": "API Key"}, + } + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/update_tool", + method="POST", + json={ + "id": str(tool_id), + "config": {}, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateTool().post() + + assert response.status_code == 400 + assert response.json["message"] == "Validation failed" + + def test_updates_multiple_fields(self, app): + from application.api.user.tools.routes import UpdateTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_manager = Mock() + mock_manager.tools = {} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/update_tool", + method="POST", + json={ + "id": str(tool_id), + "name": "new_name", + "displayName": "New Display", + "customName": "Custom", + "description": "New desc", + "actions": [{"name": "a1"}], + "status": False, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateTool().post() + + assert response.status_code == 200 + call_args = mock_collection.update_one.call_args[0][1]["$set"] + assert call_args["name"] == "new_name" + assert call_args["displayName"] == "New Display" + assert call_args["customName"] == "Custom" + assert call_args["description"] == "New desc" + assert call_args["status"] is False + + def test_returns_400_on_exception(self, app): + from application.api.user.tools.routes import UpdateTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.side_effect = Exception("db error") + mock_manager = Mock() + mock_manager.tools = {} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/update_tool", + method="POST", + json={ + "id": str(tool_id), + "config": {"base_url": "http://new.com"}, + }, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateTool().post() + + assert response.status_code == 400 + + +# --------------------------------------------------------------------------- +# Route: UpdateToolConfig +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestUpdateToolConfig: + + def test_updates_config_successfully(self, app): + from application.api.user.tools.routes import UpdateToolConfig + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": tool_id, + "name": "my_tool", + "config": {"base_url": "http://old.com"}, + } + + tool_instance = Mock() + tool_instance.get_config_requirements.return_value = {} + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/update_tool_config", + method="POST", + json={"id": str(tool_id), "config": {"base_url": "http://new.com"}}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolConfig().post() + + assert response.status_code == 200 + assert response.json["success"] is True + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import UpdateToolConfig + + with app.test_request_context( + "/api/update_tool_config", + method="POST", + json={"id": "x", "config": {}}, + ): + from flask import request + + request.decoded_token = None + response = UpdateToolConfig().post() + + assert response.status_code == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.tools.routes import UpdateToolConfig + + with app.test_request_context( + "/api/update_tool_config", + method="POST", + json={"id": "x"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolConfig().post() + + assert response.status_code == 400 + + def test_returns_404_tool_not_found(self, app): + from application.api.user.tools.routes import UpdateToolConfig + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_tool_config", + method="POST", + json={"id": str(tool_id), "config": {"base_url": "x"}}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolConfig().post() + + assert response.status_code == 404 + + def test_returns_400_on_validation_error(self, app): + from application.api.user.tools.routes import UpdateToolConfig + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": tool_id, + "name": "my_tool", + "config": {}, + } + + tool_instance = Mock() + tool_instance.get_config_requirements.return_value = { + "api_key": {"required": True, "label": "API Key"}, + } + mock_manager = Mock() + mock_manager.tools = {"my_tool": tool_instance} + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ), patch( + "application.api.user.tools.routes.tool_manager", mock_manager + ): + with app.test_request_context( + "/api/update_tool_config", + method="POST", + json={"id": str(tool_id), "config": {}}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolConfig().post() + + assert response.status_code == 400 + assert response.json["message"] == "Validation failed" + + def test_returns_400_on_exception(self, app): + from application.api.user.tools.routes import UpdateToolConfig + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.side_effect = Exception("db error") + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_tool_config", + method="POST", + json={"id": str(tool_id), "config": {"a": "b"}}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolConfig().post() + + assert response.status_code == 400 + + +# --------------------------------------------------------------------------- +# Route: UpdateToolActions +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestUpdateToolActions: + + def test_updates_actions_successfully(self, app): + from application.api.user.tools.routes import UpdateToolActions + + tool_id = ObjectId() + mock_collection = Mock() + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_tool_actions", + method="POST", + json={"id": str(tool_id), "actions": [{"name": "a1"}]}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolActions().post() + + assert response.status_code == 200 + assert response.json["success"] is True + mock_collection.update_one.assert_called_once() + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import UpdateToolActions + + with app.test_request_context( + "/api/update_tool_actions", + method="POST", + json={"id": "x", "actions": []}, + ): + from flask import request + + request.decoded_token = None + response = UpdateToolActions().post() + + assert response.status_code == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.tools.routes import UpdateToolActions + + with app.test_request_context( + "/api/update_tool_actions", + method="POST", + json={"id": "x"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolActions().post() + + assert response.status_code == 400 + + def test_returns_400_on_exception(self, app): + from application.api.user.tools.routes import UpdateToolActions + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.update_one.side_effect = Exception("db error") + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_tool_actions", + method="POST", + json={"id": str(tool_id), "actions": [{"name": "a1"}]}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolActions().post() + + assert response.status_code == 400 + + +# --------------------------------------------------------------------------- +# Route: UpdateToolStatus +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestUpdateToolStatus: + + def test_updates_status_successfully(self, app): + from application.api.user.tools.routes import UpdateToolStatus + + tool_id = ObjectId() + mock_collection = Mock() + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_tool_status", + method="POST", + json={"id": str(tool_id), "status": False}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolStatus().post() + + assert response.status_code == 200 + assert response.json["success"] is True + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import UpdateToolStatus + + with app.test_request_context( + "/api/update_tool_status", + method="POST", + json={"id": "x", "status": True}, + ): + from flask import request + + request.decoded_token = None + response = UpdateToolStatus().post() + + assert response.status_code == 401 + + def test_returns_400_missing_fields(self, app): + from application.api.user.tools.routes import UpdateToolStatus + + with app.test_request_context( + "/api/update_tool_status", + method="POST", + json={"id": "x"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolStatus().post() + + assert response.status_code == 400 + + def test_returns_400_on_exception(self, app): + from application.api.user.tools.routes import UpdateToolStatus + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.update_one.side_effect = Exception("db error") + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/update_tool_status", + method="POST", + json={"id": str(tool_id), "status": True}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = UpdateToolStatus().post() + + assert response.status_code == 400 + + +# --------------------------------------------------------------------------- +# Route: DeleteTool +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestDeleteTool: + + def test_deletes_tool_successfully(self, app): + from application.api.user.tools.routes import DeleteTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.delete_one.return_value = Mock(deleted_count=1) + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/delete_tool", + method="POST", + json={"id": str(tool_id)}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteTool().post() + + assert response.status_code == 200 + assert response.json["success"] is True + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import DeleteTool + + with app.test_request_context( + "/api/delete_tool", method="POST", json={"id": "x"} + ): + from flask import request + + request.decoded_token = None + response = DeleteTool().post() + + assert response.status_code == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.tools.routes import DeleteTool + + with app.test_request_context( + "/api/delete_tool", method="POST", json={} + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteTool().post() + + assert response.status_code == 400 + + def test_returns_404_not_found(self, app): + from application.api.user.tools.routes import DeleteTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.delete_one.return_value = Mock(deleted_count=0) + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/delete_tool", + method="POST", + json={"id": str(tool_id)}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteTool().post() + + assert response.status_code == 404 + + def test_returns_400_on_exception(self, app): + from application.api.user.tools.routes import DeleteTool + + tool_id = ObjectId() + mock_collection = Mock() + mock_collection.delete_one.side_effect = Exception("db error") + + with patch( + "application.api.user.tools.routes.user_tools_collection", + mock_collection, + ): + with app.test_request_context( + "/api/delete_tool", + method="POST", + json={"id": str(tool_id)}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = DeleteTool().post() + + assert response.status_code == 400 + + +# --------------------------------------------------------------------------- +# Route: ParseSpec +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestParseSpec: + + def test_parses_json_spec_successfully(self, app): + from application.api.user.tools.routes import ParseSpec + + metadata = {"title": "Pet API"} + actions = [{"name": "listPets"}] + + with patch( + "application.api.user.tools.routes.parse_spec", + return_value=(metadata, actions), + ): + with app.test_request_context( + "/api/parse_spec", + method="POST", + json={"spec_content": "openapi: 3.0.0"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 200 + assert response.json["success"] is True + assert response.json["metadata"] == metadata + assert response.json["actions"] == actions + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import ParseSpec + + with app.test_request_context( + "/api/parse_spec", + method="POST", + json={"spec_content": "openapi: 3.0.0"}, + ): + from flask import request + + request.decoded_token = None + response = ParseSpec().post() + + assert response.status_code == 401 + + def test_returns_400_empty_spec(self, app): + from application.api.user.tools.routes import ParseSpec + + with app.test_request_context( + "/api/parse_spec", + method="POST", + json={"spec_content": ""}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 400 + assert "Empty spec content" in response.json["message"] + + def test_returns_400_whitespace_only_spec(self, app): + from application.api.user.tools.routes import ParseSpec + + with app.test_request_context( + "/api/parse_spec", + method="POST", + json={"spec_content": " "}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 400 + + def test_returns_400_no_spec_provided(self, app): + from application.api.user.tools.routes import ParseSpec + + with app.test_request_context( + "/api/parse_spec", + method="POST", + content_type="text/plain", + data="hello", + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 400 + assert "No spec provided" in response.json["message"] + + def test_parses_file_upload(self, app): + from application.api.user.tools.routes import ParseSpec + from io import BytesIO + + metadata = {"title": "API"} + actions = [{"name": "a1"}] + + with patch( + "application.api.user.tools.routes.parse_spec", + return_value=(metadata, actions), + ): + with app.test_request_context( + "/api/parse_spec", + method="POST", + content_type="multipart/form-data", + data={"file": (BytesIO(b"openapi: 3.0.0"), "spec.yaml")}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 200 + assert response.json["success"] is True + + def test_returns_400_file_no_filename(self, app): + from application.api.user.tools.routes import ParseSpec + from io import BytesIO + + with app.test_request_context( + "/api/parse_spec", + method="POST", + content_type="multipart/form-data", + data={"file": (BytesIO(b"content"), "")}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 400 + assert "No file selected" in response.json["message"] + + def test_returns_400_on_value_error(self, app): + from application.api.user.tools.routes import ParseSpec + + with patch( + "application.api.user.tools.routes.parse_spec", + side_effect=ValueError("bad spec"), + ): + with app.test_request_context( + "/api/parse_spec", + method="POST", + json={"spec_content": "bad spec content"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 400 + assert "Invalid specification format" in response.json["error"] + + def test_returns_500_on_generic_error(self, app): + from application.api.user.tools.routes import ParseSpec + + with patch( + "application.api.user.tools.routes.parse_spec", + side_effect=RuntimeError("unexpected"), + ): + with app.test_request_context( + "/api/parse_spec", + method="POST", + json={"spec_content": "some spec"}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 500 + assert "Failed to parse specification" in response.json["error"] + + def test_returns_400_invalid_file_encoding(self, app): + from application.api.user.tools.routes import ParseSpec + from io import BytesIO + + bad_bytes = b"\x80\x81\x82\x83" + + with app.test_request_context( + "/api/parse_spec", + method="POST", + content_type="multipart/form-data", + data={"file": (BytesIO(bad_bytes), "spec.bin")}, + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = ParseSpec().post() + + assert response.status_code == 400 + assert "Invalid file encoding" in response.json["message"] + + +# --------------------------------------------------------------------------- +# Route: GetArtifact +# --------------------------------------------------------------------------- +@pytest.mark.unit +class TestGetArtifact: + + def test_returns_401_unauthenticated(self, app): + from application.api.user.tools.routes import GetArtifact + + with app.test_request_context("/api/artifact/abc"): + from flask import request + + request.decoded_token = None + response = GetArtifact().get("abc") + + assert response.status_code == 401 + + def test_returns_400_invalid_artifact_id(self, app): + from application.api.user.tools.routes import GetArtifact + + with app.test_request_context("/api/artifact/not-valid-oid"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetArtifact().get("not-valid-oid") + + assert response.status_code == 400 + assert "Invalid artifact ID" in response.json["message"] + + def test_returns_note_artifact(self, app): + from application.api.user.tools.routes import GetArtifact + + artifact_id = ObjectId() + updated_at = datetime(2025, 1, 15, 10, 30) + mock_notes = Mock() + mock_notes.find_one.return_value = { + "_id": artifact_id, + "user_id": "user1", + "note": "Line1\nLine2\nLine3", + "updated_at": updated_at, + } + mock_todos = Mock() + mock_todos.find_one.return_value = None + + mock_db = {"notes": mock_notes, "todos": mock_todos} + + with patch( + "application.core.mongo_db.MongoDB.get_client", + return_value={"test_db": mock_db}, + ), patch( + "application.core.settings.settings" + ) as mock_settings: + mock_settings.MONGO_DB_NAME = "test_db" + + with app.test_request_context(f"/api/artifact/{artifact_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetArtifact().get(str(artifact_id)) + + assert response.status_code == 200 + artifact = response.json["artifact"] + assert artifact["artifact_type"] == "note" + assert artifact["data"]["content"] == "Line1\nLine2\nLine3" + assert artifact["data"]["line_count"] == 3 + assert artifact["data"]["updated_at"] == updated_at.isoformat() + + def test_returns_note_with_no_updated_at(self, app): + from application.api.user.tools.routes import GetArtifact + + artifact_id = ObjectId() + mock_notes = Mock() + mock_notes.find_one.return_value = { + "_id": artifact_id, + "user_id": "user1", + "note": "Content", + } + + mock_db = {"notes": mock_notes, "todos": Mock()} + + with patch( + "application.core.mongo_db.MongoDB.get_client", + return_value={"test_db": mock_db}, + ), patch( + "application.core.settings.settings" + ) as mock_settings: + mock_settings.MONGO_DB_NAME = "test_db" + + with app.test_request_context(f"/api/artifact/{artifact_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetArtifact().get(str(artifact_id)) + + assert response.status_code == 200 + assert response.json["artifact"]["data"]["updated_at"] is None + + def test_returns_empty_note_line_count_zero(self, app): + from application.api.user.tools.routes import GetArtifact + + artifact_id = ObjectId() + mock_notes = Mock() + mock_notes.find_one.return_value = { + "_id": artifact_id, + "user_id": "user1", + "note": "", + } + mock_db = {"notes": mock_notes, "todos": Mock()} + + with patch( + "application.core.mongo_db.MongoDB.get_client", + return_value={"test_db": mock_db}, + ), patch( + "application.core.settings.settings" + ) as mock_settings: + mock_settings.MONGO_DB_NAME = "test_db" + + with app.test_request_context(f"/api/artifact/{artifact_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetArtifact().get(str(artifact_id)) + + assert response.status_code == 200 + assert response.json["artifact"]["data"]["line_count"] == 0 + + def test_returns_todo_artifact(self, app): + from application.api.user.tools.routes import GetArtifact + + artifact_id = ObjectId() + created_at = datetime(2025, 1, 15, 10, 0) + updated_at = datetime(2025, 1, 15, 12, 0) + mock_notes = Mock() + mock_notes.find_one.return_value = None + mock_todos = Mock() + mock_todos.find_one.return_value = { + "_id": artifact_id, + "user_id": "user1", + "tool_id": "tool123", + } + mock_todos.find.return_value = [ + { + "todo_id": "t1", + "title": "Task 1", + "status": "open", + "created_at": created_at, + "updated_at": updated_at, + }, + { + "todo_id": "t2", + "title": "Task 2", + "status": "completed", + "created_at": created_at, + "updated_at": updated_at, + }, + ] + + mock_db = {"notes": mock_notes, "todos": mock_todos} + + with patch( + "application.core.mongo_db.MongoDB.get_client", + return_value={"test_db": mock_db}, + ), patch( + "application.core.settings.settings" + ) as mock_settings: + mock_settings.MONGO_DB_NAME = "test_db" + + with app.test_request_context(f"/api/artifact/{artifact_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetArtifact().get(str(artifact_id)) + + assert response.status_code == 200 + artifact = response.json["artifact"] + assert artifact["artifact_type"] == "todo_list" + assert artifact["data"]["total_count"] == 2 + assert artifact["data"]["open_count"] == 1 + assert artifact["data"]["completed_count"] == 1 + + def test_returns_todo_with_no_dates(self, app): + from application.api.user.tools.routes import GetArtifact + + artifact_id = ObjectId() + mock_notes = Mock() + mock_notes.find_one.return_value = None + mock_todos = Mock() + mock_todos.find_one.return_value = { + "_id": artifact_id, + "user_id": "user1", + "tool_id": "tool123", + } + mock_todos.find.return_value = [ + {"todo_id": "t1", "title": "Task 1"}, + ] + + mock_db = {"notes": mock_notes, "todos": mock_todos} + + with patch( + "application.core.mongo_db.MongoDB.get_client", + return_value={"test_db": mock_db}, + ), patch( + "application.core.settings.settings" + ) as mock_settings: + mock_settings.MONGO_DB_NAME = "test_db" + + with app.test_request_context(f"/api/artifact/{artifact_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetArtifact().get(str(artifact_id)) + + assert response.status_code == 200 + item = response.json["artifact"]["data"]["items"][0] + assert item["created_at"] is None + assert item["updated_at"] is None + + def test_returns_404_not_found(self, app): + from application.api.user.tools.routes import GetArtifact + + artifact_id = ObjectId() + mock_notes = Mock() + mock_notes.find_one.return_value = None + mock_todos = Mock() + mock_todos.find_one.return_value = None + + mock_db = {"notes": mock_notes, "todos": mock_todos} + + with patch( + "application.core.mongo_db.MongoDB.get_client", + return_value={"test_db": mock_db}, + ), patch( + "application.core.settings.settings" + ) as mock_settings: + mock_settings.MONGO_DB_NAME = "test_db" + + with app.test_request_context(f"/api/artifact/{artifact_id}"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = GetArtifact().get(str(artifact_id)) + + assert response.status_code == 404 + assert "Artifact not found" in response.json["message"] diff --git a/tests/api/user/test_utils.py b/tests/api/user/test_utils.py new file mode 100644 index 00000000..c27f2fd3 --- /dev/null +++ b/tests/api/user/test_utils.py @@ -0,0 +1,411 @@ +from unittest.mock import Mock + +import pytest +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +@pytest.mark.unit +class TestGetUserId: + + def test_returns_user_id_from_decoded_token(self, app): + from application.api.user.utils import get_user_id + + with app.test_request_context(): + from flask import request + + request.decoded_token = {"sub": "user_123"} + assert get_user_id() == "user_123" + + def test_returns_none_when_no_decoded_token(self, app): + from application.api.user.utils import get_user_id + + with app.test_request_context(): + assert get_user_id() is None + + def test_returns_none_when_decoded_token_has_no_sub(self, app): + from application.api.user.utils import get_user_id + + with app.test_request_context(): + from flask import request + + request.decoded_token = {} + assert get_user_id() is None + + +@pytest.mark.unit +class TestRequireAuth: + + def test_allows_authenticated_request(self, app): + from application.api.user.utils import require_auth + + @require_auth + def protected(): + return "ok" + + with app.test_request_context(): + from flask import request + + request.decoded_token = {"sub": "user_123"} + assert protected() == "ok" + + def test_returns_401_when_unauthenticated(self, app): + from application.api.user.utils import require_auth + + @require_auth + def protected(): + return "ok" + + with app.test_request_context(): + result = protected() + assert result.status_code == 401 + + +@pytest.mark.unit +class TestSuccessResponse: + + def test_default_success_response(self, app): + from application.api.user.utils import success_response + + with app.app_context(): + resp = success_response() + assert resp.status_code == 200 + assert resp.json["success"] is True + + def test_success_response_with_data(self, app): + from application.api.user.utils import success_response + + with app.app_context(): + resp = success_response({"items": [1, 2], "total": 2}) + assert resp.status_code == 200 + assert resp.json["success"] is True + assert resp.json["items"] == [1, 2] + assert resp.json["total"] == 2 + + def test_success_response_custom_status(self, app): + from application.api.user.utils import success_response + + with app.app_context(): + resp = success_response({"id": "new"}, 201) + assert resp.status_code == 201 + + +@pytest.mark.unit +class TestErrorResponse: + + def test_default_error_response(self, app): + from application.api.user.utils import error_response + + with app.app_context(): + resp = error_response("Something went wrong") + assert resp.status_code == 400 + assert resp.json["success"] is False + assert resp.json["message"] == "Something went wrong" + + def test_error_response_custom_status(self, app): + from application.api.user.utils import error_response + + with app.app_context(): + resp = error_response("Not found", 404) + assert resp.status_code == 404 + + def test_error_response_extra_kwargs(self, app): + from application.api.user.utils import error_response + + with app.app_context(): + resp = error_response("Bad", 400, errors=["field1", "field2"]) + assert resp.json["errors"] == ["field1", "field2"] + + +@pytest.mark.unit +class TestValidateObjectId: + + def test_valid_object_id(self, app): + from application.api.user.utils import validate_object_id + + with app.app_context(): + oid = ObjectId() + result, error = validate_object_id(str(oid)) + assert result == oid + assert error is None + + def test_invalid_object_id(self, app): + from application.api.user.utils import validate_object_id + + with app.app_context(): + result, error = validate_object_id("not-a-valid-id") + assert result is None + assert error.status_code == 400 + assert "Invalid" in error.json["message"] + + def test_custom_resource_name(self, app): + from application.api.user.utils import validate_object_id + + with app.app_context(): + _, error = validate_object_id("bad", "Workflow") + assert "Workflow" in error.json["message"] + + +@pytest.mark.unit +class TestValidatePagination: + + def test_default_pagination(self, app): + from application.api.user.utils import validate_pagination + + with app.test_request_context("/?limit=10&skip=0"): + limit, skip, error = validate_pagination() + assert limit == 10 + assert skip == 0 + assert error is None + + def test_uses_defaults_when_no_params(self, app): + from application.api.user.utils import validate_pagination + + with app.test_request_context("/"): + limit, skip, error = validate_pagination() + assert limit == 20 + assert skip == 0 + assert error is None + + def test_enforces_max_limit(self, app): + from application.api.user.utils import validate_pagination + + with app.test_request_context("/?limit=500"): + limit, _, _ = validate_pagination(max_limit=100) + assert limit == 100 + + def test_invalid_limit(self, app): + from application.api.user.utils import validate_pagination + + with app.test_request_context("/?limit=-1"): + _, _, error = validate_pagination() + assert error is not None + assert error.status_code == 400 + + def test_invalid_skip(self, app): + from application.api.user.utils import validate_pagination + + with app.test_request_context("/?skip=-1"): + _, _, error = validate_pagination() + assert error is not None + + def test_non_numeric_values(self, app): + from application.api.user.utils import validate_pagination + + with app.test_request_context("/?limit=abc"): + _, _, error = validate_pagination() + assert error is not None + + +@pytest.mark.unit +class TestCheckResourceOwnership: + + def test_returns_resource_when_owned(self, app): + from application.api.user.utils import check_resource_ownership + + with app.app_context(): + collection = Mock() + oid = ObjectId() + doc = {"_id": oid, "user": "user1", "name": "test"} + collection.find_one.return_value = doc + + resource, error = check_resource_ownership(collection, oid, "user1") + assert resource == doc + assert error is None + + def test_returns_404_when_not_found(self, app): + from application.api.user.utils import check_resource_ownership + + with app.app_context(): + collection = Mock() + collection.find_one.return_value = None + + resource, error = check_resource_ownership( + collection, ObjectId(), "user1", "Workflow" + ) + assert resource is None + assert error.status_code == 404 + assert "Workflow" in error.json["message"] + + +@pytest.mark.unit +class TestSerializeObjectId: + + def test_converts_id_to_string(self): + from application.api.user.utils import serialize_object_id + + oid = ObjectId() + obj = {"_id": oid, "name": "test"} + result = serialize_object_id(obj) + assert result["id"] == str(oid) + assert "_id" not in result + + def test_custom_field_names(self): + from application.api.user.utils import serialize_object_id + + oid = ObjectId() + obj = {"custom_id": oid} + result = serialize_object_id(obj, id_field="custom_id", new_field="uid") + assert result["uid"] == str(oid) + assert "custom_id" not in result + + def test_no_id_field_present(self): + from application.api.user.utils import serialize_object_id + + obj = {"name": "test"} + result = serialize_object_id(obj) + assert "id" not in result + + +@pytest.mark.unit +class TestSerializeList: + + def test_applies_serializer_to_all_items(self): + from application.api.user.utils import serialize_list + + items = [{"_id": ObjectId()}, {"_id": ObjectId()}] + + def serializer(item): + return {"id": str(item["_id"])} + + result = serialize_list(items, serializer) + assert len(result) == 2 + assert all("id" in r for r in result) + + def test_empty_list(self): + from application.api.user.utils import serialize_list + + assert serialize_list([], lambda x: x) == [] + + +@pytest.mark.unit +class TestRequireFields: + + def test_allows_valid_request(self, app): + from application.api.user.utils import require_fields + + @require_fields(["name", "email"]) + def handler(): + return "ok" + + with app.test_request_context( + "/", method="POST", json={"name": "Alice", "email": "a@b.com"} + ): + assert handler() == "ok" + + def test_rejects_missing_fields(self, app): + from application.api.user.utils import require_fields + + @require_fields(["name", "email"]) + def handler(): + return "ok" + + with app.test_request_context("/", method="POST", json={"name": "Alice"}): + result = handler() + assert result.status_code == 400 + assert "email" in result.json["message"] + + def test_rejects_empty_body(self, app): + from application.api.user.utils import require_fields + + @require_fields(["name"]) + def handler(): + return "ok" + + with app.test_request_context( + "/", method="POST", json={} + ): + result = handler() + assert result.status_code == 400 + + +@pytest.mark.unit +class TestSafeDbOperation: + + def test_returns_result_on_success(self, app): + from application.api.user.utils import safe_db_operation + + with app.app_context(): + result, error = safe_db_operation(lambda: {"inserted": True}) + assert result == {"inserted": True} + assert error is None + + def test_returns_error_on_exception(self, app): + from application.api.user.utils import safe_db_operation + + with app.app_context(): + result, error = safe_db_operation( + lambda: (_ for _ in ()).throw(RuntimeError("db error")), + "Operation failed", + ) + assert result is None + assert error.status_code == 400 + assert error.json["message"] == "Operation failed" + + def test_hides_exception_details(self, app): + from application.api.user.utils import safe_db_operation + + with app.app_context(): + _, error = safe_db_operation( + lambda: (_ for _ in ()).throw(RuntimeError("secret credentials")), + "Failed", + ) + assert "credentials" not in error.json["message"] + + +@pytest.mark.unit +class TestValidateEnum: + + def test_valid_value(self, app): + from application.api.user.utils import validate_enum + + with app.app_context(): + assert validate_enum("draft", ["draft", "published"], "status") is None + + def test_invalid_value(self, app): + from application.api.user.utils import validate_enum + + with app.app_context(): + error = validate_enum("unknown", ["draft", "published"], "status") + assert error.status_code == 400 + assert "status" in error.json["message"] + + +@pytest.mark.unit +class TestExtractSortParams: + + def test_defaults(self, app): + from application.api.user.utils import extract_sort_params + + with app.test_request_context("/"): + field, order = extract_sort_params() + assert field == "created_at" + assert order == -1 + + def test_custom_params(self, app): + from application.api.user.utils import extract_sort_params + + with app.test_request_context("/?sort=name&order=asc"): + field, order = extract_sort_params() + assert field == "name" + assert order == 1 + + def test_enforces_allowed_fields(self, app): + from application.api.user.utils import extract_sort_params + + with app.test_request_context("/?sort=forbidden_field"): + field, _ = extract_sort_params(allowed_fields=["name", "date"]) + assert field == "created_at" + + def test_desc_order(self, app): + from application.api.user.utils import extract_sort_params + + with app.test_request_context("/?order=desc"): + _, order = extract_sort_params() + assert order == -1 diff --git a/tests/api/user/test_webhooks.py b/tests/api/user/test_webhooks.py new file mode 100644 index 00000000..304a675b --- /dev/null +++ b/tests/api/user/test_webhooks.py @@ -0,0 +1,225 @@ +from unittest.mock import Mock, patch + +import pytest +from bson import ObjectId +from flask import Flask + + +@pytest.fixture +def app(): + app = Flask(__name__) + return app + + +@pytest.mark.unit +class TestAgentWebhook: + + def test_returns_existing_webhook_url(self, app): + from application.api.user.agents.webhooks import AgentWebhook + + agent_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": agent_id, + "user": "user1", + "incoming_webhook_token": "existing_token", + } + + with patch( + "application.api.user.agents.webhooks.agents_collection", + mock_collection, + ), patch( + "application.api.user.agents.webhooks.settings", + Mock(API_URL="https://api.example.com"), + ): + with app.test_request_context( + f"/api/agent_webhook?id={agent_id}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentWebhook().get() + + assert response.status_code == 200 + assert response.json["success"] is True + assert "existing_token" in response.json["webhook_url"] + mock_collection.update_one.assert_not_called() + + def test_generates_new_webhook_token(self, app): + from application.api.user.agents.webhooks import AgentWebhook + + agent_id = ObjectId() + mock_collection = Mock() + mock_collection.find_one.return_value = { + "_id": agent_id, + "user": "user1", + "incoming_webhook_token": None, + } + + with patch( + "application.api.user.agents.webhooks.agents_collection", + mock_collection, + ), patch( + "application.api.user.agents.webhooks.settings", + Mock(API_URL="https://api.example.com"), + ), patch( + "application.api.user.agents.webhooks.secrets.token_urlsafe", + return_value="new_generated_token", + ): + with app.test_request_context( + f"/api/agent_webhook?id={agent_id}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentWebhook().get() + + assert response.status_code == 200 + assert "new_generated_token" in response.json["webhook_url"] + mock_collection.update_one.assert_called_once() + + def test_returns_401_unauthenticated(self, app): + from application.api.user.agents.webhooks import AgentWebhook + + with app.test_request_context(f"/api/agent_webhook?id={ObjectId()}"): + from flask import request + + request.decoded_token = None + response = AgentWebhook().get() + + assert response.status_code == 401 + + def test_returns_400_missing_id(self, app): + from application.api.user.agents.webhooks import AgentWebhook + + with app.test_request_context("/api/agent_webhook"): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentWebhook().get() + + assert response.status_code == 400 + + def test_returns_404_agent_not_found(self, app): + from application.api.user.agents.webhooks import AgentWebhook + + mock_collection = Mock() + mock_collection.find_one.return_value = None + + with patch( + "application.api.user.agents.webhooks.agents_collection", + mock_collection, + ): + with app.test_request_context( + f"/api/agent_webhook?id={ObjectId()}" + ): + from flask import request + + request.decoded_token = {"sub": "user1"} + response = AgentWebhook().get() + + assert response.status_code == 404 + + +@pytest.mark.unit +class TestAgentWebhookListenerPost: + + def test_enqueues_task_on_valid_post(self, app): + from application.api.user.agents.webhooks import AgentWebhookListener + + mock_task = Mock() + mock_task.id = "task_abc" + + with patch( + "application.api.user.agents.webhooks.process_agent_webhook" + ) as mock_process: + mock_process.delay.return_value = mock_task + with app.test_request_context( + "/api/webhooks/agents/tok", + method="POST", + json={"event": "new_message"}, + ): + listener = AgentWebhookListener() + response = listener._enqueue_webhook_task( + "agent123", {"event": "new_message"}, "POST" + ) + + assert response.status_code == 200 + assert response.json["task_id"] == "task_abc" + mock_process.delay.assert_called_once_with( + agent_id="agent123", payload={"event": "new_message"} + ) + + def test_returns_400_on_missing_json(self, app): + from application.api.user.agents.webhooks import AgentWebhookListener + + with app.test_request_context( + "/api/webhooks/agents/tok", + method="POST", + json=None, + content_type="application/json", + data="", + ): + from flask import request as flask_request + + # Force get_json to return None (simulating empty/missing body) + with patch.object( + flask_request, "get_json", return_value=None + ): + listener = AgentWebhookListener() + response = listener.post( + webhook_token="tok", + agent={"_id": ObjectId()}, + agent_id_str="agent123", + ) + + assert response.status_code == 400 + + def test_handles_enqueue_error(self, app): + from application.api.user.agents.webhooks import AgentWebhookListener + + with patch( + "application.api.user.agents.webhooks.process_agent_webhook" + ) as mock_process: + mock_process.delay.side_effect = Exception("Queue down") + with app.test_request_context( + "/api/webhooks/agents/tok", + method="POST", + json={"event": "test"}, + ): + listener = AgentWebhookListener() + response = listener._enqueue_webhook_task( + "agent123", {"event": "test"}, "POST" + ) + + assert response.status_code == 500 + + +@pytest.mark.unit +class TestAgentWebhookListenerGet: + + def test_uses_query_params_as_payload(self, app): + from application.api.user.agents.webhooks import AgentWebhookListener + + mock_task = Mock() + mock_task.id = "task_xyz" + + with patch( + "application.api.user.agents.webhooks.process_agent_webhook" + ) as mock_process: + mock_process.delay.return_value = mock_task + with app.test_request_context( + "/api/webhooks/agents/tok?event=ping&source=test", + method="GET", + ): + listener = AgentWebhookListener() + response = listener.get( + webhook_token="tok", + agent={"_id": ObjectId()}, + agent_id_str="agent456", + ) + + assert response.status_code == 200 + call_kwargs = mock_process.delay.call_args[1] + assert call_kwargs["payload"]["event"] == "ping" + assert call_kwargs["payload"]["source"] == "test" diff --git a/tests/api/user/test_workflows.py b/tests/api/user/test_workflows.py new file mode 100644 index 00000000..e3506ed0 --- /dev/null +++ b/tests/api/user/test_workflows.py @@ -0,0 +1,406 @@ +from datetime import datetime, timezone +from unittest.mock import Mock, patch + +import pytest +from bson import ObjectId + + +@pytest.mark.unit +class TestSerializeWorkflow: + + def test_serializes_full_workflow(self): + from application.api.user.workflows.routes import serialize_workflow + + now = datetime(2024, 6, 15, 10, 30, 0, tzinfo=timezone.utc) + doc = { + "_id": ObjectId(), + "name": "My Workflow", + "description": "A test workflow", + "created_at": now, + "updated_at": now, + } + result = serialize_workflow(doc) + assert result["id"] == str(doc["_id"]) + assert result["name"] == "My Workflow" + assert result["description"] == "A test workflow" + assert result["created_at"] == now.isoformat() + + def test_handles_missing_optional_fields(self): + from application.api.user.workflows.routes import serialize_workflow + + doc = {"_id": ObjectId()} + result = serialize_workflow(doc) + assert result["name"] is None + assert result["created_at"] is None + + +@pytest.mark.unit +class TestSerializeNode: + + def test_serializes_node(self): + from application.api.user.workflows.routes import serialize_node + + node = { + "id": "node-1", + "type": "agent", + "title": "Agent Node", + "description": "Does things", + "position": {"x": 100, "y": 200}, + "config": {"model": "gpt-4"}, + } + result = serialize_node(node) + assert result["id"] == "node-1" + assert result["type"] == "agent" + assert result["title"] == "Agent Node" + assert result["data"] == {"model": "gpt-4"} + assert result["position"] == {"x": 100, "y": 200} + + def test_defaults_for_missing_fields(self): + from application.api.user.workflows.routes import serialize_node + + node = {"id": "n1", "type": "start"} + result = serialize_node(node) + assert result["data"] == {} + assert result["title"] is None + + +@pytest.mark.unit +class TestSerializeEdge: + + def test_serializes_edge(self): + from application.api.user.workflows.routes import serialize_edge + + edge = { + "id": "edge-1", + "source_id": "node-1", + "target_id": "node-2", + "source_handle": "output", + "target_handle": "input", + } + result = serialize_edge(edge) + assert result["id"] == "edge-1" + assert result["source"] == "node-1" + assert result["target"] == "node-2" + assert result["sourceHandle"] == "output" + assert result["targetHandle"] == "input" + + +@pytest.mark.unit +class TestGetWorkflowGraphVersion: + + def test_returns_version(self): + from application.api.user.workflows.routes import get_workflow_graph_version + + assert get_workflow_graph_version({"current_graph_version": 3}) == 3 + + def test_defaults_to_1(self): + from application.api.user.workflows.routes import get_workflow_graph_version + + assert get_workflow_graph_version({}) == 1 + + def test_handles_invalid_version(self): + from application.api.user.workflows.routes import get_workflow_graph_version + + assert get_workflow_graph_version({"current_graph_version": "bad"}) == 1 + + def test_handles_zero_version(self): + from application.api.user.workflows.routes import get_workflow_graph_version + + assert get_workflow_graph_version({"current_graph_version": 0}) == 1 + + def test_handles_negative_version(self): + from application.api.user.workflows.routes import get_workflow_graph_version + + assert get_workflow_graph_version({"current_graph_version": -1}) == 1 + + +@pytest.mark.unit +class TestFetchGraphDocuments: + + def test_returns_versioned_docs(self): + from application.api.user.workflows.routes import fetch_graph_documents + + collection = Mock() + docs = [{"id": "n1", "graph_version": 2}] + collection.find.return_value = docs + + result = fetch_graph_documents(collection, "wf1", 2) + assert result == docs + collection.find.assert_called_once_with( + {"workflow_id": "wf1", "graph_version": 2} + ) + + def test_falls_back_to_unversioned_for_v1(self): + from application.api.user.workflows.routes import fetch_graph_documents + + collection = Mock() + unversioned_docs = [{"id": "n1"}] + collection.find.side_effect = [[], unversioned_docs] + + result = fetch_graph_documents(collection, "wf1", 1) + assert result == unversioned_docs + assert collection.find.call_count == 2 + + def test_no_fallback_for_higher_versions(self): + from application.api.user.workflows.routes import fetch_graph_documents + + collection = Mock() + collection.find.return_value = [] + + result = fetch_graph_documents(collection, "wf1", 3) + assert result == [] + assert collection.find.call_count == 1 + + +@pytest.mark.unit +class TestValidateWorkflowStructure: + + def _make_minimal_workflow(self): + nodes = [ + {"id": "start", "type": "start"}, + {"id": "end", "type": "end"}, + ] + edges = [{"id": "e1", "source": "start", "target": "end"}] + return nodes, edges + + def test_valid_minimal_workflow(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes, edges = self._make_minimal_workflow() + errors = validate_workflow_structure(nodes, edges) + assert errors == [] + + def test_empty_nodes(self): + from application.api.user.workflows.routes import validate_workflow_structure + + errors = validate_workflow_structure([], []) + assert any("at least one node" in e for e in errors) + + def test_missing_start_node(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes = [{"id": "end", "type": "end"}] + edges = [] + errors = validate_workflow_structure(nodes, edges) + assert any("start node" in e for e in errors) + + def test_missing_end_node(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes = [{"id": "start", "type": "start"}] + edges = [{"id": "e1", "source": "start", "target": "somewhere"}] + errors = validate_workflow_structure(nodes, edges) + assert any("end node" in e for e in errors) + + def test_start_node_without_outgoing_edge(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes = [ + {"id": "start", "type": "start"}, + {"id": "end", "type": "end"}, + ] + edges = [] + errors = validate_workflow_structure(nodes, edges) + assert any("outgoing edge" in e for e in errors) + + def test_edge_references_nonexistent_node(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes = [ + {"id": "start", "type": "start"}, + {"id": "end", "type": "end"}, + ] + edges = [{"id": "e1", "source": "start", "target": "ghost"}] + errors = validate_workflow_structure(nodes, edges) + assert any("non-existent target" in e for e in errors) + + def test_node_without_id(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes = [ + {"id": "start", "type": "start"}, + {"type": "end"}, + ] + edges = [{"id": "e1", "source": "start", "target": None}] + errors = validate_workflow_structure(nodes, edges) + assert any("must have an id" in e for e in errors) + + def test_node_without_type(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes = [ + {"id": "start", "type": "start"}, + {"id": "end"}, + ] + edges = [{"id": "e1", "source": "start", "target": "end"}] + errors = validate_workflow_structure(nodes, edges) + assert any("must have a type" in e for e in errors) + + def test_condition_node_needs_two_outgoing_edges(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes = [ + {"id": "start", "type": "start"}, + { + "id": "cond", + "type": "condition", + "title": "Check", + "data": { + "cases": [ + {"expression": "x > 1", "sourceHandle": "case1"}, + ] + }, + }, + {"id": "end", "type": "end"}, + ] + edges = [ + {"id": "e1", "source": "start", "target": "cond"}, + {"id": "e2", "source": "cond", "target": "end", "sourceHandle": "else"}, + ] + errors = validate_workflow_structure(nodes, edges) + assert any("at least 2 outgoing edges" in e for e in errors) + + def test_condition_node_needs_else_branch(self): + from application.api.user.workflows.routes import validate_workflow_structure + + nodes = [ + {"id": "start", "type": "start"}, + { + "id": "cond", + "type": "condition", + "title": "Check", + "data": { + "cases": [ + {"expression": "x > 1", "sourceHandle": "case1"}, + ] + }, + }, + {"id": "end1", "type": "end"}, + {"id": "end2", "type": "end"}, + ] + edges = [ + {"id": "e1", "source": "start", "target": "cond"}, + {"id": "e2", "source": "cond", "target": "end1", "sourceHandle": "case1"}, + {"id": "e3", "source": "cond", "target": "end2", "sourceHandle": "case1"}, + ] + errors = validate_workflow_structure(nodes, edges) + assert any("else" in e for e in errors) + + +@pytest.mark.unit +class TestCanReachEnd: + + def test_direct_end_node(self): + from application.api.user.workflows.routes import _can_reach_end + + node_map = {"end": {"id": "end", "type": "end"}} + assert _can_reach_end("end", [], node_map, {"end"}) is True + + def test_reachable_through_chain(self): + from application.api.user.workflows.routes import _can_reach_end + + node_map = { + "a": {"id": "a"}, + "b": {"id": "b"}, + "end": {"id": "end", "type": "end"}, + } + edges = [ + {"source": "a", "target": "b"}, + {"source": "b", "target": "end"}, + ] + assert _can_reach_end("a", edges, node_map, {"end"}) is True + + def test_unreachable(self): + from application.api.user.workflows.routes import _can_reach_end + + node_map = { + "a": {"id": "a"}, + "b": {"id": "b"}, + } + edges = [{"source": "a", "target": "b"}] + assert _can_reach_end("a", edges, node_map, {"end"}) is False + + def test_handles_cycles(self): + from application.api.user.workflows.routes import _can_reach_end + + node_map = {"a": {"id": "a"}, "b": {"id": "b"}} + edges = [ + {"source": "a", "target": "b"}, + {"source": "b", "target": "a"}, + ] + assert _can_reach_end("a", edges, node_map, {"end"}) is False + + +@pytest.mark.unit +class TestValidateJsonSchemaPayload: + + def test_none_input(self): + from application.api.user.workflows.routes import validate_json_schema_payload + + result, error = validate_json_schema_payload(None) + assert result is None + assert error is None + + @patch("application.api.user.workflows.routes.normalize_json_schema_payload") + def test_valid_schema(self, mock_normalize): + from application.api.user.workflows.routes import validate_json_schema_payload + + mock_normalize.return_value = {"type": "object"} + result, error = validate_json_schema_payload({"type": "object"}) + assert result == {"type": "object"} + assert error is None + + @patch("application.api.user.workflows.routes.normalize_json_schema_payload") + def test_invalid_schema(self, mock_normalize): + from application.api.user.workflows.routes import validate_json_schema_payload + from application.core.json_schema_utils import JsonSchemaValidationError + + mock_normalize.side_effect = JsonSchemaValidationError("bad schema") + result, error = validate_json_schema_payload({"bad": True}) + assert result is None + assert "bad schema" in error + + +@pytest.mark.unit +class TestNormalizeAgentNodeJsonSchemas: + + def test_non_agent_nodes_pass_through(self): + from application.api.user.workflows.routes import ( + normalize_agent_node_json_schemas, + ) + + nodes = [ + {"id": "n1", "type": "start"}, + {"id": "n2", "type": "end"}, + ] + result = normalize_agent_node_json_schemas(nodes) + assert result == nodes + + @patch("application.api.user.workflows.routes.normalize_json_schema_payload") + def test_normalizes_agent_node_schema(self, mock_normalize): + from application.api.user.workflows.routes import ( + normalize_agent_node_json_schemas, + ) + + mock_normalize.return_value = {"type": "object", "properties": {}} + nodes = [ + { + "id": "a1", + "type": "agent", + "data": {"json_schema": {"type": "object"}}, + } + ] + result = normalize_agent_node_json_schemas(nodes) + assert result[0]["data"]["json_schema"] == { + "type": "object", + "properties": {}, + } + + def test_agent_node_without_schema(self): + from application.api.user.workflows.routes import ( + normalize_agent_node_json_schemas, + ) + + nodes = [{"id": "a1", "type": "agent", "data": {"model": "gpt-4"}}] + result = normalize_agent_node_json_schemas(nodes) + assert result[0]["data"] == {"model": "gpt-4"} diff --git a/tests/core/__init__.py b/tests/core/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/core/test_model_settings.py b/tests/core/test_model_settings.py index 7257df4a..cc405539 100644 --- a/tests/core/test_model_settings.py +++ b/tests/core/test_model_settings.py @@ -337,3 +337,98 @@ class TestModelRegistry: reg = ModelRegistry() # Should have at least docsgpt-local assert reg.default_model_id is not None + + @pytest.mark.unit + def test_default_model_from_provider_fallback(self): + """When LLM_NAME is not set but LLM_PROVIDER and API_KEY are, + default should be first model of that provider.""" + mock_settings = MagicMock() + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_API_KEY = "sk-test" + mock_settings.OPENAI_API_BASE = None + mock_settings.ANTHROPIC_API_KEY = None + mock_settings.GOOGLE_API_KEY = None + mock_settings.GROQ_API_KEY = None + mock_settings.OPEN_ROUTER_API_KEY = None + mock_settings.NOVITA_API_KEY = None + mock_settings.HUGGINGFACE_API_KEY = None + mock_settings.LLM_PROVIDER = "openai" + mock_settings.LLM_NAME = None + mock_settings.API_KEY = "sk-test" + + with patch("application.core.settings.settings", mock_settings): + reg = ModelRegistry() + assert reg.default_model_id is not None + + @pytest.mark.unit + def test_add_google_models_no_key_with_provider(self): + with patch.object(ModelRegistry, "_load_models"): + reg = ModelRegistry() + reg.models = {} + mock_settings = MagicMock() + mock_settings.GOOGLE_API_KEY = None + mock_settings.LLM_PROVIDER = "google" + mock_settings.LLM_NAME = "nonexistent" + reg._add_google_models(mock_settings) + assert len(reg.models) > 0 + + @pytest.mark.unit + def test_add_groq_models_no_key_with_provider(self): + with patch.object(ModelRegistry, "_load_models"): + reg = ModelRegistry() + reg.models = {} + mock_settings = MagicMock() + mock_settings.GROQ_API_KEY = None + mock_settings.LLM_PROVIDER = "groq" + mock_settings.LLM_NAME = "nonexistent" + reg._add_groq_models(mock_settings) + assert len(reg.models) > 0 + + @pytest.mark.unit + def test_add_openrouter_models_no_key_with_provider(self): + with patch.object(ModelRegistry, "_load_models"): + reg = ModelRegistry() + reg.models = {} + mock_settings = MagicMock() + mock_settings.OPEN_ROUTER_API_KEY = None + mock_settings.LLM_PROVIDER = "openrouter" + mock_settings.LLM_NAME = "nonexistent" + reg._add_openrouter_models(mock_settings) + assert len(reg.models) > 0 + + @pytest.mark.unit + def test_add_novita_models_no_key_with_provider(self): + with patch.object(ModelRegistry, "_load_models"): + reg = ModelRegistry() + reg.models = {} + mock_settings = MagicMock() + mock_settings.NOVITA_API_KEY = None + mock_settings.LLM_PROVIDER = "novita" + mock_settings.LLM_NAME = "nonexistent" + reg._add_novita_models(mock_settings) + assert len(reg.models) > 0 + + @pytest.mark.unit + def test_to_dict_disabled_model(self): + model = AvailableModel( + id="disabled", + provider=ModelProvider.OPENAI, + display_name="Disabled", + enabled=False, + ) + d = model.to_dict() + assert d["enabled"] is False + + @pytest.mark.unit + def test_to_dict_with_attachment_types(self): + caps = ModelCapabilities( + supported_attachment_types=["image/png", "application/pdf"], + ) + model = AvailableModel( + id="vision", + provider=ModelProvider.OPENAI, + display_name="Vision", + capabilities=caps, + ) + d = model.to_dict() + assert d["supported_attachment_types"] == ["image/png", "application/pdf"] diff --git a/tests/core/test_url_validation.py b/tests/core/test_url_validation.py index 924e5cde..59d040c6 100644 --- a/tests/core/test_url_validation.py +++ b/tests/core/test_url_validation.py @@ -195,3 +195,67 @@ class TestValidateUrlSafe: is_valid, url, error = validate_url_safe("http://192.168.1.1") assert is_valid is False assert "private" in error.lower() or "internal" in error.lower() + + def test_adds_scheme_when_missing(self): + with patch("application.core.url_validation.resolve_hostname") as mock_resolve: + mock_resolve.return_value = "93.184.216.34" + is_valid, url, error = validate_url_safe("example.com") + assert is_valid is True + assert url == "http://example.com" + + +class TestIsPrivateIPExtended: + """Additional edge cases for IP classification.""" + + def test_multicast_ip(self): + assert is_private_ip("224.0.0.1") is True + + def test_unspecified_ip(self): + assert is_private_ip("0.0.0.0") is True + + def test_ipv6_loopback(self): + assert is_private_ip("::1") is True + + def test_ipv6_private(self): + assert is_private_ip("fc00::1") is True + + def test_ipv6_public(self): + assert is_private_ip("2607:f8b0:4004:800::200e") is False + + def test_reserved_ip(self): + # 240.0.0.0/4 is reserved (future use), Python's ipaddress marks it as such + assert is_private_ip("240.0.0.1") is True + + +class TestValidateUrlExtended: + """Additional URL validation tests.""" + + def test_blocks_metadata_hostname(self): + with pytest.raises(SSRFError): + validate_url("http://metadata") + + def test_allows_localhost_with_flag(self): + with patch("application.core.url_validation.resolve_hostname") as mock_resolve: + mock_resolve.return_value = "192.168.1.1" + result = validate_url( + "http://internal.local", allow_localhost=True + ) + assert result == "http://internal.local" + + def test_blocks_aws_ecs_metadata_ip(self): + with pytest.raises(SSRFError, match="metadata"): + validate_url("http://169.254.170.2") + + def test_blocks_aws_ipv6_metadata(self): + with pytest.raises(SSRFError, match="metadata"): + validate_url("http://[fd00:ec2::254]") + + def test_blocks_hostname_resolving_to_loopback(self): + with patch("application.core.url_validation.resolve_hostname") as mock_resolve: + mock_resolve.return_value = "127.0.0.1" + with pytest.raises(SSRFError): + validate_url("http://sneaky.example.com") + + def test_allows_localhost_ip_with_flag(self): + result = validate_url("http://10.0.0.1", allow_localhost=True) + assert result == "http://10.0.0.1" diff --git a/tests/llm/__init__.py b/tests/llm/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/llm/test_anthropic.py b/tests/llm/test_anthropic.py new file mode 100644 index 00000000..12c7b9cd --- /dev/null +++ b/tests/llm/test_anthropic.py @@ -0,0 +1,323 @@ +"""Unit tests for application/llm/anthropic.py — AnthropicLLM. + +Extends coverage beyond test_anthropic_llm.py: + - Constructor: api_key priority, base_url support + - get_supported_attachment_types + - prepare_messages_with_attachments: various scenarios + - _get_base64_image: error paths + - _raw_gen_stream: close called on response +""" + +import sys +import types + +import pytest + + +# --------------------------------------------------------------------------- +# Fake anthropic module +# --------------------------------------------------------------------------- + + +class _FakeCompletion: + def __init__(self, text): + self.completion = text + + +class _FakeCompletions: + def __init__(self): + self.last_kwargs = None + self._stream_items = [_FakeCompletion("s1"), _FakeCompletion("s2")] + + def create(self, **kwargs): + self.last_kwargs = kwargs + if kwargs.get("stream"): + return self._stream_items + return _FakeCompletion("final") + + +class _FakeAnthropic: + def __init__(self, api_key=None, base_url=None): + self.api_key = api_key + self.base_url = base_url + self.completions = _FakeCompletions() + + +@pytest.fixture(autouse=True) +def patch_anthropic(monkeypatch): + fake = types.ModuleType("anthropic") + fake.Anthropic = _FakeAnthropic + fake.HUMAN_PROMPT = "" + fake.AI_PROMPT = "" + + modules_to_remove = [key for key in sys.modules if key.startswith("anthropic")] + for key in modules_to_remove: + sys.modules.pop(key, None) + sys.modules["anthropic"] = fake + + if "application.llm.anthropic" in sys.modules: + del sys.modules["application.llm.anthropic"] + yield + sys.modules.pop("anthropic", None) + if "application.llm.anthropic" in sys.modules: + del sys.modules["application.llm.anthropic"] + + +@pytest.fixture +def llm(): + from application.llm.anthropic import AnthropicLLM + + instance = AnthropicLLM(api_key="test-key") + instance.storage = types.SimpleNamespace( + get_file=lambda path: _ctx_manager(b"img_bytes"), + ) + return instance + + +def _ctx_manager(data): + """Create a simple context manager returning an object with .read().""" + import contextlib + + @contextlib.contextmanager + def cm(): + yield types.SimpleNamespace(read=lambda: data) + + return cm() + + +# --------------------------------------------------------------------------- +# Constructor +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestAnthropicConstructor: + + def test_api_key_set(self): + from application.llm.anthropic import AnthropicLLM + + instance = AnthropicLLM(api_key="custom-key") + assert instance.api_key == "custom-key" + + def test_base_url_passed(self): + from application.llm.anthropic import AnthropicLLM + + instance = AnthropicLLM(api_key="k", base_url="https://custom.api") + assert instance.anthropic.base_url == "https://custom.api" + + def test_no_base_url(self): + from application.llm.anthropic import AnthropicLLM + + instance = AnthropicLLM(api_key="k") + assert instance.anthropic.base_url is None + + def test_human_and_ai_prompts_set(self): + from application.llm.anthropic import AnthropicLLM + + instance = AnthropicLLM(api_key="k") + assert instance.HUMAN_PROMPT == "" + assert instance.AI_PROMPT == "" + + +# --------------------------------------------------------------------------- +# _raw_gen +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGen: + + def test_returns_completion(self, llm): + msgs = [{"content": "context"}, {"content": "question"}] + result = llm._raw_gen(llm, model="claude-2", messages=msgs) + assert result == "final" + + def test_prompt_contains_context_and_question(self, llm): + msgs = [{"content": "my context"}, {"content": "my question"}] + llm._raw_gen(llm, model="claude-2", messages=msgs) + prompt = llm.anthropic.completions.last_kwargs["prompt"] + assert "my context" in prompt + assert "my question" in prompt + + def test_max_tokens_passed(self, llm): + msgs = [{"content": "c"}, {"content": "q"}] + llm._raw_gen(llm, model="claude-2", messages=msgs, max_tokens=200) + assert llm.anthropic.completions.last_kwargs["max_tokens_to_sample"] == 200 + + +# --------------------------------------------------------------------------- +# _raw_gen_stream +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGenStream: + + def test_yields_all_completions(self, llm): + msgs = [{"content": "c"}, {"content": "q"}] + chunks = list( + llm._raw_gen_stream(llm, model="claude", messages=msgs, max_tokens=10) + ) + assert chunks == ["s1", "s2"] + + def test_calls_close_on_response(self, llm): + closed = {"called": False} + original = llm.anthropic.completions._stream_items + + class ClosableList(list): + def close(self): + closed["called"] = True + + closable = ClosableList(original) + llm.anthropic.completions._stream_items = closable + llm.anthropic.completions.create = lambda **kw: closable + + msgs = [{"content": "c"}, {"content": "q"}] + list(llm._raw_gen_stream(llm, model="claude", messages=msgs)) + assert closed["called"] + + def test_prompt_format(self, llm): + msgs = [{"content": "ctx"}, {"content": "q"}] + list(llm._raw_gen_stream(llm, model="claude", messages=msgs)) + prompt = llm.anthropic.completions.last_kwargs["prompt"] + assert prompt.startswith("") + assert prompt.endswith("") + + +# --------------------------------------------------------------------------- +# get_supported_attachment_types +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestGetSupportedAttachmentTypes: + + def test_returns_image_types(self, llm): + result = llm.get_supported_attachment_types() + assert "image/png" in result + assert "image/jpeg" in result + assert "image/webp" in result + assert "image/gif" in result + + def test_no_pdf_support(self, llm): + result = llm.get_supported_attachment_types() + assert "application/pdf" not in result + + +# --------------------------------------------------------------------------- +# prepare_messages_with_attachments +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestPrepareMessagesWithAttachments: + + def test_no_attachments_returns_same(self, llm): + msgs = [{"role": "user", "content": "hi"}] + result = llm.prepare_messages_with_attachments(msgs) + assert result == msgs + + def test_empty_attachments_returns_same(self, llm): + msgs = [{"role": "user", "content": "hi"}] + result = llm.prepare_messages_with_attachments(msgs, []) + assert result == msgs + + def test_image_with_preconverted_data(self, llm): + msgs = [{"role": "user", "content": "look"}] + attachments = [{"mime_type": "image/png", "data": "AABBCC"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + img_part = next( + p for p in user_msg["content"] if p.get("type") == "image" + ) + assert img_part["source"]["data"] == "AABBCC" + assert img_part["source"]["type"] == "base64" + assert img_part["source"]["media_type"] == "image/png" + + def test_image_from_storage(self, llm): + llm.storage = types.SimpleNamespace( + get_file=lambda p: _ctx_manager(b"raw_image_bytes"), + ) + msgs = [{"role": "user", "content": "look"}] + attachments = [{"mime_type": "image/jpeg", "path": "/tmp/img.jpg"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + img_part = next( + p for p in user_msg["content"] if p.get("type") == "image" + ) + assert img_part["source"]["media_type"] == "image/jpeg" + assert len(img_part["source"]["data"]) > 0 + + def test_no_user_message_creates_one(self, llm): + msgs = [{"role": "system", "content": "sys"}] + attachments = [{"mime_type": "image/png", "data": "AAA"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msgs = [m for m in result if m["role"] == "user"] + assert len(user_msgs) == 1 + + def test_image_error_adds_text_fallback(self, llm): + def bad_storage(path): + raise Exception("storage error") + + llm.storage = types.SimpleNamespace(get_file=bad_storage) + msgs = [{"role": "user", "content": "look"}] + attachments = [ + {"mime_type": "image/png", "path": "/bad.png", "content": "fb"}, + ] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + text_parts = [ + p for p in user_msg["content"] + if p.get("type") == "text" and "could not" in p.get("text", "").lower() + ] + assert len(text_parts) == 1 + + def test_non_image_attachment_ignored(self, llm): + msgs = [{"role": "user", "content": "look"}] + attachments = [{"mime_type": "application/pdf"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + # content becomes list with just original text + assert isinstance(user_msg["content"], list) + assert len(user_msg["content"]) == 1 + + def test_content_not_list_becomes_empty(self, llm): + msgs = [{"role": "user", "content": 999}] + attachments = [{"mime_type": "image/png", "data": "AAA"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + assert isinstance(user_msg["content"], list) + + +# --------------------------------------------------------------------------- +# _get_base64_image +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestGetBase64Image: + + def test_raises_for_no_path(self, llm): + with pytest.raises(ValueError, match="No file path"): + llm._get_base64_image({}) + + def test_raises_for_file_not_found(self, llm): + import contextlib + + @contextlib.contextmanager + def bad_file(path): + raise FileNotFoundError("not found") + + llm.storage = types.SimpleNamespace(get_file=bad_file) + with pytest.raises(FileNotFoundError): + llm._get_base64_image({"path": "/nonexistent"}) + + def test_returns_base64_encoded(self, llm): + import base64 + + llm.storage = types.SimpleNamespace( + get_file=lambda p: _ctx_manager(b"test_data"), + ) + result = llm._get_base64_image({"path": "/tmp/img.png"}) + decoded = base64.b64decode(result) + assert decoded == b"test_data" diff --git a/tests/llm/test_base.py b/tests/llm/test_base.py new file mode 100644 index 00000000..e6d07429 --- /dev/null +++ b/tests/llm/test_base.py @@ -0,0 +1,269 @@ +"""Unit tests for application/llm/base.py — BaseLLM. + +Extends coverage beyond test_base_llm.py: + - gen / gen_stream: decorator application, argument forwarding + - _execute_with_fallback: non-streaming fallback + - _stream_with_fallback: mid-stream fallback + - fallback_llm: backup model resolution, global fallback +""" + +from unittest.mock import MagicMock, Mock, patch + +import pytest + +from application.llm.base import BaseLLM + + +# --------------------------------------------------------------------------- +# Concrete stubs +# --------------------------------------------------------------------------- + + +class StubLLM(BaseLLM): + def __init__(self, raw_gen_return="gen_result", raw_gen_stream_items=None, **kwargs): + super().__init__(**kwargs) + self._raw_gen_return = raw_gen_return + self._raw_gen_stream_items = raw_gen_stream_items or ["s1", "s2"] + + def _raw_gen(self, baseself, model, messages, stream=False, tools=None, **kw): + return self._raw_gen_return + + def _raw_gen_stream(self, baseself, model, messages, stream=True, tools=None, **kw): + yield from self._raw_gen_stream_items + + +class FailingLLM(BaseLLM): + def __init__(self, **kwargs): + super().__init__(**kwargs) + + def _raw_gen(self, baseself, model, messages, stream=False, tools=None, **kw): + raise RuntimeError("primary_failed") + + def _raw_gen_stream(self, baseself, model, messages, stream=True, tools=None, **kw): + raise RuntimeError("primary_stream_failed") + + +class FallbackLLM(BaseLLM): + def __init__(self, **kwargs): + super().__init__(**kwargs) + self.gen_called = False + self.gen_stream_called = False + + def _raw_gen(self, baseself, model, messages, stream=False, tools=None, **kw): + self.gen_called = True + return "fallback_result" + + def _raw_gen_stream(self, baseself, model, messages, stream=True, tools=None, **kw): + self.gen_stream_called = True + yield "fallback_chunk" + + def gen(self, *args, **kwargs): + self.gen_called = True + return "fallback_gen_result" + + def gen_stream(self, *args, **kwargs): + self.gen_stream_called = True + yield "fallback_stream_chunk" + + +# --------------------------------------------------------------------------- +# gen / gen_stream decorator application +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestGenMethods: + + @patch("application.llm.base.gen_cache", lambda f: f) + @patch("application.llm.base.gen_token_usage", lambda f: f) + def test_gen_returns_result(self): + llm = StubLLM(raw_gen_return="hello") + result = llm.gen(model="m", messages=[{"role": "user", "content": "hi"}]) + assert result == "hello" + + @patch("application.llm.base.stream_cache", lambda f: f) + @patch("application.llm.base.stream_token_usage", lambda f: f) + def test_gen_stream_yields_results(self): + llm = StubLLM(raw_gen_stream_items=["a", "b"]) + result = list( + llm.gen_stream(model="m", messages=[{"role": "user", "content": "hi"}]) + ) + assert result == ["a", "b"] + + @patch("application.llm.base.gen_cache", lambda f: f) + @patch("application.llm.base.gen_token_usage", lambda f: f) + def test_gen_passes_tools(self): + tools = [{"type": "function", "function": {"name": "t"}}] + + class ToolCaptureLLM(BaseLLM): + def __init__(self): + super().__init__() + self.captured_tools = None + + def _raw_gen(self, baseself, model, messages, stream=False, tools=None, **kw): + self.captured_tools = tools + return "ok" + + def _raw_gen_stream(self, baseself, model, messages, stream=True, tools=None, **kw): + yield "x" + + llm = ToolCaptureLLM() + llm.gen(model="m", messages=[], tools=tools) + assert llm.captured_tools == tools + + +# --------------------------------------------------------------------------- +# _execute_with_fallback: non-streaming +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestExecuteWithFallbackNonStreaming: + + @patch("application.llm.base.gen_cache", lambda f: f) + @patch("application.llm.base.gen_token_usage", lambda f: f) + def test_no_fallback_raises(self): + llm = FailingLLM() + with pytest.raises(RuntimeError, match="primary_failed"): + llm.gen(model="m", messages=[]) + + @patch("application.llm.base.gen_cache", lambda f: f) + @patch("application.llm.base.gen_token_usage", lambda f: f) + def test_fallback_called_on_failure(self): + fallback = FallbackLLM(model_id="fallback-model") + llm = FailingLLM() + llm._fallback_llm = fallback + + result = llm.gen(model="m", messages=[]) + assert result == "fallback_gen_result" + assert fallback.gen_called + + +# --------------------------------------------------------------------------- +# _stream_with_fallback +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestStreamWithFallback: + + @patch("application.llm.base.stream_cache", lambda f: f) + @patch("application.llm.base.stream_token_usage", lambda f: f) + def test_no_fallback_raises(self): + llm = FailingLLM() + with pytest.raises(RuntimeError, match="primary_stream_failed"): + list(llm.gen_stream(model="m", messages=[])) + + @patch("application.llm.base.stream_cache", lambda f: f) + @patch("application.llm.base.stream_token_usage", lambda f: f) + def test_fallback_called_on_stream_failure(self): + fallback = FallbackLLM(model_id="fallback-model") + llm = FailingLLM() + llm._fallback_llm = fallback + + result = list(llm.gen_stream(model="m", messages=[])) + assert "fallback_stream_chunk" in result + assert fallback.gen_stream_called + + +# --------------------------------------------------------------------------- +# fallback_llm property: backup model resolution +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestFallbackLLMResolution: + + def test_returns_cached_fallback(self): + sentinel = StubLLM() + llm = StubLLM() + llm._fallback_llm = sentinel + assert llm.fallback_llm is sentinel + + def test_none_without_config(self, monkeypatch): + monkeypatch.setattr( + "application.llm.base.settings", + MagicMock(FALLBACK_LLM_PROVIDER=None), + ) + llm = StubLLM(backup_models=[]) + assert llm.fallback_llm is None + + def test_backup_model_resolved(self, monkeypatch): + mock_fallback = StubLLM() + monkeypatch.setattr( + "application.core.model_utils.get_provider_from_model_id", + lambda mid: "openai", + ) + monkeypatch.setattr( + "application.core.model_utils.get_api_key_for_provider", + lambda p: "key", + ) + monkeypatch.setattr( + "application.llm.llm_creator.LLMCreator.create_llm", + Mock(return_value=mock_fallback), + ) + + llm = StubLLM(backup_models=["backup-model-id"]) + result = llm.fallback_llm + assert result is mock_fallback + + def test_backup_model_failure_tries_next(self, monkeypatch): + call_count = {"n": 0} + + def mock_create(*a, **kw): + call_count["n"] += 1 + if call_count["n"] == 1: + raise RuntimeError("first fail") + return StubLLM() + + monkeypatch.setattr( + "application.core.model_utils.get_provider_from_model_id", + lambda mid: "openai", + ) + monkeypatch.setattr( + "application.core.model_utils.get_api_key_for_provider", + lambda p: "key", + ) + monkeypatch.setattr( + "application.llm.llm_creator.LLMCreator.create_llm", + mock_create, + ) + + llm = StubLLM(backup_models=["bad-model", "good-model"]) + result = llm.fallback_llm + assert result is not None + assert call_count["n"] == 2 + + def test_global_fallback_used_when_no_backup(self, monkeypatch): + mock_fallback = StubLLM() + monkeypatch.setattr( + "application.llm.base.settings", + MagicMock( + FALLBACK_LLM_PROVIDER="openai", + FALLBACK_LLM_NAME="gpt-4", + FALLBACK_LLM_API_KEY="key", + API_KEY="key", + ), + ) + monkeypatch.setattr( + "application.llm.llm_creator.LLMCreator.create_llm", + Mock(return_value=mock_fallback), + ) + + llm = StubLLM(backup_models=[]) + result = llm.fallback_llm + assert result is mock_fallback + + def test_backup_provider_not_found_skipped(self, monkeypatch): + monkeypatch.setattr( + "application.core.model_utils.get_provider_from_model_id", + lambda mid: None, + ) + monkeypatch.setattr( + "application.llm.base.settings", + MagicMock(FALLBACK_LLM_PROVIDER=None), + ) + + llm = StubLLM(backup_models=["unknown-model"]) + result = llm.fallback_llm + assert result is None diff --git a/tests/llm/test_google_ai.py b/tests/llm/test_google_ai.py new file mode 100644 index 00000000..01160cb7 --- /dev/null +++ b/tests/llm/test_google_ai.py @@ -0,0 +1,755 @@ +"""Unit tests for application/llm/google_ai.py — GoogleLLM. + +Extends coverage beyond test_google_llm.py: + - _clean_messages_google: system instructions, function responses, errors + - _clean_schema: field filtering, type uppercasing, required validation + - _clean_tools_format: empty properties, required fields + - _extract_preview_from_message: various message shapes + - _summarize_messages_for_log + - _get_text_value / _is_thought_part: dict vs object forms + - _raw_gen with tools and response_schema + - _raw_gen_stream: function_call parts, thought parts, error handling + - prepare_structured_output_format: comprehensive type mapping + - prepare_messages_with_attachments: error handling + - _upload_file_to_google + - get_supported_attachment_types +""" + +import types + +import pytest + +from application.llm.google_ai import GoogleLLM + + +# --------------------------------------------------------------------------- +# Fake types module for Google AI +# --------------------------------------------------------------------------- + + +class _FakePart: + def __init__(self, text=None, function_call=None, file_data=None, thought=False): + self.text = text + self.function_call = function_call + self.file_data = file_data + self.thought = thought + + @staticmethod + def from_text(text): + return _FakePart(text=text) + + @staticmethod + def from_function_call(name, args): + return _FakePart(function_call=types.SimpleNamespace(name=name, args=args)) + + @staticmethod + def from_function_response(name, response): + return _FakePart(text=str(response)) + + @staticmethod + def from_uri(file_uri, mime_type): + return _FakePart( + file_data=types.SimpleNamespace(file_uri=file_uri, mime_type=mime_type) + ) + + +class _FakeContent: + def __init__(self, role, parts): + self.role = role + self.parts = parts + + +class FakeTypesModule: + Part = _FakePart + Content = _FakeContent + + class GenerateContentConfig: + def __init__(self): + self.system_instruction = None + self.tools = None + self.thinking_config = None + self.response_schema = None + self.response_mime_type = None + + class Tool: + def __init__(self, function_declarations=None): + self.function_declarations = function_declarations or [] + + class FunctionCall: + def __init__(self, name=None, args=None): + self.name = name + self.args = args + + +class FakeModels: + def __init__(self): + self.last_kwargs = None + + class _Resp: + def __init__(self, text=None, candidates=None): + self.text = text + self.candidates = candidates or [] + + def generate_content(self, *args, **kwargs): + self.last_kwargs = kwargs + return FakeModels._Resp(text="ok") + + def generate_content_stream(self, *args, **kwargs): + self.last_kwargs = kwargs + return [] + + +class FakeClientFiles: + def upload(self, file=None): + return types.SimpleNamespace(uri="gs://fake-uri") + + +class FakeClient: + def __init__(self, *a, **kw): + self.models = FakeModels() + self.files = FakeClientFiles() + + +@pytest.fixture(autouse=True) +def patch_google(monkeypatch): + import application.llm.google_ai as gmod + + monkeypatch.setattr(gmod, "types", FakeTypesModule) + monkeypatch.setattr(gmod.genai, "Client", FakeClient) + + +@pytest.fixture +def llm(): + instance = GoogleLLM(api_key="test-key") + instance.storage = types.SimpleNamespace( + file_exists=lambda p: True, + process_file=lambda path, fn, **kw: fn(path), + ) + return instance + + +# --------------------------------------------------------------------------- +# _clean_messages_google +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestCleanMessagesGoogle: + + def test_system_message_extracted_as_instruction(self, llm): + msgs = [ + {"role": "system", "content": "You are helpful"}, + {"role": "user", "content": "hi"}, + ] + cleaned, sys_instr = llm._clean_messages_google(msgs) + assert sys_instr == "You are helpful" + assert all(c.role != "system" for c in cleaned) + + def test_multiple_system_messages_joined(self, llm): + msgs = [ + {"role": "system", "content": "Rule 1"}, + {"role": "system", "content": "Rule 2"}, + {"role": "user", "content": "hi"}, + ] + _, sys_instr = llm._clean_messages_google(msgs) + assert "Rule 1" in sys_instr + assert "Rule 2" in sys_instr + + def test_system_list_content(self, llm): + msgs = [ + {"role": "system", "content": [{"text": "A"}, {"text": "B"}]}, + {"role": "user", "content": "hi"}, + ] + _, sys_instr = llm._clean_messages_google(msgs) + assert "A" in sys_instr and "B" in sys_instr + + def test_assistant_role_becomes_model(self, llm): + msgs = [{"role": "assistant", "content": "hi"}] + cleaned, _ = llm._clean_messages_google(msgs) + assert cleaned[0].role == "model" + + def test_tool_role_becomes_model(self, llm): + msgs = [{"role": "tool", "content": "result"}] + cleaned, _ = llm._clean_messages_google(msgs) + assert cleaned[0].role == "model" + + def test_function_call_in_content_list(self, llm): + msgs = [ + { + "role": "assistant", + "content": [ + {"function_call": {"name": "fn", "args": {"x": 1}}}, + ], + } + ] + cleaned, _ = llm._clean_messages_google(msgs) + assert len(cleaned) == 1 + assert any( + hasattr(p, "function_call") and p.function_call is not None + for p in cleaned[0].parts + ) + + def test_function_response_in_content_list(self, llm): + msgs = [ + { + "role": "assistant", + "content": [ + { + "function_response": { + "name": "fn", + "response": {"result": 42}, + } + }, + ], + } + ] + cleaned, _ = llm._clean_messages_google(msgs) + assert len(cleaned) == 1 + + def test_files_in_content_list(self, llm): + msgs = [ + { + "role": "user", + "content": [ + {"files": [{"file_uri": "gs://f", "mime_type": "image/png"}]}, + ], + } + ] + cleaned, _ = llm._clean_messages_google(msgs) + assert len(cleaned) == 1 + assert any( + hasattr(p, "file_data") and p.file_data is not None + for p in cleaned[0].parts + ) + + def test_unexpected_list_item_raises(self, llm): + msgs = [{"role": "user", "content": [{"unknown_key": "val"}]}] + with pytest.raises(ValueError, match="Unexpected content dictionary"): + llm._clean_messages_google(msgs) + + def test_unexpected_content_type_raises(self, llm): + msgs = [{"role": "user", "content": 12345}] + with pytest.raises(ValueError, match="Unexpected content type"): + llm._clean_messages_google(msgs) + + def test_no_system_instruction_returns_none(self, llm): + msgs = [{"role": "user", "content": "hi"}] + _, sys_instr = llm._clean_messages_google(msgs) + assert sys_instr is None + + def test_empty_parts_skipped(self, llm): + msgs = [{"role": "user", "content": None}] + cleaned, _ = llm._clean_messages_google(msgs) + assert len(cleaned) == 0 + + +# --------------------------------------------------------------------------- +# _clean_schema +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestCleanSchema: + + def test_type_uppercased(self, llm): + result = llm._clean_schema({"type": "string"}) + assert result["type"] == "STRING" + + def test_unsupported_fields_removed(self, llm): + result = llm._clean_schema({"type": "string", "title": "Name", "$ref": "#/x"}) + assert "title" not in result + assert "$ref" not in result + assert result["type"] == "STRING" + + def test_nested_properties_cleaned(self, llm): + # _clean_schema recursively cleans the properties dict value. + # Property names that happen to match allowed_fields survive. + # This tests the recursive cleaning on schema values. + schema = { + "type": "object", + "properties": { + "type": {"type": "string"}, + }, + } + result = llm._clean_schema(schema) + # "type" is in allowed_fields, so the property survives as a key + # Its value gets uppercased since it's a type field + assert "properties" in result + assert result["properties"]["type"]["type"] == "STRING" + + def test_required_validated_against_properties(self, llm): + # Property names must be in allowed_fields to survive _clean_schema + # "type" is in allowed_fields so it survives as a property key + schema = { + "type": "object", + "properties": {"type": {"type": "string"}}, + "required": ["type", "nonexistent"], + } + result = llm._clean_schema(schema) + assert result["required"] == ["type"] + + def test_required_removed_when_no_valid_entries(self, llm): + schema = { + "type": "object", + "properties": {"type": {"type": "string"}}, + "required": ["nonexistent"], + } + result = llm._clean_schema(schema) + assert "required" not in result + + def test_required_removed_when_no_properties(self, llm): + schema = {"type": "string", "required": ["x"]} + result = llm._clean_schema(schema) + assert "required" not in result + + def test_non_dict_passthrough(self, llm): + assert llm._clean_schema("hello") == "hello" + assert llm._clean_schema(42) == 42 + + def test_list_items_cleaned(self, llm): + schema = { + "type": "array", + "items": {"type": "string", "title": "ignored"}, + } + result = llm._clean_schema(schema) + assert "title" not in result["items"] + + +# --------------------------------------------------------------------------- +# _clean_tools_format +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestCleanToolsFormat: + + def test_basic_tool_conversion(self, llm): + tools = [ + { + "type": "function", + "function": { + "name": "search", + "description": "Search the web", + "parameters": { + "type": "object", + "properties": { + "query": {"type": "string"}, + }, + "required": ["query"], + }, + }, + } + ] + result = llm._clean_tools_format(tools) + assert len(result) == 1 + assert hasattr(result[0], "function_declarations") + + def test_tool_without_properties(self, llm): + tools = [ + { + "type": "function", + "function": { + "name": "ping", + "description": "Ping server", + "parameters": {"type": "object", "properties": {}}, + }, + } + ] + result = llm._clean_tools_format(tools) + assert len(result) == 1 + + +# --------------------------------------------------------------------------- +# _extract_preview_from_message / _summarize_messages_for_log +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestMessagePreviewAndSummary: + + def test_preview_from_parts_text(self, llm): + msg = types.SimpleNamespace( + parts=[_FakePart(text="hello world")] + ) + preview = llm._extract_preview_from_message(msg) + assert preview == "hello world" + + def test_preview_from_function_call_part(self, llm): + fc = types.SimpleNamespace(name="search") + msg = types.SimpleNamespace( + parts=[_FakePart(function_call=fc)] + ) + preview = llm._extract_preview_from_message(msg) + assert "search" in preview + + def test_preview_from_dict_string_content(self, llm): + msg = {"content": "dict content"} + preview = llm._extract_preview_from_message(msg) + assert preview == "dict content" + + def test_preview_from_dict_list_content(self, llm): + msg = {"content": [{"text": "list text"}]} + preview = llm._extract_preview_from_message(msg) + assert preview == "list text" + + def test_preview_from_dict_function_call(self, llm): + msg = {"content": [{"function_call": {"name": "fn"}}]} + preview = llm._extract_preview_from_message(msg) + assert "fn" in preview + + def test_preview_from_dict_function_response(self, llm): + msg = {"content": [{"function_response": {"name": "fn_resp"}}]} + preview = llm._extract_preview_from_message(msg) + assert "fn_resp" in preview + + def test_preview_fallback_to_str(self, llm): + msg = 42 + preview = llm._extract_preview_from_message(msg) + assert preview == "42" + + def test_summarize_messages_empty(self, llm): + result = llm._summarize_messages_for_log([]) + assert "count=0" in result + + def test_summarize_messages_truncates(self, llm): + msgs = [ + types.SimpleNamespace(parts=[_FakePart(text="a" * 100)]) + ] + result = llm._summarize_messages_for_log(msgs, preview_chars=10) + assert "..." in result + + +# --------------------------------------------------------------------------- +# _get_text_value / _is_thought_part +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestStaticHelpers: + + def test_get_text_value_dict(self): + assert GoogleLLM._get_text_value({"text": "hi"}) == "hi" + + def test_get_text_value_dict_no_text(self): + assert GoogleLLM._get_text_value({"other": "x"}) == "" + + def test_get_text_value_dict_non_string(self): + assert GoogleLLM._get_text_value({"text": 42}) == "" + + def test_get_text_value_object(self): + obj = types.SimpleNamespace(text="obj_text") + assert GoogleLLM._get_text_value(obj) == "obj_text" + + def test_get_text_value_object_no_text(self): + obj = types.SimpleNamespace() + assert GoogleLLM._get_text_value(obj) == "" + + def test_is_thought_part_dict_true(self): + assert GoogleLLM._is_thought_part({"thought": True}) is True + + def test_is_thought_part_dict_false(self): + assert GoogleLLM._is_thought_part({"thought": False}) is False + + def test_is_thought_part_object(self): + obj = types.SimpleNamespace(thought=True) + assert GoogleLLM._is_thought_part(obj) is True + + +# --------------------------------------------------------------------------- +# _raw_gen +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGen: + + def test_returns_text(self, llm): + msgs = [{"role": "user", "content": "hi"}] + result = llm._raw_gen(llm, model="gemini-2.0", messages=msgs) + assert result == "ok" + + def test_with_tools_returns_response(self, llm): + tools = [ + { + "type": "function", + "function": { + "name": "t", + "description": "d", + "parameters": {"type": "object", "properties": {}}, + }, + } + ] + msgs = [{"role": "user", "content": "hi"}] + result = llm._raw_gen(llm, model="gemini", messages=msgs, tools=tools) + assert hasattr(result, "text") + + def test_with_response_schema(self, llm): + msgs = [{"role": "user", "content": "hi"}] + llm._raw_gen( + llm, + model="gemini", + messages=msgs, + response_schema={"type": "OBJECT"}, + ) + # Should not raise + + +# --------------------------------------------------------------------------- +# _raw_gen_stream +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGenStream: + + def test_yields_text_from_candidates(self, llm, monkeypatch): + part = types.SimpleNamespace( + text="chunk1", function_call=None, thought=False + ) + candidate = types.SimpleNamespace( + content=types.SimpleNamespace(parts=[part]) + ) + chunk = types.SimpleNamespace(candidates=[candidate]) + + monkeypatch.setattr( + FakeModels, + "generate_content_stream", + lambda self, *a, **kw: [chunk], + ) + + msgs = [{"role": "user", "content": "hi"}] + result = list(llm._raw_gen_stream(llm, model="gemini", messages=msgs)) + assert "chunk1" in result + + def test_yields_function_call_part(self, llm, monkeypatch): + fc = types.SimpleNamespace(name="search") + part = types.SimpleNamespace( + text=None, function_call=fc, thought=False + ) + candidate = types.SimpleNamespace( + content=types.SimpleNamespace(parts=[part]) + ) + chunk = types.SimpleNamespace(candidates=[candidate]) + + monkeypatch.setattr( + FakeModels, + "generate_content_stream", + lambda self, *a, **kw: [chunk], + ) + + msgs = [{"role": "user", "content": "hi"}] + result = list(llm._raw_gen_stream(llm, model="gemini", messages=msgs)) + assert any(hasattr(r, "function_call") for r in result) + + def test_yields_thought_event(self, llm, monkeypatch): + part = types.SimpleNamespace( + text="thinking", function_call=None, thought=True + ) + candidate = types.SimpleNamespace( + content=types.SimpleNamespace(parts=[part]) + ) + chunk = types.SimpleNamespace(candidates=[candidate]) + + monkeypatch.setattr( + FakeModels, + "generate_content_stream", + lambda self, *a, **kw: [chunk], + ) + + msgs = [{"role": "user", "content": "hi"}] + result = list(llm._raw_gen_stream(llm, model="gemini", messages=msgs)) + assert {"type": "thought", "thought": "thinking"} in result + + def test_text_only_chunk_via_hasattr(self, llm, monkeypatch): + chunk = types.SimpleNamespace(text="fallback", candidates=None, thought=False) + + monkeypatch.setattr( + FakeModels, + "generate_content_stream", + lambda self, *a, **kw: [chunk], + ) + + msgs = [{"role": "user", "content": "hi"}] + result = list(llm._raw_gen_stream(llm, model="gemini", messages=msgs)) + assert "fallback" in result + + def test_stream_error_propagates(self, llm, monkeypatch): + def error_stream(self, *a, **kw): + raise RuntimeError("stream_err") + + monkeypatch.setattr(FakeModels, "generate_content_stream", error_stream) + + msgs = [{"role": "user", "content": "hi"}] + with pytest.raises(RuntimeError, match="stream_err"): + list(llm._raw_gen_stream(llm, model="gemini", messages=msgs)) + + def test_skips_empty_text_parts(self, llm, monkeypatch): + part = types.SimpleNamespace( + text="", function_call=None, thought=False + ) + candidate = types.SimpleNamespace( + content=types.SimpleNamespace(parts=[part]) + ) + chunk = types.SimpleNamespace(candidates=[candidate]) + + monkeypatch.setattr( + FakeModels, + "generate_content_stream", + lambda self, *a, **kw: [chunk], + ) + + msgs = [{"role": "user", "content": "hi"}] + result = list(llm._raw_gen_stream(llm, model="gemini", messages=msgs)) + assert result == [] + + +# --------------------------------------------------------------------------- +# _supports_tools / _supports_structured_output +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestSupports: + + def test_supports_tools(self, llm): + assert llm._supports_tools() is True + + def test_supports_structured_output(self, llm): + assert llm._supports_structured_output() is True + + +# --------------------------------------------------------------------------- +# prepare_structured_output_format +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestPrepareStructuredOutputFormat: + + def test_none_returns_none(self, llm): + assert llm.prepare_structured_output_format(None) is None + + def test_type_mapping(self, llm): + schema = { + "type": "object", + "properties": { + "name": {"type": "string"}, + "count": {"type": "integer"}, + "score": {"type": "number"}, + "active": {"type": "boolean"}, + "items": {"type": "array", "items": {"type": "string"}}, + }, + } + result = llm.prepare_structured_output_format(schema) + assert result["type"] == "OBJECT" + assert result["properties"]["name"]["type"] == "STRING" + assert result["properties"]["count"]["type"] == "INTEGER" + assert result["properties"]["score"]["type"] == "NUMBER" + assert result["properties"]["active"]["type"] == "BOOLEAN" + assert result["properties"]["items"]["type"] == "ARRAY" + + def test_property_ordering_added(self, llm): + schema = { + "type": "object", + "properties": {"a": {"type": "string"}, "b": {"type": "string"}}, + } + result = llm.prepare_structured_output_format(schema) + assert "propertyOrdering" in result + assert set(result["propertyOrdering"]) == {"a", "b"} + + def test_format_date_converted(self, llm): + schema = {"type": "string", "format": "date"} + result = llm.prepare_structured_output_format(schema) + assert result["format"] == "date-time" + + def test_format_datetime_preserved(self, llm): + schema = {"type": "string", "format": "date-time"} + result = llm.prepare_structured_output_format(schema) + assert result["format"] == "date-time" + + def test_anyof_processed(self, llm): + schema = { + "anyOf": [ + {"type": "string"}, + {"type": "integer"}, + ] + } + result = llm.prepare_structured_output_format(schema) + assert len(result["anyOf"]) == 2 + assert result["anyOf"][0]["type"] == "STRING" + + +# --------------------------------------------------------------------------- +# get_supported_attachment_types +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestGetSupportedAttachmentTypes: + + def test_returns_list_with_expected_types(self, llm): + result = llm.get_supported_attachment_types() + assert "application/pdf" in result + assert "image/png" in result + assert "image/jpeg" in result + + +# --------------------------------------------------------------------------- +# prepare_messages_with_attachments +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestPrepareMessagesWithAttachments: + + def test_no_attachments_returns_same(self, llm): + msgs = [{"role": "user", "content": "hi"}] + result = llm.prepare_messages_with_attachments(msgs) + assert result == msgs + + def test_upload_error_adds_text_fallback(self, llm, monkeypatch): + monkeypatch.setattr( + llm, "_upload_file_to_google", lambda a: (_ for _ in ()).throw(Exception("fail")) + ) + msgs = [{"role": "user", "content": "hi"}] + attachments = [ + {"mime_type": "image/png", "path": "/tmp/img.png", "content": "fallback"}, + ] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + text_parts = [ + p for p in user_msg["content"] + if isinstance(p, dict) and p.get("type") == "text" and "could not" in p.get("text", "").lower() + ] + assert len(text_parts) == 1 + + def test_no_user_message_creates_one(self, llm, monkeypatch): + monkeypatch.setattr(llm, "_upload_file_to_google", lambda a: "gs://uri") + msgs = [{"role": "system", "content": "sys"}] + attachments = [{"mime_type": "image/png", "path": "/img.png"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msgs = [m for m in result if m["role"] == "user"] + assert len(user_msgs) == 1 + + +# --------------------------------------------------------------------------- +# _upload_file_to_google +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestUploadFileToGoogle: + + def test_returns_cached_uri(self, llm): + attachment = {"google_file_uri": "gs://cached"} + result = llm._upload_file_to_google(attachment) + assert result == "gs://cached" + + def test_raises_for_no_path(self, llm): + with pytest.raises(ValueError, match="No file path"): + llm._upload_file_to_google({}) + + def test_raises_for_missing_file(self, llm): + llm.storage = types.SimpleNamespace(file_exists=lambda p: False) + with pytest.raises(FileNotFoundError): + llm._upload_file_to_google({"path": "/nonexistent"}) diff --git a/tests/llm/test_llama_cpp.py b/tests/llm/test_llama_cpp.py new file mode 100644 index 00000000..0a4af43f --- /dev/null +++ b/tests/llm/test_llama_cpp.py @@ -0,0 +1,193 @@ +"""Unit tests for application/llm/llama_cpp.py — LlamaCpp and LlamaSingleton. + +Covers: + - LlamaSingleton: get_instance, query_model (thread-safe) + - LlamaCpp constructor + - _raw_gen: prompt format and result extraction + - _raw_gen_stream: streaming iteration +""" + +import sys +import types + +import pytest + + +# --------------------------------------------------------------------------- +# Fake llama_cpp module +# --------------------------------------------------------------------------- + + +class FakeLlama: + def __init__(self, model_path=None, n_ctx=None): + self.model_path = model_path + self.n_ctx = n_ctx + self.last_call = None + + def __call__(self, prompt, **kwargs): + self.last_call = {"prompt": prompt, **kwargs} + if kwargs.get("stream"): + return iter( + [ + {"choices": [{"text": "chunk1"}]}, + {"choices": [{"text": "chunk2"}]}, + ] + ) + return {"choices": [{"text": "prefix ### Answer \nthe answer"}]} + + +@pytest.fixture(autouse=True) +def patch_llama_cpp(monkeypatch): + fake_mod = types.ModuleType("llama_cpp") + fake_mod.Llama = FakeLlama + sys.modules["llama_cpp"] = fake_mod + + # Clear any cached instances + if "application.llm.llama_cpp" in sys.modules: + del sys.modules["application.llm.llama_cpp"] + + yield + + sys.modules.pop("llama_cpp", None) + if "application.llm.llama_cpp" in sys.modules: + del sys.modules["application.llm.llama_cpp"] + + +@pytest.fixture +def fresh_singleton(): + from application.llm.llama_cpp import LlamaSingleton + + LlamaSingleton._instances = {} + return LlamaSingleton + + +@pytest.fixture +def llm(fresh_singleton): + from application.llm.llama_cpp import LlamaCpp + + instance = LlamaCpp(api_key="k", user_api_key=None, llm_name="/path/to/model") + return instance + + +# --------------------------------------------------------------------------- +# LlamaSingleton +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestLlamaSingleton: + + def test_get_instance_creates_llama(self, fresh_singleton): + instance = fresh_singleton.get_instance("/model/path") + assert isinstance(instance, FakeLlama) + assert instance.model_path == "/model/path" + + def test_get_instance_caches(self, fresh_singleton): + inst1 = fresh_singleton.get_instance("/model") + inst2 = fresh_singleton.get_instance("/model") + assert inst1 is inst2 + + def test_different_names_different_instances(self, fresh_singleton): + inst1 = fresh_singleton.get_instance("/model_a") + inst2 = fresh_singleton.get_instance("/model_b") + assert inst1 is not inst2 + + def test_query_model_thread_safe(self, fresh_singleton): + instance = fresh_singleton.get_instance("/model") + result = fresh_singleton.query_model(instance, "prompt", max_tokens=10) + assert "choices" in result + + def test_import_error_raised(self, fresh_singleton, monkeypatch): + # Remove the fake module to simulate import failure + sys.modules.pop("llama_cpp", None) + fresh_singleton._instances = {} + + with pytest.raises(ImportError, match="llama_cpp"): + fresh_singleton.get_instance("/new_model") + + +# --------------------------------------------------------------------------- +# LlamaCpp constructor +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestLlamaCppConstructor: + + def test_sets_api_key(self, llm): + assert llm.api_key == "k" + + def test_sets_user_api_key(self): + from application.llm.llama_cpp import LlamaCpp, LlamaSingleton + + LlamaSingleton._instances = {} + instance = LlamaCpp( + api_key="k", user_api_key="uk", llm_name="/path/model" + ) + assert instance.user_api_key == "uk" + + def test_creates_llama_instance(self, llm): + assert isinstance(llm.llama, FakeLlama) + + +# --------------------------------------------------------------------------- +# _raw_gen +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGen: + + def test_returns_answer(self, llm): + msgs = [ + {"content": "context text"}, + {"content": "user question"}, + ] + result = llm._raw_gen(llm, model="local", messages=msgs) + assert result == "the answer" + + def test_prompt_contains_instruction_and_context(self, llm): + msgs = [ + {"content": "my context"}, + {"content": "my question"}, + ] + llm._raw_gen(llm, model="local", messages=msgs) + prompt = llm.llama.last_call["prompt"] + assert "### Instruction" in prompt + assert "### Context" in prompt + assert "my question" in prompt + assert "my context" in prompt + + def test_max_tokens_passed(self, llm): + msgs = [{"content": "c"}, {"content": "q"}] + llm._raw_gen(llm, model="local", messages=msgs) + assert llm.llama.last_call["max_tokens"] == 150 + assert llm.llama.last_call["echo"] is False + + +# --------------------------------------------------------------------------- +# _raw_gen_stream +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGenStream: + + def test_yields_text_chunks(self, llm): + msgs = [{"content": "c"}, {"content": "q"}] + chunks = list(llm._raw_gen_stream(llm, model="local", messages=msgs)) + assert chunks == ["chunk1", "chunk2"] + + def test_prompt_format(self, llm): + msgs = [{"content": "ctx"}, {"content": "question"}] + list(llm._raw_gen_stream(llm, model="local", messages=msgs)) + prompt = llm.llama.last_call["prompt"] + assert "### Instruction" in prompt + assert "### Answer" in prompt + + def test_stream_flag_passed(self, llm): + msgs = [{"content": "c"}, {"content": "q"}] + list( + llm._raw_gen_stream(llm, model="local", messages=msgs, stream=True) + ) + assert llm.llama.last_call["stream"] is True diff --git a/tests/llm/test_openai.py b/tests/llm/test_openai.py new file mode 100644 index 00000000..7dcf7b8f --- /dev/null +++ b/tests/llm/test_openai.py @@ -0,0 +1,717 @@ +"""Unit tests for application/llm/openai.py — OpenAILLM. + +Extends coverage beyond test_openai_llm.py: + - _truncate_base64_for_logging helper + - _normalize_reasoning_value edge cases + - _extract_reasoning_text edge cases + - _clean_messages_openai: file type, legacy format, unexpected content type + - _raw_gen with tools and response_format + - _raw_gen_stream tool_calls yielding + - prepare_structured_output_format nested schemas + - AzureOpenAILLM constructor + - _supports_tools / _supports_structured_output + - get_supported_attachment_types + - prepare_messages_with_attachments edge cases + - _get_base64_image / _upload_file_to_openai +""" + +import types + +import pytest + +from application.llm.openai import OpenAILLM, _truncate_base64_for_logging + + +# --------------------------------------------------------------------------- +# Fake client helpers +# --------------------------------------------------------------------------- + +class _Msg: + def __init__(self, content=None, tool_calls=None): + self.content = content + self.tool_calls = tool_calls + + +class _Delta: + def __init__(self, content=None, reasoning_content=None, tool_calls=None): + self.content = content + self.reasoning_content = reasoning_content + self.tool_calls = tool_calls + + +class _Choice: + def __init__(self, content=None, delta=None, finish_reason="stop"): + if isinstance(delta, _Delta): + self.delta = delta + else: + self.delta = _Delta(content=delta) + self.message = _Msg(content=content) + self.finish_reason = finish_reason + + +class _StreamLine: + def __init__(self, choices): + self.choices = choices + + +class _Response: + def __init__(self, choices=None, lines=None): + self._choices = choices or [] + self._lines = lines or [] + + @property + def choices(self): + return self._choices + + def __iter__(self): + yield from self._lines + + def close(self): + pass + + +class FakeChatCompletions: + def __init__(self): + self.last_kwargs = None + self._response = None + + def create(self, **kwargs): + self.last_kwargs = kwargs + if self._response: + return self._response + if not kwargs.get("stream"): + return _Response(choices=[_Choice(content="hello world")]) + return _Response( + lines=[ + _StreamLine([_Choice(delta="part1")]), + _StreamLine([_Choice(delta="part2")]), + ] + ) + + +class FakeFiles: + def create(self, file=None, purpose=None): + return types.SimpleNamespace(id="file_id_uploaded") + + +class FakeClient: + def __init__(self): + self.chat = types.SimpleNamespace(completions=FakeChatCompletions()) + self.files = FakeFiles() + + +@pytest.fixture +def llm(): + instance = OpenAILLM(api_key="sk-test", user_api_key=None) + instance.storage = types.SimpleNamespace( + get_file=lambda path: types.SimpleNamespace( + __enter__=lambda s: types.SimpleNamespace(read=lambda: b"img_bytes"), + __exit__=lambda s, *a: None, + ), + file_exists=lambda path: True, + process_file=lambda path, processor_func, **kw: processor_func(path), + ) + instance.client = FakeClient() + return instance + + +# --------------------------------------------------------------------------- +# _truncate_base64_for_logging +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestTruncateBase64ForLogging: + + def test_truncates_data_url_in_content_string(self): + msgs = [{"role": "user", "content": "data:image/png;base64," + "A" * 200}] + result = _truncate_base64_for_logging(msgs) + assert "BASE64_DATA_TRUNCATED" in result[0]["content"] + assert "A" * 200 not in result[0]["content"] + + def test_truncates_url_key_in_list_content(self): + msgs = [ + { + "role": "user", + "content": [ + {"url": "data:image/png;base64," + "B" * 300}, + ], + } + ] + result = _truncate_base64_for_logging(msgs) + item = result[0]["content"][0] + assert "BASE64_DATA_TRUNCATED" in item["url"] + + def test_truncates_data_key_with_long_value(self): + msgs = [{"role": "user", "content": [{"data": "X" * 200}]}] + result = _truncate_base64_for_logging(msgs) + item = result[0]["content"][0] + assert "BASE64_DATA_TRUNCATED" in item["data"] + + def test_preserves_non_base64_content(self): + msgs = [{"role": "user", "content": "normal text"}] + result = _truncate_base64_for_logging(msgs) + assert result[0]["content"] == "normal text" + + def test_handles_message_without_content_key(self): + msgs = [{"role": "system"}] + result = _truncate_base64_for_logging(msgs) + assert "content" not in result[0] + + def test_nested_dict_truncation(self): + msgs = [ + { + "role": "user", + "content": {"nested": "data:image/jpeg;base64," + "C" * 100}, + } + ] + result = _truncate_base64_for_logging(msgs) + assert "BASE64_DATA_TRUNCATED" in result[0]["content"]["nested"] + + +# --------------------------------------------------------------------------- +# _normalize_reasoning_value +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestNormalizeReasoningValue: + + def test_none_returns_empty(self): + assert OpenAILLM._normalize_reasoning_value(None) == "" + + def test_string_passthrough(self): + assert OpenAILLM._normalize_reasoning_value("hello") == "hello" + + def test_list_concatenation(self): + assert OpenAILLM._normalize_reasoning_value(["a", "b"]) == "ab" + + def test_dict_text_key(self): + assert OpenAILLM._normalize_reasoning_value({"text": "t"}) == "t" + + def test_dict_content_key(self): + assert OpenAILLM._normalize_reasoning_value({"content": "c"}) == "c" + + def test_dict_reasoning_content_key(self): + assert OpenAILLM._normalize_reasoning_value({"reasoning_content": "rc"}) == "rc" + + def test_dict_empty_returns_empty(self): + assert OpenAILLM._normalize_reasoning_value({}) == "" + + def test_object_with_text_attribute(self): + obj = types.SimpleNamespace(text="from_attr") + assert OpenAILLM._normalize_reasoning_value(obj) == "from_attr" + + def test_object_with_content_attribute(self): + obj = types.SimpleNamespace(content="content_attr") + assert OpenAILLM._normalize_reasoning_value(obj) == "content_attr" + + def test_nested_list_of_dicts(self): + val = [{"text": "a"}, {"content": "b"}] + assert OpenAILLM._normalize_reasoning_value(val) == "ab" + + +# --------------------------------------------------------------------------- +# _extract_reasoning_text +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestExtractReasoningText: + + def test_none_delta_returns_empty(self): + assert OpenAILLM._extract_reasoning_text(None) == "" + + def test_extracts_reasoning_content_attr(self): + delta = types.SimpleNamespace(reasoning_content="thought!") + assert OpenAILLM._extract_reasoning_text(delta) == "thought!" + + def test_extracts_thinking_attr(self): + delta = types.SimpleNamespace(thinking="deep thought") + assert OpenAILLM._extract_reasoning_text(delta) == "deep thought" + + def test_extracts_from_dict_delta(self): + delta = {"reasoning_content": "dict_thought"} + assert OpenAILLM._extract_reasoning_text(delta) == "dict_thought" + + def test_no_reasoning_returns_empty(self): + delta = types.SimpleNamespace() + assert OpenAILLM._extract_reasoning_text(delta) == "" + + +# --------------------------------------------------------------------------- +# _clean_messages_openai +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestCleanMessagesOpenai: + + def test_string_content(self, llm): + msgs = [{"role": "user", "content": "hello"}] + cleaned = llm._clean_messages_openai(msgs) + assert cleaned == [{"role": "user", "content": "hello"}] + + def test_model_role_converted_to_assistant(self, llm): + msgs = [{"role": "model", "content": "hi"}] + cleaned = llm._clean_messages_openai(msgs) + assert cleaned[0]["role"] == "assistant" + + def test_file_type_in_list_content(self, llm): + msgs = [ + { + "role": "user", + "content": [ + {"type": "file", "file": {"file_id": "f1"}}, + ], + } + ] + cleaned = llm._clean_messages_openai(msgs) + content = cleaned[0]["content"] + assert any(p.get("type") == "file" for p in content) + + def test_image_url_type(self, llm): + msgs = [ + { + "role": "user", + "content": [ + {"type": "image_url", "image_url": {"url": "http://img.png"}}, + ], + } + ] + cleaned = llm._clean_messages_openai(msgs) + assert any(p.get("type") == "image_url" for p in cleaned[0]["content"]) + + def test_legacy_text_format(self, llm): + msgs = [{"role": "user", "content": [{"text": "legacy"}]}] + cleaned = llm._clean_messages_openai(msgs) + part = cleaned[0]["content"][0] + assert part["type"] == "text" + assert part["text"] == "legacy" + + def test_function_call_args_json_string(self, llm): + msgs = [ + { + "role": "assistant", + "content": [ + { + "function_call": { + "call_id": "c1", + "name": "fn", + "args": '{"a": 1}', + } + }, + ], + } + ] + cleaned = llm._clean_messages_openai(msgs) + tc_msg = next(m for m in cleaned if m.get("tool_calls")) + assert tc_msg["tool_calls"][0]["function"]["name"] == "fn" + + def test_function_response_becomes_tool_message(self, llm): + msgs = [ + { + "role": "user", + "content": [ + { + "function_response": { + "call_id": "c1", + "name": "fn", + "response": {"result": 42}, + } + }, + ], + } + ] + cleaned = llm._clean_messages_openai(msgs) + tool_msg = next(m for m in cleaned if m["role"] == "tool") + assert tool_msg["tool_call_id"] == "c1" + assert "42" in tool_msg["content"] + + def test_skips_none_content(self, llm): + msgs = [{"role": "user", "content": None}] + cleaned = llm._clean_messages_openai(msgs) + assert cleaned == [] + + def test_raises_for_unexpected_content_type(self, llm): + msgs = [{"role": "user", "content": 12345}] + with pytest.raises(ValueError, match="Unexpected content type"): + llm._clean_messages_openai(msgs) + + +# --------------------------------------------------------------------------- +# _raw_gen +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGen: + + def test_returns_content(self, llm): + msgs = [{"role": "user", "content": "hi"}] + result = llm._raw_gen(llm, model="gpt-4o", messages=msgs, stream=False) + assert result == "hello world" + + def test_with_tools_returns_choice(self, llm): + tools = [{"type": "function", "function": {"name": "t"}}] + msgs = [{"role": "user", "content": "hi"}] + result = llm._raw_gen( + llm, model="gpt-4o", messages=msgs, stream=False, tools=tools + ) + assert hasattr(result, "message") + + def test_with_response_format(self, llm): + msgs = [{"role": "user", "content": "hi"}] + llm._raw_gen( + llm, + model="gpt-4o", + messages=msgs, + stream=False, + response_format={"type": "json_object"}, + ) + kwargs = llm.client.chat.completions.last_kwargs + assert kwargs["response_format"] == {"type": "json_object"} + + def test_max_tokens_converted(self, llm): + msgs = [{"role": "user", "content": "hi"}] + llm._raw_gen( + llm, model="gpt-4o", messages=msgs, stream=False, max_tokens=100 + ) + kwargs = llm.client.chat.completions.last_kwargs + assert "max_completion_tokens" in kwargs + assert "max_tokens" not in kwargs + + def test_tools_passed_to_client(self, llm): + tools = [{"type": "function", "function": {"name": "t"}}] + msgs = [{"role": "user", "content": "hi"}] + llm._raw_gen( + llm, model="gpt-4o", messages=msgs, stream=False, tools=tools + ) + kwargs = llm.client.chat.completions.last_kwargs + assert kwargs["tools"] == tools + + +# --------------------------------------------------------------------------- +# _raw_gen_stream +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGenStream: + + def test_yields_content_chunks(self, llm): + msgs = [{"role": "user", "content": "hi"}] + chunks = list(llm._raw_gen_stream(llm, model="gpt", messages=msgs)) + assert "part1" in chunks + assert "part2" in chunks + + def test_yields_tool_call_choices(self, llm): + tool_calls_obj = [types.SimpleNamespace(id="tc1")] + delta = _Delta(content=None, tool_calls=tool_calls_obj) + choice = _Choice(delta=delta, finish_reason="tool_calls") + choice.delta = delta + line = _StreamLine([choice]) + resp = _Response(lines=[line]) + llm.client.chat.completions._response = resp + llm.client.chat.completions.create = lambda **kw: resp + + msgs = [{"role": "user", "content": "hi"}] + chunks = list(llm._raw_gen_stream(llm, model="gpt", messages=msgs)) + assert any(hasattr(c, "finish_reason") for c in chunks) + + def test_skips_empty_choices(self, llm): + line = types.SimpleNamespace(choices=None) + resp = _Response(lines=[line]) + llm.client.chat.completions.create = lambda **kw: resp + + msgs = [{"role": "user", "content": "hi"}] + chunks = list(llm._raw_gen_stream(llm, model="gpt", messages=msgs)) + assert chunks == [] + + def test_calls_close_on_response(self, llm): + closed = {"called": False} + resp = _Response(lines=[]) + + def mark_closed(): + closed["called"] = True + + resp.close = mark_closed + llm.client.chat.completions.create = lambda **kw: resp + + msgs = [{"role": "user", "content": "hi"}] + list(llm._raw_gen_stream(llm, model="gpt", messages=msgs)) + assert closed["called"] + + +# --------------------------------------------------------------------------- +# _supports_tools / _supports_structured_output +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestSupports: + + def test_supports_tools(self, llm): + assert llm._supports_tools() is True + + def test_supports_structured_output(self, llm): + assert llm._supports_structured_output() is True + + +# --------------------------------------------------------------------------- +# prepare_structured_output_format +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestPrepareStructuredOutputFormat: + + def test_none_schema_returns_none(self, llm): + assert llm.prepare_structured_output_format(None) is None + + def test_empty_schema_returns_none(self, llm): + assert llm.prepare_structured_output_format({}) is None + + def test_nested_object_gets_additional_properties_false(self, llm): + schema = { + "type": "object", + "properties": { + "inner": { + "type": "object", + "properties": { + "x": {"type": "string"}, + }, + } + }, + } + result = llm.prepare_structured_output_format(schema) + inner = result["json_schema"]["schema"]["properties"]["inner"] + assert inner["additionalProperties"] is False + assert "x" in inner["required"] + + def test_array_items_processed(self, llm): + schema = { + "type": "object", + "properties": { + "items_list": { + "type": "array", + "items": { + "type": "object", + "properties": {"name": {"type": "string"}}, + }, + } + }, + } + result = llm.prepare_structured_output_format(schema) + items_schema = result["json_schema"]["schema"]["properties"]["items_list"][ + "items" + ] + assert items_schema["additionalProperties"] is False + + def test_anyof_schemas_processed(self, llm): + schema = { + "type": "object", + "properties": { + "val": { + "anyOf": [ + {"type": "object", "properties": {"a": {"type": "string"}}}, + {"type": "string"}, + ] + } + }, + } + result = llm.prepare_structured_output_format(schema) + any_of = result["json_schema"]["schema"]["properties"]["val"]["anyOf"] + assert any_of[0]["additionalProperties"] is False + + def test_uses_schema_name_and_description(self, llm): + schema = { + "type": "object", + "name": "MySchema", + "description": "My custom schema", + "properties": {"a": {"type": "string"}}, + } + result = llm.prepare_structured_output_format(schema) + assert result["json_schema"]["name"] == "MySchema" + assert result["json_schema"]["description"] == "My custom schema" + + def test_default_name_and_description(self, llm): + schema = { + "type": "object", + "properties": {"a": {"type": "string"}}, + } + result = llm.prepare_structured_output_format(schema) + assert result["json_schema"]["name"] == "response" + assert result["json_schema"]["description"] == "Structured response" + + +# --------------------------------------------------------------------------- +# get_supported_attachment_types +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestGetSupportedAttachmentTypes: + + def test_returns_list(self, llm): + result = llm.get_supported_attachment_types() + assert isinstance(result, list) + assert len(result) > 0 + + +# --------------------------------------------------------------------------- +# prepare_messages_with_attachments +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestPrepareMessagesWithAttachments: + + def test_no_attachments_returns_same(self, llm): + msgs = [{"role": "user", "content": "hi"}] + result = llm.prepare_messages_with_attachments(msgs) + assert result == msgs + + def test_empty_attachments_returns_same(self, llm): + msgs = [{"role": "user", "content": "hi"}] + result = llm.prepare_messages_with_attachments(msgs, []) + assert result == msgs + + def test_image_with_preconverted_data(self, llm): + msgs = [{"role": "user", "content": "look at this"}] + attachments = [{"mime_type": "image/png", "data": "AABBCC"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + assert isinstance(user_msg["content"], list) + img_part = next( + p for p in user_msg["content"] if p.get("type") == "image_url" + ) + assert "AABBCC" in img_part["image_url"]["url"] + + def test_no_user_message_creates_one(self, llm): + msgs = [{"role": "system", "content": "sys"}] + attachments = [{"mime_type": "image/png", "data": "AAA"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msgs = [m for m in result if m["role"] == "user"] + assert len(user_msgs) == 1 + + def test_unsupported_mime_type_skipped(self, llm): + msgs = [{"role": "user", "content": "hi"}] + attachments = [{"mime_type": "application/octet-stream"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + # Content should still be the original string (no list conversion) + # since unsupported type is skipped but user message content is + # converted to list + assert isinstance(user_msg["content"], list) + # Only the text part should exist + assert len(user_msg["content"]) == 1 + + def test_image_error_adds_text_fallback(self, llm): + llm.storage = types.SimpleNamespace( + get_file=lambda path: (_ for _ in ()).throw(Exception("storage err")), + ) + msgs = [{"role": "user", "content": "hi"}] + attachments = [ + { + "mime_type": "image/png", + "path": "/tmp/bad.png", + "content": "fallback text", + } + ] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + text_parts = [ + p for p in user_msg["content"] if p.get("type") == "text" and "could not" in p.get("text", "").lower() + ] + assert len(text_parts) == 1 + + def test_pdf_error_adds_content_fallback(self, llm): + llm.storage = types.SimpleNamespace( + file_exists=lambda p: False, + ) + msgs = [{"role": "user", "content": "hi"}] + attachments = [ + { + "mime_type": "application/pdf", + "path": "/tmp/bad.pdf", + "content": "pdf fallback", + } + ] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + text_parts = [ + p for p in user_msg["content"] if p.get("type") == "text" and "pdf fallback" in p.get("text", "") + ] + assert len(text_parts) == 1 + + def test_content_not_list_becomes_empty_list(self, llm): + msgs = [{"role": "user", "content": 42}] + attachments = [{"mime_type": "image/png", "data": "AAA"}] + result = llm.prepare_messages_with_attachments(msgs, attachments) + user_msg = next(m for m in result if m["role"] == "user") + assert isinstance(user_msg["content"], list) + + +# --------------------------------------------------------------------------- +# _get_base64_image +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestGetBase64Image: + + def test_raises_for_no_path(self, llm): + with pytest.raises(ValueError, match="No file path"): + llm._get_base64_image({}) + + def test_raises_for_file_not_found(self, llm): + import contextlib + + @contextlib.contextmanager + def fake_get_file(path): + raise FileNotFoundError("not found") + + llm.storage = types.SimpleNamespace(get_file=fake_get_file) + with pytest.raises(FileNotFoundError): + llm._get_base64_image({"path": "/nonexistent"}) + + +# --------------------------------------------------------------------------- +# AzureOpenAILLM +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestAzureOpenAILLM: + + def test_constructor(self, monkeypatch): + monkeypatch.setattr( + "application.llm.openai.settings", + types.SimpleNamespace( + OPENAI_API_KEY="k", + API_KEY="k", + OPENAI_BASE_URL="", + OPENAI_API_BASE="https://my.azure.endpoint", + OPENAI_API_VERSION="2024-02-01", + AZURE_DEPLOYMENT_NAME="my-deployment", + ), + ) + monkeypatch.setattr( + "application.llm.openai.StorageCreator", + types.SimpleNamespace(get_storage=lambda: None), + ) + from unittest.mock import MagicMock + + monkeypatch.setattr("application.llm.openai.OpenAI", MagicMock()) + mock_azure = MagicMock() + monkeypatch.setattr("openai.AzureOpenAI", mock_azure, raising=False) + + # We need to reimport to get fresh class with mocked module + import importlib + import application.llm.openai as oai_mod + + importlib.reload(oai_mod) + + # Just verify the class exists and inherits from OpenAILLM + assert issubclass(oai_mod.AzureOpenAILLM, oai_mod.OpenAILLM) diff --git a/tests/llm/test_premai.py b/tests/llm/test_premai.py new file mode 100644 index 00000000..06f57e36 --- /dev/null +++ b/tests/llm/test_premai.py @@ -0,0 +1,190 @@ +"""Unit tests for application/llm/premai.py — PremAILLM. + +Covers: + - Constructor + - _raw_gen: API call and return value + - _raw_gen_stream: streaming with delta content filtering +""" + +import sys +import types + +import pytest + + +# --------------------------------------------------------------------------- +# Fake premai module +# --------------------------------------------------------------------------- + + +class _FakeMessage: + def __init__(self, content): + self.message = {"content": content} + + +class _FakeDelta: + def __init__(self, content): + self.delta = {"content": content} + + +class _FakeChoice: + def __init__(self, content): + self.message = {"content": content} + + +class _FakeStreamChoice: + def __init__(self, content): + self.delta = {"content": content} + + +class _FakeResponse: + def __init__(self, content="result_text"): + self.choices = [_FakeChoice(content)] + + +class _FakeStreamLine: + def __init__(self, content): + self.choices = [_FakeStreamChoice(content)] + + +class _FakeChatCompletions: + def __init__(self): + self.last_kwargs = None + + def create(self, **kwargs): + self.last_kwargs = kwargs + if kwargs.get("stream"): + return [ + _FakeStreamLine("chunk1"), + _FakeStreamLine("chunk2"), + _FakeStreamLine(None), # None content should be filtered + ] + return _FakeResponse() + + +class _FakeChat: + def __init__(self): + self.completions = _FakeChatCompletions() + + +class _FakePrem: + def __init__(self, api_key=None): + self.api_key = api_key + self.chat = _FakeChat() + + +@pytest.fixture(autouse=True) +def patch_premai(monkeypatch): + fake_mod = types.ModuleType("premai") + fake_mod.Prem = _FakePrem + sys.modules["premai"] = fake_mod + + if "application.llm.premai" in sys.modules: + del sys.modules["application.llm.premai"] + yield + sys.modules.pop("premai", None) + if "application.llm.premai" in sys.modules: + del sys.modules["application.llm.premai"] + + +@pytest.fixture +def llm(): + from application.llm.premai import PremAILLM + + return PremAILLM(api_key="test-key") + + +# --------------------------------------------------------------------------- +# Constructor +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestPremAIConstructor: + + def test_sets_api_key(self, llm): + assert llm.api_key == "test-key" + + def test_sets_user_api_key_none(self, llm): + assert llm.user_api_key is None + + def test_client_created(self, llm): + assert isinstance(llm.client, _FakePrem) + + def test_project_id_from_settings(self, llm): + from application.core.settings import settings + + assert llm.project_id == settings.PREMAI_PROJECT_ID + + +# --------------------------------------------------------------------------- +# _raw_gen +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGen: + + def test_returns_content(self, llm): + msgs = [{"role": "user", "content": "hi"}] + result = llm._raw_gen(llm, model="model-1", messages=msgs) + assert result == "result_text" + + def test_passes_model_and_project_id(self, llm): + msgs = [{"role": "user", "content": "hi"}] + llm._raw_gen(llm, model="my-model", messages=msgs) + kwargs = llm.client.chat.completions.last_kwargs + assert kwargs["model"] == "my-model" + assert kwargs["project_id"] == llm.project_id + assert kwargs["stream"] is False + + def test_passes_messages(self, llm): + msgs = [{"role": "user", "content": "hello"}] + llm._raw_gen(llm, model="m", messages=msgs) + kwargs = llm.client.chat.completions.last_kwargs + assert kwargs["messages"] == msgs + + def test_extra_kwargs_forwarded(self, llm): + msgs = [{"role": "user", "content": "hi"}] + llm._raw_gen(llm, model="m", messages=msgs, temperature=0.5) + kwargs = llm.client.chat.completions.last_kwargs + assert kwargs["temperature"] == 0.5 + + +# --------------------------------------------------------------------------- +# _raw_gen_stream +# --------------------------------------------------------------------------- + + +@pytest.mark.unit +class TestRawGenStream: + + def test_yields_non_none_content(self, llm): + msgs = [{"role": "user", "content": "hi"}] + chunks = list( + llm._raw_gen_stream(llm, model="m", messages=msgs, stream=True) + ) + assert chunks == ["chunk1", "chunk2"] + + def test_filters_none_content(self, llm): + msgs = [{"role": "user", "content": "hi"}] + chunks = list( + llm._raw_gen_stream(llm, model="m", messages=msgs, stream=True) + ) + assert None not in chunks + + def test_passes_stream_true(self, llm): + msgs = [{"role": "user", "content": "hi"}] + list(llm._raw_gen_stream(llm, model="m", messages=msgs)) + kwargs = llm.client.chat.completions.last_kwargs + assert kwargs["stream"] is True + + def test_passes_extra_kwargs(self, llm): + msgs = [{"role": "user", "content": "hi"}] + list( + llm._raw_gen_stream( + llm, model="m", messages=msgs, max_tokens=100 + ) + ) + kwargs = llm.client.chat.completions.last_kwargs + assert kwargs["max_tokens"] == 100 diff --git a/tests/parser/file/__init__.py b/tests/parser/file/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/parser/file/test_bulk.py b/tests/parser/file/test_bulk.py new file mode 100644 index 00000000..2d501878 --- /dev/null +++ b/tests/parser/file/test_bulk.py @@ -0,0 +1,367 @@ +"""Comprehensive tests for application/parser/file/bulk.py + +Covers: SimpleDirectoryReader (init, file discovery, load_data, directory +structure building), get_default_file_extractor. +""" + +from unittest.mock import MagicMock, patch + +import pytest + +from application.parser.schema.base import Document + + +# ===================================================================== +# Helpers +# ===================================================================== + + +@pytest.fixture +def temp_dir(tmp_path): + """Create a temporary directory with test files.""" + (tmp_path / "file1.md").write_text("# Heading\n\nContent 1") + (tmp_path / "file2.txt").write_text("Plain text content") + (tmp_path / ".hidden").write_text("hidden file") + sub = tmp_path / "subdir" + sub.mkdir() + (sub / "file3.md").write_text("Nested content") + return tmp_path + + +@pytest.fixture +def temp_dir_with_types(tmp_path): + """Directory with multiple file types.""" + (tmp_path / "doc.md").write_text("markdown") + (tmp_path / "data.json").write_text('{"key": "value"}') + (tmp_path / "notes.txt").write_text("text") + return tmp_path + + +# ===================================================================== +# SimpleDirectoryReader - Init +# ===================================================================== + + +@pytest.mark.unit +class TestSimpleDirectoryReaderInit: + + def test_init_with_dir(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + reader = SimpleDirectoryReader(input_dir=str(temp_dir)) + assert len(reader.input_files) >= 2 + + def test_init_with_files(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + files = [str(temp_dir / "file1.md")] + reader = SimpleDirectoryReader(input_files=files) + assert len(reader.input_files) == 1 + + def test_init_requires_input(self): + from application.parser.file.bulk import SimpleDirectoryReader + + with pytest.raises(ValueError, match="Must provide"): + SimpleDirectoryReader() + + def test_exclude_hidden(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + reader = SimpleDirectoryReader(input_dir=str(temp_dir), exclude_hidden=True) + filenames = [f.name for f in reader.input_files] + assert ".hidden" not in filenames + + def test_include_hidden(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + reader = SimpleDirectoryReader(input_dir=str(temp_dir), exclude_hidden=False) + filenames = [f.name for f in reader.input_files] + assert ".hidden" in filenames + + def test_recursive(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + reader = SimpleDirectoryReader(input_dir=str(temp_dir), recursive=True) + filenames = [f.name for f in reader.input_files] + assert "file3.md" in filenames + + def test_non_recursive(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + reader = SimpleDirectoryReader(input_dir=str(temp_dir), recursive=False) + filenames = [f.name for f in reader.input_files] + assert "file3.md" not in filenames + + def test_required_exts(self, temp_dir_with_types): + from application.parser.file.bulk import SimpleDirectoryReader + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir_with_types), required_exts=[".md"] + ) + filenames = [f.name for f in reader.input_files] + assert "doc.md" in filenames + assert "data.json" not in filenames + assert "notes.txt" not in filenames + + def test_required_exts_case_insensitive(self, tmp_path): + from application.parser.file.bulk import SimpleDirectoryReader + + (tmp_path / "FILE.MD").write_text("content") + reader = SimpleDirectoryReader( + input_dir=str(tmp_path), required_exts=[".md"] + ) + assert len(reader.input_files) == 1 + + def test_num_files_limit(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), num_files_limit=1, recursive=False + ) + assert len(reader.input_files) <= 1 + + def test_custom_file_extractor(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + mock_parser = MagicMock() + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), + file_extractor={".md": mock_parser}, + ) + assert ".md" in reader.file_extractor + + +# ===================================================================== +# SimpleDirectoryReader - load_data +# ===================================================================== + + +@pytest.mark.unit +class TestSimpleDirectoryReaderLoadData: + + def test_load_data_returns_documents(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = "parsed content" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), + file_extractor={".md": mock_parser, ".txt": mock_parser}, + recursive=False, + exclude_hidden=True, + ) + docs = reader.load_data() + assert len(docs) >= 1 + for doc in docs: + assert isinstance(doc, Document) + + def test_load_data_concatenate(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = "content" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), + file_extractor={".md": mock_parser, ".txt": mock_parser}, + recursive=False, + exclude_hidden=True, + ) + docs = reader.load_data(concatenate=True) + assert len(docs) == 1 + + def test_load_data_with_file_metadata(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + def custom_metadata(filename): + return {"custom_key": f"meta_{filename}"} + + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = "parsed" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), + file_extractor={".md": mock_parser, ".txt": mock_parser}, + file_metadata=custom_metadata, + recursive=False, + exclude_hidden=True, + ) + docs = reader.load_data() + assert len(docs) >= 1 + for doc in docs: + assert doc.extra_info is not None + assert "custom_key" in doc.extra_info + + def test_load_data_inits_parser_if_not_set(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + mock_parser = MagicMock() + mock_parser.parser_config_set = False + mock_parser.parse_file.return_value = "content" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), + file_extractor={".md": mock_parser, ".txt": mock_parser}, + recursive=False, + exclude_hidden=True, + ) + reader.load_data() + mock_parser.init_parser.assert_called() + + def test_load_data_standard_read_for_unknown_ext(self, tmp_path): + from application.parser.file.bulk import SimpleDirectoryReader + + (tmp_path / "file.xyz").write_text("xyz content") + reader = SimpleDirectoryReader( + input_dir=str(tmp_path), + file_extractor={}, + ) + docs = reader.load_data() + assert len(docs) == 1 + assert "xyz content" in docs[0].text + + def test_load_data_list_return_from_parser(self, tmp_path): + from application.parser.file.bulk import SimpleDirectoryReader + + (tmp_path / "multi.md").write_text("content") + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = ["part1", "part2"] + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(tmp_path), + file_extractor={".md": mock_parser}, + ) + docs = reader.load_data() + assert len(docs) == 2 + + def test_load_data_tracks_token_counts(self, tmp_path): + from application.parser.file.bulk import SimpleDirectoryReader + + (tmp_path / "test.md").write_text("hello world") + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = "hello world" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(tmp_path), + file_extractor={".md": mock_parser}, + ) + reader.load_data() + assert hasattr(reader, "file_token_counts") + assert len(reader.file_token_counts) >= 1 + + +# ===================================================================== +# Directory Structure Building +# ===================================================================== + + +@pytest.mark.unit +class TestBuildDirectoryStructure: + + def test_builds_structure(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = "content" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), + file_extractor={".md": mock_parser, ".txt": mock_parser}, + exclude_hidden=True, + ) + reader.load_data() + assert hasattr(reader, "directory_structure") + assert isinstance(reader.directory_structure, dict) + + def test_structure_contains_files_and_dirs(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = "content" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), + file_extractor={".md": mock_parser, ".txt": mock_parser}, + exclude_hidden=True, + ) + reader.load_data() + struct = reader.directory_structure + # Should contain subdir + assert "subdir" in struct + # Files should have metadata + for key, val in struct.items(): + if isinstance(val, dict) and "type" in val: + assert "size_bytes" in val + + def test_structure_excludes_hidden(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = "c" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_dir=str(temp_dir), + file_extractor={".md": mock_parser, ".txt": mock_parser}, + exclude_hidden=True, + ) + reader.load_data() + assert ".hidden" not in reader.directory_structure + + def test_no_structure_without_input_dir(self, temp_dir): + from application.parser.file.bulk import SimpleDirectoryReader + + files = [str(temp_dir / "file1.md")] + mock_parser = MagicMock() + mock_parser.parser_config_set = True + mock_parser.parse_file.return_value = "content" + mock_parser.get_file_metadata.return_value = {} + + reader = SimpleDirectoryReader( + input_files=files, + file_extractor={".md": mock_parser}, + ) + reader.load_data() + assert reader.directory_structure == {} + + +# ===================================================================== +# get_default_file_extractor +# ===================================================================== + + +@pytest.mark.unit +class TestGetDefaultFileExtractor: + + def test_returns_dict(self): + from application.parser.file.bulk import get_default_file_extractor + + with patch.dict("sys.modules", {"docling": None, "docling.document_converter": None}): + result = get_default_file_extractor() + assert isinstance(result, dict) + assert ".pdf" in result + + def test_fallback_parsers_on_import_error(self): + with patch( + "application.parser.file.bulk.get_default_file_extractor" + ) as mock_fn: + mock_fn.return_value = {".pdf": MagicMock(), ".md": MagicMock()} + result = mock_fn() + assert ".pdf" in result diff --git a/tests/parser/file/test_docling_parser.py b/tests/parser/file/test_docling_parser.py new file mode 100644 index 00000000..16582b71 --- /dev/null +++ b/tests/parser/file/test_docling_parser.py @@ -0,0 +1,382 @@ +"""Comprehensive tests for application/parser/file/docling_parser.py + +Covers: DoclingParser (init, _init_parser, _get_ocr_options, _export_content, +parse_file), subclass initialization, error handling. +""" + +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + + +# ===================================================================== +# DoclingParser - Init +# ===================================================================== + + +@pytest.mark.unit +class TestDoclingParserInit: + + def test_default_init(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser() + assert parser.ocr_enabled is True + assert parser.table_structure is True + assert parser.export_format == "markdown" + assert parser.use_rapidocr is True + assert parser.ocr_languages == ["english"] + assert parser.force_full_page_ocr is False + assert parser._converter is None + + def test_custom_init(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser( + ocr_enabled=False, + table_structure=False, + export_format="text", + use_rapidocr=False, + ocr_languages=["german"], + force_full_page_ocr=True, + ) + assert parser.ocr_enabled is False + assert parser.table_structure is False + assert parser.export_format == "text" + assert parser.use_rapidocr is False + assert parser.ocr_languages == ["german"] + assert parser.force_full_page_ocr is True + + +# ===================================================================== +# Init Parser +# ===================================================================== + + +@pytest.mark.unit +class TestDoclingParserInitParser: + + def test_init_parser_raises_without_docling(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser() + + with patch("importlib.util.find_spec", return_value=None): + with pytest.raises(ImportError, match="docling is required"): + parser._init_parser() + + def test_init_parser_success(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser() + + mock_converter = MagicMock() + with patch("importlib.util.find_spec", return_value=MagicMock()), \ + patch.object(parser, "_create_converter", return_value=mock_converter): + result = parser._init_parser() + + assert isinstance(result, dict) + assert result["ocr_enabled"] is True + assert result["table_structure"] is True + assert parser._converter is mock_converter + + +# ===================================================================== +# Get OCR Options +# ===================================================================== + + +@pytest.mark.unit +class TestGetOCROptions: + + def test_returns_none_when_rapidocr_disabled(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(use_rapidocr=False) + assert parser._get_ocr_options() is None + + def test_returns_options_when_available(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(use_rapidocr=True, ocr_languages=["english"]) + + mock_options = MagicMock() + with patch( + "application.parser.file.docling_parser.DoclingParser._get_ocr_options", + return_value=mock_options, + ): + result = parser._get_ocr_options() + assert result is mock_options + + def test_returns_none_on_import_error(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(use_rapidocr=True) + + # Simulate the ImportError path + original = parser._get_ocr_options + + def patched_get_ocr(): + try: + raise ImportError("No RapidOcrOptions") + except ImportError: + return None + + parser._get_ocr_options = patched_get_ocr + assert parser._get_ocr_options() is None + parser._get_ocr_options = original + + +# ===================================================================== +# Export Content +# ===================================================================== + + +@pytest.mark.unit +class TestExportContent: + + def test_export_markdown(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(export_format="markdown") + mock_doc = MagicMock() + mock_doc.export_to_markdown.return_value = "# Title\n\nContent here" + mock_doc.texts = [] + + result = parser._export_content(mock_doc) + assert "# Title" in result + mock_doc.export_to_markdown.assert_called_once() + + def test_export_html(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(export_format="html") + mock_doc = MagicMock() + mock_doc.export_to_html.return_value = "

Title

" + mock_doc.texts = [] + + result = parser._export_content(mock_doc) + assert "

" in result + + def test_export_text(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(export_format="text") + mock_doc = MagicMock() + mock_doc.export_to_text.return_value = "Plain text content" + mock_doc.texts = [] + + result = parser._export_content(mock_doc) + assert "Plain text" in result + + def test_fallback_to_texts_on_minimal_content(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(export_format="markdown") + mock_doc = MagicMock() + mock_doc.export_to_markdown.return_value = "" + + text1 = MagicMock() + text1.text = "OCR extracted text 1" + text2 = MagicMock() + text2.text = "OCR extracted text 2" + mock_doc.texts = [text1, text2] + + result = parser._export_content(mock_doc) + assert "OCR extracted text 1" in result + assert "OCR extracted text 2" in result + + def test_no_fallback_for_substantial_content(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(export_format="markdown") + mock_doc = MagicMock() + mock_doc.export_to_markdown.return_value = "A" * 100 + mock_doc.texts = [] + + result = parser._export_content(mock_doc) + assert result == "A" * 100 + + def test_fallback_skipped_when_no_texts(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(export_format="markdown") + mock_doc = MagicMock() + mock_doc.export_to_markdown.return_value = "short" + mock_doc.texts = [] + + result = parser._export_content(mock_doc) + assert result == "short" + + def test_fallback_skips_empty_texts(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser(export_format="markdown") + mock_doc = MagicMock() + mock_doc.export_to_markdown.return_value = "" + + empty_text = MagicMock() + empty_text.text = "" + mock_doc.texts = [empty_text] + + result = parser._export_content(mock_doc) + assert result == "" + + +# ===================================================================== +# Parse File +# ===================================================================== + + +@pytest.mark.unit +class TestDoclingParserParseFile: + + def test_parse_file_success(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser() + + mock_converter = MagicMock() + mock_result = MagicMock() + mock_doc = MagicMock() + mock_doc.export_to_markdown.return_value = "Parsed document content" + mock_doc.texts = [] + mock_result.document = mock_doc + mock_converter.convert.return_value = mock_result + parser._converter = mock_converter + + result = parser.parse_file(Path("test.pdf")) + assert "Parsed document content" in result + + def test_parse_file_inits_converter_on_first_call(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser() + parser._converter = None + + mock_converter = MagicMock() + mock_result = MagicMock() + mock_doc = MagicMock() + mock_doc.export_to_markdown.return_value = "content" + mock_doc.texts = [] + mock_result.document = mock_doc + mock_converter.convert.return_value = mock_result + + with patch.object(parser, "_init_parser") as mock_init: + parser._converter = mock_converter + mock_init.return_value = {} + result = parser.parse_file(Path("test.pdf")) + assert "content" in result + + def test_parse_file_error_ignore(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser() + mock_converter = MagicMock() + mock_converter.convert.side_effect = Exception("Parse failed") + parser._converter = mock_converter + + result = parser.parse_file(Path("bad.pdf"), errors="ignore") + assert "Error" in result + + def test_parse_file_error_raise(self): + from application.parser.file.docling_parser import DoclingParser + + parser = DoclingParser() + mock_converter = MagicMock() + mock_converter.convert.side_effect = Exception("Parse failed") + parser._converter = mock_converter + + with pytest.raises(Exception, match="Parse failed"): + parser.parse_file(Path("bad.pdf"), errors="strict") + + +# ===================================================================== +# Subclass Init +# ===================================================================== + + +@pytest.mark.unit +class TestDoclingSubclasses: + + def test_pdf_parser_init(self): + from application.parser.file.docling_parser import DoclingPDFParser + + parser = DoclingPDFParser() + assert parser.ocr_enabled is True + assert parser.export_format == "markdown" + + def test_pdf_parser_custom_ocr(self): + from application.parser.file.docling_parser import DoclingPDFParser + + parser = DoclingPDFParser(ocr_enabled=False, force_full_page_ocr=True) + assert parser.ocr_enabled is False + assert parser.force_full_page_ocr is True + + def test_docx_parser_init(self): + from application.parser.file.docling_parser import DoclingDocxParser + + parser = DoclingDocxParser() + assert parser.export_format == "markdown" + + def test_pptx_parser_init(self): + from application.parser.file.docling_parser import DoclingPPTXParser + + parser = DoclingPPTXParser() + assert parser.export_format == "markdown" + + def test_xlsx_parser_init(self): + from application.parser.file.docling_parser import DoclingXLSXParser + + parser = DoclingXLSXParser() + assert parser.table_structure is True + + def test_html_parser_init(self): + from application.parser.file.docling_parser import DoclingHTMLParser + + parser = DoclingHTMLParser() + assert parser.export_format == "markdown" + + def test_image_parser_init(self): + from application.parser.file.docling_parser import DoclingImageParser + + parser = DoclingImageParser() + assert parser.ocr_enabled is True + assert parser.force_full_page_ocr is True + + def test_image_parser_custom(self): + from application.parser.file.docling_parser import DoclingImageParser + + parser = DoclingImageParser(ocr_enabled=False) + assert parser.ocr_enabled is False + + def test_csv_parser_init(self): + from application.parser.file.docling_parser import DoclingCSVParser + + parser = DoclingCSVParser() + assert parser.table_structure is True + + def test_markdown_parser_init(self): + from application.parser.file.docling_parser import DoclingMarkdownParser + + parser = DoclingMarkdownParser() + assert parser.export_format == "markdown" + + def test_asciidoc_parser_init(self): + from application.parser.file.docling_parser import DoclingAsciiDocParser + + parser = DoclingAsciiDocParser() + assert parser.export_format == "markdown" + + def test_vtt_parser_init(self): + from application.parser.file.docling_parser import DoclingVTTParser + + parser = DoclingVTTParser() + assert parser.export_format == "markdown" + + def test_xml_parser_init(self): + from application.parser.file.docling_parser import DoclingXMLParser + + parser = DoclingXMLParser() + assert parser.export_format == "markdown" diff --git a/tests/parser/file/test_docs_parser.py b/tests/parser/file/test_docs_parser.py index c0de52ec..c30ee40d 100644 --- a/tests/parser/file/test_docs_parser.py +++ b/tests/parser/file/test_docs_parser.py @@ -1,117 +1,187 @@ -import pytest +"""Comprehensive tests for application/parser/file/docs_parser.py + +Covers: PDFParser (init, parse with pypdf, parse as image, import error), +DocxParser (init, parse, import error). +""" + from pathlib import Path -from unittest.mock import patch, MagicMock +from unittest.mock import MagicMock, patch, mock_open + +import pytest from application.parser.file.docs_parser import PDFParser, DocxParser -@pytest.fixture -def pdf_parser(): - return PDFParser() +# ===================================================================== +# PDFParser - Init +# ===================================================================== -@pytest.fixture -def docx_parser(): - return DocxParser() +@pytest.mark.unit +class TestPDFParserInit: + + def test_init_parser(self): + parser = PDFParser() + result = parser._init_parser() + assert isinstance(result, dict) + assert result == {} + + def test_parser_config_not_set_initially(self): + parser = PDFParser() + assert not parser.parser_config_set + + def test_parser_config_set_after_init(self): + parser = PDFParser() + parser.init_parser() + assert parser.parser_config_set -def test_pdf_init_parser(): - parser = PDFParser() - assert isinstance(parser._init_parser(), dict) - assert not parser.parser_config_set - parser.init_parser() - assert parser.parser_config_set +# ===================================================================== +# PDFParser - Parse File +# ===================================================================== -def test_docx_init_parser(): - parser = DocxParser() - assert isinstance(parser._init_parser(), dict) - assert not parser.parser_config_set - parser.init_parser() - assert parser.parser_config_set +@pytest.mark.unit +class TestPDFParserParse: + + @patch("application.parser.file.docs_parser.settings") + def test_parse_with_pypdf(self, mock_settings): + mock_settings.PARSE_PDF_AS_IMAGE = False + + parser = PDFParser() + + mock_page1 = MagicMock() + mock_page1.extract_text.return_value = "Page 1 content" + mock_page2 = MagicMock() + mock_page2.extract_text.return_value = "Page 2 content" + + mock_reader = MagicMock() + mock_reader.pages = [mock_page1, mock_page2] + + with patch("application.parser.file.docs_parser.PdfReader", + create=True), \ + patch("builtins.open", mock_open()): + # Need to patch the import inside the function + import sys + mock_pypdf = MagicMock() + mock_pypdf.PdfReader = MagicMock(return_value=mock_reader) + sys.modules["pypdf"] = mock_pypdf + + try: + result = parser.parse_file(Path("test.pdf")) + assert "Page 1 content" in result + assert "Page 2 content" in result + finally: + del sys.modules["pypdf"] + + @patch("application.parser.file.docs_parser.settings") + @patch("application.parser.file.docs_parser.requests") + def test_parse_as_image(self, mock_requests, mock_settings): + mock_settings.PARSE_PDF_AS_IMAGE = True + + mock_response = MagicMock() + mock_response.json.return_value = {"markdown": "# OCR Result"} + mock_requests.post.return_value = mock_response + + parser = PDFParser() + + with patch("builtins.open", mock_open(read_data=b"fake pdf")): + result = parser.parse_file(Path("test.pdf")) + assert result == "# OCR Result" + + @patch("application.parser.file.docs_parser.settings") + def test_parse_raises_on_missing_pypdf(self, mock_settings): + mock_settings.PARSE_PDF_AS_IMAGE = False + + parser = PDFParser() + + # Simulate the import error path + original = parser.parse_file + + def mock_parse(*args, **kwargs): + raise ValueError("pypdf is required to read PDF files.") + + parser.parse_file = mock_parse + + try: + with pytest.raises(ValueError, match="pypdf is required"): + parser.parse_file(Path("test.pdf")) + finally: + parser.parse_file = original -@patch("application.parser.file.docs_parser.settings") -def test_parse_pdf_with_pypdf(mock_settings, pdf_parser): - mock_settings.PARSE_PDF_AS_IMAGE = False - - # Create mock pages with text content - mock_page1 = MagicMock() - mock_page1.extract_text.return_value = "Test PDF content page 1" - mock_page2 = MagicMock() - mock_page2.extract_text.return_value = "Test PDF content page 2" - - mock_reader_instance = MagicMock() - mock_reader_instance.pages = [mock_page1, mock_page2] - - original_parse_file = pdf_parser.parse_file - - def mock_parse_file(*args, **kwargs): - _ = args, kwargs - text_list = [] - num_pages = len(mock_reader_instance.pages) - for page_index in range(num_pages): - page = mock_reader_instance.pages[page_index] - page_text = page.extract_text() - text_list.append(page_text) - text = "\n".join(text_list) - return text - - pdf_parser.parse_file = mock_parse_file - - try: - result = pdf_parser.parse_file(Path("test.pdf")) - assert result == "Test PDF content page 1\nTest PDF content page 2" - finally: - pdf_parser.parse_file = original_parse_file +# ===================================================================== +# DocxParser - Init +# ===================================================================== -@patch("application.parser.file.docs_parser.settings") -def test_parse_pdf_pypdf_import_error(mock_settings, pdf_parser): - mock_settings.PARSE_PDF_AS_IMAGE = False +@pytest.mark.unit +class TestDocxParserInit: - original_parse_file = pdf_parser.parse_file + def test_init_parser(self): + parser = DocxParser() + result = parser._init_parser() + assert isinstance(result, dict) + assert result == {} - def mock_parse_file(*args, **kwargs): - _ = args, kwargs - raise ValueError("pypdf is required to read PDF files.") + def test_parser_config_not_set_initially(self): + parser = DocxParser() + assert not parser.parser_config_set - pdf_parser.parse_file = mock_parse_file - - try: - with pytest.raises(ValueError, match="pypdf is required to read PDF files"): - pdf_parser.parse_file(Path("test.pdf")) - finally: - pdf_parser.parse_file = original_parse_file + def test_parser_config_set_after_init(self): + parser = DocxParser() + parser.init_parser() + assert parser.parser_config_set -def test_parse_docx(docx_parser): - original_parse_file = docx_parser.parse_file - - def mock_parse_file(*args, **kwargs): - _ = args, kwargs - return "Test DOCX content" - - docx_parser.parse_file = mock_parse_file - - try: - result = docx_parser.parse_file(Path("test.docx")) - assert result == "Test DOCX content" - finally: - docx_parser.parse_file = original_parse_file +# ===================================================================== +# DocxParser - Parse File +# ===================================================================== -def test_parse_docx_import_error(docx_parser): - original_parse_file = docx_parser.parse_file +@pytest.mark.unit +class TestDocxParserParse: - def mock_parse_file(*args, **kwargs): - _ = args, kwargs - raise ValueError("docx2txt is required to read Microsoft Word files.") + def test_parse_file_success(self): + parser = DocxParser() - docx_parser.parse_file = mock_parse_file + import sys + mock_docx2txt = MagicMock() + mock_docx2txt.process.return_value = "DOCX content here" + sys.modules["docx2txt"] = mock_docx2txt - try: - with pytest.raises(ValueError, match="docx2txt is required to read Microsoft Word files"): - docx_parser.parse_file(Path("test.docx")) - finally: - docx_parser.parse_file = original_parse_file \ No newline at end of file + try: + result = parser.parse_file(Path("test.docx")) + assert result == "DOCX content here" + finally: + del sys.modules["docx2txt"] + + def test_parse_raises_on_missing_docx2txt(self): + parser = DocxParser() + + original = parser.parse_file + + def mock_parse(*args, **kwargs): + raise ValueError("docx2txt is required to read Microsoft Word files.") + + parser.parse_file = mock_parse + + try: + with pytest.raises(ValueError, match="docx2txt is required"): + parser.parse_file(Path("test.docx")) + finally: + parser.parse_file = original + + +# ===================================================================== +# BaseParser properties +# ===================================================================== + + +@pytest.mark.unit +class TestBaseParserProperties: + + def test_get_file_metadata_default(self): + parser = PDFParser() + meta = parser.get_file_metadata(Path("test.pdf")) + assert meta == {} diff --git a/tests/parser/remote/__init__.py b/tests/parser/remote/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/parser/remote/test_github_loader.py b/tests/parser/remote/test_github_loader.py index f52003c1..bd0a5fc0 100644 --- a/tests/parser/remote/test_github_loader.py +++ b/tests/parser/remote/test_github_loader.py @@ -116,6 +116,159 @@ class TestGitHubLoaderLoadData: +class TestGitHubLoaderIsTextFile: + def test_known_extension(self): + loader = GitHubLoader() + assert loader.is_text_file("app.py") is True + assert loader.is_text_file("data.json") is True + + def test_unknown_extension_with_text_mime(self): + loader = GitHubLoader() + assert loader.is_text_file("file.xml") is True + + def test_binary_file(self): + loader = GitHubLoader() + assert loader.is_text_file("image.png") is False + + @patch("application.parser.remote.github_loader.mimetypes.guess_type") + def test_mime_fallback_text(self, mock_mime): + mock_mime.return_value = ("text/plain", None) + loader = GitHubLoader() + assert loader.is_text_file("unknownfile.xyz") is True + + +class TestGitHubLoaderMakeRequest: + @patch("application.parser.remote.github_loader.requests.get") + def test_success(self, mock_get): + loader = GitHubLoader() + mock_get.return_value = make_response({"ok": True}, 200) + resp = loader._make_request("http://example.com") + assert resp.status_code == 200 + + @patch("application.parser.remote.github_loader.time.sleep") + @patch("application.parser.remote.github_loader.requests.get") + def test_rate_limit_retry(self, mock_get, mock_sleep): + loader = GitHubLoader() + rate_resp = MagicMock() + rate_resp.status_code = 403 + rate_resp.json.return_value = {"message": "API rate limit exceeded"} + rate_resp.headers = { + "X-RateLimit-Remaining": "0", + "X-RateLimit-Reset": "9999999", + } + ok_resp = make_response({"ok": True}, 200) + mock_get.side_effect = [rate_resp, ok_resp] + + resp = loader._make_request("http://example.com", max_retries=2) + assert resp.status_code == 200 + mock_sleep.assert_called_once() + + @patch("application.parser.remote.github_loader.requests.get") + def test_rate_limit_exhausted(self, mock_get): + loader = GitHubLoader() + rate_resp = MagicMock() + rate_resp.status_code = 403 + rate_resp.json.return_value = {"message": "API rate limit exceeded"} + rate_resp.headers = { + "X-RateLimit-Remaining": "0", + "X-RateLimit-Reset": "9999", + } + mock_get.return_value = rate_resp + + with pytest.raises(Exception, match="rate limit exceeded"): + loader._make_request("http://example.com", max_retries=1) + + @patch("application.parser.remote.github_loader.requests.get") + def test_403_non_rate_limit(self, mock_get): + loader = GitHubLoader() + resp = MagicMock() + resp.status_code = 403 + resp.json.return_value = {"message": "Forbidden - need auth"} + resp.headers = {"X-RateLimit-Remaining": "50", "X-RateLimit-Reset": "9999"} + mock_get.return_value = resp + + with pytest.raises(Exception, match="GitHub API error"): + loader._make_request("http://example.com", max_retries=1) + + @patch("application.parser.remote.github_loader.requests.get") + def test_other_error_raises(self, mock_get): + loader = GitHubLoader() + resp = make_response( + status_code=500, + raise_error=requests.HTTPError("Server Error"), + ) + mock_get.return_value = resp + + with pytest.raises(requests.HTTPError): + loader._make_request("http://example.com", max_retries=1) + + +class TestGitHubLoaderFetchRepoFilesErrors: + @patch("application.parser.remote.github_loader.requests.get") + def test_api_error_message_in_dict(self, mock_get): + loader = GitHubLoader() + mock_get.return_value = make_response( + {"message": "Not Found"}, 200 + ) + + with pytest.raises(Exception, match="GitHub API error"): + loader.fetch_repo_files("owner/repo") + + @patch("application.parser.remote.github_loader.requests.get") + def test_non_list_response(self, mock_get): + loader = GitHubLoader() + mock_get.return_value = make_response("not a list", 200) + + with pytest.raises(TypeError, match="Expected list"): + loader.fetch_repo_files("owner/repo") + + +class TestGitHubLoaderFetchFileContentEdgeCases: + @patch("application.parser.remote.github_loader.requests.get") + def test_empty_base64_text_returns_none(self, mock_get): + loader = GitHubLoader() + b64 = base64.b64encode(b"").decode("utf-8") + mock_get.return_value = make_response( + {"encoding": "base64", "content": b64} + ) + result = loader.fetch_file_content("owner/repo", "empty.py") + assert result is None + + @patch("application.parser.remote.github_loader.requests.get") + def test_empty_non_base64_returns_none(self, mock_get): + loader = GitHubLoader() + mock_get.return_value = make_response( + {"encoding": "none", "content": " "} + ) + result = loader.fetch_file_content("owner/repo", "empty.txt") + assert result is None + + @patch("application.parser.remote.github_loader.requests.get") + def test_decode_failure_returns_none(self, mock_get): + loader = GitHubLoader() + mock_get.return_value = make_response( + {"encoding": "base64", "content": "invalid!!base64"} + ) + result = loader.fetch_file_content("owner/repo", "broken.py") + assert result is None + + +class TestGitHubLoaderLoadDataSkipsNone: + def test_skips_binary_files(self, monkeypatch): + loader = GitHubLoader() + monkeypatch.setattr( + loader, "fetch_repo_files", lambda repo, path="": ["a.py", "b.png"] + ) + + def fake_content(repo, fp): + return "code" if fp == "a.py" else None + + monkeypatch.setattr(loader, "fetch_file_content", fake_content) + docs = loader.load_data("https://github.com/o/r") + assert len(docs) == 1 + assert docs[0].doc_id == "a.py" + + class TestGitHubLoaderRobustness: @patch("application.parser.remote.github_loader.requests.get") def test_fetch_repo_files_non_json_raises(self, mock_get): diff --git a/tests/parser/remote/test_sitemap_loader.py b/tests/parser/remote/test_sitemap_loader.py new file mode 100644 index 00000000..b593bd35 --- /dev/null +++ b/tests/parser/remote/test_sitemap_loader.py @@ -0,0 +1,306 @@ +"""Comprehensive tests for application/parser/remote/sitemap_loader.py + +Covers: SitemapLoader (init, load_data, _extract_urls, _is_sitemap, +_parse_sitemap, URL validation, error handling). +""" + +from unittest.mock import MagicMock, patch + +import pytest +import requests + +from application.parser.remote.sitemap_loader import SitemapLoader + + +# ===================================================================== +# SitemapLoader - Init +# ===================================================================== + + +@pytest.mark.unit +class TestSitemapLoaderInit: + + def test_default_limit(self): + loader = SitemapLoader() + assert loader.limit == 20 + + def test_custom_limit(self): + loader = SitemapLoader(limit=5) + assert loader.limit == 5 + + def test_has_loader_class(self): + loader = SitemapLoader() + assert loader.loader is not None + + +# ===================================================================== +# _is_sitemap +# ===================================================================== + + +@pytest.mark.unit +class TestIsSitemap: + + def test_xml_content_type(self): + loader = SitemapLoader() + response = MagicMock() + response.headers = {"Content-Type": "application/xml"} + response.url = "https://example.com/sitemap.xml" + response.text = "" + assert loader._is_sitemap(response) is True + + def test_xml_url_extension(self): + loader = SitemapLoader() + response = MagicMock() + response.headers = {"Content-Type": "text/html"} + response.url = "https://example.com/sitemap.xml" + response.text = "" + assert loader._is_sitemap(response) is True + + def test_sitemapindex_in_body(self): + loader = SitemapLoader() + response = MagicMock() + response.headers = {"Content-Type": "text/html"} + response.url = "https://example.com/sitemap" + response.text = "" + assert loader._is_sitemap(response) is True + + def test_urlset_in_body(self): + loader = SitemapLoader() + response = MagicMock() + response.headers = {"Content-Type": "text/html"} + response.url = "https://example.com/page" + response.text = "" + assert loader._is_sitemap(response) is True + + def test_regular_page(self): + loader = SitemapLoader() + response = MagicMock() + response.headers = {"Content-Type": "text/html"} + response.url = "https://example.com/about" + response.text = "About us" + assert loader._is_sitemap(response) is False + + def test_text_xml_content_type(self): + loader = SitemapLoader() + response = MagicMock() + response.headers = {"Content-Type": "text/xml; charset=utf-8"} + response.url = "https://example.com/feed" + response.text = "" + assert loader._is_sitemap(response) is True + + +# ===================================================================== +# _parse_sitemap +# ===================================================================== + + +@pytest.mark.unit +class TestParseSitemap: + + def test_parse_basic_sitemap(self): + loader = SitemapLoader() + sitemap_xml = b""" + + https://example.com/page1 + https://example.com/page2 + """ + + urls = loader._parse_sitemap(sitemap_xml) + assert "https://example.com/page1" in urls + assert "https://example.com/page2" in urls + + def test_parse_nested_sitemap(self): + loader = SitemapLoader() + + parent_xml = b""" + + + https://example.com/sitemap-child.xml + + """ + + with patch.object( + loader, "_extract_urls", + return_value=["https://example.com/page1"] + ): + urls = loader._parse_sitemap(parent_xml) + assert "https://example.com/page1" in urls + + def test_parse_empty_sitemap(self): + loader = SitemapLoader() + sitemap_xml = b""" + + """ + + urls = loader._parse_sitemap(sitemap_xml) + assert urls == [] + + +# ===================================================================== +# _extract_urls +# ===================================================================== + + +@pytest.mark.unit +class TestExtractUrls: + + @patch("application.parser.remote.sitemap_loader.validate_url") + @patch("application.parser.remote.sitemap_loader.requests.get") + def test_extract_urls_from_sitemap(self, mock_get, mock_validate): + loader = SitemapLoader() + + response = MagicMock() + response.headers = {"Content-Type": "application/xml"} + response.url = "https://example.com/sitemap.xml" + response.text = "https://example.com/p" + response.content = b""" + + https://example.com/p + """ + mock_get.return_value = response + + urls = loader._extract_urls("https://example.com/sitemap.xml") + assert "https://example.com/p" in urls + + @patch("application.parser.remote.sitemap_loader.validate_url") + @patch("application.parser.remote.sitemap_loader.requests.get") + def test_extract_urls_not_sitemap(self, mock_get, mock_validate): + loader = SitemapLoader() + + response = MagicMock() + response.headers = {"Content-Type": "text/html"} + response.url = "https://example.com/page" + response.text = "Normal page" + mock_get.return_value = response + + urls = loader._extract_urls("https://example.com/page") + assert urls == ["https://example.com/page"] + + @patch("application.parser.remote.sitemap_loader.validate_url") + @patch("application.parser.remote.sitemap_loader.requests.get") + def test_extract_urls_http_error(self, mock_get, mock_validate): + loader = SitemapLoader() + mock_get.side_effect = requests.exceptions.HTTPError("404") + + urls = loader._extract_urls("https://example.com/missing") + assert urls == [] + + @patch("application.parser.remote.sitemap_loader.validate_url") + @patch("application.parser.remote.sitemap_loader.requests.get") + def test_extract_urls_connection_error(self, mock_get, mock_validate): + loader = SitemapLoader() + mock_get.side_effect = requests.exceptions.ConnectionError() + + urls = loader._extract_urls("https://example.com/bad") + assert urls == [] + + def test_extract_urls_ssrf_blocked(self): + from application.core.url_validation import SSRFError + + loader = SitemapLoader() + + with patch( + "application.parser.remote.sitemap_loader.validate_url", + side_effect=SSRFError("blocked"), + ): + urls = loader._extract_urls("http://169.254.169.254/") + assert urls == [] + + +# ===================================================================== +# load_data +# ===================================================================== + + +@pytest.mark.unit +class TestSitemapLoaderLoadData: + + @patch("application.parser.remote.sitemap_loader.validate_url") + def test_load_data_success(self, mock_validate): + loader = SitemapLoader(limit=10) + + mock_doc = MagicMock() + mock_loader_instance = MagicMock() + mock_loader_instance.load.return_value = [mock_doc] + loader.loader = MagicMock(return_value=mock_loader_instance) + + with patch.object( + loader, "_extract_urls", + return_value=["https://example.com/page1"] + ): + docs = loader.load_data("https://example.com/sitemap.xml") + assert len(docs) == 1 + + @patch("application.parser.remote.sitemap_loader.validate_url") + def test_load_data_no_urls(self, mock_validate): + loader = SitemapLoader() + + with patch.object(loader, "_extract_urls", return_value=[]): + docs = loader.load_data("https://example.com/empty-sitemap.xml") + assert docs == [] + + @patch("application.parser.remote.sitemap_loader.validate_url") + def test_load_data_list_input(self, mock_validate): + loader = SitemapLoader() + + mock_loader_instance = MagicMock() + mock_loader_instance.load.return_value = [MagicMock()] + loader.loader = MagicMock(return_value=mock_loader_instance) + + with patch.object( + loader, "_extract_urls", + return_value=["https://example.com/page1"] + ): + docs = loader.load_data(["https://example.com/sitemap.xml"]) + assert len(docs) == 1 + + @patch("application.parser.remote.sitemap_loader.validate_url") + def test_load_data_respects_limit(self, mock_validate): + loader = SitemapLoader(limit=2) + + mock_loader_instance = MagicMock() + mock_loader_instance.load.return_value = [MagicMock()] + loader.loader = MagicMock(return_value=mock_loader_instance) + + urls = [f"https://example.com/page{i}" for i in range(10)] + with patch.object(loader, "_extract_urls", return_value=urls): + docs = loader.load_data("https://example.com/sitemap.xml") + assert len(docs) == 2 + + @patch("application.parser.remote.sitemap_loader.validate_url") + def test_load_data_handles_url_error(self, mock_validate): + loader = SitemapLoader() + loader.loader = MagicMock(side_effect=Exception("Load failed")) + + with patch.object( + loader, "_extract_urls", + return_value=["https://example.com/broken"] + ): + docs = loader.load_data("https://example.com/sitemap.xml") + assert docs == [] + + def test_load_data_ssrf_blocked(self): + from application.core.url_validation import SSRFError + + loader = SitemapLoader() + + with patch( + "application.parser.remote.sitemap_loader.validate_url", + side_effect=SSRFError("blocked"), + ): + docs = loader.load_data("http://169.254.169.254/") + assert docs == [] + + @patch("application.parser.remote.sitemap_loader.validate_url") + def test_load_data_no_limit(self, mock_validate): + loader = SitemapLoader(limit=None) + + mock_loader_instance = MagicMock() + mock_loader_instance.load.return_value = [MagicMock()] + loader.loader = MagicMock(return_value=mock_loader_instance) + + urls = [f"https://example.com/page{i}" for i in range(5)] + with patch.object(loader, "_extract_urls", return_value=urls): + docs = loader.load_data("https://example.com/sitemap.xml") + assert len(docs) == 5 diff --git a/tests/parser/test_chunking.py b/tests/parser/test_chunking.py new file mode 100644 index 00000000..fd5c0951 --- /dev/null +++ b/tests/parser/test_chunking.py @@ -0,0 +1,279 @@ +"""Comprehensive tests for application/parser/chunking.py + +Covers: Chunker (init, separate_header_and_body, split_document, +classic_chunk, chunk), edge cases, token counting. +""" + +import pytest + +from application.parser.chunking import Chunker +from application.parser.schema.base import Document + + +# ===================================================================== +# Chunker - Init +# ===================================================================== + + +@pytest.mark.unit +class TestChunkerInit: + + def test_default_init(self): + chunker = Chunker() + assert chunker.chunking_strategy == "classic_chunk" + assert chunker.max_tokens == 2000 + assert chunker.min_tokens == 150 + assert chunker.duplicate_headers is False + + def test_custom_init(self): + chunker = Chunker( + chunking_strategy="classic_chunk", + max_tokens=1000, + min_tokens=50, + duplicate_headers=True, + ) + assert chunker.max_tokens == 1000 + assert chunker.min_tokens == 50 + assert chunker.duplicate_headers is True + + def test_invalid_strategy_raises(self): + with pytest.raises(ValueError, match="Unsupported chunking strategy"): + Chunker(chunking_strategy="unknown_strategy") + + +# ===================================================================== +# Separate Header and Body +# ===================================================================== + + +@pytest.mark.unit +class TestSeparateHeaderAndBody: + + def test_with_header(self): + chunker = Chunker() + text = "line1\nline2\nline3\nbody content here" + header, body = chunker.separate_header_and_body(text) + assert "line1" in header + assert "line2" in header + assert "line3" in header + assert "body content here" in body + + def test_without_header(self): + chunker = Chunker() + text = "short" + header, body = chunker.separate_header_and_body(text) + assert header == "" + assert body == "short" + + def test_empty_text(self): + chunker = Chunker() + header, body = chunker.separate_header_and_body("") + assert header == "" + assert body == "" + + def test_exactly_three_lines(self): + chunker = Chunker() + text = "line1\nline2\nline3\n" + header, body = chunker.separate_header_and_body(text) + assert header == "line1\nline2\nline3\n" + assert body == "" + + +# ===================================================================== +# Split Document +# ===================================================================== + + +@pytest.mark.unit +class TestSplitDocument: + + def test_split_large_document(self): + chunker = Chunker(max_tokens=50, min_tokens=5) + long_text = "word " * 200 + doc = Document(text=long_text, doc_id="doc1") + + result = chunker.split_document(doc) + assert len(result) > 1 + for split_doc in result: + assert split_doc.doc_id.startswith("doc1-") + assert split_doc.extra_info is not None + assert "token_count" in split_doc.extra_info + + def test_split_preserves_header_on_first(self): + chunker = Chunker(max_tokens=50, min_tokens=5, duplicate_headers=False) + text = "h1\nh2\nh3\n" + "word " * 200 + doc = Document(text=text, doc_id="doc1") + + result = chunker.split_document(doc) + assert len(result) > 1 + # First chunk should contain header + assert "h1" in result[0].text + + def test_split_duplicates_header(self): + chunker = Chunker(max_tokens=50, min_tokens=5, duplicate_headers=True) + text = "h1\nh2\nh3\n" + "word " * 200 + doc = Document(text=text, doc_id="doc1") + + result = chunker.split_document(doc) + assert len(result) > 1 + # First chunk should contain header + assert "h1" in result[0].text + + def test_split_preserves_embedding(self): + chunker = Chunker(max_tokens=50, min_tokens=5) + doc = Document( + text="word " * 200, + doc_id="doc1", + embedding=[0.1, 0.2], + ) + + result = chunker.split_document(doc) + for split_doc in result: + assert split_doc.embedding == [0.1, 0.2] + + def test_split_preserves_extra_info(self): + chunker = Chunker(max_tokens=50, min_tokens=5) + doc = Document( + text="word " * 200, + doc_id="doc1", + extra_info={"source": "test"}, + ) + + result = chunker.split_document(doc) + for split_doc in result: + assert split_doc.extra_info["source"] == "test" + assert "token_count" in split_doc.extra_info + + +# ===================================================================== +# Classic Chunk +# ===================================================================== + + +@pytest.mark.unit +class TestClassicChunk: + + def test_small_doc_passes_through(self): + chunker = Chunker(max_tokens=2000, min_tokens=1) + doc = Document(text="Short text", doc_id="d1") + + result = chunker.classic_chunk([doc]) + assert len(result) == 1 + assert result[0].extra_info is not None + assert "token_count" in result[0].extra_info + + def test_large_doc_gets_split(self): + chunker = Chunker(max_tokens=50, min_tokens=5) + doc = Document(text="word " * 200, doc_id="d1") + + result = chunker.classic_chunk([doc]) + assert len(result) > 1 + + def test_medium_doc_within_range(self): + chunker = Chunker(max_tokens=2000, min_tokens=5) + doc = Document(text="Hello " * 50, doc_id="d1") + + result = chunker.classic_chunk([doc]) + assert len(result) == 1 + + def test_multiple_docs(self): + chunker = Chunker(max_tokens=2000, min_tokens=1) + docs = [ + Document(text="Doc 1 content", doc_id="d1"), + Document(text="Doc 2 content", doc_id="d2"), + ] + + result = chunker.classic_chunk(docs) + assert len(result) == 2 + + def test_empty_docs_list(self): + chunker = Chunker() + result = chunker.classic_chunk([]) + assert result == [] + + def test_very_small_doc_below_min(self): + chunker = Chunker(max_tokens=2000, min_tokens=500) + doc = Document(text="tiny", doc_id="d1") + + result = chunker.classic_chunk([doc]) + assert len(result) == 1 + assert result[0].extra_info["token_count"] < 500 + + def test_existing_extra_info_preserved(self): + chunker = Chunker(max_tokens=2000, min_tokens=1) + doc = Document( + text="Hello world", + doc_id="d1", + extra_info={"source": "test"}, + ) + + result = chunker.classic_chunk([doc]) + assert result[0].extra_info["source"] == "test" + assert "token_count" in result[0].extra_info + + def test_none_extra_info_initialized(self): + chunker = Chunker(max_tokens=2000, min_tokens=1) + doc = Document(text="Hello", doc_id="d1", extra_info=None) + + result = chunker.classic_chunk([doc]) + assert result[0].extra_info is not None + assert "token_count" in result[0].extra_info + + +# ===================================================================== +# Chunk (dispatcher) +# ===================================================================== + + +@pytest.mark.unit +class TestChunkDispatcher: + + def test_dispatch_classic_chunk(self): + chunker = Chunker(chunking_strategy="classic_chunk") + doc = Document(text="content", doc_id="d1") + + result = chunker.chunk([doc]) + assert len(result) == 1 + + def test_dispatch_unknown_raises(self): + chunker = Chunker() + chunker.chunking_strategy = "nonexistent" + + with pytest.raises(ValueError, match="Unsupported chunking strategy"): + chunker.chunk([Document(text="x", doc_id="d")]) + + +# ===================================================================== +# Integration-like test +# ===================================================================== + + +@pytest.mark.unit +class TestChunkerIntegration: + + def test_mixed_document_sizes(self): + chunker = Chunker(max_tokens=50, min_tokens=5) + docs = [ + Document(text="small text", doc_id="small"), + Document(text="word " * 200, doc_id="large"), + Document(text="medium " * 20, doc_id="medium"), + ] + + result = chunker.chunk(docs) + # Small and medium should pass through, large should be split + assert len(result) >= 3 + doc_ids = [d.doc_id for d in result] + assert "small" in doc_ids + + def test_all_chunks_have_token_counts(self): + chunker = Chunker(max_tokens=50, min_tokens=1) + docs = [ + Document(text="word " * 200, doc_id="big"), + Document(text="tiny", doc_id="small"), + ] + + result = chunker.chunk(docs) + for doc in result: + assert doc.extra_info is not None + assert "token_count" in doc.extra_info + assert doc.extra_info["token_count"] > 0 diff --git a/tests/parser/test_schema.py b/tests/parser/test_schema.py new file mode 100644 index 00000000..a2f50531 --- /dev/null +++ b/tests/parser/test_schema.py @@ -0,0 +1,58 @@ +import pytest + +from application.parser.schema.schema import BaseDocument + + +class ConcreteDoc(BaseDocument): + @classmethod + def get_type(cls) -> str: + return "test" + + +@pytest.mark.unit +class TestBaseDocument: + + def test_get_text(self): + doc = ConcreteDoc(text="hello") + assert doc.get_text() == "hello" + + def test_get_text_raises_when_none(self): + doc = ConcreteDoc() + with pytest.raises(ValueError, match="text field not set"): + doc.get_text() + + def test_get_doc_id(self): + doc = ConcreteDoc(text="x", doc_id="doc1") + assert doc.get_doc_id() == "doc1" + + def test_get_doc_id_raises_when_none(self): + doc = ConcreteDoc(text="x") + with pytest.raises(ValueError, match="doc_id not set"): + doc.get_doc_id() + + def test_is_doc_id_none(self): + doc = ConcreteDoc(text="x") + assert doc.is_doc_id_none is True + + def test_is_doc_id_not_none(self): + doc = ConcreteDoc(text="x", doc_id="y") + assert doc.is_doc_id_none is False + + def test_get_embedding(self): + doc = ConcreteDoc(text="x", embedding=[1.0, 2.0]) + assert doc.get_embedding() == [1.0, 2.0] + + def test_get_embedding_raises_when_none(self): + doc = ConcreteDoc(text="x") + with pytest.raises(ValueError, match="embedding not set"): + doc.get_embedding() + + def test_extra_info_str(self): + doc = ConcreteDoc(text="x", extra_info={"key": "value", "num": 42}) + result = doc.extra_info_str + assert "key: value" in result + assert "num: 42" in result + + def test_extra_info_str_none(self): + doc = ConcreteDoc(text="x") + assert doc.extra_info_str is None diff --git a/tests/security/__init__.py b/tests/security/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/security/test_encryption.py b/tests/security/test_encryption.py index 2d7e4efd..e6c28eec 100644 --- a/tests/security/test_encryption.py +++ b/tests/security/test_encryption.py @@ -97,3 +97,83 @@ def test_pad_and_unpad_are_inverse(): assert len(padded) % 16 == 0 assert encryption._unpad_data(padded) == original + + +@pytest.mark.unit +def test_pad_data_exact_block_size(): + # When input is exactly 16 bytes, a full block of padding is added + original = b"0123456789abcdef" + assert len(original) == 16 + + padded = encryption._pad_data(original) + + # Should be 32 bytes (16 + 16 padding) + assert len(padded) == 32 + assert encryption._unpad_data(padded) == original + + +@pytest.mark.unit +def test_pad_data_various_sizes(): + for size in range(1, 33): + data = b"x" * size + padded = encryption._pad_data(data) + assert len(padded) % 16 == 0 + assert encryption._unpad_data(padded) == data + + +@pytest.mark.unit +def test_encrypt_decrypt_complex_credentials(monkeypatch): + monkeypatch.setattr(encryption.settings, "ENCRYPTION_SECRET_KEY", "complex-secret") + + credentials = { + "token": "abc123", + "refresh": "xyz789", + "nested": {"key": "value"}, + "list_field": [1, 2, 3], + "unicode": "\u4f60\u597d\u4e16\u754c", + } + + encrypted = encryption.encrypt_credentials(credentials, "user-456") + decrypted = encryption.decrypt_credentials(encrypted, "user-456") + + assert decrypted == credentials + + +@pytest.mark.unit +def test_decrypt_with_wrong_user_returns_empty(monkeypatch): + monkeypatch.setattr(encryption.settings, "ENCRYPTION_SECRET_KEY", "test-secret") + + credentials = {"token": "abc123"} + encrypted = encryption.encrypt_credentials(credentials, "user-1") + + # Decrypting with wrong user should fail gracefully + result = encryption.decrypt_credentials(encrypted, "user-2") + assert result == {} + + +@pytest.mark.unit +def test_decrypt_with_wrong_secret_returns_empty(monkeypatch): + monkeypatch.setattr(encryption.settings, "ENCRYPTION_SECRET_KEY", "secret-1") + credentials = {"token": "abc123"} + encrypted = encryption.encrypt_credentials(credentials, "user-1") + + # Change the secret key + monkeypatch.setattr(encryption.settings, "ENCRYPTION_SECRET_KEY", "secret-2") + result = encryption.decrypt_credentials(encrypted, "user-1") + assert result == {} + + +@pytest.mark.unit +def test_encrypt_credentials_empty_dict(monkeypatch): + monkeypatch.setattr(encryption.settings, "ENCRYPTION_SECRET_KEY", "test-secret") + assert encryption.encrypt_credentials({}, "user-1") == "" + + +@pytest.mark.unit +def test_decrypt_credentials_truncated_payload(monkeypatch): + monkeypatch.setattr(encryption.settings, "ENCRYPTION_SECRET_KEY", "test-secret") + # base64 of only 10 bytes - not enough for salt+iv + import base64 + + short = base64.b64encode(b"0123456789").decode() + assert encryption.decrypt_credentials(short, "user-1") == {} diff --git a/tests/storage/__init__.py b/tests/storage/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/storage/test_s3_storage.py b/tests/storage/test_s3_storage.py index da2afac2..967da1f5 100644 --- a/tests/storage/test_s3_storage.py +++ b/tests/storage/test_s3_storage.py @@ -413,3 +413,100 @@ class TestS3StorageRemoveDirectory: result = s3_storage.remove_directory(directory) assert result is False + + @pytest.mark.unit + def test_remove_directory_returns_false_on_delete_errors( + self, s3_storage, mock_boto3_client + ): + """Should return False when delete_objects response contains Errors.""" + directory = "documents/" + + paginator_mock = MagicMock() + mock_boto3_client.get_paginator.return_value = paginator_mock + paginator_mock.paginate.return_value = [ + {"Contents": [{"Key": "documents/file1.txt"}]} + ] + + mock_boto3_client.delete_objects.return_value = { + "Errors": [{"Key": "documents/file1.txt", "Code": "InternalError"}] + } + + result = s3_storage.remove_directory(directory) + + assert result is False + + +class TestS3StorageDirectorySlashHandling: + """Test that directories without trailing slashes get them added.""" + + @pytest.mark.unit + def test_list_files_adds_trailing_slash(self, s3_storage, mock_boto3_client): + """Should add trailing slash when listing directory without one.""" + paginator_mock = MagicMock() + mock_boto3_client.get_paginator.return_value = paginator_mock + paginator_mock.paginate.return_value = [{}] + + s3_storage.list_files("documents") + + paginator_mock.paginate.assert_called_once_with( + Bucket="test-bucket", Prefix="documents/" + ) + + @pytest.mark.unit + def test_list_files_empty_directory_string(self, s3_storage, mock_boto3_client): + """Empty string directory should not get a slash added.""" + paginator_mock = MagicMock() + mock_boto3_client.get_paginator.return_value = paginator_mock + paginator_mock.paginate.return_value = [{}] + + s3_storage.list_files("") + + paginator_mock.paginate.assert_called_once_with( + Bucket="test-bucket", Prefix="" + ) + + @pytest.mark.unit + def test_is_directory_adds_trailing_slash(self, s3_storage, mock_boto3_client): + """Should add trailing slash for is_directory check.""" + mock_boto3_client.list_objects_v2.return_value = {} + + s3_storage.is_directory("docs") + + mock_boto3_client.list_objects_v2.assert_called_once_with( + Bucket="test-bucket", Prefix="docs/", MaxKeys=1 + ) + + @pytest.mark.unit + def test_remove_directory_adds_trailing_slash(self, s3_storage, mock_boto3_client): + """Should add trailing slash for remove_directory.""" + paginator_mock = MagicMock() + mock_boto3_client.get_paginator.return_value = paginator_mock + paginator_mock.paginate.return_value = [{}] + + s3_storage.remove_directory("docs") + + paginator_mock.paginate.assert_called_once() + call_kwargs = paginator_mock.paginate.call_args[1] + assert call_kwargs["Prefix"] == "docs/" + + +class TestS3StorageProcessFileError: + """Test error handling in process_file.""" + + @pytest.mark.unit + def test_process_file_propagates_processor_error( + self, s3_storage, mock_boto3_client + ): + """Should propagate errors from the processor function.""" + path = "documents/test.txt" + mock_boto3_client.head_object.return_value = {} + + with patch("tempfile.NamedTemporaryFile") as mock_temp: + mock_file = MagicMock() + mock_file.name = "/tmp/test_file" + mock_temp.return_value.__enter__.return_value = mock_file + + processor_func = MagicMock(side_effect=RuntimeError("Process failed")) + + with pytest.raises(RuntimeError, match="Process failed"): + s3_storage.process_file(path, processor_func) diff --git a/tests/stt/__init__.py b/tests/stt/__init__.py new file mode 100644 index 00000000..e69de29b diff --git a/tests/stt/test_faster_whisper.py b/tests/stt/test_faster_whisper.py new file mode 100644 index 00000000..969c6bd8 --- /dev/null +++ b/tests/stt/test_faster_whisper.py @@ -0,0 +1,234 @@ +"""Tests for application/stt/faster_whisper_stt.py""" + +from pathlib import Path +from unittest.mock import MagicMock, patch + +import pytest + +from application.stt.faster_whisper_stt import FasterWhisperSTT + + +@pytest.mark.unit +class TestFasterWhisperSTTInit: + + def test_init_defaults(self): + stt = FasterWhisperSTT() + assert stt.model_size == "base" + assert stt.device == "auto" + assert stt.compute_type == "int8" + assert stt._model is None + + def test_init_custom_params(self): + stt = FasterWhisperSTT( + model_size="large-v2", + device="cuda", + compute_type="float16", + ) + assert stt.model_size == "large-v2" + assert stt.device == "cuda" + assert stt.compute_type == "float16" + + +@pytest.mark.unit +class TestFasterWhisperSTTGetModel: + + def test_get_model_lazy_init(self): + stt = FasterWhisperSTT() + + mock_whisper_model = MagicMock() + mock_module = MagicMock() + mock_module.WhisperModel.return_value = mock_whisper_model + + with patch.dict("sys.modules", {"faster_whisper": mock_module}): + model = stt._get_model() + + assert model is mock_whisper_model + mock_module.WhisperModel.assert_called_once_with( + "base", + device="auto", + compute_type="int8", + ) + + def test_get_model_caches(self): + stt = FasterWhisperSTT() + + mock_whisper_model = MagicMock() + mock_module = MagicMock() + mock_module.WhisperModel.return_value = mock_whisper_model + + with patch.dict("sys.modules", {"faster_whisper": mock_module}): + model1 = stt._get_model() + model2 = stt._get_model() + + assert model1 is model2 + assert mock_module.WhisperModel.call_count == 1 + + def test_get_model_raises_import_error(self): + stt = FasterWhisperSTT() + + with patch.dict("sys.modules", {"faster_whisper": None}): + with pytest.raises(ImportError, match="faster-whisper is required"): + stt._get_model() + + +@pytest.mark.unit +class TestFasterWhisperSTTTranscribe: + + def _make_stt_with_mock_model(self): + stt = FasterWhisperSTT() + mock_model = MagicMock() + stt._model = mock_model + return stt, mock_model + + def test_transcribe_basic(self): + stt, mock_model = self._make_stt_with_mock_model() + + seg1 = MagicMock() + seg1.text = " Hello world " + seg1.start = 0.0 + seg1.end = 1.5 + + seg2 = MagicMock() + seg2.text = " How are you " + seg2.start = 1.5 + seg2.end = 3.0 + + info = MagicMock() + info.language = "en" + info.duration = 3.0 + + mock_model.transcribe.return_value = (iter([seg1, seg2]), info) + + result = stt.transcribe(Path("/tmp/audio.wav")) + + assert result["text"] == "Hello world How are you" + assert result["language"] == "en" + assert result["duration_s"] == 3.0 + assert result["segments"] == [] # timestamps=False by default + assert result["provider"] == "faster_whisper" + + mock_model.transcribe.assert_called_once_with( + "/tmp/audio.wav", + language=None, + word_timestamps=False, + ) + + def test_transcribe_with_language(self): + stt, mock_model = self._make_stt_with_mock_model() + + info = MagicMock() + info.language = "fr" + info.duration = 1.0 + + mock_model.transcribe.return_value = (iter([]), info) + + result = stt.transcribe(Path("/tmp/audio.wav"), language="fr") + + mock_model.transcribe.assert_called_once_with( + "/tmp/audio.wav", + language="fr", + word_timestamps=False, + ) + assert result["language"] == "fr" + + def test_transcribe_with_timestamps(self): + stt, mock_model = self._make_stt_with_mock_model() + + seg = MagicMock() + seg.text = " Hello " + seg.start = 0.0 + seg.end = 1.0 + + info = MagicMock() + info.language = "en" + info.duration = 1.0 + + mock_model.transcribe.return_value = (iter([seg]), info) + + result = stt.transcribe(Path("/tmp/audio.wav"), timestamps=True) + + assert len(result["segments"]) == 1 + assert result["segments"][0]["start"] == 0.0 + assert result["segments"][0]["end"] == 1.0 + assert result["segments"][0]["text"] == "Hello" + + mock_model.transcribe.assert_called_once_with( + "/tmp/audio.wav", + language=None, + word_timestamps=True, + ) + + def test_transcribe_empty_segments(self): + stt, mock_model = self._make_stt_with_mock_model() + + info = MagicMock() + info.language = "en" + info.duration = 0.0 + + mock_model.transcribe.return_value = (iter([]), info) + + result = stt.transcribe(Path("/tmp/audio.wav")) + + assert result["text"] == "" + assert result["segments"] == [] + + def test_transcribe_segment_with_empty_text(self): + stt, mock_model = self._make_stt_with_mock_model() + + seg = MagicMock() + seg.text = " " + seg.start = 0.0 + seg.end = 0.5 + + info = MagicMock() + info.language = "en" + info.duration = 0.5 + + mock_model.transcribe.return_value = (iter([seg]), info) + + result = stt.transcribe(Path("/tmp/audio.wav")) + + # Empty text stripped should not be included in text_parts + assert result["text"] == "" + + def test_transcribe_diarize_is_ignored(self): + stt, mock_model = self._make_stt_with_mock_model() + + info = MagicMock() + info.language = "en" + info.duration = 1.0 + + mock_model.transcribe.return_value = (iter([]), info) + + # diarize param should be accepted but ignored + result = stt.transcribe( + Path("/tmp/audio.wav"), + diarize=True, + ) + + assert result["provider"] == "faster_whisper" + + def test_transcribe_missing_attrs_use_none(self): + stt, mock_model = self._make_stt_with_mock_model() + + seg = MagicMock(spec=[]) # No attributes + seg.text = "" # Override to avoid AttributeError on text + + # Create a segment that uses getattr fallbacks + class MinimalSegment: + pass + + minimal = MinimalSegment() + + info_cls = type("Info", (), {})() + + mock_model.transcribe.return_value = (iter([minimal]), info_cls) + + result = stt.transcribe(Path("/tmp/audio.wav"), timestamps=True) + + assert result["language"] is None + assert result["duration_s"] is None + # Segment should have None for start/end + assert len(result["segments"]) == 1 + assert result["segments"][0]["start"] is None + assert result["segments"][0]["end"] is None diff --git a/tests/stt/test_live_session.py b/tests/stt/test_live_session.py index 67b7b26b..e0c5315a 100644 --- a/tests/stt/test_live_session.py +++ b/tests/stt/test_live_session.py @@ -1,8 +1,22 @@ +import json +from unittest.mock import MagicMock + +import pytest + from application.stt.live_session import ( apply_live_stt_hypothesis, + create_live_stt_session, + delete_live_stt_session, finalize_live_stt_session, + get_live_stt_session_key, get_live_stt_transcript_text, + join_transcript_parts, + load_live_stt_session, + normalize_transcript_text, + save_live_stt_session, strip_committed_prefix, + LIVE_STT_SESSION_PREFIX, + LIVE_STT_SESSION_TTL_SECONDS, ) @@ -140,3 +154,241 @@ def test_apply_live_stt_hypothesis_rejects_older_chunks(): assert "older" in str(exc) else: raise AssertionError("Expected older chunk to raise ValueError") + + +# ── normalize_transcript_text ─────────────────────────────────────────────── + + +def test_normalize_transcript_text_strips_and_collapses_whitespace(): + assert normalize_transcript_text(" hello world ") == "hello world" + + +def test_normalize_transcript_text_empty(): + assert normalize_transcript_text("") == "" + + +def test_normalize_transcript_text_none(): + assert normalize_transcript_text(None) == "" + + +def test_normalize_transcript_text_tabs_and_newlines(): + assert normalize_transcript_text("hello\t\nworld") == "hello world" + + +# ── join_transcript_parts ─────────────────────────────────────────────────── + + +def test_join_transcript_parts_multiple(): + assert join_transcript_parts("hello", "world") == "hello world" + + +def test_join_transcript_parts_empty_parts(): + assert join_transcript_parts("hello", "", "world") == "hello world" + + +def test_join_transcript_parts_all_empty(): + assert join_transcript_parts("", "", "") == "" + + +def test_join_transcript_parts_single(): + assert join_transcript_parts("hello") == "hello" + + +def test_join_transcript_parts_whitespace_only(): + assert join_transcript_parts(" ", " hello ") == "hello" + + +# ── create_live_stt_session ───────────────────────────────────────────────── + + +def test_create_live_stt_session_basic(): + session = create_live_stt_session("user1", language="en") + assert session["user"] == "user1" + assert session["language"] == "en" + assert session["committed_text"] == "" + assert session["mutable_text"] == "" + assert session["previous_hypothesis"] == "" + assert session["latest_hypothesis"] == "" + assert session["last_chunk_index"] == -1 + assert "session_id" in session + assert len(session["session_id"]) == 36 # UUID format + + +def test_create_live_stt_session_no_language(): + session = create_live_stt_session("user1") + assert session["language"] is None + + +# ── get_live_stt_session_key ──────────────────────────────────────────────── + + +def test_get_live_stt_session_key(): + key = get_live_stt_session_key("abc-123") + assert key == f"{LIVE_STT_SESSION_PREFIX}abc-123" + + +# ── save_live_stt_session ────────────────────────────────────────────────── + + +def test_save_live_stt_session(): + mock_redis = MagicMock() + session_state = { + "session_id": "test-session-id", + "user": "user1", + "committed_text": "hello", + "mutable_text": "world", + } + + save_live_stt_session(mock_redis, session_state) + + expected_key = f"{LIVE_STT_SESSION_PREFIX}test-session-id" + mock_redis.setex.assert_called_once_with( + expected_key, + LIVE_STT_SESSION_TTL_SECONDS, + json.dumps(session_state), + ) + + +# ── load_live_stt_session ────────────────────────────────────────────────── + + +def test_load_live_stt_session_found(): + mock_redis = MagicMock() + session_data = {"session_id": "test-id", "committed_text": "hello"} + mock_redis.get.return_value = json.dumps(session_data).encode("utf-8") + + result = load_live_stt_session(mock_redis, "test-id") + + assert result == session_data + + +def test_load_live_stt_session_not_found(): + mock_redis = MagicMock() + mock_redis.get.return_value = None + + result = load_live_stt_session(mock_redis, "nonexistent") + + assert result is None + + +def test_load_live_stt_session_string_response(): + mock_redis = MagicMock() + session_data = {"session_id": "test-id"} + # Some redis clients return strings instead of bytes + mock_redis.get.return_value = json.dumps(session_data) + + result = load_live_stt_session(mock_redis, "test-id") + + assert result == session_data + + +# ── delete_live_stt_session ───────────────────────────────────────────────── + + +def test_delete_live_stt_session(): + mock_redis = MagicMock() + + delete_live_stt_session(mock_redis, "test-id") + + expected_key = f"{LIVE_STT_SESSION_PREFIX}test-id" + mock_redis.delete.assert_called_once_with(expected_key) + + +# ── strip_committed_prefix edge cases ────────────────────────────────────── + + +def test_strip_committed_prefix_empty_committed(): + result = strip_committed_prefix("", "hello world") + assert result == "hello world" + + +def test_strip_committed_prefix_empty_hypothesis(): + result = strip_committed_prefix("hello", "") + assert result == "" + + +def test_strip_committed_prefix_both_empty(): + result = strip_committed_prefix("", "") + assert result == "" + + +def test_strip_committed_prefix_no_overlap(): + result = strip_committed_prefix( + "completely different text", + "no overlap here at all", + ) + assert result == "no overlap here at all" + + +# ── apply_live_stt_hypothesis edge cases ─────────────────────────────────── + + +def test_apply_live_stt_hypothesis_negative_chunk_index(): + session_state = create_live_stt_session("user1") + + with pytest.raises(ValueError, match="non-negative"): + apply_live_stt_hypothesis(session_state, "hello", -1) + + +def test_apply_live_stt_hypothesis_same_chunk_index_is_noop(): + session_state = create_live_stt_session("user1") + session_state["last_chunk_index"] = 5 + + original_state = dict(session_state) + result = apply_live_stt_hypothesis(session_state, "hello", 5) + + assert result["last_chunk_index"] == 5 + assert result["committed_text"] == original_state["committed_text"] + + +def test_apply_live_stt_hypothesis_silence_commits_all_previous(): + session_state = create_live_stt_session("user1") + session_state["last_chunk_index"] = 0 + session_state["latest_hypothesis"] = "previous words here" + + apply_live_stt_hypothesis(session_state, "", 1, is_silence=True) + + # Silence with empty current hypothesis should commit previous + assert "previous words here" in session_state["committed_text"] + + +# ── get_live_stt_transcript_text ──────────────────────────────────────────── + + +def test_get_live_stt_transcript_text_both_parts(): + state = {"committed_text": "hello", "mutable_text": "world"} + assert get_live_stt_transcript_text(state) == "hello world" + + +def test_get_live_stt_transcript_text_committed_only(): + state = {"committed_text": "hello", "mutable_text": ""} + assert get_live_stt_transcript_text(state) == "hello" + + +def test_get_live_stt_transcript_text_mutable_only(): + state = {"committed_text": "", "mutable_text": "world"} + assert get_live_stt_transcript_text(state) == "world" + + +def test_get_live_stt_transcript_text_empty(): + state = {"committed_text": "", "mutable_text": ""} + assert get_live_stt_transcript_text(state) == "" + + +# ── finalize_live_stt_session edge cases ──────────────────────────────────── + + +def test_finalize_live_stt_session_empty(): + state = { + "committed_text": "", + "latest_hypothesis": "", + } + assert finalize_live_stt_session(state) == "" + + +def test_finalize_live_stt_session_committed_only(): + state = { + "committed_text": "all committed", + "latest_hypothesis": "", + } + assert finalize_live_stt_session(state) == "all committed" diff --git a/tests/stt/test_openai_stt.py b/tests/stt/test_openai_stt.py new file mode 100644 index 00000000..f934f16a --- /dev/null +++ b/tests/stt/test_openai_stt.py @@ -0,0 +1,275 @@ +"""Tests for application/stt/openai_stt.py""" + +from pathlib import Path +from unittest.mock import MagicMock, patch, mock_open + +import pytest + + +@pytest.mark.unit +class TestOpenAISTTInit: + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_init_defaults_from_settings(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = "sk-from-settings" + mock_settings.API_KEY = "sk-fallback" + mock_settings.OPENAI_BASE_URL = "https://custom.api.com/v1" + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT() + + assert stt.api_key == "sk-from-settings" + assert stt.base_url == "https://custom.api.com/v1" + assert stt.model == "whisper-1" + mock_openai_cls.assert_called_once_with( + api_key="sk-from-settings", + base_url="https://custom.api.com/v1", + ) + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_init_explicit_params_override_settings(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = "sk-settings" + mock_settings.API_KEY = None + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT( + api_key="sk-explicit", + base_url="https://explicit.api.com", + model="whisper-2", + ) + + assert stt.api_key == "sk-explicit" + assert stt.base_url == "https://explicit.api.com" + assert stt.model == "whisper-2" + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_init_falls_back_to_api_key(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = None + mock_settings.API_KEY = "sk-fallback-key" + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT() + + assert stt.api_key == "sk-fallback-key" + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_init_default_base_url(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = "sk-test" + mock_settings.API_KEY = None + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT() + + assert stt.base_url == "https://api.openai.com/v1" + + +@pytest.mark.unit +class TestOpenAISTTTranscribe: + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_transcribe_basic(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = "sk-test" + mock_settings.API_KEY = None + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + mock_client = MagicMock() + mock_openai_cls.return_value = mock_client + + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "text": "Hello world", + "language": "en", + "duration": 2.5, + "segments": [], + } + mock_client.audio.transcriptions.create.return_value = mock_response + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT() + + file_path = Path("/tmp/test_audio.wav") + with patch("builtins.open", mock_open(read_data=b"audio_data")): + result = stt.transcribe(file_path) + + assert result["text"] == "Hello world" + assert result["language"] == "en" + assert result["duration_s"] == 2.5 + assert result["segments"] == [] + assert result["provider"] == "openai" + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_transcribe_with_language(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = "sk-test" + mock_settings.API_KEY = None + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + mock_client = MagicMock() + mock_openai_cls.return_value = mock_client + + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "text": "Bonjour", + "language": "fr", + "duration": 1.0, + "segments": [], + } + mock_client.audio.transcriptions.create.return_value = mock_response + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT() + + file_path = Path("/tmp/test_audio.wav") + with patch("builtins.open", mock_open(read_data=b"audio_data")): + result = stt.transcribe(file_path, language="fr") + + assert result["language"] == "fr" + call_kwargs = mock_client.audio.transcriptions.create.call_args[1] + assert call_kwargs["language"] == "fr" + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_transcribe_with_timestamps(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = "sk-test" + mock_settings.API_KEY = None + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + mock_client = MagicMock() + mock_openai_cls.return_value = mock_client + + segment_obj = MagicMock() + segment_obj.model_dump.return_value = { + "start": 0.0, + "end": 1.5, + "text": "Hello", + } + + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "text": "Hello", + "language": "en", + "duration": 1.5, + "segments": [segment_obj], + } + mock_client.audio.transcriptions.create.return_value = mock_response + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT() + + file_path = Path("/tmp/test_audio.wav") + with patch("builtins.open", mock_open(read_data=b"audio_data")): + result = stt.transcribe(file_path, timestamps=True) + + call_kwargs = mock_client.audio.transcriptions.create.call_args[1] + assert call_kwargs["timestamp_granularities"] == ["segment"] + assert len(result["segments"]) == 1 + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_transcribe_no_segments_key(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = "sk-test" + mock_settings.API_KEY = None + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + mock_client = MagicMock() + mock_openai_cls.return_value = mock_client + + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "text": "Hello", + } + mock_client.audio.transcriptions.create.return_value = mock_response + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT() + + file_path = Path("/tmp/test_audio.wav") + with patch("builtins.open", mock_open(read_data=b"audio_data")): + result = stt.transcribe(file_path) + + assert result["text"] == "Hello" + assert result["segments"] == [] + + @patch("application.stt.openai_stt.OpenAI") + @patch("application.stt.openai_stt.settings") + def test_transcribe_language_fallback_to_param(self, mock_settings, mock_openai_cls): + mock_settings.OPENAI_API_KEY = "sk-test" + mock_settings.API_KEY = None + mock_settings.OPENAI_BASE_URL = None + mock_settings.OPENAI_STT_MODEL = "whisper-1" + + mock_client = MagicMock() + mock_openai_cls.return_value = mock_client + + mock_response = MagicMock() + mock_response.model_dump.return_value = { + "text": "Test", + "language": None, + "duration": 1.0, + } + mock_client.audio.transcriptions.create.return_value = mock_response + + from application.stt.openai_stt import OpenAISTT + + stt = OpenAISTT() + + file_path = Path("/tmp/test_audio.wav") + with patch("builtins.open", mock_open(read_data=b"audio_data")): + result = stt.transcribe(file_path, language="de") + + assert result["language"] == "de" + + +@pytest.mark.unit +class TestOpenAISTTToDict: + + def test_to_dict_with_model_dump(self): + from application.stt.openai_stt import OpenAISTT + + obj = MagicMock() + obj.model_dump.return_value = {"key": "value"} + + result = OpenAISTT._to_dict(obj) + assert result == {"key": "value"} + + def test_to_dict_with_dict(self): + from application.stt.openai_stt import OpenAISTT + + result = OpenAISTT._to_dict({"key": "value"}) + assert result == {"key": "value"} + + def test_to_dict_with_other_type(self): + from application.stt.openai_stt import OpenAISTT + + result = OpenAISTT._to_dict("string_value") + assert result == {} + + def test_to_dict_with_none(self): + from application.stt.openai_stt import OpenAISTT + + result = OpenAISTT._to_dict(None) + assert result == {} diff --git a/tests/test_auth.py b/tests/test_auth.py new file mode 100644 index 00000000..8cd16abe --- /dev/null +++ b/tests/test_auth.py @@ -0,0 +1,83 @@ +from unittest.mock import Mock, patch + +import pytest + + +@pytest.mark.unit +class TestHandleAuth: + + def test_returns_local_when_no_auth_type(self): + from application.auth import handle_auth + + mock_request = Mock() + with patch("application.auth.settings") as mock_settings: + mock_settings.AUTH_TYPE = "none" + result = handle_auth(mock_request) + + assert result == {"sub": "local"} + + def test_returns_none_when_no_jwt_header(self): + from application.auth import handle_auth + + mock_request = Mock() + mock_request.headers.get.return_value = None + with patch("application.auth.settings") as mock_settings: + mock_settings.AUTH_TYPE = "simple_jwt" + result = handle_auth(mock_request) + + assert result is None + + def test_decodes_valid_jwt(self): + from application.auth import handle_auth + + mock_request = Mock() + mock_request.headers.get.return_value = "Bearer valid_token" + + with patch("application.auth.settings") as mock_settings, patch( + "application.auth.jwt" + ) as mock_jwt: + mock_settings.AUTH_TYPE = "simple_jwt" + mock_settings.JWT_SECRET_KEY = "secret" + mock_jwt.decode.return_value = {"sub": "user123"} + result = handle_auth(mock_request) + + assert result == {"sub": "user123"} + mock_jwt.decode.assert_called_once_with( + "valid_token", + "secret", + algorithms=["HS256"], + options={"verify_exp": False}, + ) + + def test_returns_error_on_invalid_jwt(self): + from application.auth import handle_auth + + mock_request = Mock() + mock_request.headers.get.return_value = "Bearer bad_token" + + with patch("application.auth.settings") as mock_settings, patch( + "application.auth.jwt" + ) as mock_jwt: + mock_settings.AUTH_TYPE = "session_jwt" + mock_settings.JWT_SECRET_KEY = "secret" + mock_jwt.decode.side_effect = Exception("Invalid token") + result = handle_auth(mock_request) + + assert result["error"] == "invalid_token" + + def test_strips_bearer_prefix(self): + from application.auth import handle_auth + + mock_request = Mock() + mock_request.headers.get.return_value = "Bearer my_token" + + with patch("application.auth.settings") as mock_settings, patch( + "application.auth.jwt" + ) as mock_jwt: + mock_settings.AUTH_TYPE = "simple_jwt" + mock_settings.JWT_SECRET_KEY = "secret" + mock_jwt.decode.return_value = {"sub": "user1"} + handle_auth(mock_request) + + mock_jwt.decode.assert_called_once() + assert mock_jwt.decode.call_args[0][0] == "my_token" diff --git a/tests/test_cache.py b/tests/test_cache.py index dc802c64..d8d28998 100644 --- a/tests/test_cache.py +++ b/tests/test_cache.py @@ -2,7 +2,12 @@ import json from unittest.mock import MagicMock, patch import pytest -from application.cache import gen_cache, gen_cache_key, stream_cache +from application.cache import ( + gen_cache, + gen_cache_key, + get_redis_instance, + stream_cache, +) from application.utils import get_hash @@ -120,3 +125,296 @@ def test_stream_cache_miss(mock_make_redis): assert result == ["new_chunk"] mock_redis_instance.get.assert_called_once() mock_redis_instance.set.assert_called_once() + + +# ── get_redis_instance ────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestGetRedisInstance: + + def setup_method(self): + """Reset module-level redis state between tests.""" + import application.cache as cache_mod + + cache_mod._redis_instance = None + cache_mod._redis_creation_failed = False + + def teardown_method(self): + import application.cache as cache_mod + + cache_mod._redis_instance = None + cache_mod._redis_creation_failed = False + + @patch("application.cache.redis.Redis.from_url") + @patch("application.cache.settings") + def test_creates_redis_instance(self, mock_settings, mock_from_url): + mock_settings.CACHE_REDIS_URL = "redis://localhost:6379/0" + mock_instance = MagicMock() + mock_from_url.return_value = mock_instance + + result = get_redis_instance() + + assert result is mock_instance + mock_from_url.assert_called_once_with( + "redis://localhost:6379/0", socket_connect_timeout=2 + ) + + @patch("application.cache.redis.Redis.from_url") + @patch("application.cache.settings") + def test_returns_cached_instance(self, mock_settings, mock_from_url): + mock_settings.CACHE_REDIS_URL = "redis://localhost:6379/0" + mock_instance = MagicMock() + mock_from_url.return_value = mock_instance + + result1 = get_redis_instance() + result2 = get_redis_instance() + + assert result1 is result2 + assert mock_from_url.call_count == 1 + + @patch("application.cache.redis.Redis.from_url") + @patch("application.cache.settings") + def test_value_error_stops_retries(self, mock_settings, mock_from_url): + import application.cache as cache_mod + + mock_settings.CACHE_REDIS_URL = "invalid://url" + mock_from_url.side_effect = ValueError("Invalid Redis URL") + + result = get_redis_instance() + + assert result is None + assert cache_mod._redis_creation_failed is True + + # Subsequent calls should not retry + mock_from_url.reset_mock() + result2 = get_redis_instance() + assert result2 is None + mock_from_url.assert_not_called() + + @patch("application.cache.redis.Redis.from_url") + @patch("application.cache.settings") + def test_connection_error_allows_retries(self, mock_settings, mock_from_url): + import application.cache as cache_mod + import redis as redis_mod + + mock_settings.CACHE_REDIS_URL = "redis://unreachable:6379/0" + mock_from_url.side_effect = redis_mod.ConnectionError("Connection refused") + + result = get_redis_instance() + + assert result is None + assert cache_mod._redis_creation_failed is False + + # Subsequent calls should retry + mock_from_url.side_effect = None + mock_from_url.return_value = MagicMock() + result2 = get_redis_instance() + assert result2 is not None + + +# ── gen_cache_key edge cases ──────────────────────────────────────────────── + + +@pytest.mark.unit +def test_gen_cache_key_with_tools(): + messages = [{"role": "user", "content": "test"}] + tools = [{"type": "function", "function": {"name": "test"}}] + + key = gen_cache_key(messages, model="docgpt", tools=tools) + assert isinstance(key, str) + assert len(key) == 32 + + +@pytest.mark.unit +def test_gen_cache_key_default_model(): + messages = [{"role": "user", "content": "test"}] + key = gen_cache_key(messages) + assert isinstance(key, str) + assert len(key) == 32 + + +@pytest.mark.unit +def test_gen_cache_key_deterministic(): + messages = [{"role": "user", "content": "test"}] + key1 = gen_cache_key(messages, model="m1") + key2 = gen_cache_key(messages, model="m1") + assert key1 == key2 + + +@pytest.mark.unit +def test_gen_cache_key_different_models(): + messages = [{"role": "user", "content": "test"}] + key1 = gen_cache_key(messages, model="m1") + key2 = gen_cache_key(messages, model="m2") + assert key1 != key2 + + +# ── gen_cache with tools bypass ───────────────────────────────────────────── + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_gen_cache_bypasses_when_tools_provided(mock_make_redis): + """When tools are provided, caching is bypassed.""" + mock_redis_instance = MagicMock() + mock_make_redis.return_value = mock_redis_instance + + @gen_cache + def mock_function(self, model, messages, stream, tools): + return "direct_result" + + messages = [{"role": "user", "content": "test"}] + tools = [{"type": "function"}] + result = mock_function(None, "model", messages, stream=False, tools=tools) + + assert result == "direct_result" + mock_redis_instance.get.assert_not_called() + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_gen_cache_no_redis(mock_make_redis): + """When redis is unavailable, function runs without caching.""" + mock_make_redis.return_value = None + + @gen_cache + def mock_function(self, model, messages, stream, tools): + return "no_cache_result" + + messages = [{"role": "user", "content": "test"}] + result = mock_function(None, "model", messages, stream=False, tools=None) + + assert result == "no_cache_result" + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_gen_cache_redis_get_error(mock_make_redis): + """When redis.get raises, function falls through gracefully.""" + mock_redis_instance = MagicMock() + mock_make_redis.return_value = mock_redis_instance + mock_redis_instance.get.side_effect = Exception("Redis error") + + @gen_cache + def mock_function(self, model, messages, stream, tools): + return "fallback_result" + + messages = [{"role": "user", "content": "test"}] + result = mock_function(None, "model", messages, stream=False, tools=None) + + assert result == "fallback_result" + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_gen_cache_redis_set_error(mock_make_redis): + """When redis.set raises, the result is still returned.""" + mock_redis_instance = MagicMock() + mock_make_redis.return_value = mock_redis_instance + mock_redis_instance.get.return_value = None + mock_redis_instance.set.side_effect = Exception("Redis write error") + + @gen_cache + def mock_function(self, model, messages, stream, tools): + return "result_str" + + messages = [{"role": "user", "content": "test"}] + result = mock_function(None, "model", messages, stream=False, tools=None) + + assert result == "result_str" + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_gen_cache_non_string_result_not_cached(mock_make_redis): + """Non-string results should not be cached.""" + mock_redis_instance = MagicMock() + mock_make_redis.return_value = mock_redis_instance + mock_redis_instance.get.return_value = None + + @gen_cache + def mock_function(self, model, messages, stream, tools): + return {"key": "value"} # not a string + + messages = [{"role": "user", "content": "test"}] + result = mock_function(None, "model", messages, stream=False, tools=None) + + assert result == {"key": "value"} + mock_redis_instance.set.assert_not_called() + + +# ── stream_cache edge cases ───────────────────────────────────────────────── + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_stream_cache_bypasses_when_tools_provided(mock_make_redis): + """When tools are provided, streaming cache is bypassed.""" + mock_redis_instance = MagicMock() + mock_make_redis.return_value = mock_redis_instance + + @stream_cache + def mock_function(self, model, messages, stream, tools): + yield "direct_chunk" + + messages = [{"role": "user", "content": "test"}] + tools = [{"type": "function"}] + result = list(mock_function(None, "model", messages, stream=True, tools=tools)) + + assert result == ["direct_chunk"] + mock_redis_instance.get.assert_not_called() + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_stream_cache_no_redis(mock_make_redis): + """When redis is unavailable, streaming works without caching.""" + mock_make_redis.return_value = None + + @stream_cache + def mock_function(self, model, messages, stream, tools): + yield "chunk1" + yield "chunk2" + + messages = [{"role": "user", "content": "test"}] + result = list(mock_function(None, "model", messages, stream=True, tools=None)) + + assert result == ["chunk1", "chunk2"] + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_stream_cache_redis_get_error(mock_make_redis): + """When redis.get raises during stream, falls through gracefully.""" + mock_redis_instance = MagicMock() + mock_make_redis.return_value = mock_redis_instance + mock_redis_instance.get.side_effect = Exception("Redis error") + + @stream_cache + def mock_function(self, model, messages, stream, tools): + yield "fallback_chunk" + + messages = [{"role": "user", "content": "test"}] + result = list(mock_function(None, "model", messages, stream=True, tools=None)) + + assert result == ["fallback_chunk"] + + +@pytest.mark.unit +@patch("application.cache.get_redis_instance") +def test_stream_cache_redis_set_error(mock_make_redis): + """When redis.set raises during stream save, chunks are still yielded.""" + mock_redis_instance = MagicMock() + mock_make_redis.return_value = mock_redis_instance + mock_redis_instance.get.return_value = None + mock_redis_instance.set.side_effect = Exception("Redis write error") + + @stream_cache + def mock_function(self, model, messages, stream, tools): + yield "chunk" + + messages = [{"role": "user", "content": "test"}] + result = list(mock_function(None, "model", messages, stream=True, tools=None)) + + assert result == ["chunk"] diff --git a/tests/test_error.py b/tests/test_error.py index b86de00a..8913ae6d 100644 --- a/tests/test_error.py +++ b/tests/test_error.py @@ -1,5 +1,5 @@ import pytest -from application.error import bad_request, response_error +from application.error import bad_request, response_error, sanitize_api_error from flask import Flask @@ -41,3 +41,43 @@ def test_response_error_without_message(app): response = response_error(code_status=500) assert response.status_code == 500 assert response.json == {"error": "Internal Server Error"} + + +@pytest.mark.unit +class TestSanitizeApiError: + + def test_503_unavailable(self): + assert "temporarily unavailable" in sanitize_api_error("503 Service Unavailable") + + def test_high_demand(self): + assert "temporarily unavailable" in sanitize_api_error("high demand") + + def test_429_rate_limit(self): + assert "Rate limit" in sanitize_api_error("429 Too Many Requests") + + def test_quota_exceeded(self): + assert "Rate limit" in sanitize_api_error("Quota exceeded") + + def test_401_unauthorized(self): + assert "Authentication" in sanitize_api_error("401 Unauthorized") + + def test_invalid_api_key(self): + assert "Authentication" in sanitize_api_error("Invalid API key provided") + + def test_timeout(self): + assert "timed out" in sanitize_api_error("Request timed out") + + def test_connection_error(self): + assert "Network" in sanitize_api_error("Connection refused") + + def test_long_message_sanitized(self): + assert "error occurred" in sanitize_api_error("x" * 201) + + def test_traceback_sanitized(self): + assert "error occurred" in sanitize_api_error("Traceback (most recent call)") + + def test_json_sanitized(self): + assert "error occurred" in sanitize_api_error('{"error": "something"}') + + def test_short_safe_message_passed_through(self): + assert sanitize_api_error("Something broke") == "Something broke" diff --git a/tests/test_logging.py b/tests/test_logging.py new file mode 100644 index 00000000..a9329284 --- /dev/null +++ b/tests/test_logging.py @@ -0,0 +1,202 @@ +from unittest.mock import Mock, patch + +import pytest + +from application.logging import build_stack_data + + +@pytest.mark.unit +class TestBuildStackData: + + def test_raises_on_none_obj(self): + with pytest.raises(ValueError, match="cannot be None"): + build_stack_data(None) + + def test_auto_discovers_attributes(self): + class Obj: + name = "test" + count = 5 + + result = build_stack_data(Obj()) + assert result["name"] == "test" + assert result["count"] == 5 + + def test_include_attributes(self): + class Obj: + name = "test" + count = 5 + hidden = "secret" + + result = build_stack_data(Obj(), include_attributes=["name"]) + assert result["name"] == "test" + assert "hidden" not in result + + def test_exclude_attributes(self): + class Obj: + pass + + obj = Obj() + obj.name = "test" + obj.secret = "hidden" + result = build_stack_data( + obj, + include_attributes=["name", "secret"], + exclude_attributes=["secret"], + ) + assert "secret" not in result + assert result["name"] == "test" + + def test_list_of_dicts(self): + class Obj: + pass + + obj = Obj() + obj.items = [{"a": 1}, {"b": 2}] + result = build_stack_data(obj, include_attributes=["items"]) + assert result["items"] == [{"a": 1}, {"b": 2}] + + def test_list_of_objects(self): + class Inner: + def __init__(self, v): + self.val = v + + class Obj: + pass + + obj = Obj() + obj.items = [Inner(1), Inner(2)] + result = build_stack_data(obj, include_attributes=["items"]) + assert result["items"] == [{"val": 1}, {"val": 2}] + + def test_list_of_strings(self): + class Obj: + pass + + obj = Obj() + obj.tags = [1, 2, 3] + result = build_stack_data(obj, include_attributes=["tags"]) + assert result["tags"] == ["1", "2", "3"] + + def test_dict_attribute(self): + class Obj: + pass + + obj = Obj() + obj.meta = {"key": 123} + result = build_stack_data(obj, include_attributes=["meta"]) + assert result["meta"] == {"key": "123"} + + def test_none_attribute_skipped(self): + class Obj: + pass + + obj = Obj() + obj.empty = None + result = build_stack_data(obj, include_attributes=["empty"]) + assert "empty" not in result + + def test_custom_data_merged(self): + class Obj: + name = "test" + + result = build_stack_data( + Obj(), + include_attributes=["name"], + custom_data={"extra": "val"}, + ) + assert result["extra"] == "val" + + def test_attribute_error_handled(self): + class Obj: + pass + + result = build_stack_data(Obj(), include_attributes=["nonexistent"]) + assert result == {} + + +@pytest.mark.unit +class TestLogActivity: + + def test_log_activity_decorator_yields(self): + from application.logging import log_activity + + class FakeAgent: + endpoint = "test" + user = "user1" + user_api_key = "key1" + query = "hi" + + @log_activity() + def my_gen(agent, log_context=None): + yield "chunk1" + yield "chunk2" + + with patch("application.logging._log_to_mongodb"): + result = list(my_gen(FakeAgent())) + assert result == ["chunk1", "chunk2"] + + def test_log_activity_handles_exception(self): + from application.logging import log_activity + + class FakeAgent: + endpoint = "test" + user = "user1" + user_api_key = "" + + @log_activity() + def failing_gen(agent, log_context=None): + yield "ok" + raise RuntimeError("boom") + + with patch("application.logging._log_to_mongodb"), pytest.raises( + RuntimeError, match="boom" + ): + list(failing_gen(FakeAgent())) + + +@pytest.mark.unit +class TestLogToMongoDB: + + def test_logs_entry(self): + from application.logging import _log_to_mongodb + + mock_collection = Mock() + mock_db = {"stack_logs": mock_collection} + + with patch( + "application.logging.MongoDB.get_client", + return_value={"docsgpt": mock_db}, + ), patch("application.logging.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + _log_to_mongodb("ep", "aid", "user", "key", "q", [], "info") + + mock_collection.insert_one.assert_called_once() + doc = mock_collection.insert_one.call_args[0][0] + assert doc["endpoint"] == "ep" + assert doc["level"] == "info" + + def test_truncates_long_strings(self): + from application.logging import _log_to_mongodb + + mock_collection = Mock() + mock_db = {"stack_logs": mock_collection} + + with patch( + "application.logging.MongoDB.get_client", + return_value={"docsgpt": mock_db}, + ), patch("application.logging.settings") as mock_settings: + mock_settings.MONGO_DB_NAME = "docsgpt" + _log_to_mongodb("ep", "aid", "user", "key", "x" * 20000, [], "info") + + doc = mock_collection.insert_one.call_args[0][0] + assert len(doc["query"]) == 10000 + + def test_handles_mongo_error(self): + from application.logging import _log_to_mongodb + + with patch( + "application.logging.MongoDB.get_client", + side_effect=Exception("DB down"), + ): + # Should not raise + _log_to_mongodb("ep", "aid", "user", "key", "q", [], "info") diff --git a/tests/test_usage.py b/tests/test_usage.py index 10185fef..f8d8106f 100644 --- a/tests/test_usage.py +++ b/tests/test_usage.py @@ -3,7 +3,9 @@ import sys import pytest from application.usage import ( + _count_prompt_tokens, _count_tokens, + _serialize_for_token_count, gen_token_usage, stream_token_usage, update_token_usage, @@ -324,3 +326,264 @@ def test_update_token_usage_skips_when_all_ids_missing(monkeypatch): ) assert inserted_docs == [] + + +# ── _serialize_for_token_count ────────────────────────────────────────────── + + +@pytest.mark.unit +class TestSerializeForTokenCount: + + def test_string_passthrough(self): + assert _serialize_for_token_count("hello") == "hello" + + def test_data_url_returns_empty(self): + data_url = "data:image/png;base64,iVBORw0KGgoAAAA..." + assert _serialize_for_token_count(data_url) == "" + + def test_none_returns_empty(self): + assert _serialize_for_token_count(None) == "" + + def test_list_recursion(self): + result = _serialize_for_token_count(["hello", "world"]) + assert result == ["hello", "world"] + + def test_dict_skips_binary_fields(self): + data = { + "text": "hello", + "data": "binary_stuff", + "base64": "encoded_data", + "image_data": "img_bytes", + } + result = _serialize_for_token_count(data) + assert "text" in result + assert "data" not in result + assert "base64" not in result + assert "image_data" not in result + + def test_dict_skips_base64_url(self): + data = {"url": "data:image/png;base64,abc123"} + result = _serialize_for_token_count(data) + assert "url" not in result + + def test_dict_keeps_normal_url(self): + data = {"url": "https://example.com/image.png"} + result = _serialize_for_token_count(data) + assert "url" in result + + def test_object_with_model_dump(self): + class PydanticLike: + def model_dump(self): + return {"key": "value"} + + result = _serialize_for_token_count(PydanticLike()) + assert result == {"key": "value"} + + def test_object_with_to_dict(self): + class DictLike: + def to_dict(self): + return {"key": "value"} + + result = _serialize_for_token_count(DictLike()) + assert result == {"key": "value"} + + def test_object_with_dict_attr(self): + class SimpleObj: + def __init__(self): + self.name = "test" + + result = _serialize_for_token_count(SimpleObj()) + assert result == {"name": "test"} + + def test_number_to_string(self): + assert _serialize_for_token_count(42) == "42" + + def test_nested_dict_with_list(self): + data = {"items": ["a", "b"], "nested": {"key": "val"}} + result = _serialize_for_token_count(data) + assert result["items"] == ["a", "b"] + assert result["nested"] == {"key": "val"} + + +# ── _count_tokens ─────────────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestCountTokens: + + def test_none_returns_zero(self): + assert _count_tokens(None) == 0 + + def test_empty_string_returns_zero(self): + assert _count_tokens("") == 0 + + def test_data_url_returns_zero(self): + data_url = "data:image/png;base64,iVBORw0KGgoAAAA..." + assert _count_tokens(data_url) == 0 + + def test_dict_counts(self): + assert _count_tokens({"key": "some text here"}) > 0 + + def test_list_counts(self): + assert _count_tokens(["some text", "more text"]) > 0 + + +# ── _count_prompt_tokens ──────────────────────────────────────────────────── + + +@pytest.mark.unit +class TestCountPromptTokens: + + def test_empty_messages(self): + assert _count_prompt_tokens([], tools=None) == 0 + + def test_none_messages(self): + assert _count_prompt_tokens(None, tools=None) == 0 + + def test_dict_messages(self): + messages = [{"content": "Hello world"}] + tokens = _count_prompt_tokens(messages, tools=None) + assert tokens > 0 + + def test_non_dict_messages(self): + class MessageObj: + def __init__(self): + self.content = "Hello world" + + messages = [MessageObj()] + tokens = _count_prompt_tokens(messages, tools=None) + assert tokens > 0 + + def test_with_tools(self): + messages = [{"content": "Hello"}] + tools = [ + { + "type": "function", + "function": { + "name": "search", + "parameters": {"type": "object"}, + }, + } + ] + tokens_without = _count_prompt_tokens(messages, tools=None) + tokens_with = _count_prompt_tokens(messages, tools=tools) + assert tokens_with > tokens_without + + def test_with_usage_attachments(self): + messages = [{"content": "Hello"}] + attachments = [{"mime_type": "text/plain", "content": "file data"}] + tokens_without = _count_prompt_tokens(messages, tools=None) + tokens_with = _count_prompt_tokens( + messages, tools=None, usage_attachments=attachments + ) + assert tokens_with > tokens_without + + def test_with_response_format(self): + messages = [{"content": "Hello"}] + tokens_without = _count_prompt_tokens(messages, tools=None) + tokens_with = _count_prompt_tokens( + messages, tools=None, response_format={"type": "json_object"} + ) + assert tokens_with > tokens_without + + def test_message_with_tool_calls_field(self): + messages = [ + { + "content": "Hello", + "tool_calls": [ + {"id": "call_1", "function": {"name": "test", "arguments": "{}"}} + ], + } + ] + tokens = _count_prompt_tokens(messages, tools=None) + assert tokens > 0 + + def test_message_with_tool_call_id(self): + messages = [ + { + "content": "Result of tool", + "tool_call_id": "call_1", + } + ] + tokens = _count_prompt_tokens(messages, tools=None) + assert tokens > 0 + + +# ── update_token_usage edge cases ─────────────────────────────────────────── + + +@pytest.mark.unit +def test_update_token_usage_with_user_api_key(monkeypatch): + inserted_docs = [] + + class FakeCollection: + def insert_one(self, doc): + inserted_docs.append(doc) + + modules_without_pytest = dict(sys.modules) + modules_without_pytest.pop("pytest", None) + + monkeypatch.setattr("application.usage.sys.modules", modules_without_pytest) + monkeypatch.setattr("application.usage.usage_collection", FakeCollection()) + + update_token_usage( + decoded_token=None, + user_api_key="api-key-123", + token_usage={"prompt_tokens": 10, "generated_tokens": 5}, + agent_id=None, + ) + + assert len(inserted_docs) == 1 + assert inserted_docs[0]["api_key"] == "api-key-123" + assert inserted_docs[0]["user_id"] is None + assert "agent_id" not in inserted_docs[0] + + +@pytest.mark.unit +def test_update_token_usage_with_decoded_token(monkeypatch): + inserted_docs = [] + + class FakeCollection: + def insert_one(self, doc): + inserted_docs.append(doc) + + modules_without_pytest = dict(sys.modules) + modules_without_pytest.pop("pytest", None) + + monkeypatch.setattr("application.usage.sys.modules", modules_without_pytest) + monkeypatch.setattr("application.usage.usage_collection", FakeCollection()) + + update_token_usage( + decoded_token={"sub": "user-abc"}, + user_api_key=None, + token_usage={"prompt_tokens": 20, "generated_tokens": 10}, + agent_id=None, + ) + + assert len(inserted_docs) == 1 + assert inserted_docs[0]["user_id"] == "user-abc" + + +@pytest.mark.unit +def test_update_token_usage_non_dict_decoded_token(monkeypatch): + inserted_docs = [] + + class FakeCollection: + def insert_one(self, doc): + inserted_docs.append(doc) + + modules_without_pytest = dict(sys.modules) + modules_without_pytest.pop("pytest", None) + + monkeypatch.setattr("application.usage.sys.modules", modules_without_pytest) + monkeypatch.setattr("application.usage.usage_collection", FakeCollection()) + + update_token_usage( + decoded_token="not-a-dict", + user_api_key="key", + token_usage={"prompt_tokens": 5, "generated_tokens": 3}, + agent_id=None, + ) + + assert len(inserted_docs) == 1 + assert inserted_docs[0]["user_id"] is None diff --git a/tests/test_utils.py b/tests/test_utils.py index ed1a4977..0190b44c 100644 --- a/tests/test_utils.py +++ b/tests/test_utils.py @@ -531,3 +531,134 @@ class TestCleanTextForTts: def test_removes_double_colons(self): result = clean_text_for_tts("module::function") assert "::" not in result + + @pytest.mark.unit + def test_removes_non_ascii(self): + result = clean_text_for_tts("hello \U0001f600 world") + assert "\U0001f600" not in result + assert "hello" in result + assert "world" in result + + @pytest.mark.unit + def test_empty_string(self): + result = clean_text_for_tts("") + assert result == "" + + @pytest.mark.unit + def test_removes_underscore_bold(self): + result = clean_text_for_tts("__bold text__") + assert "bold text" in result + assert "__" not in result + + @pytest.mark.unit + def test_removes_underscore_italic(self): + result = clean_text_for_tts("_italic text_") + assert "italic text" in result + + +class TestLimitChatHistoryEdgeCases: + + @pytest.mark.unit + def test_max_token_limit_caps_at_model_limit(self): + """When max_token_limit exceeds model limit, model limit is used.""" + with patch("application.utils.get_token_limit", return_value=100): + history = [ + {"prompt": "q", "response": "a"}, + ] + result = limit_chat_history(history, max_token_limit=999999) + assert len(result) <= 1 + + @pytest.mark.unit + def test_max_token_limit_none_uses_model_limit(self): + with patch("application.utils.get_token_limit", return_value=100000): + history = [{"prompt": "q", "response": "a"}] + result = limit_chat_history(history, max_token_limit=None) + assert len(result) == 1 + + @pytest.mark.unit + def test_messages_without_prompt_response_keys(self): + """Messages lacking prompt/response should still be included.""" + with patch("application.utils.get_token_limit", return_value=100000): + history = [{"custom_key": "value"}] + result = limit_chat_history(history, max_token_limit=100000) + assert len(result) == 1 + + @pytest.mark.unit + def test_single_message_exceeds_limit(self): + """If the most recent message exceeds the limit, it's excluded.""" + history = [ + {"prompt": "x" * 50000, "response": "y" * 50000}, + ] + result = limit_chat_history(history, max_token_limit=10) + assert len(result) == 0 + + +class TestSafeFilenameEdgeCases: + + @pytest.mark.unit + def test_filename_with_spaces(self): + result = safe_filename("my document.pdf") + assert result == "my_document.pdf" + + @pytest.mark.unit + def test_filename_with_special_chars(self): + result = safe_filename("file@#$.txt") + # secure_filename strips special chars + assert result.endswith(".txt") + + @pytest.mark.unit + def test_chinese_filename_gets_uuid(self): + result = safe_filename("\u6587\u4ef6.pdf") + # secure_filename strips non-latin, so UUID is generated + assert result.endswith(".pdf") + assert len(result) > 5 + + +class TestGenerateImageUrlEdgeCases: + + @pytest.mark.unit + def test_non_string_input(self): + result = generate_image_url(123) + # Not a string, not starting with http, uses default strategy + assert "/api/images/" in result or "s3" in result + + @pytest.mark.unit + def test_default_strategy_is_backend(self): + with patch("application.utils.settings") as s: + # Simulate missing URL_STRATEGY attribute + del s.URL_STRATEGY + s.API_URL = "http://localhost:7091" + result = generate_image_url("img.png") + assert "localhost:7091" in result + + +class TestGetHashEdgeCases: + + @pytest.mark.unit + def test_empty_string(self): + h = get_hash("") + assert len(h) == 32 + + @pytest.mark.unit + def test_unicode_string(self): + h = get_hash("\u4f60\u597d\u4e16\u754c") + assert len(h) == 32 + + +class TestValidateFunctionNameEdgeCases: + + @pytest.mark.unit + def test_single_char(self): + assert validate_function_name("a") is True + + @pytest.mark.unit + def test_only_numbers(self): + assert validate_function_name("123") is True + + @pytest.mark.unit + def test_with_dots(self): + assert validate_function_name("func.name") is False + + @pytest.mark.unit + def test_with_slash(self): + assert validate_function_name("path/to") is False diff --git a/tests/vectorstore/test_faiss.py b/tests/vectorstore/test_faiss.py index 36ca8542..c9fd40e3 100644 --- a/tests/vectorstore/test_faiss.py +++ b/tests/vectorstore/test_faiss.py @@ -357,3 +357,158 @@ class TestGetVectorstore: assert get_vectorstore("") == "indexes" assert get_vectorstore(None) == "indexes" + + def test_with_nested_path(self): + from application.vectorstore.faiss import get_vectorstore + + assert get_vectorstore("user/source123") == "indexes/user/source123" + + +@pytest.mark.unit +class TestFaissStoreAddChunk: + @patch("application.vectorstore.faiss.StorageCreator") + @patch("application.vectorstore.faiss.FAISS") + @patch.object( + __import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore, + "_get_embeddings", + ) + @patch("application.vectorstore.faiss.settings") + def test_add_chunk_with_metadata( + self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator + ): + mock_settings.EMBEDDINGS_NAME = "test_model" + mock_emb = Mock(dimension=3) + mock_get_emb.return_value = mock_emb + mock_ds = Mock() + mock_ds.index = Mock(d=3) + mock_ds.add_documents.return_value = ["new_id"] + mock_faiss.from_documents.return_value = mock_ds + mock_storage = Mock() + mock_storage_creator.get_storage.return_value = mock_storage + + from application.vectorstore.faiss import FaissStore + + store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()]) + store._save_to_storage = Mock(return_value=True) + + doc_id = store.add_chunk("new text", metadata={"source": "test"}) + + assert doc_id == ["new_id"] + mock_ds.add_documents.assert_called_once() + store._save_to_storage.assert_called_once() + + @patch("application.vectorstore.faiss.StorageCreator") + @patch("application.vectorstore.faiss.FAISS") + @patch.object( + __import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore, + "_get_embeddings", + ) + @patch("application.vectorstore.faiss.settings") + def test_add_chunk_default_metadata( + self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator + ): + mock_settings.EMBEDDINGS_NAME = "test_model" + mock_emb = Mock(dimension=3) + mock_get_emb.return_value = mock_emb + mock_ds = Mock() + mock_ds.index = Mock(d=3) + mock_ds.add_documents.return_value = ["new_id"] + mock_faiss.from_documents.return_value = mock_ds + mock_storage = Mock() + mock_storage_creator.get_storage.return_value = mock_storage + + from application.vectorstore.faiss import FaissStore + + store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()]) + store._save_to_storage = Mock(return_value=True) + + doc_id = store.add_chunk("new text") + + assert doc_id == ["new_id"] + + +@pytest.mark.unit +class TestFaissStoreSaveLocalNoPath: + @patch("application.vectorstore.faiss.StorageCreator") + @patch("application.vectorstore.faiss.FAISS") + @patch.object( + __import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore, + "_get_embeddings", + ) + @patch("application.vectorstore.faiss.settings") + def test_save_local_without_path( + self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator + ): + mock_settings.EMBEDDINGS_NAME = "test_model" + mock_emb = Mock(dimension=3) + mock_get_emb.return_value = mock_emb + mock_ds = Mock() + mock_ds.index = Mock(d=3) + mock_faiss.from_documents.return_value = mock_ds + mock_storage = Mock() + mock_storage_creator.get_storage.return_value = mock_storage + + from application.vectorstore.faiss import FaissStore + + store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()]) + store._save_to_storage = Mock(return_value=True) + + result = store.save_local() + + # Should NOT call docsearch.save_local with a path + mock_ds.save_local.assert_not_called() + store._save_to_storage.assert_called_once() + assert result is True + + +@pytest.mark.unit +class TestFaissStoreAssertEmbeddingDimensionsMatch: + @patch("application.vectorstore.faiss.StorageCreator") + @patch("application.vectorstore.faiss.FAISS") + @patch.object( + __import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore, + "_get_embeddings", + ) + @patch("application.vectorstore.faiss.settings") + def test_dimension_match_passes( + self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator + ): + mock_settings.EMBEDDINGS_NAME = ( + "huggingface_sentence-transformers/all-mpnet-base-v2" + ) + mock_emb = Mock(dimension=768) + mock_get_emb.return_value = mock_emb + mock_ds = Mock() + mock_ds.index = Mock(d=768) # Matching dimension + mock_faiss.from_documents.return_value = mock_ds + mock_storage_creator.get_storage.return_value = Mock() + + from application.vectorstore.faiss import FaissStore + + # Should not raise + store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()]) + assert store is not None + + @patch("application.vectorstore.faiss.StorageCreator") + @patch("application.vectorstore.faiss.FAISS") + @patch.object( + __import__("application.vectorstore.base", fromlist=["BaseVectorStore"]).BaseVectorStore, + "_get_embeddings", + ) + @patch("application.vectorstore.faiss.settings") + def test_non_huggingface_skips_dimension_check( + self, mock_settings, mock_get_emb, mock_faiss, mock_storage_creator + ): + mock_settings.EMBEDDINGS_NAME = "openai_text-embedding-ada-002" + mock_emb = Mock(dimension=1536) + mock_get_emb.return_value = mock_emb + mock_ds = Mock() + mock_ds.index = Mock(d=999) # Mismatched but doesn't matter + mock_faiss.from_documents.return_value = mock_ds + mock_storage_creator.get_storage.return_value = Mock() + + from application.vectorstore.faiss import FaissStore + + # Should not raise since embedding name is not the huggingface one + store = FaissStore(source_id="t", embeddings_key="k", docs_init=[Mock()]) + assert store is not None