Files
mcp-nextcloud/tests/unit/providers/test_gateway_provider.py
T
Chris CoutinhoandClaude Opus 4.8 64318f0b25 feat(usage): meter embedding tokens as embeddings_queries on both paths
embeddings_queries now records the embedding request's token count (the unit
upstream providers bill on) instead of an operation count, and fires on the
indexing path too. Previously only semantic search recorded it (value=1), so a
re-indexing run produced no embeddings_queries events at all — only pages_chunks.

- Provider layer: additive embed_with_usage / embed_batch_with_usage surface the
  per-request token count (Mistral/OpenAI usage.total_tokens, Bedrock Titan
  inputTextTokenCount, Ollama prompt_eval_count); a char-based estimate is the
  fallback (Simple, and any provider/response without a token field). Gateway and
  EmbeddingService forward through. The count travels as a return value / a
  per-request SearchAlgorithm attribute — never on the singleton — so concurrent
  indexing + search can't mis-attribute bills.
- Indexing (vector/processor.py): records embeddings_queries (value=batch tokens)
  alongside the existing pages_chunks event.
- Search (server/semantic.py): value is now the query embedding's token count,
  relayed from BM25HybridSearchAlgorithm via query_token_count.

The astrolabe_embeddings_queries Stripe meter (sum aggregation) now sums tokens
with no CP/Terraform change. The meter "queries"->tokens naming/unit
clarification (homelab-terraform #254) + CP rollup/portal copy is a follow-up.

Deck #67.

Co-Authored-By: Claude Opus 4.8 (1M context) <noreply@anthropic.com>
2026-06-08 00:53:58 +02:00

412 lines
14 KiB
Python

"""Gateway provider registration + M2M OIDC auth (design §10.2).
The gateway is manual-only: selected by EMBEDDING_PROVIDER=gateway and never by
the autodetect chain. Auth is the gateway's own M2M OIDC realm (parallel to the
tenant realm); creds are all-or-nothing.
"""
import time
from unittest.mock import AsyncMock, MagicMock
import httpx
import pytest
from nextcloud_mcp_server.config import Settings
from nextcloud_mcp_server.embedding.gateway_client import (
GatewayProvider,
GatewayTokenProvider,
)
from nextcloud_mcp_server.providers.registry import ProviderRegistry, reset_provider
from nextcloud_mcp_server.providers.simple import SimpleProvider
def _patch_settings(monkeypatch, settings):
monkeypatch.setattr(
"nextcloud_mcp_server.providers.registry.get_settings", lambda: settings
)
reset_provider()
def test_gateway_selected_unauthenticated(monkeypatch):
settings = Settings(
embedding_provider="gateway",
embedding_gateway_url="https://gateway:8083",
embedding_gateway_model="mistral/mistral-embed",
)
_patch_settings(monkeypatch, settings)
provider = ProviderRegistry.create_provider()
assert isinstance(provider, GatewayProvider)
assert provider.embedding_model == "mistral/mistral-embed"
assert provider.supports_embeddings is True
assert provider.supports_generation is False
assert provider._token_provider is None # unauthenticated
def test_gateway_selected_with_m2m_oidc(monkeypatch):
settings = Settings(
embedding_provider="gateway",
embedding_gateway_url="https://gateway:8083",
embedding_gateway_token_url="https://idp.example/oauth2/token",
embedding_gateway_client_id="mcp-server",
embedding_gateway_client_secret="shh",
embedding_gateway_scope="astrolabe-embedding-gateway/embed",
)
_patch_settings(monkeypatch, settings)
provider = ProviderRegistry.create_provider()
assert isinstance(provider, GatewayProvider)
assert isinstance(provider._token_provider, GatewayTokenProvider)
def test_partial_m2m_creds_rejected():
with pytest.raises(ValueError, match="must be set together"):
Settings(
embedding_provider="gateway",
embedding_gateway_url="https://gateway:8083",
embedding_gateway_client_id="mcp-server", # missing token_url/secret
)
def test_autodetect_default_does_not_pick_gateway(monkeypatch):
settings = Settings()
_patch_settings(monkeypatch, settings)
assert isinstance(ProviderRegistry.create_provider(), SimpleProvider)
def test_openai_creds_do_not_trigger_gateway(monkeypatch):
settings = Settings(openai_api_key="sk-test")
_patch_settings(monkeypatch, settings)
assert not isinstance(ProviderRegistry.create_provider(), GatewayProvider)
async def test_token_provider_caches_and_refreshes(monkeypatch):
calls = {"n": 0}
def handler(request: httpx.Request) -> httpx.Response:
calls["n"] += 1
assert request.headers["Authorization"].startswith("Basic ")
body = dict(httpx.QueryParams(request.content.decode()))
assert body["grant_type"] == "client_credentials"
assert body["scope"] == "embed"
return httpx.Response(
200, json={"access_token": f"tok{calls['n']}", "expires_in": 3600}
)
transport = httpx.MockTransport(handler)
orig_async_client = httpx.AsyncClient
def _client(*args, **kwargs):
kwargs["transport"] = transport
return orig_async_client(*args, **kwargs)
monkeypatch.setattr(httpx, "AsyncClient", _client)
tp = GatewayTokenProvider(
token_url="https://idp.example/oauth2/token",
client_id="cid",
client_secret="sec",
scope="embed",
)
t1 = await tp.get_token()
t2 = await tp.get_token() # cached → no new HTTP call
assert t1 == t2 == "tok1"
assert calls["n"] == 1
# Expire the cache → next call refreshes.
assert tp._cache is not None
tp._cache = (tp._cache[0], time.time() - 1)
t3 = await tp.get_token()
assert t3 == "tok2"
assert calls["n"] == 2
async def test_token_provider_concurrent_callers_issue_single_request(monkeypatch):
"""Two concurrent get_token() calls must share one token request, not race."""
import anyio
calls = {"n": 0}
async def handler(request: httpx.Request) -> httpx.Response:
calls["n"] += 1
# Hold the "network" open so a second caller arrives mid-flight.
await anyio.sleep(0.05)
return httpx.Response(
200, json={"access_token": f"tok{calls['n']}", "expires_in": 3600}
)
transport = httpx.MockTransport(handler)
orig_async_client = httpx.AsyncClient
def _client(*args, **kwargs):
kwargs["transport"] = transport
return orig_async_client(*args, **kwargs)
monkeypatch.setattr(httpx, "AsyncClient", _client)
tp = GatewayTokenProvider(
token_url="https://idp.example/oauth2/token",
client_id="cid",
client_secret="sec",
)
results: list[str] = []
async def _fetch():
results.append(await tp.get_token())
async with anyio.create_task_group() as tg:
tg.start_soon(_fetch)
tg.start_soon(_fetch)
# The lock serialises the check-then-fetch cycle: only one HTTP request,
# and both callers observe the same cached token.
assert calls["n"] == 1
assert results == ["tok1", "tok1"]
# --- Dimension discovery via gateway GET /v1/models -------------------------
def _mock_async_client(monkeypatch, handler):
"""Route every httpx.AsyncClient through a MockTransport (mirrors the
token-provider tests above)."""
transport = httpx.MockTransport(handler)
orig = httpx.AsyncClient
def _client(*args, **kwargs):
kwargs["transport"] = transport
return orig(*args, **kwargs)
monkeypatch.setattr(httpx, "AsyncClient", _client)
async def test_detect_dimension_from_models_endpoint(monkeypatch):
"""_detect_dimension() resolves the dimension from /v1/models with no embed
call — the regression that crashed external-mode startup."""
seen = {}
def handler(request: httpx.Request) -> httpx.Response:
seen["url"] = str(request.url)
seen["auth"] = request.headers.get("Authorization")
return httpx.Response(
200,
json={
"object": "list",
"data": [
{
"id": "mistral/mistral-embed",
"object": "model",
"dimension": 1024,
},
{
"id": "text-embedding-3-large",
"object": "model",
"dimension": 3072,
},
],
},
)
_mock_async_client(monkeypatch, handler)
provider = GatewayProvider(
base_url="http://gw:8083/v1", embedding_model="mistral/mistral-embed"
)
await provider._detect_dimension()
assert provider.get_dimension() == 1024
assert seen["url"].endswith("/v1/models")
assert seen["auth"] is None # unauthenticated gateway
async def test_detect_dimension_sends_bearer(monkeypatch):
"""When a token provider is configured, discovery presents the M2M bearer."""
captured = {}
def handler(request: httpx.Request) -> httpx.Response:
if request.url.path.endswith("/token"):
return httpx.Response(200, json={"access_token": "tok", "expires_in": 3600})
captured["auth"] = request.headers.get("Authorization")
return httpx.Response(
200, json={"data": [{"id": "mistral/mistral-embed", "dimension": 1024}]}
)
_mock_async_client(monkeypatch, handler)
tp = GatewayTokenProvider(
token_url="http://idp.example/token", client_id="c", client_secret="s"
)
provider = GatewayProvider(
base_url="http://gw:8083/v1",
embedding_model="mistral/mistral-embed",
token_provider=tp,
)
await provider._detect_dimension()
assert provider.get_dimension() == 1024
assert captured["auth"] == "Bearer tok"
async def test_detect_dimension_non_fatal_on_http_error(monkeypatch):
"""An old gateway without /v1/models (404) must not crash startup —
dimension stays unknown so lazy detect-on-first-embed still applies."""
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(404, json={"detail": "not found"})
_mock_async_client(monkeypatch, handler)
provider = GatewayProvider(
base_url="http://gw:8083/v1", embedding_model="mistral/mistral-embed"
)
await provider._detect_dimension() # must not raise
with pytest.raises(RuntimeError):
provider.get_dimension() # still unknown
async def test_detect_dimension_model_absent(monkeypatch):
"""Gateway reachable but doesn't list our model → no dimension set,
no raise."""
def handler(request: httpx.Request) -> httpx.Response:
return httpx.Response(200, json={"data": [{"id": "other", "dimension": 99}]})
_mock_async_client(monkeypatch, handler)
provider = GatewayProvider(
base_url="http://gw:8083/v1", embedding_model="mistral/mistral-embed"
)
await provider._detect_dimension()
with pytest.raises(RuntimeError):
provider.get_dimension()
async def test_detect_dimension_skips_when_already_known(monkeypatch):
"""If the dimension is already known, discovery makes no HTTP call."""
called = {"n": 0}
def handler(request: httpx.Request) -> httpx.Response:
called["n"] += 1
return httpx.Response(200, json={"data": []})
_mock_async_client(monkeypatch, handler)
provider = GatewayProvider(
base_url="http://gw:8083/v1", embedding_model="mistral/mistral-embed"
)
provider._dimension = 1024 # pre-set (e.g. explicit override / OpenAI model)
await provider._detect_dimension()
assert called["n"] == 0
assert provider.get_dimension() == 1024
# --- /v1 base-path normalization --------------------------------------------
# EMBEDDING_GATEWAY_URL is configured as a bare origin (scheme://host:port);
# the provider appends the gateway's /v1 base path so both the OpenAI SDK's
# embed posts ({base}/embeddings) and discovery ({base}/models) land under /v1.
def _client_base(provider: GatewayProvider) -> str:
return str(provider.client.base_url).rstrip("/")
def test_bare_base_url_gets_v1_base_path():
provider = GatewayProvider(
base_url="http://gw:8083", embedding_model="mistral/mistral-embed"
)
assert _client_base(provider).endswith("/v1")
def test_v1_base_url_is_idempotent():
# A URL that already carries /v1 (e.g. legacy config) is not doubled.
provider = GatewayProvider(
base_url="http://gw:8083/v1", embedding_model="mistral/mistral-embed"
)
base = _client_base(provider)
assert base.endswith("/v1")
assert not base.endswith("/v1/v1")
def test_trailing_slash_base_url_normalized():
provider = GatewayProvider(
base_url="http://gw:8083/", embedding_model="mistral/mistral-embed"
)
base = _client_base(provider)
assert base.endswith("/v1")
assert not base.endswith("/v1/v1")
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."""
provider = GatewayProvider(
base_url="http://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
async def test_gateway_embed_batch_with_usage_forwards_after_bearer(monkeypatch):
"""embed_batch_with_usage also refreshes the bearer before delegating."""
provider = GatewayProvider(
base_url="http://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
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."""
seen = {}
def handler(request: httpx.Request) -> httpx.Response:
seen["url"] = str(request.url)
return httpx.Response(
200, json={"data": [{"id": "mistral/mistral-embed", "dimension": 1024}]}
)
_mock_async_client(monkeypatch, handler)
provider = GatewayProvider(
base_url="http://gw:8083", embedding_model="mistral/mistral-embed"
)
await provider._detect_dimension()
assert provider.get_dimension() == 1024
assert seen["url"].endswith("/v1/models")