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:
Chris Coutinho
2026-06-08 15:18:36 +02:00
committed by GitHub
27 changed files with 1284 additions and 136 deletions
@@ -55,7 +55,7 @@ def upgrade() -> None:
postgresql.TIMESTAMP(timezone=True) if is_pg else sa.TIMESTAMP(), postgresql.TIMESTAMP(timezone=True) if is_pg else sa.TIMESTAMP(),
nullable=False, nullable=False,
), ),
# Catalog metric: 'embeddings_queries' or 'pages_chunks'. Deliberately # Catalog metric: 'tokens_embedded' or 'pages_embedded'. Deliberately
# an unconstrained Text (no CHECK/enum) — the metric catalog lives in # an unconstrained Text (no CHECK/enum) — the metric catalog lives in
# control-plane config, not the app-DB schema. If a third metric is # control-plane config, not the app-DB schema. If a third metric is
# ever added, the CP-side catalog must learn it too, or its rollup will # ever added, the CP-side catalog must learn it too, or its rollup will
@@ -217,10 +217,21 @@ class GatewayProvider(OpenAIProvider):
exc, exc,
) )
# Bearer-refresh override topology. OpenAIProvider routes embed_batch(),
# embed_with_usage() and embed_batch_with_usage() all through
# embed_batch_with_usage(); only embed() (single) is self-contained. So we
# override exactly two methods to refresh the bearer exactly once on every
# path: embed() (its own entrypoint) and embed_batch_with_usage() (the
# shared funnel for the other three). Overriding embed_batch() as well would
# double-call _ensure_bearer() (override → super().embed_batch() →
# self.embed_batch_with_usage() → override again).
async def embed(self, text: str) -> list[float]: async def embed(self, text: str) -> list[float]:
await self._ensure_bearer() await self._ensure_bearer()
return await super().embed(text) return await super().embed(text)
async def embed_batch(self, texts: list[str]) -> list[list[float]]: async def embed_batch_with_usage(
self, texts: list[str]
) -> tuple[list[list[float]], int]:
await self._ensure_bearer() await self._ensure_bearer()
return await super().embed_batch(texts) return await super().embed_batch_with_usage(texts)
+16
View File
@@ -52,6 +52,22 @@ class EmbeddingService:
""" """
return await self.provider.embed_batch(texts) return await self.provider.embed_batch(texts)
async def embed_with_usage(self, text: str) -> tuple[list[float], int]:
"""Embed one text and report the request's token count.
Returns ``(embedding, token_count)`` for usage metering (Deck #67).
"""
return await self.provider.embed_with_usage(text)
async def embed_batch_with_usage(
self, texts: list[str]
) -> tuple[list[list[float]], int]:
"""Embed multiple texts and report the total token count.
Returns ``(embeddings, token_count)`` for usage metering (Deck #67).
"""
return await self.provider.embed_batch_with_usage(texts)
def get_dimension(self) -> int: def get_dimension(self) -> int:
""" """
Get embedding dimension. Get embedding dimension.
@@ -338,6 +338,22 @@ embedding_chars_total = Counter(
["kind", "provider"], ["kind", "provider"],
) )
# Token consumption — the billed cost unit (mirrors the tokens_embedded billing
# measure, Deck #67). On a dedicated counter (not folded into the chunk/request
# metrics above) so query embeds don't inflate indexing dashboards; labelled by
# operation = index | query. Always emitted, independent of USAGE_METERING_ENABLED.
#
# Dashboard note: operation="query" is recorded at embed time (before Qdrant /
# verify-on-read), whereas the billing-store tokens_embedded row is written only
# after the search fully succeeds. So this counter can legitimately exceed the
# billing aggregate when a search fails post-embed — don't alert on that gap as
# a divergence bug.
embedding_tokens_total = Counter(
"astrolabe_embedding_tokens_total",
"Total embedding tokens consumed (provider-reported or estimated)",
["provider", "operation"], # operation: index | query
)
# --- Chunking & indexed-by-type ----------------------------------------------- # --- Chunking & indexed-by-type -----------------------------------------------
document_chunks_total = Counter( document_chunks_total = Counter(
@@ -727,6 +743,25 @@ def record_embedding(
embedding_chars_total.labels(kind=kind, provider=provider).inc(chars) embedding_chars_total.labels(kind=kind, provider=provider).inc(chars)
def record_embedding_tokens(provider: str, operation: str, tokens: int) -> None:
"""Export embedding token consumption to Prometheus.
Mirrors the ``tokens_embedded`` billing measure (Deck #67) as an always-on
observability signal — emitted regardless of ``USAGE_METERING_ENABLED`` so
OSS/self-host deployments still see token cost in Grafana.
Args:
provider: Provider family (mistral | openai | bedrock | ollama | simple).
operation: ``"index"`` (chunk-batch embedding) or ``"query"`` (search
query embedding).
tokens: Token count for this embedding request (no-op when ``<= 0``).
"""
if tokens > 0:
embedding_tokens_total.labels(provider=provider, operation=operation).inc(
tokens
)
def record_document_chunks(doc_type: str, count: int) -> None: def record_document_chunks(doc_type: str, count: int) -> None:
""" """
Record the number of chunks produced for a document. Record the number of chunks produced for a document.
+46
View File
@@ -1,5 +1,6 @@
"""Unified provider interface for embeddings and text generation.""" """Unified provider interface for embeddings and text generation."""
import math
from abc import ABC, abstractmethod from abc import ABC, abstractmethod
@@ -55,6 +56,51 @@ class Provider(ABC):
""" """
pass pass
@staticmethod
def _estimate_tokens(texts: list[str]) -> int:
"""Best-effort token estimate when a provider returns no usage data.
Uses a coarse ~4-chars-per-token heuristic so the billable token
value stays non-zero and monotone with input size for local/dev
providers (Simple, Ollama without ``prompt_eval_count``). Real
providers override ``*_with_usage`` to report exact counts.
"""
return math.ceil(sum(len(t) for t in texts) / 4)
async def embed_with_usage(self, text: str) -> tuple[list[float], int]:
"""Embed one text and report the request's token count.
Returns ``(embedding, token_count)``. The default delegates to
:meth:`embed` and estimates the tokens; providers that surface real
usage from their embedding response override this. Used by the
usage-metering hooks (Deck #67) to bill ``tokens_embedded`` by
tokens rather than by operation count.
IMPORTANT (recursion invariant): this default calls ``self.embed``. A
provider that overrides ``embed()`` to delegate to ``embed_with_usage()``
(to avoid duplicating request logic) MUST also override this method, or
the two will call each other forever. The shipped providers that use
that delegation (Bedrock) do override both — keep that pairing.
"""
embedding = await self.embed(text)
return embedding, self._estimate_tokens([text])
async def embed_batch_with_usage(
self, texts: list[str]
) -> tuple[list[list[float]], int]:
"""Embed multiple texts and report the total token count.
Returns ``(embeddings, token_count)``; the default estimates. See
:meth:`embed_with_usage`.
IMPORTANT (recursion invariant): this default calls ``self.embed_batch``.
A provider that overrides ``embed_batch()`` to delegate to
``embed_batch_with_usage()`` (Mistral, OpenAI, Ollama do) MUST also
override this method, or the two recurse infinitely. Keep the pairing.
"""
embeddings = await self.embed_batch(texts)
return embeddings, self._estimate_tokens(texts)
@abstractmethod @abstractmethod
def get_dimension(self) -> int: def get_dimension(self) -> int:
""" """
+34 -4
View File
@@ -164,6 +164,16 @@ class BedrockProvider(Provider):
NotImplementedError: If embeddings not enabled (no embedding_model) NotImplementedError: If embeddings not enabled (no embedding_model)
ClientError: If Bedrock API call fails ClientError: If Bedrock API call fails
""" """
embedding, _ = await self.embed_with_usage(text)
return embedding
async def embed_with_usage(self, text: str) -> tuple[list[float], int]:
"""Embed one text, reporting the request's token count.
Titan Embed responses carry ``inputTextTokenCount``; for Cohere /
unknown models (no token field) this falls back to a char-based
estimate. Used by the usage-metering hooks (Deck #67).
"""
if not self.supports_embeddings: if not self.supports_embeddings:
raise NotImplementedError( raise NotImplementedError(
"Embedding not supported - no embedding_model configured" "Embedding not supported - no embedding_model configured"
@@ -182,7 +192,13 @@ class BedrockProvider(Provider):
response_body = json.loads(response["body"].read()) response_body = json.loads(response["body"].read())
embedding = self._parse_embedding_response(response_body) embedding = self._parse_embedding_response(response_body)
return embedding token_count = response_body.get("inputTextTokenCount")
tokens = (
round(token_count)
if isinstance(token_count, (int, float))
else self._estimate_tokens([text])
)
return embedding, tokens
except (BotoCoreError, ClientError) as e: except (BotoCoreError, ClientError) as e:
logger.error("Bedrock embedding error: %s", e) logger.error("Bedrock embedding error: %s", e)
@@ -205,16 +221,30 @@ class BedrockProvider(Provider):
NotImplementedError: If embeddings not enabled (no embedding_model) NotImplementedError: If embeddings not enabled (no embedding_model)
ClientError: If Bedrock API call fails ClientError: If Bedrock API call fails
""" """
embeddings, _ = await self.embed_batch_with_usage(texts)
return embeddings
async def embed_batch_with_usage(
self, texts: list[str]
) -> tuple[list[list[float]], int]:
"""Embed multiple texts, summing the per-call token counts.
Bedrock has no batch embedding API, so requests run sequentially and
the token total is the sum of each call's ``inputTextTokenCount``
(Titan) or estimate (Cohere/unknown).
"""
if not self.supports_embeddings: if not self.supports_embeddings:
raise NotImplementedError( raise NotImplementedError(
"Embedding not supported - no embedding_model configured" "Embedding not supported - no embedding_model configured"
) )
embeddings = [] embeddings: list[list[float]] = []
total_tokens = 0
for text in texts: for text in texts:
embedding = await self.embed(text) embedding, tokens = await self.embed_with_usage(text)
embeddings.append(embedding) embeddings.append(embedding)
return embeddings total_tokens += tokens
return embeddings, total_tokens
async def _detect_dimension(self): async def _detect_dimension(self):
""" """
+49 -6
View File
@@ -122,17 +122,42 @@ class MistralProvider(Provider):
async def embed_batch(self, texts: list[str]) -> list[list[float]]: async def embed_batch(self, texts: list[str]) -> list[list[float]]:
"""Generate embeddings for multiple texts, chunking by ``BATCH_SIZE``.""" """Generate embeddings for multiple texts, chunking by ``BATCH_SIZE``."""
embeddings, _ = await self.embed_batch_with_usage(texts)
return embeddings
async def embed_with_usage(self, text: str) -> tuple[list[float], int]:
"""Embed one text, reporting the Mistral request's token count."""
embeddings, tokens = await self.embed_batch_with_usage([text])
if not embeddings:
raise RuntimeError(
f"Mistral embeddings API returned no embedding for model "
f"{self.embedding_model}"
)
return embeddings[0], tokens
async def embed_batch_with_usage(
self, texts: list[str]
) -> tuple[list[list[float]], int]:
"""Embed multiple texts, summing the Mistral-reported token usage.
Returns ``(embeddings, total_tokens)`` where ``total_tokens`` is the
sum of ``response.usage.total_tokens`` across the ``BATCH_SIZE`` sub-
requests (the unit Mistral bills on). Used by the usage-metering hooks
to record ``tokens_embedded`` by tokens (Deck #67).
"""
if not self.supports_embeddings: if not self.supports_embeddings:
raise NotImplementedError(_NO_EMBEDDING_MODEL_MSG) raise NotImplementedError(_NO_EMBEDDING_MODEL_MSG)
if not texts: if not texts:
return [] return [], 0
all_embeddings: list[list[float]] = [] all_embeddings: list[list[float]] = []
total_tokens = 0
for i in range(0, len(texts), BATCH_SIZE): for i in range(0, len(texts), BATCH_SIZE):
batch = texts[i : i + BATCH_SIZE] batch = texts[i : i + BATCH_SIZE]
batch_embeddings = await self._embed_batch_request(batch) batch_embeddings, batch_tokens = await self._embed_batch_request(batch)
all_embeddings.extend(batch_embeddings) all_embeddings.extend(batch_embeddings)
total_tokens += batch_tokens
if self._dimension is None and batch_embeddings: if self._dimension is None and batch_embeddings:
self._dimension = len(batch_embeddings[0]) self._dimension = len(batch_embeddings[0])
@@ -142,11 +167,18 @@ class MistralProvider(Provider):
self.embedding_model, self.embedding_model,
) )
return all_embeddings return all_embeddings, total_tokens
@_retry_429 @_retry_429
async def _embed_batch_request(self, batch: list[str]) -> list[list[float]]: async def _embed_batch_request(
"""Single batch request with rate-limit retry.""" self, batch: list[str]
) -> tuple[list[list[float]], int]:
"""Single batch request with rate-limit retry.
Returns ``(embeddings, token_count)``; ``token_count`` comes from the
response's ``usage.total_tokens`` and falls back to a char-based
estimate if the API omits usage.
"""
assert self.embedding_model is not None assert self.embedding_model is not None
response = await self.client.embeddings.create_async( response = await self.client.embeddings.create_async(
model=self.embedding_model, model=self.embedding_model,
@@ -170,7 +202,18 @@ class MistralProvider(Provider):
f"Mistral embeddings API returned {len(result)} embeddings " f"Mistral embeddings API returned {len(result)} embeddings "
f"for {len(batch)} inputs" f"for {len(batch)} inputs"
) )
return result
usage = getattr(response, "usage", None)
total_tokens = getattr(usage, "total_tokens", None) if usage else None
# Guard on numeric type (not just ``is not None``): a real response
# gives an int, but test doubles / partial responses can surface a
# non-numeric attribute — fall back to the estimate there.
tokens = (
round(total_tokens)
if isinstance(total_tokens, (int, float))
else self._estimate_tokens(batch)
)
return result, tokens
def get_dimension(self) -> int: def get_dimension(self) -> int:
if not self.supports_embeddings: if not self.supports_embeddings:
+61 -14
View File
@@ -82,17 +82,12 @@ class OllamaProvider(Provider):
Raises: Raises:
NotImplementedError: If embeddings not enabled (no embedding_model) NotImplementedError: If embeddings not enabled (no embedding_model)
""" """
if not self.supports_embeddings: # Delegate to embed_with_usage so single and batch embeds use the same
raise NotImplementedError( # /api/embed endpoint (the legacy /api/embeddings differs in payload and
"Embedding not supported - no embedding_model configured" # omits prompt_eval_count). _detect_dimension() and other embed() callers
) # therefore stay consistent with the search/indexing path.
embedding, _ = await self.embed_with_usage(text)
response = await self.client.post( return embedding
f"{self.base_url}/api/embeddings",
json={"model": self.embedding_model, "prompt": text},
)
response.raise_for_status()
return response.json()["embedding"]
async def embed_batch( async def embed_batch(
self, texts: list[str], batch_size: int = 32 self, texts: list[str], batch_size: int = 32
@@ -116,12 +111,44 @@ class OllamaProvider(Provider):
Raises: Raises:
NotImplementedError: If embeddings not enabled (no embedding_model) NotImplementedError: If embeddings not enabled (no embedding_model)
""" """
embeddings, _ = await self.embed_batch_with_usage(texts, batch_size=batch_size)
return embeddings
async def embed_with_usage(self, text: str) -> tuple[list[float], int]:
"""Embed one text, reporting the request's token count.
Routes through ``/api/embed`` (which carries ``prompt_eval_count``)
rather than the legacy ``/api/embeddings`` so a token count is
available; falls back to a char-based estimate when the field is
absent. Used by the usage-metering hooks (Deck #67).
"""
embeddings, tokens = await self.embed_batch_with_usage([text])
if not embeddings:
raise RuntimeError(
"Ollama embeddings API returned no embedding for model "
f"{self.embedding_model}"
)
return embeddings[0], tokens
async def embed_batch_with_usage(
self, texts: list[str], batch_size: int = 32
) -> tuple[list[list[float]], int]:
"""Embed multiple texts, summing ``prompt_eval_count`` token usage.
Returns ``(embeddings, total_tokens)``. Ollama's ``/api/embed`` may
omit ``prompt_eval_count`` (older versions); a char-based estimate is
used per batch when it does.
"""
if not self.supports_embeddings: if not self.supports_embeddings:
raise NotImplementedError( raise NotImplementedError(
"Embedding not supported - no embedding_model configured" "Embedding not supported - no embedding_model configured"
) )
all_embeddings = [] if not texts:
return [], 0
all_embeddings: list[list[float]] = []
total_tokens = 0
for i in range(0, len(texts), batch_size): for i in range(0, len(texts), batch_size):
batch = texts[i : i + batch_size] batch = texts[i : i + batch_size]
response = await self.client.post( response = await self.client.post(
@@ -129,9 +156,29 @@ class OllamaProvider(Provider):
json={"model": self.embedding_model, "input": batch}, json={"model": self.embedding_model, "input": batch},
) )
response.raise_for_status() response.raise_for_status()
all_embeddings.extend(response.json()["embeddings"]) data = response.json()
all_embeddings.extend(data["embeddings"])
return all_embeddings # Cache the dimension inline (mirrors OpenAI/Mistral) so it is set
# via any embed path, not only an explicit _detect_dimension() call.
if self._dimension is None and data["embeddings"]:
self._dimension = len(data["embeddings"][0])
# ``prompt_eval_count`` is assumed to be the batch-level total for a
# multi-input /api/embed call. Ollama's API docs aren't explicit
# about batch aggregation; if a version reports only the last
# input's tokens this understates the batch. Unverified against a
# live instance — Ollama isn't the Cloud billing provider (Mistral
# is). If it proves last-item-only, switch to per-item requests and
# sum. The char-based estimate covers versions that omit the field.
prompt_eval = data.get("prompt_eval_count")
total_tokens += (
round(prompt_eval)
if isinstance(prompt_eval, (int, float))
else self._estimate_tokens(batch)
)
return all_embeddings, total_tokens
async def _detect_dimension(self): async def _detect_dimension(self):
""" """
+55 -6
View File
@@ -153,19 +153,49 @@ class OpenAIProvider(Provider):
"Embedding not supported - no embedding_model configured" "Embedding not supported - no embedding_model configured"
) )
embeddings, _ = await self.embed_batch_with_usage(texts)
return embeddings
async def embed_with_usage(self, text: str) -> tuple[list[float], int]:
"""Embed one text, reporting the request's token count."""
embeddings, tokens = await self.embed_batch_with_usage([text])
if not embeddings:
raise RuntimeError(
"OpenAI embeddings API returned no embedding for model "
f"{self.embedding_model}"
)
return embeddings[0], tokens
async def embed_batch_with_usage(
self, texts: list[str]
) -> tuple[list[list[float]], int]:
"""Embed multiple texts, summing the API-reported token usage.
Returns ``(embeddings, total_tokens)`` where ``total_tokens`` sums
``response.usage.total_tokens`` across the sub-requests (the unit the
provider bills on). Used by the usage-metering hooks (Deck #67). Also
serves the gateway path via :class:`GatewayProvider`.
"""
if not self.supports_embeddings:
raise NotImplementedError(
"Embedding not supported - no embedding_model configured"
)
if not texts: if not texts:
return [] return [], 0
# OpenAI supports batches up to 2048, but use smaller batches for safety # OpenAI supports batches up to 2048, but use smaller batches for safety
batch_size = 100 batch_size = 100
all_embeddings: list[list[float]] = [] all_embeddings: list[list[float]] = []
total_tokens = 0
for i in range(0, len(texts), batch_size): for i in range(0, len(texts), batch_size):
batch = texts[i : i + batch_size] batch = texts[i : i + batch_size]
# Use helper method with retry logic for each batch # Use helper method with retry logic for each batch
batch_embeddings = await self._embed_batch_request(batch) batch_embeddings, batch_tokens = await self._embed_batch_request(batch)
all_embeddings.extend(batch_embeddings) all_embeddings.extend(batch_embeddings)
total_tokens += batch_tokens
# Update dimension if not set # Update dimension if not set
if self._dimension is None and batch_embeddings: if self._dimension is None and batch_embeddings:
@@ -176,11 +206,18 @@ class OpenAIProvider(Provider):
self.embedding_model, self.embedding_model,
) )
return all_embeddings return all_embeddings, total_tokens
@_retry_429 @_retry_429
async def _embed_batch_request(self, batch: list[str]) -> list[list[float]]: async def _embed_batch_request(
"""Make a single batch embedding request with retry logic.""" self, batch: list[str]
) -> tuple[list[list[float]], int]:
"""Make a single batch embedding request with retry logic.
Returns ``(embeddings, token_count)``; ``token_count`` comes from the
response's ``usage.total_tokens`` and falls back to a char-based
estimate if the API omits usage.
"""
assert self.embedding_model is not None # Type narrowing assert self.embedding_model is not None # Type narrowing
response = await self.client.embeddings.create( response = await self.client.embeddings.create(
input=batch, input=batch,
@@ -188,7 +225,19 @@ class OpenAIProvider(Provider):
) )
# Sort by index to maintain order # Sort by index to maintain order
sorted_data = sorted(response.data, key=lambda x: x.index) sorted_data = sorted(response.data, key=lambda x: x.index)
return [item.embedding for item in sorted_data] embeddings = [item.embedding for item in sorted_data]
usage = getattr(response, "usage", None)
total_tokens = getattr(usage, "total_tokens", None) if usage else None
# Guard on numeric type (not just ``is not None``): a real response
# gives an int, but test doubles / partial responses can surface a
# non-numeric attribute — fall back to the estimate there.
tokens = (
round(total_tokens)
if isinstance(total_tokens, (int, float))
else self._estimate_tokens(batch)
)
return embeddings, tokens
def get_dimension(self) -> int: def get_dimension(self) -> int:
""" """
+16
View File
@@ -285,9 +285,25 @@ class SearchAlgorithm(ABC):
query_embedding: The query embedding generated during the last search. query_embedding: The query embedding generated during the last search.
Available after search() completes for algorithms that use embeddings. Available after search() completes for algorithms that use embeddings.
Can be reused by callers to avoid redundant embedding generation. Can be reused by callers to avoid redundant embedding generation.
query_token_count: Token count of the query embedding request from the
last search (provider-reported, or estimated). Set by algorithms
that embed the query so the usage-metering hook can bill
``tokens_embedded`` by tokens (Deck #67). The instance is
per-request, so this side-channel is concurrency-safe.
""" """
# Class-level defaults are a safety net; __init__ shadows them per instance.
query_embedding: list[float] | None = None query_embedding: list[float] | None = None
query_token_count: int | None = None
def __init__(self) -> None:
# Set the query-embedding side-channel as instance attributes so
# concurrent SearchAlgorithm instances never share it through the
# class-level defaults above — per-request isolation by construction,
# not just by the convention that each subclass redeclares them.
# Subclasses with their own __init__ should call super().__init__().
self.query_embedding: list[float] | None = None
self.query_token_count: int | None = None
@abstractmethod @abstractmethod
async def search( async def search(
+37 -5
View File
@@ -9,7 +9,10 @@ from qdrant_client.models import Filter
from nextcloud_mcp_server.config import get_settings from nextcloud_mcp_server.config import get_settings
from nextcloud_mcp_server.embedding import get_bm25_service, get_embedding_service from nextcloud_mcp_server.embedding import get_bm25_service, get_embedding_service
from nextcloud_mcp_server.observability.metrics import record_qdrant_operation from nextcloud_mcp_server.observability.metrics import (
record_embedding_tokens,
record_qdrant_operation,
)
from nextcloud_mcp_server.observability.tracing import trace_operation from nextcloud_mcp_server.observability.tracing import trace_operation
from nextcloud_mcp_server.search.access_filter import build_base_filter_conditions from nextcloud_mcp_server.search.access_filter import build_base_filter_conditions
from nextcloud_mcp_server.search.algorithms import ( from nextcloud_mcp_server.search.algorithms import (
@@ -53,9 +56,16 @@ class BM25HybridSearchAlgorithm(SearchAlgorithm):
f"Invalid fusion algorithm '{fusion}'. Must be 'rrf' or 'dbsf'" f"Invalid fusion algorithm '{fusion}'. Must be 'rrf' or 'dbsf'"
) )
# super() sets the per-instance query_embedding / query_token_count
# side-channel; this adds the cache key for it.
super().__init__()
self.score_threshold = score_threshold self.score_threshold = score_threshold
self.fusion = models.Fusion.RRF if fusion == "rrf" else models.Fusion.DBSF self.fusion = models.Fusion.RRF if fusion == "rrf" else models.Fusion.DBSF
self.fusion_name = fusion self.fusion_name = fusion
# ``_embedded_query`` is the query string whose dense embedding is held
# in ``query_embedding`` — repeated search() calls on this per-request
# instance (the doc_types loop) reuse it instead of re-embedding.
self._embedded_query: str | None = None
@property @property
def name(self) -> str: def name(self) -> str:
@@ -128,13 +138,35 @@ class BM25HybridSearchAlgorithm(SearchAlgorithm):
self.fusion_name, self.fusion_name,
) )
# Generate dense embedding for semantic search # Generate dense embedding for semantic search. Cache it per query on
# this (per-request) instance: nc_semantic_search calls search() once
# per doc_type with the same query, so re-embedding each time would make
# N redundant API calls and bill the query's tokens N times (Deck #67).
# Reuse the first call's embedding + token count so the query is embedded
# — and metered — exactly once.
with trace_operation("search.get_embedding_service"): with trace_operation("search.get_embedding_service"):
embedding_service = get_embedding_service() embedding_service = get_embedding_service()
with trace_operation("search.dense_embedding"): with trace_operation("search.dense_embedding"):
dense_embedding = await embedding_service.embed(query) if self.query_embedding is not None and self._embedded_query == query:
# Store for reuse by callers (e.g., viz_routes PCA visualization) dense_embedding = self.query_embedding
self.query_embedding = dense_embedding else:
(
dense_embedding,
query_tokens,
) = await embedding_service.embed_with_usage(query)
# Store for reuse by callers (e.g., viz_routes PCA
# visualization) and for the usage-metering hook in
# server/semantic.py (token count).
self.query_embedding = dense_embedding
self.query_token_count = query_tokens
self._embedded_query = query
# Export query-embedding token cost to Prometheus
# (operation=query), mirroring the per-search billing record in
# server/semantic.py. Inside the cache-miss branch so a reused
# embedding isn't double-counted.
record_embedding_tokens(
settings.get_embedding_provider_family(), "query", query_tokens
)
logger.debug("Generated dense embedding (dimension=%s)", len(dense_embedding)) logger.debug("Generated dense embedding (dimension=%s)", len(dense_embedding))
# Generate sparse embedding for BM25 keyword search # Generate sparse embedding for BM25 keyword search
+1
View File
@@ -35,6 +35,7 @@ class SemanticSearchAlgorithm(SearchAlgorithm):
Args: Args:
score_threshold: Minimum similarity score (0-1, default: 0.7) score_threshold: Minimum similarity score (0-1, default: 0.7)
""" """
super().__init__()
self.score_threshold = score_threshold self.score_threshold = score_threshold
@property @property
+80 -44
View File
@@ -56,6 +56,64 @@ logger = logging.getLogger(__name__)
_USAGE_METADATA_MAX_DOC_TYPES = 16 _USAGE_METADATA_MAX_DOC_TYPES = 16
async def record_search_usage(
*,
enabled: bool,
user_id: str,
fusion: str,
doc_types: list[str] | None,
token_count: int | None,
) -> None:
"""Record the billable ``tokens_embedded`` event for one semantic search.
The value is the query embedding's token count (provider-reported or
estimated) — the unit upstream providers bill on, and the same metric the
indexing path records for chunk embeddings (Deck #67). ``nc_semantic_search``
and ``nc_semantic_search_answer`` (which reuses it) both flow through here —
do not add a second hook. ``nc_semantic_search_answer`` exposes no
``doc_types`` parameter, so its searches always meter with
``doc_types=None``.
Best-effort and flag-gated: a metering failure is logged and never breaks
the search. Unlike the indexing path's chunk-count guard, a 0-token query is
still recorded (the query embedding ran); a zero-value row is a no-op at the
Stripe ``sum`` aggregation.
Privacy note: ``user_id`` stays tenant-local — the CP rollup aggregates
GROUP BY (day, metric) into ``usage_daily`` (no metadata column), so nothing
here propagates to Stripe; it is retained only to keep Deck #67's future
per-user attribution derivable from app-DB metadata without a re-migration.
"""
if not enabled:
return
try:
store = await UsageEventStore.shared()
await store.record_usage_event(
metric="tokens_embedded",
value=token_count or 0,
metadata={
"user_id": user_id,
"fusion": fusion,
# Bounded copy — see _USAGE_METADATA_MAX_DOC_TYPES. Both None and
# [] normalize to null so a future metadata->'doc_types' IS NULL
# query counts the all-types case consistently.
"doc_types": (
doc_types[:_USAGE_METADATA_MAX_DOC_TYPES] if doc_types else None
),
},
# The caller already confirmed the flag, so pass enabled=True
# directly — the store then skips a second uncached Settings build on
# this hot query path (ADR-024).
enabled=True,
)
except Exception:
# Reached only when shared()/store construction itself raises
# (record_usage_event swallows its own write failures). Metering is on,
# so warn — a silent DEBUG line would hide "operator enabled metering
# but gets no data".
logger.warning("usage metering hook (tokens_embedded) skipped", exc_info=True)
def configure_semantic_tools(mcp: FastMCP): def configure_semantic_tools(mcp: FastMCP):
"""Configure semantic search tools for MCP server.""" """Configure semantic search tools for MCP server."""
@@ -527,51 +585,29 @@ def configure_semantic_tools(mcp: FastMCP):
logger.info("Returning %d results from BM25 hybrid search", len(results)) logger.info("Returning %d results from BM25 hybrid search", len(results))
# Usage metering (Deck #67): one billable 'embeddings_queries' # Usage metering (Deck #67): record the query embedding's token
# event per successful search (the query embedding is the metered # count as a billable 'tokens_embedded' event. query_token_count
# cost). Best-effort and gated on the flag so the off-path touches # is set by BM25HybridSearchAlgorithm during the search() above; the
# no storage. nc_semantic_search_answer reuses this tool, so it # doc_types loop reuses one search_algo instance for the same query
# records here too — do not add a second hook there. # and the algorithm caches the dense embedding per query, so the
# query is embedded — and metered — exactly once regardless of how
# many doc_types were searched. See record_search_usage for the
# metric/privacy details.
# #
# Privacy note: user_id stays tenant-local. The CP rollup # NOTE (v1 billing gap): this fires only on a fully successful
# aggregates GROUP BY (day, metric) into usage_daily, which has no # search. If the query embed succeeded (provider billed the tokens,
# metadata column, so nothing here propagates to Stripe; the value # and Prometheus recorded them via record_embedding_tokens) but a
# is retained only so Deck #67's "per-user attribution derivable # later step (Qdrant/verify) raised, no tokens_embedded row is
# from app-DB metadata later" stays possible without a re-migration. # written — the embed cost is real but absent from the billing
if settings.usage_metering_enabled: # ledger. Acceptable for v1 (search failures are rare and the meter
try: # is not billed today); revisit if billing accuracy needs it.
store = await UsageEventStore.shared() await record_search_usage(
await store.record_usage_event( enabled=settings.usage_metering_enabled,
metric="embeddings_queries", user_id=username,
value=1, fusion=fusion,
metadata={ doc_types=doc_types,
"user_id": username, token_count=search_algo.query_token_count,
"fusion": fusion, )
# Bounded copy — see _USAGE_METADATA_MAX_DOC_TYPES.
# Both None and [] normalize to null so a future
# metadata->'doc_types' IS NULL query counts the
# all-types case consistently.
"doc_types": (
doc_types[:_USAGE_METADATA_MAX_DOC_TYPES]
if doc_types
else None
),
},
# The outer guard already confirmed the flag, so pass
# enabled=True directly — the store then skips a second
# uncached Settings build on this hot query path
# (ADR-024).
enabled=True,
)
except Exception:
# Reached only when shared()/store construction itself
# raises (record_usage_event swallows its own write
# failures). Metering is on, so warn — a silent DEBUG line
# would hide "operator enabled metering but gets no data".
logger.warning(
"usage metering hook (embeddings_queries) skipped",
exc_info=True,
)
return SemanticSearchResponse( return SemanticSearchResponse(
results=results, results=results,
+2 -2
View File
@@ -102,8 +102,8 @@ class UsageEventStore:
logged and swallowed — this must never break the caller's operation. logged and swallowed — this must never break the caller's operation.
Args: Args:
metric: Catalog metric, e.g. ``"embeddings_queries"`` or metric: Catalog metric, e.g. ``"tokens_embedded"`` or
``"pages_chunks"``. ``"pages_embedded"``.
value: Count/quantity for this event. value: Count/quantity for this event.
occurred_at: Operation completion time; defaults to now (UTC). occurred_at: Operation completion time; defaults to now (UTC).
metadata: Optional rawest-unit context (provider, model, tokens, metadata: Optional rawest-unit context (provider, model, tokens,
+102 -38
View File
@@ -23,6 +23,7 @@ from nextcloud_mcp_server.observability.metrics import (
record_document_chunks, record_document_chunks,
record_document_parse_failed, record_document_parse_failed,
record_embedding, record_embedding,
record_embedding_tokens,
record_qdrant_operation, record_qdrant_operation,
record_vector_sync_processing, record_vector_sync_processing,
update_vector_sync_queue_size, update_vector_sync_queue_size,
@@ -112,6 +113,80 @@ def should_use_page_aware(
return page_aware_enabled and doc_type == "file" and bool(page_boundaries) return page_aware_enabled and doc_type == "file" and bool(page_boundaries)
async def record_indexing_usage(
*,
enabled: bool,
provider: str,
model: str,
doc_type: str,
user_id: str,
chunk_count: int,
token_count: int,
total_chars: int,
) -> None:
"""Record the two billable usage events for one embedded document.
``pages_embedded`` is the buyer-facing "pages indexed" dimension;
``tokens_embedded`` is the embedding request's token count — the same metric
search records, so the meter bills embedding tokens whether they were
incurred indexing a document or embedding a query (Deck #67).
TODO(#282): ``pages_embedded`` currently carries the raw chunk count
(``len(chunk_texts)``) as an interim value. The real normalized "pages
indexed" count — real pages for paginated types (PDF/DOCX/PPT), a fixed
chars/tokens-per-page constant otherwise — is deferred to instrumentation
card #282; this code (card #284) only lands the metric name/contract.
Best-effort and flag-gated: a metering failure is logged and never breaks
indexing. No-op when metering is disabled or the document produced no chunks
(an empty batch embeds nothing and would only write zero-value rows).
Privacy note: ``user_id`` stays tenant-local — the CP rollup aggregates
GROUP BY (day, metric) into ``usage_daily`` (no metadata column), so nothing
here reaches Stripe; it is retained only to keep Deck #67's future per-user
attribution derivable from the app DB without a re-migration.
"""
if not enabled or chunk_count == 0:
return
metadata = {
"provider": provider,
"model": model,
"doc_type": doc_type,
"user_id": user_id,
"total_chars": total_chars,
}
try:
store = await UsageEventStore.shared()
# enabled=True: the guard above already confirmed the flag, so the store
# skips a second uncached Settings build per record (ADR-024).
# record_usage_event swallows its own write failures, so the two records
# are independent; if pages_embedded somehow raised mid-way,
# tokens_embedded would be skipped, leaving an unmatched pages_embedded
# row — acceptable under the (day, metric) SUM-aggregation billing model.
await store.record_usage_event(
# TODO(#282): value is the interim chunk count; switch to normalized
# real-page count when the per-page constant lands.
metric="pages_embedded",
value=chunk_count,
metadata=metadata,
enabled=True,
)
await store.record_usage_event(
metric="tokens_embedded",
value=token_count,
metadata=metadata,
enabled=True,
)
except Exception:
# Reached only when shared()/store construction itself raises
# (record_usage_event swallows its own write failures). Metering is on,
# so warn rather than hide the "enabled but no billing data" case.
logger.warning(
"usage metering hook (indexing embeddings) skipped", exc_info=True
)
async def processor_task( async def processor_task(
worker_id: int, worker_id: int,
receive_stream: MemoryObjectReceiveStream[DocumentTask], receive_stream: MemoryObjectReceiveStream[DocumentTask],
@@ -598,9 +673,13 @@ async def _index_document(
doc_type="file", doc_type="file",
user_id=doc_task.user_id, user_id=doc_task.user_id,
) )
# No embedding ran, so no usage is recorded here — stated
# explicitly so a "fewer tokens_embedded rows than expected"
# audit lands on the dedup path rather than reconstructing it
# from Qdrant claim logs.
logger.info( logger.info(
"Dedup hit for file %s (etag=%s); claimed for user %s " "Dedup hit for file %s (etag=%s); claimed for user %s "
"without reprocessing", "without reprocessing (no embedding/usage recorded)",
doc_task.doc_id, doc_task.doc_id,
doc_task.etag, doc_task.etag,
doc_task.user_id, doc_task.user_id,
@@ -809,7 +888,10 @@ async def _index_document(
embedding_service = get_embedding_service() embedding_service = get_embedding_service()
embed_start = time.time() embed_start = time.time()
try: try:
dense_embeddings = await embedding_service.embed_batch(chunk_texts) (
dense_embeddings,
embed_tokens,
) = await embedding_service.embed_batch_with_usage(chunk_texts)
except Exception: except Exception:
record_embedding( record_embedding(
"dense", provider, time.time() - embed_start, status="error" "dense", provider, time.time() - embed_start, status="error"
@@ -822,42 +904,24 @@ async def _index_document(
chunks=len(chunk_texts), chunks=len(chunk_texts),
chars=total_chars, chars=total_chars,
) )
# Usage metering (Deck #67): record chunks embedded as a billable # Export token consumption to Prometheus (always-on, independent of
# 'pages_chunks' event. Best-effort and gated on the flag so the # the billing flag) so Grafana sees indexing token cost.
# off-path (OSS default) touches no storage; placed after the record_embedding_tokens(provider, "index", embed_tokens)
# embedding succeeds so it can never affect the indexing path. # Usage metering (Deck #67): record the chunk volume +
# # embedding-token count for this document. Best-effort and
# Privacy note: user_id stays tenant-local — the CP rollup # flag-gated; placed after the embedding succeeds so it can never
# aggregates GROUP BY (day, metric) into usage_daily (no metadata # affect the indexing path. See record_indexing_usage for the
# column), so nothing here reaches Stripe; it is retained only to # metric/privacy details.
# keep Deck #67's future per-user attribution derivable from the await record_indexing_usage(
# app DB without a re-migration. enabled=settings.usage_metering_enabled,
if settings.usage_metering_enabled: provider=provider,
try: model=settings.get_embedding_model_name(),
store = await UsageEventStore.shared() doc_type=doc_task.doc_type,
await store.record_usage_event( user_id=doc_task.user_id,
metric="pages_chunks", chunk_count=len(chunk_texts),
value=len(chunk_texts), token_count=embed_tokens,
metadata={ total_chars=total_chars,
"provider": provider, )
"model": settings.get_embedding_model_name(),
"doc_type": doc_task.doc_type,
"user_id": doc_task.user_id,
"total_chars": total_chars,
},
# The outer guard already confirmed the flag, so pass
# enabled=True directly — the store then skips a second
# uncached Settings build here (ADR-024).
enabled=True,
)
except Exception:
# Reached only when shared()/store construction itself
# raises (record_usage_event swallows its own write
# failures). Metering is on, so warn rather than hide the
# "enabled but no billing data" case in DEBUG logs.
logger.warning(
"usage metering hook (pages_chunks) skipped", exc_info=True
)
async def generate_sparse_embeddings(): async def generate_sparse_embeddings():
"""Generate sparse embeddings (BM25 for keyword matching).""" """Generate sparse embeddings (BM25 for keyword matching)."""
+67
View File
@@ -255,6 +255,73 @@ async def test_bedrock_dimension_detection(mock_bedrock_client):
assert provider.get_dimension() == 1536 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 @pytest.mark.unit
async def test_bedrock_cohere_embedding(mock_bedrock_client): async def test_bedrock_cohere_embedding(mock_bedrock_client):
"""Test Bedrock with Cohere embedding model.""" """Test Bedrock with Cohere embedding model."""
@@ -6,6 +6,7 @@ tenant realm); creds are all-or-nothing.
""" """
import time import time
from unittest.mock import AsyncMock, MagicMock
import httpx import httpx
import pytest import pytest
@@ -327,6 +328,105 @@ def test_trailing_slash_base_url_normalized():
assert not base.endswith("/v1/v1") 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): 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 """End-to-end of the fix: a bare-origin base_url still resolves the
dimension because discovery lands on /v1/models.""" dimension because discovery lands on /v1/models."""
+62
View File
@@ -255,6 +255,68 @@ async def test_mistral_batch_raises_on_count_mismatch(mock_mistral_client):
await provider.embed_batch(["a", "b"]) 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 @pytest.mark.unit
def test_mistral_is_rate_limit_predicate(): def test_mistral_is_rate_limit_predicate():
"""_is_rate_limit returns True only for SDKErrors with status_code == 429.""" """_is_rate_limit returns True only for SDKErrors with status_code == 429."""
+69
View File
@@ -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()
+40
View File
@@ -280,6 +280,46 @@ async def test_openai_empty_batch():
assert embeddings == [] 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 @pytest.mark.unit
async def test_openai_close(mock_openai_client): async def test_openai_close(mock_openai_client):
"""Test OpenAI client close.""" """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
+72
View File
@@ -1,5 +1,7 @@
"""Unit tests for BM25 hybrid search algorithm.""" """Unit tests for BM25 hybrid search algorithm."""
from unittest.mock import AsyncMock, MagicMock
import pytest import pytest
from qdrant_client import models from qdrant_client import models
@@ -52,3 +54,73 @@ def test_bm25_hybrid_requires_vector_db():
"""Test BM25HybridSearchAlgorithm reports it requires vector database.""" """Test BM25HybridSearchAlgorithm reports it requires vector database."""
algo = BM25HybridSearchAlgorithm() algo = BM25HybridSearchAlgorithm()
assert algo.requires_vector_db is True 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
+1
View File
@@ -0,0 +1 @@
"""Unit tests for server-layer MCP tools."""
+124
View File
@@ -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,
)
+32 -1
View File
@@ -12,7 +12,10 @@ from __future__ import annotations
import pytest import pytest
from nextcloud_mcp_server.config import Settings 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 pytestmark = pytest.mark.unit
@@ -120,3 +123,31 @@ class TestRecordEmbedding:
assert metric_sample( assert metric_sample(
"astrolabe_embedding_requests_total", {**labels, "status": "error"} "astrolabe_embedding_requests_total", {**labels, "status": "error"}
) == pytest.approx(1.0) ) == 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)
+104
View File
@@ -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,
)
+13 -13
View File
@@ -95,7 +95,7 @@ async def test_flag_off_is_noop(storage, monkeypatch):
"""With metering disabled, nothing is written (zero DB work).""" """With metering disabled, nothing is written (zero DB work)."""
_set_metering(monkeypatch, False) _set_metering(monkeypatch, False)
store = UsageEventStore(storage) 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 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) monkeypatch.setattr(store_module, "get_settings", _boom)
store = UsageEventStore(storage) 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 assert await _count(storage) == 0
eid = str(uuid.uuid4()) eid = str(uuid.uuid4())
await store.record_usage_event( 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 assert await _count(storage) == 1
@@ -131,7 +131,7 @@ async def test_insert_roundtrip(storage, monkeypatch):
store = UsageEventStore(storage) store = UsageEventStore(storage)
eid = str(uuid.uuid4()) eid = str(uuid.uuid4())
await store.record_usage_event( await store.record_usage_event(
metric="pages_chunks", metric="pages_embedded",
value=7, value=7,
event_id=eid, event_id=eid,
metadata={"provider": "gateway"}, metadata={"provider": "gateway"},
@@ -140,7 +140,7 @@ async def test_insert_roundtrip(storage, monkeypatch):
assert row is not None assert row is not None
# Postgres returns event_id as a uuid.UUID; normalize to str for compare. # Postgres returns event_id as a uuid.UUID; normalize to str for compare.
assert str(row[0]) == eid assert str(row[0]) == eid
assert row[2] == "pages_chunks" assert row[2] == "pages_embedded"
assert row[3] == 7 assert row[3] == 7
@@ -149,11 +149,11 @@ async def test_on_conflict_dedup(storage, monkeypatch):
_set_metering(monkeypatch, True) _set_metering(monkeypatch, True)
store = UsageEventStore(storage) store = UsageEventStore(storage)
eid = str(uuid.uuid4()) eid = str(uuid.uuid4())
await store.record_usage_event(metric="pages_chunks", value=1, event_id=eid) await store.record_usage_event(metric="pages_embedded", value=1, event_id=eid)
await store.record_usage_event(metric="embeddings_queries", value=99, event_id=eid) await store.record_usage_event(metric="tokens_embedded", value=99, event_id=eid)
assert await _count(storage) == 1 assert await _count(storage) == 1
row = await _fetch(storage, eid) 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 assert row[3] == 1
@@ -164,7 +164,7 @@ async def test_metadata_json_roundtrip(storage, monkeypatch):
eid = str(uuid.uuid4()) eid = str(uuid.uuid4())
meta = {"provider": "gateway", "model": "titan", "nested": {"chunks": 3}} meta = {"provider": "gateway", "model": "titan", "nested": {"chunks": 3}}
await store.record_usage_event( 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) row = await _fetch(storage, eid)
raw = row[4] raw = row[4]
@@ -187,7 +187,7 @@ async def test_occurred_at_roundtrip(storage, monkeypatch):
eid = str(uuid.uuid4()) eid = str(uuid.uuid4())
when = datetime(2026, 1, 15, 12, 0, 0, tzinfo=timezone.utc) when = datetime(2026, 1, 15, 12, 0, 0, tzinfo=timezone.utc)
await store.record_usage_event( 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) row = await _fetch(storage, eid)
stored = row[1] stored = row[1]
@@ -207,7 +207,7 @@ async def test_metadata_none_is_null(storage, monkeypatch):
store = UsageEventStore(storage) store = UsageEventStore(storage)
eid = str(uuid.uuid4()) eid = str(uuid.uuid4())
await store.record_usage_event( 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) row = await _fetch(storage, eid)
assert row[4] is None assert row[4] is None
@@ -233,7 +233,7 @@ async def test_best_effort_swallows_db_errors(storage, monkeypatch, caplog):
# Must not raise. # Must not raise.
with caplog.at_level(logging.WARNING, logger="nextcloud_mcp_server.usage.store"): 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, "record_db_operation should be called on the error path"
assert recorded[-1][3] == "error" assert recorded[-1][3] == "error"
@@ -262,7 +262,7 @@ async def test_best_effort_swallows_unserializable_metadata(
# Must not raise. # Must not raise.
with caplog.at_level(logging.WARNING, logger="nextcloud_mcp_server.usage.store"): with caplog.at_level(logging.WARNING, logger="nextcloud_mcp_server.usage.store"):
await store.record_usage_event( 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. # Nothing was written — the encode failed before the insert.