embeddings_queries now records the embedding request's token count (the unit upstream providers bill on) instead of an operation count, and fires on the indexing path too. Previously only semantic search recorded it (value=1), so a re-indexing run produced no embeddings_queries events at all — only pages_chunks. - Provider layer: additive embed_with_usage / embed_batch_with_usage surface the per-request token count (Mistral/OpenAI usage.total_tokens, Bedrock Titan inputTextTokenCount, Ollama prompt_eval_count); a char-based estimate is the fallback (Simple, and any provider/response without a token field). Gateway and EmbeddingService forward through. The count travels as a return value / a per-request SearchAlgorithm attribute — never on the singleton — so concurrent indexing + search can't mis-attribute bills. - Indexing (vector/processor.py): records embeddings_queries (value=batch tokens) alongside the existing pages_chunks event. - Search (server/semantic.py): value is now the query embedding's token count, relayed from BM25HybridSearchAlgorithm via query_token_count. The astrolabe_embeddings_queries Stripe meter (sum aggregation) now sums tokens with no CP/Terraform change. The meter "queries"->tokens naming/unit clarification (homelab-terraform #254) + CP rollup/portal copy is a follow-up. Deck #67. Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
332 lines
12 KiB
Python
332 lines
12 KiB
Python
"""Unit tests for Mistral provider."""
|
|
|
|
from unittest.mock import AsyncMock, MagicMock
|
|
|
|
import pytest
|
|
from mistralai.client.errors import SDKError
|
|
|
|
from nextcloud_mcp_server.providers.mistral import (
|
|
BATCH_SIZE,
|
|
MISTRAL_EMBEDDING_DIMENSIONS,
|
|
MistralProvider,
|
|
_is_rate_limit,
|
|
)
|
|
|
|
|
|
def _make_data(embedding: list[float], index: int) -> MagicMock:
|
|
"""Build a mock EmbeddingResponseData entry."""
|
|
item = MagicMock()
|
|
item.embedding = embedding
|
|
item.index = index
|
|
return item
|
|
|
|
|
|
def _make_response(embeddings: list[list[float]]) -> MagicMock:
|
|
"""Build a mock EmbeddingResponse with `embeddings` indexed in order."""
|
|
response = MagicMock()
|
|
response.data = [_make_data(emb, i) for i, emb in enumerate(embeddings)]
|
|
return response
|
|
|
|
|
|
@pytest.fixture
|
|
def mock_mistral_client(mocker):
|
|
"""Mock the Mistral SDK constructor."""
|
|
mock_client = MagicMock()
|
|
mock_client.embeddings = MagicMock()
|
|
mocker.patch(
|
|
"nextcloud_mcp_server.providers.mistral.Mistral", return_value=mock_client
|
|
)
|
|
return mock_client
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embedding_single(mock_mistral_client):
|
|
"""Single text embed: round-trip through SDK with correct kwargs."""
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(
|
|
return_value=_make_response([[0.1, 0.2, 0.3]])
|
|
)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
embedding = await provider.embed("hello world")
|
|
|
|
assert embedding == [0.1, 0.2, 0.3]
|
|
mock_mistral_client.embeddings.create_async.assert_awaited_once_with(
|
|
model="mistral-embed",
|
|
inputs=["hello world"],
|
|
)
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embedding_batch_single_call(mock_mistral_client):
|
|
"""Batch smaller than BATCH_SIZE issues a single API call."""
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(
|
|
return_value=_make_response([[0.1, 0.2], [0.3, 0.4], [0.5, 0.6]])
|
|
)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
embeddings = await provider.embed_batch(["a", "b", "c"])
|
|
|
|
assert embeddings == [[0.1, 0.2], [0.3, 0.4], [0.5, 0.6]]
|
|
assert mock_mistral_client.embeddings.create_async.await_count == 1
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embedding_batch_chunking(mock_mistral_client):
|
|
"""Batches exceeding BATCH_SIZE are split into multiple API calls."""
|
|
|
|
# Each call returns one embedding per input it received; capture by side
|
|
# effect so we can inspect lengths per chunk.
|
|
def _side_effect(*, model, inputs, **_kwargs):
|
|
return _make_response([[float(i)] for i in range(len(inputs))])
|
|
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(side_effect=_side_effect)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
total = BATCH_SIZE * 2 + 5 # forces three chunks: 64, 64, 5 (with default)
|
|
embeddings = await provider.embed_batch([f"text-{i}" for i in range(total)])
|
|
|
|
assert len(embeddings) == total
|
|
assert mock_mistral_client.embeddings.create_async.await_count == 3
|
|
# Verify the chunk sizes the SDK was actually called with.
|
|
chunk_sizes = [
|
|
len(call.kwargs["inputs"])
|
|
for call in mock_mistral_client.embeddings.create_async.await_args_list
|
|
]
|
|
assert chunk_sizes == [BATCH_SIZE, BATCH_SIZE, 5]
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embedding_batch_order_preserved(mock_mistral_client):
|
|
"""Out-of-order index in response data is sorted before returning."""
|
|
response = MagicMock()
|
|
response.data = [
|
|
_make_data([0.3, 0.3], 2),
|
|
_make_data([0.1, 0.1], 0),
|
|
_make_data([0.2, 0.2], 1),
|
|
]
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(return_value=response)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
embeddings = await provider.embed_batch(["x", "y", "z"])
|
|
|
|
assert embeddings == [[0.1, 0.1], [0.2, 0.2], [0.3, 0.3]]
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_supports_capabilities(mock_mistral_client):
|
|
"""Mistral provider advertises embeddings only."""
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
assert provider.supports_embeddings is True
|
|
assert provider.supports_generation is False
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_generate_not_implemented(mock_mistral_client):
|
|
"""generate() always raises NotImplementedError."""
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
with pytest.raises(NotImplementedError, match="does not support generation"):
|
|
await provider.generate("test prompt")
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_get_dimension_known_model(mock_mistral_client):
|
|
"""Known model: dimension available without an API call."""
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
assert provider.get_dimension() == MISTRAL_EMBEDDING_DIMENSIONS["mistral-embed"]
|
|
mock_mistral_client.embeddings.create_async.assert_not_called()
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_get_dimension_unknown_model_detected(mock_mistral_client):
|
|
"""Unknown model: dimension detected on first embed() call."""
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(
|
|
return_value=_make_response([[0.1] * 768])
|
|
)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="custom-mistral")
|
|
|
|
with pytest.raises(RuntimeError, match="not detected yet"):
|
|
provider.get_dimension()
|
|
|
|
await provider.embed("test")
|
|
assert provider.get_dimension() == 768
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_no_embeddings_disabled(mock_mistral_client):
|
|
"""Setting embedding_model=None disables the embedding capability."""
|
|
provider = MistralProvider(api_key="test-key", embedding_model=None)
|
|
assert provider.supports_embeddings is False
|
|
|
|
with pytest.raises(NotImplementedError, match="no embedding_model configured"):
|
|
await provider.embed("test")
|
|
with pytest.raises(NotImplementedError, match="no embedding_model configured"):
|
|
await provider.embed_batch(["test"])
|
|
with pytest.raises(NotImplementedError, match="no embedding_model configured"):
|
|
provider.get_dimension()
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_empty_batch(mock_mistral_client):
|
|
"""An empty batch returns [] without calling the API."""
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
assert await provider.embed_batch([]) == []
|
|
mock_mistral_client.embeddings.create_async.assert_not_called()
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_close_no_error(mock_mistral_client):
|
|
"""close() is best-effort and does not raise."""
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
# No __aexit__ on the mock by default → close() should silently no-op.
|
|
await provider.close()
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_base_url_passed_to_sdk(mocker):
|
|
"""base_url is forwarded as server_url to the Mistral SDK constructor."""
|
|
mock_ctor = mocker.patch(
|
|
"nextcloud_mcp_server.providers.mistral.Mistral", return_value=MagicMock()
|
|
)
|
|
|
|
MistralProvider(
|
|
api_key="test-key",
|
|
embedding_model="mistral-embed",
|
|
base_url="https://example.com/mistral",
|
|
)
|
|
mock_ctor.assert_called_once_with(
|
|
api_key="test-key",
|
|
server_url="https://example.com/mistral",
|
|
)
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embed_raises_on_empty_response_data(mock_mistral_client):
|
|
"""embed(): empty response.data triggers the defensive RuntimeError guard."""
|
|
empty_response = MagicMock()
|
|
empty_response.data = []
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(return_value=empty_response)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
with pytest.raises(RuntimeError, match="returned no embedding"):
|
|
await provider.embed("test")
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embed_raises_on_null_embedding(mock_mistral_client):
|
|
"""embed(): a single response item with embedding=None is rejected."""
|
|
null_item = MagicMock()
|
|
null_item.embedding = None
|
|
null_item.index = 0
|
|
null_response = MagicMock()
|
|
null_response.data = [null_item]
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(return_value=null_response)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
with pytest.raises(RuntimeError, match="returned no embedding"):
|
|
await provider.embed("test")
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_batch_raises_on_null_embedding(mock_mistral_client):
|
|
"""_embed_batch_request: a null embedding inside a batch raises explicitly."""
|
|
good = _make_data([0.1, 0.2], 0)
|
|
bad = MagicMock()
|
|
bad.embedding = None
|
|
bad.index = 1
|
|
response = MagicMock()
|
|
response.data = [good, bad]
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(return_value=response)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
with pytest.raises(RuntimeError, match="null embedding"):
|
|
await provider.embed_batch(["a", "b"])
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_batch_raises_on_count_mismatch(mock_mistral_client):
|
|
"""_embed_batch_request: fewer embeddings returned than inputs sent."""
|
|
# Two inputs sent, one embedding returned.
|
|
response = _make_response([[0.1, 0.2]])
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(return_value=response)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
with pytest.raises(RuntimeError, match="returned 1 embeddings for 2 inputs"):
|
|
await provider.embed_batch(["a", "b"])
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embed_batch_with_usage_reports_tokens(mock_mistral_client):
|
|
"""embed_batch_with_usage returns the provider-reported total_tokens."""
|
|
response = _make_response([[0.1, 0.2], [0.3, 0.4]])
|
|
response.usage = MagicMock(total_tokens=11)
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(return_value=response)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
embeddings, tokens = await provider.embed_batch_with_usage(["a", "b"])
|
|
|
|
assert embeddings == [[0.1, 0.2], [0.3, 0.4]]
|
|
assert tokens == 11
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embed_batch_with_usage_sums_across_chunks(mock_mistral_client):
|
|
"""Token counts sum across the BATCH_SIZE sub-requests (1 token/input here)."""
|
|
|
|
def _side_effect(*, model, inputs, **_kwargs):
|
|
resp = _make_response([[float(i)] for i in range(len(inputs))])
|
|
resp.usage = MagicMock(total_tokens=len(inputs))
|
|
return resp
|
|
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(side_effect=_side_effect)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
total = BATCH_SIZE * 2 + 5 # three chunks
|
|
embeddings, tokens = await provider.embed_batch_with_usage(
|
|
[f"t-{i}" for i in range(total)]
|
|
)
|
|
|
|
assert len(embeddings) == total
|
|
assert tokens == total # summed across all three chunks
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_with_usage_estimates_when_usage_absent(mock_mistral_client):
|
|
"""Missing usage falls back to the char-based estimate, not a crash."""
|
|
response = _make_response([[0.1, 0.2]])
|
|
response.usage = None
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(return_value=response)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
_, tokens = await provider.embed_batch_with_usage(["abcd"]) # 4 chars → 1 token
|
|
|
|
assert tokens == 1
|
|
|
|
|
|
@pytest.mark.unit
|
|
async def test_mistral_embed_with_usage_single(mock_mistral_client):
|
|
"""embed_with_usage returns the single embedding plus its token count."""
|
|
response = _make_response([[0.5, 0.6]])
|
|
response.usage = MagicMock(total_tokens=3)
|
|
mock_mistral_client.embeddings.create_async = AsyncMock(return_value=response)
|
|
|
|
provider = MistralProvider(api_key="test-key", embedding_model="mistral-embed")
|
|
embedding, tokens = await provider.embed_with_usage("hello")
|
|
|
|
assert embedding == [0.5, 0.6]
|
|
assert tokens == 3
|
|
|
|
|
|
@pytest.mark.unit
|
|
def test_mistral_is_rate_limit_predicate():
|
|
"""_is_rate_limit returns True only for SDKErrors with status_code == 429."""
|
|
err_429 = MagicMock(spec=SDKError)
|
|
err_429.status_code = 429
|
|
err_500 = MagicMock(spec=SDKError)
|
|
err_500.status_code = 500
|
|
|
|
assert _is_rate_limit(err_429) is True
|
|
assert _is_rate_limit(err_500) is False
|
|
# ValueError has no status_code attr → getattr returns None → False.
|
|
assert _is_rate_limit(ValueError()) is False
|