Studio: cache MCP tool discovery instead of re-probing every chat send (#5828)

* Studio: cache MCP tool discovery instead of re-probing every chat send

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Stop re-probing offline/down MCP servers every time

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Add tests for mid-probe delete and OAuth cool-off paths

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Don't cool-off a server edited or deleted mid-probe

* [pre-commit.ci] auto fixes from pre-commit.com hooks

for more information, see https://pre-commit.ci

* Guard MCP refresh cache writes

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
Co-authored-by: Roland Tannous <115670425+rolandtannous@users.noreply.github.com>
Co-authored-by: Lee Jackson <130007945+Imagineer99@users.noreply.github.com>
Co-authored-by: imagineer99 <samleejackson0@gmail.com>
This commit is contained in:
oobabooga 2026-06-12 10:18:06 -03:00 committed by GitHub
commit 522f308be2
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 749 additions and 20 deletions

View file

@ -8,6 +8,7 @@ import json
import os
import shlex
import sys
import time
from typing import Any, Optional
from loggers import get_logger
@ -16,6 +17,14 @@ logger = get_logger(__name__)
MCP_TOOL_PREFIX = "mcp__"
# A failed probe isn't cached (a recovered server must come back), but it's
# recorded so a down server isn't re-probed -- and the chat send re-hung for
# the full timeout -- on every message. Cool off for this long after a failure;
# much longer for OAuth, whose probe can hang up to _OAUTH_PROBE_TIMEOUT,
# so that hang doesn't recur every minute.
FAILED_PROBE_COOLOFF_SECONDS = 60.0
OAUTH_FAILED_PROBE_COOLOFF_SECONDS = 300.0
_oauth_token_store = None
@ -227,6 +236,53 @@ async def list_tools_async(
return await asyncio.wait_for(_fetch(), timeout = timeout)
# Discovered-tool cache, keyed by MCP server id. get_enabled_mcp_tools()
# probes a server only on a cache miss, keeping MCP discovery off the chat
# send's critical path -- tool schemas are stable within a session. The
# /refresh route warms it; a URL/header/OAuth change or a delete evicts it.
# Successful probes are cached indefinitely.
_tool_cache: dict[str, list[dict]] = {}
# server_id -> monotonic time before which a failed server must not be
# re-probed (see record_probe_failure). Cleared on a successful probe or
# eviction.
_probe_cooloff_until: dict[str, float] = {}
# MCP server fields whose change invalidates a server's discovered tools: the
# endpoint/auth used to probe it (url, headers, oauth) or whether it's used at
# all (is_enabled). A rename does not. The update route's eviction and
# get_enabled_mcp_tools' mid-probe guard both key off this so they can't drift.
TOOL_CACHE_INVALIDATING_FIELDS = frozenset({"url", "headers_json", "use_oauth", "is_enabled"})
def get_cached_tools(server_id: str) -> Optional[list[dict]]:
return _tool_cache.get(server_id)
def cache_tools(server_id: str, tools: list[dict]) -> None:
_tool_cache[server_id] = tools
_probe_cooloff_until.pop(server_id, None)
def record_probe_failure(server_id: str, use_oauth: bool = False) -> None:
cooloff = OAUTH_FAILED_PROBE_COOLOFF_SECONDS if use_oauth else FAILED_PROBE_COOLOFF_SECONDS
_probe_cooloff_until[server_id] = time.monotonic() + cooloff
def in_failure_cooloff(server_id: str) -> bool:
return _probe_cooloff_until.get(server_id, 0.0) > time.monotonic()
def invalidate_tool_cache(server_id: Optional[str] = None) -> None:
"""Evict one server's cached tools, or every entry when server_id is None."""
if server_id is None:
_tool_cache.clear()
_probe_cooloff_until.clear()
else:
_tool_cache.pop(server_id, None)
_probe_cooloff_until.pop(server_id, None)
def _flatten_result(result: Any) -> str:
parts = []
for block in getattr(result, "content", None) or []:

View file

@ -24,11 +24,16 @@ import urllib.request
from core.inference.mcp_client import (
MCP_TOOL_PREFIX,
TOOL_CACHE_INVALIDATING_FIELDS,
cache_tools,
call_tool_sync,
get_cached_tools,
in_failure_cooloff,
is_stdio,
list_tools_async,
parse_server_headers,
probe_timeout,
record_probe_failure,
stdio_mcp_enabled,
)
from storage import mcp_servers_db
@ -634,28 +639,56 @@ async def get_enabled_mcp_tools() -> list[dict]:
if not servers:
return []
results = await asyncio.gather(
*(
list_tools_async(
url = s["url"],
headers = parse_server_headers(s),
timeout = probe_timeout(s["url"], bool(s.get("use_oauth"))),
use_oauth = bool(s.get("use_oauth")),
)
for s in servers
),
return_exceptions = True,
)
# Skip servers still in their post-failure cool-off, otherwise a down
# server gets re-probed -- and blocks the send for the full timeout -- on
# every message.
uncached = [
s for s in servers if get_cached_tools(s["id"]) is None and not in_failure_cooloff(s["id"])
]
if uncached:
results = await asyncio.gather(
*(
list_tools_async(
url = s["url"],
headers = parse_server_headers(s),
timeout = probe_timeout(s["url"], bool(s.get("use_oauth"))),
use_oauth = bool(s.get("use_oauth")),
)
for s in uncached
),
return_exceptions = True,
)
# An edit/delete can land while we await a probe (up to 305 s for
# OAuth); its cache eviction is a no-op against an entry we haven't
# written yet. Re-read and drop a result whose server changed or
# was removed mid-probe, else a stale tool list caches indefinitely.
current = {s["id"]: s for s in mcp_servers_db.list_servers()}
for server, payload in zip(uncached, results):
# Guard the failure branch too: a stale failure must not park a
# cool-off on the fresh config, or the server the user just fixed
# is skipped for the whole window.
fresh = current.get(server["id"])
if fresh is None or any(
fresh.get(k) != server.get(k) for k in TOOL_CACHE_INVALIDATING_FIELDS
):
continue
if isinstance(payload, BaseException):
logger.warning(
"MCP server '%s' (%s) discovery failed: %s",
server.get("display_name") or server["id"],
server.get("url"),
payload,
)
# Failures aren't cached, but record one so a down server
# isn't re-probed every send during the cool-off.
record_probe_failure(server["id"], bool(fresh.get("use_oauth")))
continue
cache_tools(server["id"], payload)
specs: list[dict] = []
for server, payload in zip(servers, results):
if isinstance(payload, BaseException):
logger.warning(
"MCP server '%s' (%s) discovery failed: %s",
server.get("display_name") or server["id"],
server.get("url"),
payload,
)
for server in servers:
payload = get_cached_tools(server["id"])
if payload is None:
continue
specs.extend(_mcp_specs_for_server(server, payload))
return specs

View file

@ -10,12 +10,16 @@ from fastapi import APIRouter, Depends, HTTPException
from auth.authentication import get_current_subject
from core.inference.mcp_client import (
TOOL_CACHE_INVALIDATING_FIELDS,
cache_tools,
clear_oauth_tokens_async,
invalidate_tool_cache,
is_stdio,
list_tools_async,
parse_server_headers,
parse_stdio_command,
probe_timeout,
record_probe_failure,
stdio_mcp_enabled,
)
from core.inference.mcp_config_import import parse_mcp_config
@ -198,6 +202,11 @@ async def update_mcp_server(
):
await clear_oauth_tokens_async(old["url"])
mcp_servers_db.update_server(server_id, changes)
# A new endpoint/auth makes cached tools wrong and disabling makes them
# unreachable, so drop them and let the next send re-probe; a rename
# leaves them valid.
if changes.keys() & TOOL_CACHE_INVALIDATING_FIELDS:
invalidate_tool_cache(server_id)
return _row_to_response(mcp_servers_db.get_server(server_id))
@ -209,6 +218,7 @@ async def delete_mcp_server(server_id: str, current_subject: str = Depends(get_c
if old.get("use_oauth"):
await clear_oauth_tokens_async(old["url"])
mcp_servers_db.delete_server(server_id)
invalidate_tool_cache(server_id)
@router.post("/{server_id}/refresh", response_model = McpServerProbeResult)
@ -238,8 +248,23 @@ async def refresh_mcp_server_tools(
error = str(exc),
exc_info = True,
)
current = mcp_servers_db.get_server(server_id)
if current is not None and not any(
current.get(k) != server.get(k) for k in TOOL_CACHE_INVALIDATING_FIELDS
):
# Start the cool-off so the next chat send doesn't immediately re-hang
# on this server's timeout. If the row changed while the probe was
# awaiting, the failure belongs to the old config and must not park
# the newly edited server.
record_probe_failure(server_id, use_oauth)
return McpServerProbeResult(ok = False, error = safe_curated_detail(exc))
# Warm the chat-path cache so the next send skips re-probing.
current = mcp_servers_db.get_server(server_id)
if current is not None and not any(
current.get(k) != server.get(k) for k in TOOL_CACHE_INVALIDATING_FIELDS
):
cache_tools(server_id, tools)
return McpServerProbeResult(ok = True, tool_count = len(tools))

View file

@ -598,3 +598,614 @@ def test_safetensors_agentic_empty_allowlist_still_means_allow_all():
)
# Empty allow-list = run anything (preserved contract).
assert calls == [("python", {"code": "1"})] or len(calls) >= 1
# ── discovery cache ─────────────────────────────────────────────────
def _one_tool(name = "echo"):
return [{"name": name, "inputSchema": {"type": "object", "properties": {}}}]
def test_get_enabled_mcp_tools_caches_discovery(tmp_path, monkeypatch):
"""A second send must serve tools from cache instead of re-probing."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
calls: list[str] = []
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
calls.append(url)
return _one_tool()
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
first = asyncio.run(tools_mod.get_enabled_mcp_tools())
second = asyncio.run(tools_mod.get_enabled_mcp_tools())
assert len(calls) == 1 # probed once, cache hit on the second send
assert [t["function"]["name"] for t in first] == ["mcp__s1__echo"]
assert first == second
def test_get_enabled_mcp_tools_does_not_cache_failures(tmp_path, monkeypatch):
"""A failed probe isn't cached: once the cool-off elapses, it's retried."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
attempts = {"n": 0}
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
attempts["n"] += 1
if attempts["n"] == 1:
raise RuntimeError("server down")
return _one_tool()
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == [] # failure -> empty
# Expire the cool-off (an until-time in the past) so the server is retried.
mcp_client._probe_cooloff_until["s1"] = 0.0
second = asyncio.run(tools_mod.get_enabled_mcp_tools())
assert attempts["n"] == 2 # retried after the cool-off, not cached
assert [t["function"]["name"] for t in second] == ["mcp__s1__echo"]
def test_refresh_warms_tool_cache(tmp_path, monkeypatch):
"""Clicking Refresh must populate the cache the chat path reads."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
import routes.mcp_servers as routes_mcp
monkeypatch.setattr(mcp_client, "_tool_cache", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
async def fake_refresh(
url,
headers = None,
timeout = None,
use_oauth = False,
):
return _one_tool()
monkeypatch.setattr(routes_mcp, "list_tools_async", fake_refresh)
res = asyncio.run(routes_mcp.refresh_mcp_server_tools("s1", current_subject = "u"))
assert res.ok and res.tool_count == 1
def boom(*a, **k):
raise AssertionError("chat path re-probed despite a warm cache")
monkeypatch.setattr(tools_mod, "list_tools_async", boom)
specs = asyncio.run(tools_mod.get_enabled_mcp_tools())
assert [t["function"]["name"] for t in specs] == ["mcp__s1__echo"]
def test_update_url_evicts_tool_cache(tmp_path, monkeypatch):
"""Re-pointing the URL must drop the old endpoint's cached tools."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from models.mcp_servers import McpServerUpdate
import routes.mcp_servers as routes_mcp
monkeypatch.setattr(mcp_client, "_tool_cache", {"s1": _one_tool("stale")})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://old/mcp", is_enabled = True)
asyncio.run(
routes_mcp.update_mcp_server(
"s1", McpServerUpdate(url = "https://new/mcp"), current_subject = "u"
)
)
assert mcp_client.get_cached_tools("s1") is None
def test_update_display_name_keeps_tool_cache(tmp_path, monkeypatch):
"""A rename touches no endpoint, so the cache must survive it."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from models.mcp_servers import McpServerUpdate
import routes.mcp_servers as routes_mcp
cached = _one_tool()
monkeypatch.setattr(mcp_client, "_tool_cache", {"s1": cached})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
asyncio.run(
routes_mcp.update_mcp_server("s1", McpServerUpdate(display_name = "B"), current_subject = "u")
)
assert mcp_client.get_cached_tools("s1") == cached
def test_update_disable_evicts_tool_cache(tmp_path, monkeypatch):
"""Disabling a server must drop its cached tools, not leave them unread."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from models.mcp_servers import McpServerUpdate
import routes.mcp_servers as routes_mcp
monkeypatch.setattr(mcp_client, "_tool_cache", {"s1": _one_tool()})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
asyncio.run(
routes_mcp.update_mcp_server("s1", McpServerUpdate(is_enabled = False), current_subject = "u")
)
assert mcp_client.get_cached_tools("s1") is None
def test_delete_evicts_tool_cache(tmp_path, monkeypatch):
"""Deleting a server must not leave its tools cached."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
import routes.mcp_servers as routes_mcp
monkeypatch.setattr(mcp_client, "_tool_cache", {"s1": _one_tool()})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
asyncio.run(routes_mcp.delete_mcp_server("s1", current_subject = "u"))
assert mcp_client.get_cached_tools("s1") is None
def test_invalidate_tool_cache_clears_all(monkeypatch):
from core.inference import mcp_client
monkeypatch.setattr(mcp_client, "_tool_cache", {"a": _one_tool(), "b": _one_tool()})
mcp_client.invalidate_tool_cache()
assert mcp_client.get_cached_tools("a") is None
assert mcp_client.get_cached_tools("b") is None
def test_get_enabled_mcp_tools_probes_only_uncached(tmp_path, monkeypatch):
"""An already-cached server must not be re-probed alongside a cold one."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {"s1": _one_tool("cached")})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://a/mcp", is_enabled = True)
mcp_servers_db.create_server(id = "s2", display_name = "B", url = "https://b/mcp", is_enabled = True)
probed: list[str] = []
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
probed.append(url)
return _one_tool("fresh")
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
specs = asyncio.run(tools_mod.get_enabled_mcp_tools())
assert probed == ["https://b/mcp"] # only the uncached server is probed
assert sorted(t["function"]["name"] for t in specs) == ["mcp__s1__cached", "mcp__s2__fresh"]
def test_get_enabled_mcp_tools_partial_failure_caches_healthy(tmp_path, monkeypatch):
"""One server failing must not stop the others from being cached/served."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://bad/mcp", is_enabled = True)
mcp_servers_db.create_server(id = "s2", display_name = "B", url = "https://good/mcp", is_enabled = True)
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
if "bad" in url:
raise RuntimeError("down")
return _one_tool("ok")
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
specs = asyncio.run(tools_mod.get_enabled_mcp_tools())
assert [t["function"]["name"] for t in specs] == ["mcp__s2__ok"]
assert mcp_client.get_cached_tools("s1") is None # failure not cached
assert mcp_client.get_cached_tools("s2") == _one_tool("ok") # healthy cached
def test_get_enabled_mcp_tools_caches_empty_tool_list(tmp_path, monkeypatch):
"""A server exposing zero tools is cached as [] (a hit), not re-probed."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
calls: list[str] = []
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
calls.append(url)
return []
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == []
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == []
assert len(calls) == 1 # [] is a cache hit, not re-probed every send
assert mcp_client.get_cached_tools("s1") == []
def test_update_headers_evicts_tool_cache(tmp_path, monkeypatch):
"""Changing auth headers must drop tools discovered under the old headers."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from models.mcp_servers import McpServerUpdate
import routes.mcp_servers as routes_mcp
monkeypatch.setattr(mcp_client, "_tool_cache", {"s1": _one_tool()})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
asyncio.run(
routes_mcp.update_mcp_server(
"s1",
McpServerUpdate(headers = {"Authorization": "Bearer new"}),
current_subject = "u",
)
)
assert mcp_client.get_cached_tools("s1") is None
def test_get_enabled_mcp_tools_skips_cache_when_config_changes_mid_probe(tmp_path, monkeypatch):
"""A config edit landing during an in-flight probe must not be clobbered
by the now-stale probe result (TOCTOU on the cache write)."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://old/mcp", is_enabled = True)
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
# Simulate a PUT landing while we are awaiting the probe.
mcp_servers_db.update_server("s1", {"url": "https://new/mcp"})
return _one_tool()
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
specs = asyncio.run(tools_mod.get_enabled_mcp_tools())
assert specs == [] # stale result is neither served...
assert mcp_client.get_cached_tools("s1") is None # ...nor cached
def test_get_enabled_mcp_tools_no_cooloff_when_config_changes_mid_failed_probe(
tmp_path, monkeypatch
):
"""An edit landing while a probe of the OLD config is failing must not park
a cool-off on the now-fresh config -- else the re-pointed server the user
just fixed is needlessly skipped for the whole cool-off window."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://old/mcp", is_enabled = True)
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
# The user re-points the server while the old endpoint's probe fails.
mcp_servers_db.update_server("s1", {"url": "https://new/mcp"})
raise RuntimeError("old endpoint down")
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == []
# The failure was for the OLD config, so the new one must stay re-probable.
assert not mcp_client.in_failure_cooloff("s1")
def test_get_enabled_mcp_tools_no_cooloff_when_server_deleted_mid_failed_probe(
tmp_path, monkeypatch
):
"""A delete landing while a probe fails must not leave an orphan cool-off
entry keyed by the since-removed server id."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
mcp_servers_db.delete_server("s1")
raise RuntimeError("down")
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == []
assert "s1" not in mcp_client._probe_cooloff_until # no orphan cool-off
def test_get_enabled_mcp_tools_skips_failed_server_during_cooloff(tmp_path, monkeypatch):
"""A down server is probed once, then skipped during the cool-off instead
of being re-probed (and re-hung) on every send."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
attempts = {"n": 0}
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
attempts["n"] += 1
raise RuntimeError("down")
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == [] # probes, fails
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == [] # within cool-off
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == [] # still skipped
assert attempts["n"] == 1 # only the first send probed
def test_cache_tools_clears_failure_cooloff(monkeypatch):
"""A successful probe lifts a server's failure cool-off."""
from core.inference import mcp_client
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_client.record_probe_failure("s1")
assert mcp_client.in_failure_cooloff("s1")
mcp_client.cache_tools("s1", _one_tool())
assert not mcp_client.in_failure_cooloff("s1")
def test_oauth_failure_cools_off_longer_than_plain(monkeypatch):
"""An OAuth server's failure cools off longer than a plain server's, so its
multi-minute probe hang doesn't recur every minute."""
from core.inference import mcp_client
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_client.record_probe_failure("plain", use_oauth = False)
mcp_client.record_probe_failure("oauth", use_oauth = True)
assert mcp_client._probe_cooloff_until["oauth"] > mcp_client._probe_cooloff_until["plain"]
def test_invalidate_clears_failure_cooloff(monkeypatch):
"""Eviction drops the failure cool-off so an edited server re-probes at once."""
from core.inference import mcp_client
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {"s1": 1.0, "s2": 2.0})
mcp_client.invalidate_tool_cache("s1")
assert "s1" not in mcp_client._probe_cooloff_until
assert "s2" in mcp_client._probe_cooloff_until
mcp_client.invalidate_tool_cache()
assert mcp_client._probe_cooloff_until == {}
def test_refresh_failure_records_cooloff(tmp_path, monkeypatch):
"""A failed manual refresh starts the cool-off so the next chat send does
not immediately hang on the down server."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
import routes.mcp_servers as routes_mcp
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
async def boom(
url,
headers = None,
timeout = None,
use_oauth = False,
):
raise RuntimeError("down")
monkeypatch.setattr(routes_mcp, "list_tools_async", boom)
res = asyncio.run(routes_mcp.refresh_mcp_server_tools("s1", current_subject = "u"))
assert res.ok is False
assert mcp_client.in_failure_cooloff("s1")
def test_refresh_drops_result_when_config_changes_mid_probe(tmp_path, monkeypatch):
"""A manual refresh must not warm the chat cache with tools discovered
under an old config if the server is edited while the probe is in flight."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
import routes.mcp_servers as routes_mcp
monkeypatch.setattr(mcp_client, "_tool_cache", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://old/mcp", is_enabled = True)
async def fake_refresh(
url,
headers = None,
timeout = None,
use_oauth = False,
):
mcp_servers_db.update_server("s1", {"url": "https://new/mcp"})
return _one_tool("stale")
monkeypatch.setattr(routes_mcp, "list_tools_async", fake_refresh)
res = asyncio.run(routes_mcp.refresh_mcp_server_tools("s1", current_subject = "u"))
assert res.ok and res.tool_count == 1
assert mcp_client.get_cached_tools("s1") is None
def test_refresh_failure_no_cooloff_when_config_changes_mid_probe(tmp_path, monkeypatch):
"""A manual refresh failure for an old config must not cool off the freshly
edited server."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
import routes.mcp_servers as routes_mcp
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://old/mcp", is_enabled = True)
async def boom(
url,
headers = None,
timeout = None,
use_oauth = False,
):
mcp_servers_db.update_server("s1", {"url": "https://new/mcp"})
raise RuntimeError("old endpoint down")
monkeypatch.setattr(routes_mcp, "list_tools_async", boom)
res = asyncio.run(routes_mcp.refresh_mcp_server_tools("s1", current_subject = "u"))
assert res.ok is False
assert not mcp_client.in_failure_cooloff("s1")
def test_get_enabled_mcp_tools_drops_result_when_server_deleted_mid_probe(tmp_path, monkeypatch):
"""A delete landing while a probe is in flight must drop the now-orphan
result -- the `fresh is None` arm of the mid-probe TOCTOU guard. The
result is neither served nor cached under the since-removed id."""
import asyncio
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(id = "s1", display_name = "A", url = "https://x/mcp", is_enabled = True)
async def fake(
url,
headers = None,
timeout = None,
use_oauth = False,
):
# Simulate a DELETE landing while we await the probe.
mcp_servers_db.delete_server("s1")
return _one_tool()
monkeypatch.setattr(tools_mod, "list_tools_async", fake)
specs = asyncio.run(tools_mod.get_enabled_mcp_tools())
assert specs == [] # orphan result not served
assert mcp_client.get_cached_tools("s1") is None # nor cached under a gone id
def test_oauth_probe_failure_in_chat_path_uses_long_cooloff(tmp_path, monkeypatch):
"""When an OAuth server fails discovery during a send, the chat path must
record the OAuth (long) cool-off, not the plain one -- otherwise its
multi-minute browser hang recurs every minute."""
import asyncio
import time
_reset_db(tmp_path, monkeypatch)
from core.inference import mcp_client
from core.inference import tools as tools_mod
monkeypatch.setattr(mcp_client, "_tool_cache", {})
monkeypatch.setattr(mcp_client, "_probe_cooloff_until", {})
mcp_servers_db.create_server(
id = "s1",
display_name = "A",
url = "https://x/mcp",
is_enabled = True,
use_oauth = True,
)
async def boom(
url,
headers = None,
timeout = None,
use_oauth = False,
):
raise RuntimeError("oauth down")
monkeypatch.setattr(tools_mod, "list_tools_async", boom)
assert asyncio.run(tools_mod.get_enabled_mcp_tools()) == []
assert mcp_client.in_failure_cooloff("s1")
# The recorded window must exceed the plain cool-off, proving the OAuth
# branch (use_oauth=True) fired -- not the 60 s default.
remaining = mcp_client._probe_cooloff_until["s1"] - time.monotonic()
assert remaining > mcp_client.FAILED_PROBE_COOLOFF_SECONDS

View file

@ -19,6 +19,10 @@ from storage import mcp_servers_db
def _reset_db(tmp_path, monkeypatch):
monkeypatch.setenv("UNSLOTH_STUDIO_HOME", str(tmp_path))
monkeypatch.setattr(mcp_servers_db, "_schema_ready", False)
# The discovered-tool cache is process-global and keyed by server id; tests
# reuse "stdio1", so clear it (and the failure cool-off) for isolation —
# otherwise a prior test's warm cache makes discovery skip its probe.
mcp_client.invalidate_tool_cache()
def _enable(monkeypatch):