Merge pull request #875 from cbcoutinho/feat/meter-embedding-tokens
feat(usage): meter embedding tokens (tokens_embedded/pages_embedded) on both paths + Prometheus export
This commit is contained in:
@@ -255,6 +255,73 @@ async def test_bedrock_dimension_detection(mock_bedrock_client):
|
||||
assert provider.get_dimension() == 1536
|
||||
|
||||
|
||||
def _titan_body(embedding, token_count=None):
|
||||
payload = {"embedding": embedding}
|
||||
if token_count is not None:
|
||||
payload["inputTextTokenCount"] = token_count
|
||||
return {
|
||||
"body": MagicMock(read=MagicMock(return_value=json.dumps(payload).encode()))
|
||||
}
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_bedrock_embed_with_usage_reports_titan_tokens(mock_bedrock_client):
|
||||
"""Titan's inputTextTokenCount is surfaced as the token count."""
|
||||
mock_bedrock_client.invoke_model.return_value = _titan_body(
|
||||
[0.1, 0.2], token_count=6
|
||||
)
|
||||
|
||||
provider = BedrockProvider(
|
||||
region_name="us-east-1",
|
||||
embedding_model="amazon.titan-embed-text-v2:0",
|
||||
generation_model=None,
|
||||
)
|
||||
embedding, tokens = await provider.embed_with_usage("test text")
|
||||
|
||||
assert embedding == [0.1, 0.2]
|
||||
assert tokens == 6
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_bedrock_embed_batch_with_usage_sums_token_counts(mock_bedrock_client):
|
||||
"""Sequential per-text calls sum their inputTextTokenCount values."""
|
||||
mock_bedrock_client.invoke_model.return_value = _titan_body(
|
||||
[0.1, 0.2], token_count=4
|
||||
)
|
||||
|
||||
provider = BedrockProvider(
|
||||
region_name="us-east-1",
|
||||
embedding_model="amazon.titan-embed-text-v2:0",
|
||||
generation_model=None,
|
||||
)
|
||||
embeddings, tokens = await provider.embed_batch_with_usage(["t1", "t2", "t3"])
|
||||
|
||||
assert len(embeddings) == 3
|
||||
assert tokens == 12 # 4 tokens per call × 3 calls
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_bedrock_with_usage_estimates_when_token_count_absent(
|
||||
mock_bedrock_client,
|
||||
):
|
||||
"""Cohere returns no inputTextTokenCount → char-based estimate."""
|
||||
mock_bedrock_client.invoke_model.return_value = {
|
||||
"body": MagicMock(
|
||||
read=MagicMock(
|
||||
return_value=json.dumps({"embeddings": [[0.1, 0.2]]}).encode()
|
||||
)
|
||||
)
|
||||
}
|
||||
|
||||
provider = BedrockProvider(
|
||||
region_name="us-east-1",
|
||||
embedding_model="cohere.embed-english-v3",
|
||||
)
|
||||
_, tokens = await provider.embed_with_usage("abcdefgh") # 8 chars → 2 tokens
|
||||
|
||||
assert tokens == 2
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_bedrock_cohere_embedding(mock_bedrock_client):
|
||||
"""Test Bedrock with Cohere embedding model."""
|
||||
|
||||
@@ -6,6 +6,7 @@ tenant realm); creds are all-or-nothing.
|
||||
"""
|
||||
|
||||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
@@ -327,6 +328,105 @@ def test_trailing_slash_base_url_normalized():
|
||||
assert not base.endswith("/v1/v1")
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_gateway_embed_with_usage_forwards_after_bearer(monkeypatch):
|
||||
"""embed_with_usage refreshes the bearer, then returns the (embedding,
|
||||
token_count) from the inherited OpenAI implementation."""
|
||||
# https mock host (never contacted — the OpenAI client is patched below).
|
||||
provider = GatewayProvider(
|
||||
base_url="https://gw:8083/v1", embedding_model="mistral/mistral-embed"
|
||||
)
|
||||
|
||||
order: list[str] = []
|
||||
|
||||
async def _ensure_bearer():
|
||||
order.append("bearer")
|
||||
|
||||
monkeypatch.setattr(provider, "_ensure_bearer", _ensure_bearer)
|
||||
|
||||
item = MagicMock()
|
||||
item.embedding = [0.1, 0.2]
|
||||
item.index = 0
|
||||
response = MagicMock()
|
||||
response.data = [item]
|
||||
response.usage = MagicMock(total_tokens=8)
|
||||
|
||||
async def _create(**_kwargs):
|
||||
order.append("embed")
|
||||
return response
|
||||
|
||||
monkeypatch.setattr(provider.client.embeddings, "create", _create)
|
||||
|
||||
embedding, tokens = await provider.embed_with_usage("hello")
|
||||
|
||||
assert embedding == [0.1, 0.2]
|
||||
assert tokens == 8
|
||||
assert order == ["bearer", "embed"] # bearer refreshed before the embed call
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_gateway_embed_batch_with_usage_forwards_after_bearer(monkeypatch):
|
||||
"""embed_batch_with_usage also refreshes the bearer before delegating."""
|
||||
# https mock host (never contacted — the OpenAI client is patched below).
|
||||
provider = GatewayProvider(
|
||||
base_url="https://gw:8083/v1", embedding_model="mistral/mistral-embed"
|
||||
)
|
||||
ensured = {"n": 0}
|
||||
|
||||
async def _ensure_bearer():
|
||||
ensured["n"] += 1
|
||||
|
||||
monkeypatch.setattr(provider, "_ensure_bearer", _ensure_bearer)
|
||||
|
||||
item = MagicMock()
|
||||
item.embedding = [0.3, 0.4]
|
||||
item.index = 0
|
||||
response = MagicMock()
|
||||
response.data = [item]
|
||||
response.usage = MagicMock(total_tokens=5)
|
||||
monkeypatch.setattr(
|
||||
provider.client.embeddings, "create", AsyncMock(return_value=response)
|
||||
)
|
||||
|
||||
embeddings, tokens = await provider.embed_batch_with_usage(["x"])
|
||||
|
||||
assert embeddings == [[0.3, 0.4]]
|
||||
assert tokens == 5
|
||||
assert ensured["n"] == 1
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_gateway_embed_batch_ensures_bearer_once(monkeypatch):
|
||||
"""embed_batch() has no override: it routes through the inherited OpenAI
|
||||
embed_batch() → embed_batch_with_usage() (overridden), so the bearer is
|
||||
refreshed exactly once — not twice."""
|
||||
# https mock host (never contacted — the OpenAI client is patched below).
|
||||
provider = GatewayProvider(
|
||||
base_url="https://gw:8083/v1", embedding_model="mistral/mistral-embed"
|
||||
)
|
||||
ensured = {"n": 0}
|
||||
|
||||
async def _ensure_bearer():
|
||||
ensured["n"] += 1
|
||||
|
||||
monkeypatch.setattr(provider, "_ensure_bearer", _ensure_bearer)
|
||||
|
||||
item = MagicMock()
|
||||
item.embedding = [0.1, 0.2]
|
||||
item.index = 0
|
||||
response = MagicMock()
|
||||
response.data = [item]
|
||||
response.usage = MagicMock(total_tokens=4)
|
||||
monkeypatch.setattr(
|
||||
provider.client.embeddings, "create", AsyncMock(return_value=response)
|
||||
)
|
||||
|
||||
embeddings = await provider.embed_batch(["x"])
|
||||
|
||||
assert embeddings == [[0.1, 0.2]]
|
||||
assert ensured["n"] == 1 # not 2 — embed_batch() must not double-refresh
|
||||
|
||||
|
||||
async def test_detect_dimension_with_bare_base_url_hits_v1_models(monkeypatch):
|
||||
"""End-to-end of the fix: a bare-origin base_url still resolves the
|
||||
dimension because discovery lands on /v1/models."""
|
||||
|
||||
@@ -255,6 +255,68 @@ async def test_mistral_batch_raises_on_count_mismatch(mock_mistral_client):
|
||||
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."""
|
||||
|
||||
@@ -0,0 +1,69 @@
|
||||
"""Unit tests for Ollama provider token-usage surfacing.
|
||||
|
||||
The provider has no other unit coverage; these focus on the ``*_with_usage``
|
||||
methods added for usage metering (Deck #67) — provider-reported
|
||||
``prompt_eval_count`` and the char-based estimate fallback when it's absent.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from nextcloud_mcp_server.providers.ollama import OllamaProvider
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def ollama_provider():
|
||||
# Construct with no models so __init__ skips _check_model_is_loaded (no
|
||||
# network call), then enable embeddings post-construction. https mock host
|
||||
# (never contacted — client.post is patched in each test).
|
||||
provider = OllamaProvider(base_url="https://ollama:11434")
|
||||
provider.embedding_model = "nomic-embed-text"
|
||||
return provider
|
||||
|
||||
|
||||
def _embed_response(embeddings, prompt_eval_count=None):
|
||||
payload = {"embeddings": embeddings}
|
||||
if prompt_eval_count is not None:
|
||||
payload["prompt_eval_count"] = prompt_eval_count
|
||||
resp = MagicMock()
|
||||
resp.json = MagicMock(return_value=payload)
|
||||
resp.raise_for_status = MagicMock()
|
||||
return resp
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_ollama_embed_batch_with_usage_reports_prompt_eval_count(ollama_provider):
|
||||
"""prompt_eval_count from /api/embed is surfaced as the token count."""
|
||||
ollama_provider.client.post = AsyncMock(
|
||||
return_value=_embed_response([[0.1, 0.2], [0.3, 0.4]], prompt_eval_count=7)
|
||||
)
|
||||
|
||||
embeddings, tokens = await ollama_provider.embed_batch_with_usage(["a", "b"])
|
||||
|
||||
assert embeddings == [[0.1, 0.2], [0.3, 0.4]]
|
||||
assert tokens == 7
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_ollama_with_usage_estimates_when_count_absent(ollama_provider):
|
||||
"""Older Ollama omits prompt_eval_count → char-based estimate."""
|
||||
ollama_provider.client.post = AsyncMock(
|
||||
return_value=_embed_response([[0.1]], prompt_eval_count=None)
|
||||
)
|
||||
|
||||
_, tokens = await ollama_provider.embed_with_usage("abcdefgh") # 8 chars → 2
|
||||
|
||||
assert tokens == 2
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_ollama_empty_batch_with_usage(ollama_provider):
|
||||
"""Empty batch returns no embeddings, zero tokens, and makes no request."""
|
||||
ollama_provider.client.post = AsyncMock()
|
||||
|
||||
embeddings, tokens = await ollama_provider.embed_batch_with_usage([])
|
||||
|
||||
assert embeddings == []
|
||||
assert tokens == 0
|
||||
ollama_provider.client.post.assert_not_called()
|
||||
@@ -280,6 +280,46 @@ async def test_openai_empty_batch():
|
||||
assert embeddings == []
|
||||
|
||||
|
||||
def _embed_item(embedding, index):
|
||||
item = MagicMock()
|
||||
item.embedding = embedding
|
||||
item.index = index
|
||||
return item
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_openai_embed_batch_with_usage_reports_tokens(mock_openai_client):
|
||||
"""embed_batch_with_usage returns the response's total_tokens."""
|
||||
response = MagicMock()
|
||||
response.data = [_embed_item([0.1, 0.2], 0), _embed_item([0.3, 0.4], 1)]
|
||||
response.usage = MagicMock(total_tokens=9)
|
||||
mock_openai_client.embeddings.create = AsyncMock(return_value=response)
|
||||
|
||||
provider = OpenAIProvider(
|
||||
api_key="test-key", embedding_model="text-embedding-3-small"
|
||||
)
|
||||
embeddings, tokens = await provider.embed_batch_with_usage(["a", "b"])
|
||||
|
||||
assert embeddings == [[0.1, 0.2], [0.3, 0.4]]
|
||||
assert tokens == 9
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_openai_with_usage_estimates_when_usage_absent(mock_openai_client):
|
||||
"""Missing usage falls back to the char-based estimate."""
|
||||
response = MagicMock()
|
||||
response.data = [_embed_item([0.1], 0)]
|
||||
response.usage = None
|
||||
mock_openai_client.embeddings.create = AsyncMock(return_value=response)
|
||||
|
||||
provider = OpenAIProvider(
|
||||
api_key="test-key", embedding_model="text-embedding-3-small"
|
||||
)
|
||||
_, tokens = await provider.embed_with_usage("abcdefgh") # 8 chars → 2 tokens
|
||||
|
||||
assert tokens == 2
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_openai_close(mock_openai_client):
|
||||
"""Test OpenAI client close."""
|
||||
|
||||
@@ -0,0 +1,52 @@
|
||||
"""Token-usage surfacing: the Provider ABC estimate default + SimpleProvider.
|
||||
|
||||
The usage-metering hooks (Deck #67) bill ``tokens_embedded`` by tokens. Real
|
||||
providers report exact counts from their API response; providers without a token
|
||||
field (Simple, and the ABC default) fall back to a char-based estimate so the
|
||||
billable value stays non-zero and monotone with input size.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from nextcloud_mcp_server.providers.base import Provider
|
||||
from nextcloud_mcp_server.providers.simple import SimpleProvider
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
def test_estimate_tokens_is_char_based():
|
||||
"""~4-chars-per-token, ceil-rounded, summed across inputs."""
|
||||
assert Provider._estimate_tokens(["abcd"]) == 1 # 4 chars
|
||||
assert Provider._estimate_tokens(["abcde"]) == 2 # 5 chars → ceil(5/4)
|
||||
assert Provider._estimate_tokens(["ab", "cd"]) == 1 # 4 chars total
|
||||
assert Provider._estimate_tokens([]) == 0
|
||||
assert Provider._estimate_tokens([""]) == 0
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_simple_provider_embed_with_usage_estimates():
|
||||
"""SimpleProvider has no real usage → estimate path via the ABC default."""
|
||||
provider = SimpleProvider(dimension=8)
|
||||
embedding, tokens = await provider.embed_with_usage("abcdefgh") # 8 chars → 2
|
||||
|
||||
assert len(embedding) == 8
|
||||
assert tokens == 2
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_simple_provider_embed_batch_with_usage_estimates():
|
||||
"""Batch estimate sums character counts across all inputs."""
|
||||
provider = SimpleProvider(dimension=8)
|
||||
embeddings, tokens = await provider.embed_batch_with_usage(["abcd", "efgh"])
|
||||
|
||||
assert len(embeddings) == 2
|
||||
assert tokens == 2 # 8 chars total → 2 tokens
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_simple_provider_empty_batch_with_usage():
|
||||
"""Empty batch returns no embeddings and zero tokens."""
|
||||
provider = SimpleProvider(dimension=8)
|
||||
embeddings, tokens = await provider.embed_batch_with_usage([])
|
||||
|
||||
assert embeddings == []
|
||||
assert tokens == 0
|
||||
@@ -1,5 +1,7 @@
|
||||
"""Unit tests for BM25 hybrid search algorithm."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
from qdrant_client import models
|
||||
|
||||
@@ -52,3 +54,73 @@ def test_bm25_hybrid_requires_vector_db():
|
||||
"""Test BM25HybridSearchAlgorithm reports it requires vector database."""
|
||||
algo = BM25HybridSearchAlgorithm()
|
||||
assert algo.requires_vector_db is True
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def patched_search(monkeypatch):
|
||||
"""Stub the embedding / BM25 / Qdrant deps of search() and return the
|
||||
embed_with_usage mock so tests can assert how often the query was embedded."""
|
||||
embed = AsyncMock(return_value=([0.1, 0.2, 0.3], 7))
|
||||
svc = MagicMock()
|
||||
svc.embed_with_usage = embed
|
||||
monkeypatch.setattr(
|
||||
"nextcloud_mcp_server.search.bm25_hybrid.get_embedding_service", lambda: svc
|
||||
)
|
||||
|
||||
bm25 = MagicMock()
|
||||
bm25.encode_async = AsyncMock(return_value={"indices": [1], "values": [0.5]})
|
||||
monkeypatch.setattr(
|
||||
"nextcloud_mcp_server.search.bm25_hybrid.get_bm25_service",
|
||||
AsyncMock(return_value=bm25),
|
||||
)
|
||||
|
||||
qdrant = MagicMock()
|
||||
empty = MagicMock()
|
||||
empty.points = []
|
||||
qdrant.query_points = AsyncMock(return_value=empty)
|
||||
monkeypatch.setattr(
|
||||
"nextcloud_mcp_server.search.bm25_hybrid.get_qdrant_client",
|
||||
AsyncMock(return_value=qdrant),
|
||||
)
|
||||
|
||||
settings = MagicMock()
|
||||
settings.get_collection_name.return_value = "test_collection"
|
||||
settings.get_embedding_provider_family.return_value = "mistral"
|
||||
monkeypatch.setattr(
|
||||
"nextcloud_mcp_server.search.bm25_hybrid.get_settings", lambda: settings
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"nextcloud_mcp_server.search.bm25_hybrid.build_base_filter_conditions",
|
||||
lambda **kwargs: [],
|
||||
)
|
||||
return embed
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_query_embedded_and_metered_once_across_doc_types(patched_search):
|
||||
"""nc_semantic_search calls search() once per doc_type on one instance with
|
||||
the same query; the dense embedding (and its billed token count) must be
|
||||
computed exactly once, not once per type."""
|
||||
embed = patched_search
|
||||
algo = BM25HybridSearchAlgorithm()
|
||||
|
||||
for dtype in ("note", "file", "deck_card"):
|
||||
await algo.search(query="hello", user_id="alice", doc_type=dtype)
|
||||
|
||||
assert embed.await_count == 1 # embedded once, not 3×
|
||||
assert (
|
||||
algo.query_token_count == 7
|
||||
) # single query's token count, not summed/overwritten
|
||||
assert algo.query_embedding == [0.1, 0.2, 0.3]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_different_query_invalidates_cache(patched_search):
|
||||
"""A different query string re-embeds (and re-meters)."""
|
||||
embed = patched_search
|
||||
algo = BM25HybridSearchAlgorithm()
|
||||
|
||||
await algo.search(query="hello", user_id="alice")
|
||||
await algo.search(query="world", user_id="alice")
|
||||
|
||||
assert embed.await_count == 2
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
"""Unit tests for server-layer MCP tools."""
|
||||
@@ -0,0 +1,124 @@
|
||||
"""Unit tests for the search-path usage-metering helper (Deck #67).
|
||||
|
||||
``record_search_usage`` records the billable ``tokens_embedded`` event for a
|
||||
semantic search. These pin the value mapping (query token count), the flag-off
|
||||
no-op, the doc_types metadata bounding, and the best-effort failure path —
|
||||
covering the server-tool metering wiring without standing up the full
|
||||
``nc_semantic_search`` tool.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from nextcloud_mcp_server.server import semantic
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def store_spy(monkeypatch):
|
||||
"""Patch UsageEventStore.shared() to return a spy store."""
|
||||
store = MagicMock()
|
||||
store.record_usage_event = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
semantic.UsageEventStore, "shared", AsyncMock(return_value=store)
|
||||
)
|
||||
return store
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_records_query_token_count(store_spy):
|
||||
"""The event value is the query embedding's token count."""
|
||||
await semantic.record_search_usage(
|
||||
enabled=True,
|
||||
user_id="alice",
|
||||
fusion="rrf",
|
||||
doc_types=["note", "file"],
|
||||
token_count=42,
|
||||
)
|
||||
|
||||
store_spy.record_usage_event.assert_awaited_once()
|
||||
kwargs = store_spy.record_usage_event.await_args.kwargs
|
||||
assert kwargs["metric"] == "tokens_embedded"
|
||||
assert kwargs["value"] == 42
|
||||
assert kwargs["enabled"] is True
|
||||
assert kwargs["metadata"]["user_id"] == "alice"
|
||||
assert kwargs["metadata"]["fusion"] == "rrf"
|
||||
assert kwargs["metadata"]["doc_types"] == ["note", "file"]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_disabled_is_noop(store_spy):
|
||||
"""Flag off → no store access, no event."""
|
||||
await semantic.record_search_usage(
|
||||
enabled=False,
|
||||
user_id="alice",
|
||||
fusion="rrf",
|
||||
doc_types=None,
|
||||
token_count=10,
|
||||
)
|
||||
store_spy.record_usage_event.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_none_token_count_records_zero(store_spy):
|
||||
"""A missing token count (pre-embed error) records value 0, not None."""
|
||||
await semantic.record_search_usage(
|
||||
enabled=True,
|
||||
user_id="alice",
|
||||
fusion="dbsf",
|
||||
doc_types=None,
|
||||
token_count=None,
|
||||
)
|
||||
kwargs = store_spy.record_usage_event.await_args.kwargs
|
||||
assert kwargs["value"] == 0
|
||||
# None and [] both normalize to null for consistent IS NULL counting.
|
||||
assert kwargs["metadata"]["doc_types"] is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_empty_doc_types_normalizes_to_null(store_spy):
|
||||
"""An empty doc_types list normalizes to None, same as a None input, so a
|
||||
metadata->'doc_types' IS NULL query counts the all-types case consistently."""
|
||||
await semantic.record_search_usage(
|
||||
enabled=True,
|
||||
user_id="alice",
|
||||
fusion="rrf",
|
||||
doc_types=[],
|
||||
token_count=5,
|
||||
)
|
||||
kwargs = store_spy.record_usage_event.await_args.kwargs
|
||||
assert kwargs["metadata"]["doc_types"] is None
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_doc_types_metadata_is_bounded(store_spy):
|
||||
"""A large doc_types list is truncated to the metadata cap."""
|
||||
many = [f"type-{i}" for i in range(40)]
|
||||
await semantic.record_search_usage(
|
||||
enabled=True,
|
||||
user_id="alice",
|
||||
fusion="rrf",
|
||||
doc_types=many,
|
||||
token_count=5,
|
||||
)
|
||||
recorded = store_spy.record_usage_event.await_args.kwargs["metadata"]["doc_types"]
|
||||
assert recorded == many[: semantic._USAGE_METADATA_MAX_DOC_TYPES]
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_store_failure_is_swallowed(monkeypatch):
|
||||
"""A store-construction failure is logged, never raised into the search."""
|
||||
monkeypatch.setattr(
|
||||
semantic.UsageEventStore,
|
||||
"shared",
|
||||
AsyncMock(side_effect=RuntimeError("boom")),
|
||||
)
|
||||
|
||||
# Must not raise.
|
||||
await semantic.record_search_usage(
|
||||
enabled=True,
|
||||
user_id="alice",
|
||||
fusion="rrf",
|
||||
doc_types=None,
|
||||
token_count=7,
|
||||
)
|
||||
@@ -12,7 +12,10 @@ from __future__ import annotations
|
||||
import pytest
|
||||
|
||||
from nextcloud_mcp_server.config import Settings
|
||||
from nextcloud_mcp_server.observability.metrics import record_embedding
|
||||
from nextcloud_mcp_server.observability.metrics import (
|
||||
record_embedding,
|
||||
record_embedding_tokens,
|
||||
)
|
||||
|
||||
pytestmark = pytest.mark.unit
|
||||
|
||||
@@ -120,3 +123,31 @@ class TestRecordEmbedding:
|
||||
assert metric_sample(
|
||||
"astrolabe_embedding_requests_total", {**labels, "status": "error"}
|
||||
) == pytest.approx(1.0)
|
||||
|
||||
|
||||
class TestRecordEmbeddingTokens:
|
||||
"""astrolabe_embedding_tokens_total — token cost split by index/query."""
|
||||
|
||||
def test_index_increments_by_token_count(self, metric_sample):
|
||||
labels = {"provider": "tok-prov", "operation": "index"}
|
||||
before = metric_sample("astrolabe_embedding_tokens_total", labels)
|
||||
record_embedding_tokens("tok-prov", "index", 4242)
|
||||
assert metric_sample(
|
||||
"astrolabe_embedding_tokens_total", labels
|
||||
) == pytest.approx(before + 4242)
|
||||
|
||||
def test_query_operation_is_separate_series(self, metric_sample):
|
||||
labels = {"provider": "tok-prov", "operation": "query"}
|
||||
before = metric_sample("astrolabe_embedding_tokens_total", labels)
|
||||
record_embedding_tokens("tok-prov", "query", 7)
|
||||
assert metric_sample(
|
||||
"astrolabe_embedding_tokens_total", labels
|
||||
) == pytest.approx(before + 7)
|
||||
|
||||
def test_zero_or_negative_is_noop(self, metric_sample):
|
||||
labels = {"provider": "tok-noop", "operation": "index"}
|
||||
record_embedding_tokens("tok-noop", "index", 0)
|
||||
record_embedding_tokens("tok-noop", "index", -3)
|
||||
assert metric_sample(
|
||||
"astrolabe_embedding_tokens_total", labels
|
||||
) == pytest.approx(0.0)
|
||||
|
||||
@@ -0,0 +1,104 @@
|
||||
"""Unit tests for the indexing-path usage-metering helper (Deck #67).
|
||||
|
||||
``record_indexing_usage`` records the two billable events (``pages_embedded`` +
|
||||
``tokens_embedded``) after a document's chunks are embedded. These cover the
|
||||
value mapping, the flag/zero-chunk no-ops, and the best-effort failure path
|
||||
without standing up the full document pipeline.
|
||||
"""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
from nextcloud_mcp_server.vector import processor
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def store_spy(monkeypatch):
|
||||
"""Patch UsageEventStore.shared() to return a spy store."""
|
||||
store = MagicMock()
|
||||
store.record_usage_event = AsyncMock()
|
||||
monkeypatch.setattr(
|
||||
processor.UsageEventStore, "shared", AsyncMock(return_value=store)
|
||||
)
|
||||
return store
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_records_pages_embedded_and_token_count(store_spy):
|
||||
"""Both events fire: pages_embedded = chunk count, tokens_embedded = tokens."""
|
||||
await processor.record_indexing_usage(
|
||||
enabled=True,
|
||||
provider="mistral",
|
||||
model="mistral-embed",
|
||||
doc_type="file",
|
||||
user_id="alice",
|
||||
chunk_count=110,
|
||||
token_count=4242,
|
||||
total_chars=170826,
|
||||
)
|
||||
|
||||
calls = store_spy.record_usage_event.await_args_list
|
||||
by_metric = {c.kwargs["metric"]: c.kwargs["value"] for c in calls}
|
||||
assert by_metric == {"pages_embedded": 110, "tokens_embedded": 4242}
|
||||
for c in calls:
|
||||
# Hot-path fast-gate + tenant-local attribution metadata.
|
||||
assert c.kwargs["enabled"] is True
|
||||
assert c.kwargs["metadata"]["provider"] == "mistral"
|
||||
assert c.kwargs["metadata"]["model"] == "mistral-embed"
|
||||
assert c.kwargs["metadata"]["user_id"] == "alice"
|
||||
assert c.kwargs["metadata"]["doc_type"] == "file"
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_disabled_is_noop(store_spy):
|
||||
"""Flag off → no store access, no events."""
|
||||
await processor.record_indexing_usage(
|
||||
enabled=False,
|
||||
provider="mistral",
|
||||
model="mistral-embed",
|
||||
doc_type="file",
|
||||
user_id="alice",
|
||||
chunk_count=10,
|
||||
token_count=20,
|
||||
total_chars=5,
|
||||
)
|
||||
store_spy.record_usage_event.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_zero_chunks_is_noop(store_spy):
|
||||
"""A document with no chunks records nothing (no zero-value rows)."""
|
||||
await processor.record_indexing_usage(
|
||||
enabled=True,
|
||||
provider="mistral",
|
||||
model="mistral-embed",
|
||||
doc_type="file",
|
||||
user_id="alice",
|
||||
chunk_count=0,
|
||||
token_count=0,
|
||||
total_chars=0,
|
||||
)
|
||||
store_spy.record_usage_event.assert_not_awaited()
|
||||
|
||||
|
||||
@pytest.mark.unit
|
||||
async def test_store_failure_is_swallowed(monkeypatch):
|
||||
"""A store-construction failure is logged, never raised into indexing."""
|
||||
monkeypatch.setattr(
|
||||
processor.UsageEventStore,
|
||||
"shared",
|
||||
AsyncMock(side_effect=RuntimeError("boom")),
|
||||
)
|
||||
|
||||
# Must not raise.
|
||||
await processor.record_indexing_usage(
|
||||
enabled=True,
|
||||
provider="mistral",
|
||||
model="mistral-embed",
|
||||
doc_type="file",
|
||||
user_id="alice",
|
||||
chunk_count=3,
|
||||
token_count=7,
|
||||
total_chars=9,
|
||||
)
|
||||
@@ -95,7 +95,7 @@ async def test_flag_off_is_noop(storage, monkeypatch):
|
||||
"""With metering disabled, nothing is written (zero DB work)."""
|
||||
_set_metering(monkeypatch, False)
|
||||
store = UsageEventStore(storage)
|
||||
await store.record_usage_event(metric="pages_chunks", value=5)
|
||||
await store.record_usage_event(metric="pages_embedded", value=5)
|
||||
assert await _count(storage) == 0
|
||||
|
||||
|
||||
@@ -115,12 +115,12 @@ async def test_enabled_param_short_circuits_without_reading_settings(
|
||||
monkeypatch.setattr(store_module, "get_settings", _boom)
|
||||
store = UsageEventStore(storage)
|
||||
|
||||
await store.record_usage_event(metric="pages_chunks", value=1, enabled=False)
|
||||
await store.record_usage_event(metric="pages_embedded", value=1, enabled=False)
|
||||
assert await _count(storage) == 0
|
||||
|
||||
eid = str(uuid.uuid4())
|
||||
await store.record_usage_event(
|
||||
metric="pages_chunks", value=1, event_id=eid, enabled=True
|
||||
metric="pages_embedded", value=1, event_id=eid, enabled=True
|
||||
)
|
||||
assert await _count(storage) == 1
|
||||
|
||||
@@ -131,7 +131,7 @@ async def test_insert_roundtrip(storage, monkeypatch):
|
||||
store = UsageEventStore(storage)
|
||||
eid = str(uuid.uuid4())
|
||||
await store.record_usage_event(
|
||||
metric="pages_chunks",
|
||||
metric="pages_embedded",
|
||||
value=7,
|
||||
event_id=eid,
|
||||
metadata={"provider": "gateway"},
|
||||
@@ -140,7 +140,7 @@ async def test_insert_roundtrip(storage, monkeypatch):
|
||||
assert row is not None
|
||||
# Postgres returns event_id as a uuid.UUID; normalize to str for compare.
|
||||
assert str(row[0]) == eid
|
||||
assert row[2] == "pages_chunks"
|
||||
assert row[2] == "pages_embedded"
|
||||
assert row[3] == 7
|
||||
|
||||
|
||||
@@ -149,11 +149,11 @@ async def test_on_conflict_dedup(storage, monkeypatch):
|
||||
_set_metering(monkeypatch, True)
|
||||
store = UsageEventStore(storage)
|
||||
eid = str(uuid.uuid4())
|
||||
await store.record_usage_event(metric="pages_chunks", value=1, event_id=eid)
|
||||
await store.record_usage_event(metric="embeddings_queries", value=99, event_id=eid)
|
||||
await store.record_usage_event(metric="pages_embedded", value=1, event_id=eid)
|
||||
await store.record_usage_event(metric="tokens_embedded", value=99, event_id=eid)
|
||||
assert await _count(storage) == 1
|
||||
row = await _fetch(storage, eid)
|
||||
assert row[2] == "pages_chunks" # DO NOTHING, not DO UPDATE
|
||||
assert row[2] == "pages_embedded" # DO NOTHING, not DO UPDATE
|
||||
assert row[3] == 1
|
||||
|
||||
|
||||
@@ -164,7 +164,7 @@ async def test_metadata_json_roundtrip(storage, monkeypatch):
|
||||
eid = str(uuid.uuid4())
|
||||
meta = {"provider": "gateway", "model": "titan", "nested": {"chunks": 3}}
|
||||
await store.record_usage_event(
|
||||
metric="pages_chunks", value=3, event_id=eid, metadata=meta
|
||||
metric="pages_embedded", value=3, event_id=eid, metadata=meta
|
||||
)
|
||||
row = await _fetch(storage, eid)
|
||||
raw = row[4]
|
||||
@@ -187,7 +187,7 @@ async def test_occurred_at_roundtrip(storage, monkeypatch):
|
||||
eid = str(uuid.uuid4())
|
||||
when = datetime(2026, 1, 15, 12, 0, 0, tzinfo=timezone.utc)
|
||||
await store.record_usage_event(
|
||||
metric="pages_chunks", value=1, event_id=eid, occurred_at=when
|
||||
metric="pages_embedded", value=1, event_id=eid, occurred_at=when
|
||||
)
|
||||
row = await _fetch(storage, eid)
|
||||
stored = row[1]
|
||||
@@ -207,7 +207,7 @@ async def test_metadata_none_is_null(storage, monkeypatch):
|
||||
store = UsageEventStore(storage)
|
||||
eid = str(uuid.uuid4())
|
||||
await store.record_usage_event(
|
||||
metric="embeddings_queries", value=1, event_id=eid, metadata=None
|
||||
metric="tokens_embedded", value=1, event_id=eid, metadata=None
|
||||
)
|
||||
row = await _fetch(storage, eid)
|
||||
assert row[4] is None
|
||||
@@ -233,7 +233,7 @@ async def test_best_effort_swallows_db_errors(storage, monkeypatch, caplog):
|
||||
|
||||
# Must not raise.
|
||||
with caplog.at_level(logging.WARNING, logger="nextcloud_mcp_server.usage.store"):
|
||||
await store.record_usage_event(metric="pages_chunks", value=1)
|
||||
await store.record_usage_event(metric="pages_embedded", value=1)
|
||||
|
||||
assert recorded, "record_db_operation should be called on the error path"
|
||||
assert recorded[-1][3] == "error"
|
||||
@@ -262,7 +262,7 @@ async def test_best_effort_swallows_unserializable_metadata(
|
||||
# Must not raise.
|
||||
with caplog.at_level(logging.WARNING, logger="nextcloud_mcp_server.usage.store"):
|
||||
await store.record_usage_event(
|
||||
metric="pages_chunks", value=1, metadata=bad_metadata
|
||||
metric="pages_embedded", value=1, metadata=bad_metadata
|
||||
)
|
||||
|
||||
# Nothing was written — the encode failed before the insert.
|
||||
|
||||
Reference in New Issue
Block a user