mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 03:12:55 +00:00
chore: more tests
This commit is contained in:
1 parent
dc6db847ca
commit
d5c0322e2a
64 files changed
+24436
-140
No files matched your search
@@ -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",
|
||||
)
|
||||
@@ -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
|
||||
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 "<script>" 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 "&" in RequestBodySerializer._escape_xml("&")
|
||||
|
||||
def test_escape_xml_lt(self):
|
||||
assert "<" in RequestBodySerializer._escape_xml("<")
|
||||
|
||||
def test_escape_xml_gt(self):
|
||||
assert ">" in RequestBodySerializer._escape_xml(">")
|
||||
|
||||
def test_escape_xml_quote(self):
|
||||
assert """ in RequestBodySerializer._escape_xml('"')
|
||||
|
||||
def test_escape_xml_apos(self):
|
||||
assert "'" in RequestBodySerializer._escape_xml("'")
|
||||
|
||||
def test_dict_to_xml_list(self):
|
||||
xml = RequestBodySerializer._dict_to_xml({"items": [1, 2, 3]})
|
||||
assert "<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,
|
||||
)
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
@@ -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
|
||||
@@ -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 == []
|
||||
@@ -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"}'
|
||||
Whitespace-only changes.
@@ -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
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
@@ -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"
|
||||
@@ -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"}
|
||||
Whitespace-only changes.
@@ -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"]
|
||||
@@ -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"
|
||||
Whitespace-only changes.
@@ -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"
|
||||
@@ -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
|
||||
@@ -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"})
|
||||
@@ -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
|
||||
@@ -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)
|
||||
@@ -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
|
||||
Whitespace-only changes.
@@ -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
|
||||
@@ -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"
|
||||
@@ -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 == {}
|
||||
Whitespace-only changes.
@@ -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):
|
||||
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Whitespace-only changes.
@@ -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") == {}
|
||||
Whitespace-only changes.
@@ -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)
|
||||
Whitespace-only changes.
@@ -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
|
||||
@@ -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"
|
||||
@@ -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 == {}
|
||||
@@ -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
@@ -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
@@ -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"
|
||||
@@ -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")
|
||||
@@ -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
|
||||
@@ -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
|
||||
@@ -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
|
||||
Reference in new issue
Block a user