diff --git a/src/chatbot_plugin/llm/rate_limit/sliding_window_strategy.py b/src/chatbot_plugin/llm/rate_limit/sliding_window_strategy.py index 5e305c8..4a9302e 100644 --- a/src/chatbot_plugin/llm/rate_limit/sliding_window_strategy.py +++ b/src/chatbot_plugin/llm/rate_limit/sliding_window_strategy.py @@ -87,4 +87,7 @@ def _tpm_wait(self, now: float, estimated_tokens: int) -> float: if current_tokens + estimated_tokens <= self._tpm: return 0 # Wait until the oldest entry exits the window + if not self._tpm_window: + # Single request exceeds TPM — no window to wait on, proceed anyway + return 0 return self._tpm_window[0][0] + 60.0 - now diff --git a/src/tests/llm/test_base_provider.py b/src/tests/llm/test_base_provider.py new file mode 100644 index 0000000..51d0999 --- /dev/null +++ b/src/tests/llm/test_base_provider.py @@ -0,0 +1,91 @@ +"""Tests for BaseProvider retry logic and _is_retryable.""" + +import pytest + +from chatbot_plugin.llm.base_provider import BaseProvider, _is_retryable +from chatbot_plugin.llm.rate_limit.quota_strategy import RateLimitExhausted + + +# ── _is_retryable() ── + + +def test_retryable_generic_exception(): + assert _is_retryable(ConnectionError("timeout")) is True + + +def test_retryable_runtime_error(): + assert _is_retryable(RuntimeError("oops")) is True + + +def test_not_retryable_rate_limit_exhausted(): + assert _is_retryable(RateLimitExhausted()) is False + + +def test_not_retryable_value_error(): + assert _is_retryable(ValueError("bad value")) is False + + +def test_not_retryable_key_error(): + assert _is_retryable(KeyError("missing")) is False + + +# ── BaseProvider.generate() ── + + +class _StubProvider(BaseProvider): + """Test double that records calls and simulates _call_api behavior.""" + + def __init__(self, side_effects: list): + super().__init__(model="stub-model") + self._side_effects = list(side_effects) + self._call_count = 0 + + async def _call_api(self, system_prompt: str, human_prompt: str) -> str: + self._call_count += 1 + if self._side_effects: + result = self._side_effects.pop(0) + if isinstance(result, Exception): + raise result + return result + return "default" + + +@pytest.mark.asyncio +async def test_generate_success(): + provider = _StubProvider(["hello"]) + result = await provider.generate("sys", "human") + assert result == "hello" + + +@pytest.mark.asyncio +async def test_generate_retries_on_transient_error(): + provider = _StubProvider([ConnectionError("timeout"), "recovered"]) + result = await provider.generate("sys", "human") + assert result == "recovered" + assert provider._call_count == 2 + + +@pytest.mark.asyncio +async def test_generate_raises_rate_limit_exhausted(): + provider = _StubProvider([RateLimitExhausted()]) + with pytest.raises(RateLimitExhausted): + await provider.generate("sys", "human") + + +@pytest.mark.asyncio +async def test_generate_returns_none_on_non_retryable_error(): + provider = _StubProvider([ValueError("bad data")]) + result = await provider.generate("sys", "human") + assert result is None + + +@pytest.mark.asyncio +async def test_generate_returns_none_after_exhausted_retries(): + provider = _StubProvider([ + ConnectionError("fail1"), + ConnectionError("fail2"), + ConnectionError("fail3"), + ]) + result = await provider.generate("sys", "human") + assert result is None + assert provider._call_count == 3 diff --git a/src/tests/llm/test_bootstrap.py b/src/tests/llm/test_bootstrap.py new file mode 100644 index 0000000..264531e --- /dev/null +++ b/src/tests/llm/test_bootstrap.py @@ -0,0 +1,114 @@ +"""Tests for LLM bootstrap factory.""" + +import os +import tempfile + +import pytest + +from chatbot_plugin.llm.bootstrap import build_llm_service, _create_provider, _create_strategy +from chatbot_plugin.llm.rate_limit import SlidingWindowStrategy, NoOpStrategy + + +# ── _create_provider() ── + + +def test_create_provider_claude(): + provider = _create_provider("claude", "key", "model") + assert provider is not None + assert provider._model == "model" + + +def test_create_provider_gemini(): + provider = _create_provider("gemini", "key", "model") + assert provider is not None + assert provider._model == "model" + + +def test_create_provider_openrouter(): + provider = _create_provider("openrouter", "key", "model") + assert provider is not None + assert provider._model == "model" + + +def test_create_provider_unknown_returns_none(): + assert _create_provider("unknown", "key", "model") is None + + +# ── _create_strategy() ── + + +def test_create_strategy_sliding_window(): + cfg = {"type": "sliding_window", "rpm": 5, "tpm": 1000, "rpd": 200} + strategy = _create_strategy(cfg) + assert isinstance(strategy, SlidingWindowStrategy) + + +def test_create_strategy_sliding_window_defaults(): + cfg = {"type": "sliding_window"} + strategy = _create_strategy(cfg) + assert isinstance(strategy, SlidingWindowStrategy) + + +def test_create_strategy_no_op_for_empty(): + strategy = _create_strategy({}) + assert isinstance(strategy, NoOpStrategy) + + +def test_create_strategy_no_op_for_unknown(): + strategy = _create_strategy({"type": "unknown"}) + assert isinstance(strategy, NoOpStrategy) + + +# ── build_llm_service() ── + + +def _write_toml(content: str) -> str: + fd, path = tempfile.mkstemp(suffix=".toml") + with os.fdopen(fd, "w") as f: + f.write(content) + return path + + +def test_build_llm_service_raises_on_no_providers(): + path = _write_toml("") + try: + with pytest.raises(ValueError, match="No valid LLM providers"): + build_llm_service(path) + finally: + os.unlink(path) + + +def test_build_llm_service_skips_missing_api_key(monkeypatch): + path = _write_toml('[[providers]]\nname = "claude"\nmodel = "m"\napi_key_env = "NO_SUCH_KEY"\npriority = 1\n') + # Ensure the env var is NOT set + monkeypatch.delenv("NO_SUCH_KEY", raising=False) + try: + with pytest.raises(ValueError, match="No valid LLM providers"): + build_llm_service(path) + finally: + os.unlink(path) + + +def test_build_llm_service_skips_unknown_provider(monkeypatch): + monkeypatch.setenv("TEST_KEY", "sk-test") + path = _write_toml( + '[[providers]]\nname = "unknown"\nmodel = "m"\napi_key_env = "TEST_KEY"\npriority = 1\n' + ) + try: + with pytest.raises(ValueError, match="No valid LLM providers"): + build_llm_service(path) + finally: + os.unlink(path) + + +def test_build_llm_service_success(monkeypatch): + monkeypatch.setenv("TEST_KEY", "sk-test") + path = _write_toml( + '[[providers]]\nname = "claude"\nmodel = "claude-sonnet-4-6-20250514"\napi_key_env = "TEST_KEY"\npriority = 1\n' + '\n[[providers]]\nname = "gemini"\nmodel = "gemini-2.5-flash"\napi_key_env = "TEST_KEY"\npriority = 2\n' + ) + try: + service = build_llm_service(path) + assert service is not None + finally: + os.unlink(path) diff --git a/src/tests/llm/test_provider_implementations.py b/src/tests/llm/test_provider_implementations.py new file mode 100644 index 0000000..77ce497 --- /dev/null +++ b/src/tests/llm/test_provider_implementations.py @@ -0,0 +1,210 @@ +"""Tests for individual LLM provider implementations.""" + +import pytest +from unittest.mock import AsyncMock, MagicMock, patch + +from chatbot_plugin.llm.rate_limit.quota_strategy import RateLimitExhausted + + +# ── ClaudeProvider ── + + +class TestClaudeProvider: + def _make_provider(self): + with patch("chatbot_plugin.llm.claude_provider.anthropic") as mock_anthropic: + mock_client = AsyncMock() + mock_anthropic.AsyncAnthropic.return_value = mock_client + from chatbot_plugin.llm.claude_provider import ClaudeProvider + provider = ClaudeProvider(api_key="sk-test", model="claude-sonnet-4-6-20250514") + return provider, mock_client + + @pytest.mark.asyncio + async def test_call_api_success(self): + provider, mock_client = self._make_provider() + mock_response = MagicMock() + mock_response.content = [MagicMock(text="Hello from Claude")] + mock_response.usage.input_tokens = 10 + mock_response.usage.output_tokens = 5 + mock_client.messages.create.return_value = mock_response + + result = await provider._call_api("sys", "human") + assert result == "Hello from Claude" + + @pytest.mark.asyncio + async def test_call_api_uses_correct_params(self): + provider, mock_client = self._make_provider() + mock_response = MagicMock() + mock_response.content = [MagicMock(text="ok")] + mock_response.usage.input_tokens = 0 + mock_response.usage.output_tokens = 0 + mock_client.messages.create.return_value = mock_response + + await provider._call_api("system-instr", "user-msg") + call_kwargs = mock_client.messages.create.call_args + assert call_kwargs.kwargs["model"] == "claude-sonnet-4-6-20250514" + assert call_kwargs.kwargs["system"] == "system-instr" + assert call_kwargs.kwargs["messages"] == [{"role": "user", "content": "user-msg"}] + + +# ── GeminiProvider ── + + +class TestGeminiProvider: + @pytest.mark.asyncio + @patch("chatbot_plugin.llm.gemini_provider.genai") + async def test_call_api_success(self, mock_genai): + mock_client = MagicMock() + mock_genai.Client.return_value = mock_client + mock_genai.GenerateContentConfig = MagicMock() + from chatbot_plugin.llm.gemini_provider import GeminiProvider + provider = GeminiProvider(api_key="test-key", model="gemini-2.5-flash") + + mock_response = MagicMock() + mock_response.candidates = [MagicMock(finish_reason="STOP")] + mock_response.text = "Hello from Gemini" + mock_response.usage_metadata = MagicMock( + prompt_token_count=10, candidates_token_count=5 + ) + mock_client.models.generate_content.return_value = mock_response + + result = await provider._call_api("sys", "human") + assert result == "Hello from Gemini" + + @pytest.mark.asyncio + @patch("chatbot_plugin.llm.gemini_provider.genai") + async def test_call_api_no_candidates_returns_empty(self, mock_genai): + mock_client = MagicMock() + mock_genai.Client.return_value = mock_client + mock_genai.GenerateContentConfig = MagicMock() + from chatbot_plugin.llm.gemini_provider import GeminiProvider + provider = GeminiProvider(api_key="test-key", model="gemini-2.5-flash") + + mock_response = MagicMock() + mock_response.candidates = [] + mock_client.models.generate_content.return_value = mock_response + + result = await provider._call_api("sys", "human") + assert result == "" + + @pytest.mark.asyncio + @patch("chatbot_plugin.llm.gemini_provider.genai") + async def test_call_api_blocked_finish_reason_returns_empty(self, mock_genai): + mock_client = MagicMock() + mock_genai.Client.return_value = mock_client + mock_genai.GenerateContentConfig = MagicMock() + from chatbot_plugin.llm.gemini_provider import GeminiProvider + provider = GeminiProvider(api_key="test-key", model="gemini-2.5-flash") + + mock_response = MagicMock() + mock_candidate = MagicMock(finish_reason="SAFETY") + mock_response.candidates = [mock_candidate] + mock_client.models.generate_content.return_value = mock_response + + result = await provider._call_api("sys", "human") + assert result == "" + + @pytest.mark.asyncio + @patch("chatbot_plugin.llm.gemini_provider.genai") + async def test_call_api_daily_quota_raises_rate_limit_exhausted(self, mock_genai): + mock_client = MagicMock() + mock_genai.Client.return_value = mock_client + mock_genai.GenerateContentConfig = MagicMock() + from chatbot_plugin.llm.gemini_provider import GeminiProvider + provider = GeminiProvider(api_key="test-key", model="gemini-2.5-flash") + + mock_client.models.generate_content.side_effect = Exception( + "429 RESOURCE_EXHAUSTED: PerDay limit exceeded" + ) + + with pytest.raises(RateLimitExhausted): + await provider._call_api("sys", "human") + + @pytest.mark.asyncio + @patch("chatbot_plugin.llm.gemini_provider.genai") + async def test_call_api_other_exception_reraises(self, mock_genai): + mock_client = MagicMock() + mock_genai.Client.return_value = mock_client + mock_genai.GenerateContentConfig = MagicMock() + from chatbot_plugin.llm.gemini_provider import GeminiProvider + provider = GeminiProvider(api_key="test-key", model="gemini-2.5-flash") + + mock_client.models.generate_content.side_effect = RuntimeError("network error") + + with pytest.raises(RuntimeError, match="network error"): + await provider._call_api("sys", "human") + + @pytest.mark.asyncio + @patch("chatbot_plugin.llm.gemini_provider.genai") + async def test_call_api_no_usage_metadata(self, mock_genai): + mock_client = MagicMock() + mock_genai.Client.return_value = mock_client + mock_genai.GenerateContentConfig = MagicMock() + from chatbot_plugin.llm.gemini_provider import GeminiProvider + provider = GeminiProvider(api_key="test-key", model="gemini-2.5-flash") + + mock_response = MagicMock() + mock_response.candidates = [MagicMock(finish_reason="STOP")] + mock_response.text = "ok" + mock_response.usage_metadata = None + mock_client.models.generate_content.return_value = mock_response + + result = await provider._call_api("sys", "human") + assert result == "ok" + + +# ── OpenRouterProvider ── + + +class TestOpenRouterProvider: + def _make_provider(self): + with patch("chatbot_plugin.llm.openrouter_provider.httpx") as mock_httpx: + mock_client = AsyncMock() + mock_httpx.AsyncClient.return_value = mock_client + from chatbot_plugin.llm.openrouter_provider import OpenRouterProvider + provider = OpenRouterProvider(api_key="sk-test", model="test-model") + return provider, mock_client + + @pytest.mark.asyncio + async def test_call_api_success(self): + provider, mock_client = self._make_provider() + mock_response = MagicMock() + mock_response.status_code = 200 + mock_response.raise_for_status = MagicMock() + mock_response.json.return_value = { + "choices": [{"message": {"content": "Hello from OpenRouter"}}], + "usage": {"prompt_tokens": 10, "completion_tokens": 5}, + } + mock_client.post.return_value = mock_response + + result = await provider._call_api("sys", "human") + assert result == "Hello from OpenRouter" + + @pytest.mark.asyncio + async def test_call_api_uses_correct_params(self): + provider, mock_client = self._make_provider() + mock_response = MagicMock() + mock_response.raise_for_status = MagicMock() + mock_response.json.return_value = { + "choices": [{"message": {"content": "ok"}}], + } + mock_client.post.return_value = mock_response + + await provider._call_api("system-instr", "user-msg") + call_kwargs = mock_client.post.call_args + body = call_kwargs.kwargs["json"] + assert body["model"] == "test-model" + assert body["messages"][0] == {"role": "system", "content": "system-instr"} + assert body["messages"][1] == {"role": "user", "content": "user-msg"} + + @pytest.mark.asyncio + async def test_call_api_missing_usage_defaults_to_zero(self): + provider, mock_client = self._make_provider() + mock_response = MagicMock() + mock_response.raise_for_status = MagicMock() + mock_response.json.return_value = { + "choices": [{"message": {"content": "ok"}}], + } + mock_client.post.return_value = mock_response + + result = await provider._call_api("sys", "human") + assert result == "ok" diff --git a/src/tests/llm/test_rate_limit.py b/src/tests/llm/test_rate_limit.py index be94b0e..727db5c 100644 --- a/src/tests/llm/test_rate_limit.py +++ b/src/tests/llm/test_rate_limit.py @@ -61,3 +61,42 @@ async def try_acquire(): # Launch 3 acquires concurrently — all should succeed (RPM=3) await asyncio.gather(*[try_acquire() for _ in range(3)]) assert results.count("ok") == 3 + + @pytest.mark.asyncio + async def test_rpd_zero_skips_daily_check(self): + """rpd=0 means no daily limit enforced.""" + strategy = SlidingWindowStrategy(rpm=100, tpm=100000, rpd=0) + for _ in range(50): + await strategy.acquire(10) + + @pytest.mark.asyncio + async def test_record_usage_on_empty_window_is_noop(self): + """record_usage when window is empty does not crash.""" + strategy = SlidingWindowStrategy(rpm=5, tpm=10000, rpd=100) + await strategy.record_usage(100) # no prior acquire + + @pytest.mark.asyncio + async def test_tpm_exceeds_single_request(self): + """A single request with estimated_tokens > tpm still acquires.""" + strategy = SlidingWindowStrategy(rpm=100, tpm=10, rpd=100) + # This should succeed — there's no prior load + await strategy.acquire(50) + + @pytest.mark.asyncio + async def test_rpm_wait_returns_positive_when_full(self): + """When RPM window is full, _rpm_wait returns a positive wait time.""" + strategy = SlidingWindowStrategy(rpm=2, tpm=100000, rpd=100) + await strategy.acquire(10) + await strategy.acquire(10) + # Now RPM is full — compute_wait should return > 0 + wait = await strategy._compute_wait(10) + assert wait > 0 + + @pytest.mark.asyncio + async def test_tpm_wait_returns_positive_when_full(self): + """When TPM is full, _compute_wait returns a positive wait time.""" + strategy = SlidingWindowStrategy(rpm=100, tpm=20, rpd=100) + await strategy.acquire(20) + # TPM is full — compute_wait should return > 0 + wait = await strategy._compute_wait(20) + assert wait > 0 diff --git a/src/tests/llm/test_resilient_service.py b/src/tests/llm/test_resilient_service.py index 9005a95..8bfa67a 100644 --- a/src/tests/llm/test_resilient_service.py +++ b/src/tests/llm/test_resilient_service.py @@ -69,3 +69,54 @@ async def test_handlers_sorted_by_priority(self): result = await service.generate("sys", "human") assert result == "gemini reply" # priority 1 tried first + + @pytest.mark.asyncio + async def test_fallback_on_generic_exception(self): + """Non-RateLimitExhausted exceptions also trigger fallback.""" + h1 = _handler("gemini", 1, generate_side_effect=RuntimeError("boom")) + h2 = _handler("claude", 2, generate_return="claude reply") + service = ResilientLLMService([h1, h2]) + + result = await service.generate("sys", "human") + assert result == "claude reply" + + @pytest.mark.asyncio + async def test_mixed_failures_then_success(self): + """First provider RateLimitExhausted, second returns None, third succeeds.""" + h1 = _handler("gemini", 1, generate_side_effect=RateLimitExhausted("daily")) + h2 = _handler("claude", 2, generate_return=None) + h3 = _handler("openrouter", 3, generate_return="or reply") + service = ResilientLLMService([h1, h2, h3]) + + result = await service.generate("sys", "human") + assert result == "or reply" + + +class TestProviderHandler: + @pytest.mark.asyncio + async def test_generate_acquires_and_records_on_success(self): + """ProviderHandler calls strategy.acquire and strategy.record_usage on success.""" + from chatbot_plugin.llm.resilient_llm_service import ProviderHandler + strategy = AsyncMock() + provider = AsyncMock() + provider.generate.return_value = "success" + handler = ProviderHandler(provider=provider, strategy=strategy, priority=1, name="test") + + result = await handler.generate("sys", "human") + assert result == "success" + strategy.acquire.assert_called_once() + strategy.record_usage.assert_called_once() + + @pytest.mark.asyncio + async def test_generate_skips_record_usage_on_none(self): + """ProviderHandler does not call record_usage when provider returns None.""" + from chatbot_plugin.llm.resilient_llm_service import ProviderHandler + strategy = AsyncMock() + provider = AsyncMock() + provider.generate.return_value = None + handler = ProviderHandler(provider=provider, strategy=strategy, priority=1, name="test") + + result = await handler.generate("sys", "human") + assert result is None + strategy.acquire.assert_called_once() + strategy.record_usage.assert_not_called() diff --git a/src/tests/routers/test_chat.py b/src/tests/routers/test_chat.py index 82e27eb..be60275 100644 --- a/src/tests/routers/test_chat.py +++ b/src/tests/routers/test_chat.py @@ -109,3 +109,60 @@ async def test_status_returns_shape(self, client: AsyncClient): assert "total_chunks" in data assert "last_indexed_at" in data assert "pending_articles" in data + + +# ── Additional router branch tests ── + + +class TestRouterAdditional: + @pytest.mark.asyncio + async def test_message_with_user_id(self, client: AsyncClient): + """Router passes user_id through to service.chat.""" + with patch("chatbot_plugin.service.ChatbotService.chat", new_callable=AsyncMock) as mock: + mock.return_value = MagicMock( + reply="hi", articles_used=[], model_dump=lambda: {"reply": "hi", "articles_used": []} + ) + resp = await client.post("/chat/message", json={"message": "hello", "user_id": "user-1"}) + assert resp.status_code == 200 + mock.assert_called_once_with("hello", "user-1") + + @pytest.mark.asyncio + async def test_search_with_topic_id(self, client: AsyncClient): + """Router passes topic_id through to service.search.""" + with patch("chatbot_plugin.service.ChatbotService.search", new_callable=AsyncMock) as mock: + mock.return_value = MagicMock( + chunks=[], model_dump=lambda: {"chunks": []} + ) + resp = await client.post("/chat/search", json={"query": "RAG", "topic_id": "topic-1"}) + assert resp.status_code == 200 + mock.assert_called_once_with("RAG", 10, "topic-1") + + @pytest.mark.asyncio + async def test_index_with_article_id(self, client: AsyncClient): + """Router passes article_id through to service.trigger_index.""" + with patch("chatbot_plugin.service.ChatbotService.trigger_index", new_callable=AsyncMock) as mock: + mock.return_value = MagicMock( + job_id="job-1", status="started", + model_dump=lambda: {"job_id": "job-1", "status": "started"} + ) + resp = await client.post("/chat/index", json={"article_id": "uuid-1"}) + assert resp.status_code == 202 + mock.assert_called_once_with("uuid-1") + + @pytest.mark.asyncio + async def test_message_503_on_llm_failure(self, client: AsyncClient): + """Router returns 503 when service.chat raises HTTPException(503).""" + from fastapi import HTTPException + with patch("chatbot_plugin.service.ChatbotService.chat", new_callable=AsyncMock) as mock: + mock.side_effect = HTTPException(status_code=503, detail="LLM provider unavailable") + resp = await client.post("/chat/message", json={"message": "hello"}) + assert resp.status_code == 503 + + @pytest.mark.asyncio + async def test_set_llm_service_overrides(self, client: AsyncClient): + """set_llm_service can override the LLM service for testing.""" + from chatbot_plugin.routers import set_llm_service, _llm_service + # The conftest already called set_llm_service, so _llm_service is not None + # Just verify it was set + from chatbot_plugin import routers + assert routers._llm_service is not None diff --git a/src/tests/test_chain.py b/src/tests/test_chain.py new file mode 100644 index 0000000..603515f --- /dev/null +++ b/src/tests/test_chain.py @@ -0,0 +1,35 @@ +"""Tests for RAG chain (rag_generate).""" + +import pytest +from unittest.mock import AsyncMock + +from chatbot_plugin.llm.rate_limit.quota_strategy import RateLimitExhausted +from chatbot_plugin.rag.chain import rag_generate + + +@pytest.mark.asyncio +async def test_rag_generate_success(): + mock_llm = AsyncMock() + mock_llm.generate.return_value = "RAG is retrieval-augmented generation." + result = await rag_generate("What is RAG?", [], mock_llm) + assert "RAG" in result + + +@pytest.mark.asyncio +async def test_rag_generate_none_raises_runtime_error(): + mock_llm = AsyncMock() + mock_llm.generate.return_value = None + with pytest.raises(RuntimeError, match="All LLM providers failed"): + await rag_generate("hello", [], mock_llm) + + +@pytest.mark.asyncio +async def test_rag_generate_passes_articles_to_prompt(): + articles = [{"title": "AI", "content": "Artificial intelligence."}] + mock_llm = AsyncMock() + mock_llm.generate.return_value = "AI response" + result = await rag_generate("What is AI?", articles, mock_llm) + assert result == "AI response" + # Verify generate was called with system and human prompts + call_args = mock_llm.generate.call_args + assert len(call_args.args) == 2 # system_prompt, human_prompt diff --git a/src/tests/test_prompt.py b/src/tests/test_prompt.py new file mode 100644 index 0000000..01c5868 --- /dev/null +++ b/src/tests/test_prompt.py @@ -0,0 +1,93 @@ +"""Tests for RAG prompt building.""" + +import pytest + +from chatbot_plugin.rag.prompt import build_context, build_messages, SYSTEM_PROMPT + + +# ── build_context() ── + + +def test_build_context_empty_articles(): + assert build_context([]) == "" + + +def test_build_context_single_article_fits(): + articles = [{"title": "AI", "content": "Artificial intelligence."}] + result = build_context(articles, max_tokens=100) + assert "[source: AI]" in result + assert "Artificial intelligence." in result + + +def test_build_context_multiple_articles(): + articles = [ + {"title": "A", "content": "Content A."}, + {"title": "B", "content": "Content B."}, + ] + result = build_context(articles, max_tokens=100) + assert "[source: A]" in result + assert "[source: B]" in result + + +def test_build_context_truncation_with_ellipsis(): + long_content = "x" * 500 + articles = [{"title": "Long", "content": long_content}] + result = build_context(articles, max_tokens=50) + assert "[source: Long]" in result + assert result.endswith("...") + + +def test_build_context_budget_exhausted_skips_remaining(): + articles = [ + {"title": "First", "content": "x" * 200}, + {"title": "Second", "content": "y" * 200}, + ] + result = build_context(articles, max_tokens=30) + assert "[source: First]" in result + assert "[source: Second]" not in result + + +def test_build_context_too_small_budget_skips_article(): + articles = [{"title": "Tiny", "content": "x" * 500}] + result = build_context(articles, max_tokens=5) + assert "[source: Tiny]" not in result or len(result) < 200 + + +def test_build_context_untitled_default(): + articles = [{"content": "No title here."}] + result = build_context(articles, max_tokens=100) + assert "[source: Untitled]" in result + + +def test_build_context_empty_content(): + articles = [{"title": "Empty", "content": ""}] + result = build_context(articles, max_tokens=100) + assert "[source: Empty]" in result + + +def test_build_context_defaults_to_settings_max_tokens(): + articles = [{"title": "Test", "content": "Hello"}] + result = build_context(articles) + assert "[source: Test]" in result + + +# ── build_messages() ── + + +def test_build_messages_returns_tuple(): + system, human = build_messages("What is AI?", []) + assert system == SYSTEM_PROMPT + assert "What is AI?" in human + + +def test_build_messages_empty_articles_fallback(): + system, human = build_messages("Hello", []) + assert "No relevant articles found." in human + + +def test_build_messages_with_articles(): + articles = [{"title": "AI", "content": "Artificial intelligence."}] + system, human = build_messages("What is AI?", articles) + assert system == SYSTEM_PROMPT + assert "[source: AI]" in human + assert "What is AI?" in human diff --git a/src/tests/test_service.py b/src/tests/test_service.py index 95858d0..a404912 100644 --- a/src/tests/test_service.py +++ b/src/tests/test_service.py @@ -106,3 +106,83 @@ async def test_get_status_returns_shape(service, mock_db, mock_llm_service): assert result.pending_articles == 42 assert result.total_chunks == 0 assert result.last_indexed_at is None + + +# ── Missing branch tests ── + + +@pytest.mark.asyncio +async def test_chat_generic_exception_raises_503(service, mock_db, mock_llm_service): + """Non-RuntimeError exceptions from rag_generate also produce 503.""" + mock_db.execute.return_value = _mock_result(rows=[]) + mock_llm_service.generate.side_effect = Exception("unexpected error") + + with pytest.raises(HTTPException) as exc_info: + await service.chat("hello") + assert exc_info.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_chat_untitled_fallback_for_none_title(service, mock_db, mock_llm_service): + """Articles with None/empty title get 'Untitled' fallback.""" + mock_db.execute.return_value = _mock_result( + rows=[{"id": "uuid-1", "title": None, "content": "Some content", "rank": 0.5}] + ) + mock_llm_service.generate.return_value = "Reply" + + result = await service.chat("hello") + assert result.articles_used[0].title == "Untitled" + + +@pytest.mark.asyncio +async def test_search_with_topic_id(service, mock_db, mock_llm_service): + """search() with topic_id passes it through to _search_articles.""" + mock_db.execute.return_value = _mock_result( + rows=[{"id": "uuid-1", "title": "Article A", "content": "Content A", "rank": 0.8}] + ) + + result = await service.search("RAG", topic_id="topic-uuid") + assert len(result.chunks) == 1 + # Verify the SQL was executed (topic_id passed to params) + mock_db.execute.assert_called_once() + + +@pytest.mark.asyncio +async def test_search_untitled_fallback(service, mock_db, mock_llm_service): + """Search results with None/empty title get 'Untitled' fallback.""" + mock_db.execute.return_value = _mock_result( + rows=[{"id": "uuid-1", "title": "", "content": "Content", "rank": 0.5}] + ) + + result = await service.search("test") + assert result.chunks[0].article_title == "Untitled" + + +@pytest.mark.asyncio +async def test_search_empty_content_fallback(service, mock_db, mock_llm_service): + """Search results with None content get empty string fallback.""" + mock_db.execute.return_value = _mock_result( + rows=[{"id": "uuid-1", "title": "Title", "content": None, "rank": 0.5}] + ) + + result = await service.search("test") + assert result.chunks[0].content == "" + + +@pytest.mark.asyncio +async def test_trigger_index_with_article_found(service, mock_db, mock_llm_service): + """trigger_index with article_id where the article exists.""" + mock_db.execute.return_value = _mock_result(scalar_val="uuid-1") + + result = await service.trigger_index(article_id="uuid-1") + assert result.status == "started" + assert result.job_id + + +@pytest.mark.asyncio +async def test_get_status_with_none_scalar(service, mock_db, mock_llm_service): + """get_status when count(*) returns None — should default to 0.""" + mock_db.execute.return_value = _mock_result(scalar_val=None) + + result = await service.get_status() + assert result.pending_articles == 0