Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions src/chatbot_plugin/llm/rate_limit/sliding_window_strategy.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
91 changes: 91 additions & 0 deletions src/tests/llm/test_base_provider.py
Original file line number Diff line number Diff line change
@@ -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
114 changes: 114 additions & 0 deletions src/tests/llm/test_bootstrap.py
Original file line number Diff line number Diff line change
@@ -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)
Loading
Loading