chore: more tests

This commit is contained in:
Alex committed 2026-03-30 16:13:08 +01:00
1 parent dc6db847ca
commit d5c0322e2a
64 files changed
+24436 -140

No files matched your search

+92
View File
@@ -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",
)
+428 -45
View File
@@ -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
View File
Whitespace-only changes.
@@ -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 '<?xml version="1.0"' in body
assert "<name>Alice</name>" in body
assert headers["Content-Type"] == "application/xml"
def test_nested_xml(self):
body, headers = RequestBodySerializer.serialize(
{"user": {"name": "Alice"}}, ContentType.XML
)
assert "<user>" in body
assert "<name>Alice</name>" in body
def test_xml_escapes_special_chars(self):
body, headers = RequestBodySerializer.serialize(
{"data": "<script>alert('xss')</script>"}, ContentType.XML
)
assert "&lt;script&gt;" in body
def test_xml_with_list(self):
body, _ = RequestBodySerializer.serialize(
{"items": [1, 2, 3]}, ContentType.XML
)
assert "<item>1</item>" in body
assert "<item>2</item>" in body
assert "<item>3</item>" 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 "&amp;" in RequestBodySerializer._escape_xml("&")
def test_escape_xml_lt(self):
assert "&lt;" in RequestBodySerializer._escape_xml("<")
def test_escape_xml_gt(self):
assert "&gt;" in RequestBodySerializer._escape_xml(">")
def test_escape_xml_quote(self):
assert "&quot;" in RequestBodySerializer._escape_xml('"')
def test_escape_xml_apos(self):
assert "&apos;" in RequestBodySerializer._escape_xml("'")
def test_dict_to_xml_list(self):
xml = RequestBodySerializer._dict_to_xml({"items": [1, 2, 3]})
assert "<item>1</item>" in xml
assert "<item>2</item>" in xml
def test_dict_to_xml_custom_root(self):
xml = RequestBodySerializer._dict_to_xml({"key": "val"}, root_name="data")
assert "<data>" in xml
assert "<key>val</key>" in xml
def test_dict_to_xml_deeply_nested(self):
xml = RequestBodySerializer._dict_to_xml({"a": {"b": {"c": "deep"}}})
assert "<c>deep</c>" 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,
)
+516
View File
@@ -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 = "<root><item>1</item></root>"
mock_resp.content = b"<root><item>1</item></root>"
result = get_tool._parse_response(mock_resp)
assert "<root>" in result
def test_html_response(self, get_tool):
mock_resp = MagicMock()
mock_resp.headers = {"Content-Type": "text/html"}
mock_resp.text = "<html><body>Hi</body></html>"
mock_resp.content = b"<html><body>Hi</body></html>"
result = get_tool._parse_response(mock_resp)
assert "<html>" 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 = "<data/>"
mock_resp.content = b"<data/>"
result = get_tool._parse_response(mock_resp)
assert result == "<data/>"
# =====================================================================
# 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
+596
View File
@@ -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
File diff suppressed because it is too large. Load diff
+449
View File
@@ -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() == {}
@@ -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
+363
View File
@@ -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
@@ -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
+184 -1
View File
@@ -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 == []
+416
View File
@@ -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"}'
View File
Whitespace-only changes.
+879
View File
@@ -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
@@ -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
File diff suppressed because it is too large. Load diff
File diff suppressed because it is too large. Load diff
+768
View File
@@ -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
+388
View File
@@ -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
+360
View File
@@ -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
+509
View File
@@ -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
+70
View File
@@ -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
+288
View File
@@ -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
+690
View File
@@ -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()
File diff suppressed because it is too large. Load diff
File diff suppressed because it is too large. Load diff
+411
View File
@@ -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
+225
View File
@@ -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"
+406
View File
@@ -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"}
View File
Whitespace-only changes.
+95
View File
@@ -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"]
+64
View File
@@ -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"
View File
Whitespace-only changes.
+323
View File
@@ -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 = "<HUMAN>"
fake.AI_PROMPT = "<AI>"
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 == "<HUMAN>"
assert instance.AI_PROMPT == "<AI>"
# ---------------------------------------------------------------------------
# _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("<HUMAN>")
assert prompt.endswith("<AI>")
# ---------------------------------------------------------------------------
# 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"
+269
View File
@@ -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
+755
View File
@@ -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"})
+193
View File
@@ -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
+717
View File
@@ -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)
+190
View File
@@ -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
View File
Whitespace-only changes.
+367
View File
@@ -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
+382
View File
@@ -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 = "<h1>Title</h1>"
mock_doc.texts = []
result = parser._export_content(mock_doc)
assert "<h1>" 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 = "<!-- image -->"
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"
+162 -92
View File
@@ -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
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 == {}
View File
Whitespace-only changes.
+153
View File
@@ -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):
+306
View File
@@ -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 = "<sitemapindex><sitemap></sitemap></sitemapindex>"
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 = "<urlset><url></url></urlset>"
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 = "<html><body>About us</body></html>"
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"""<?xml version="1.0" encoding="UTF-8"?>
<urlset xmlns="http://www.sitemaps.org/schemas/sitemap/0.9">
<url><loc>https://example.com/page1</loc></url>
<url><loc>https://example.com/page2</loc></url>
</urlset>"""
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"""<?xml version="1.0" encoding="UTF-8"?>
<sitemapindex xmlns="http://www.sitemaps.org/schemas/sitemap/0.9">
<sitemap>
<loc>https://example.com/sitemap-child.xml</loc>
</sitemap>
</sitemapindex>"""
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"""<?xml version="1.0" encoding="UTF-8"?>
<urlset xmlns="http://www.sitemaps.org/schemas/sitemap/0.9">
</urlset>"""
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 = "<urlset><url><loc>https://example.com/p</loc></url></urlset>"
response.content = b"""<?xml version="1.0" encoding="UTF-8"?>
<urlset xmlns="http://www.sitemaps.org/schemas/sitemap/0.9">
<url><loc>https://example.com/p</loc></url>
</urlset>"""
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 = "<html>Normal page</html>"
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
+279
View File
@@ -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
+58
View File
@@ -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
View File
Whitespace-only changes.
+80
View File
@@ -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") == {}
View File
Whitespace-only changes.
+97
View File
@@ -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)
View File
Whitespace-only changes.
+234
View File
@@ -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
+252
View File
@@ -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"
+275
View File
@@ -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 == {}
+83
View File
@@ -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"
+299 -1
View File
@@ -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"]
+41 -1
View File
@@ -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"
+202
View File
@@ -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")
+263
View File
@@ -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
+131
View File
@@ -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
+155
View File
@@ -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