137 lines
6.3 KiB
Python
137 lines
6.3 KiB
Python
|
|
import asyncio
|
||
|
|
import json
|
||
|
|
from contextvars import ContextVar
|
||
|
|
|
||
|
|
import httpx
|
||
|
|
import pytest
|
||
|
|
from fastmcp import FastMCP
|
||
|
|
from fastapi.testclient import TestClient
|
||
|
|
|
||
|
|
from hub_core.mcp import CORE_TOOL_NAMES, RUNTIME_TOOL_NAMES, HubCoreMCPServer
|
||
|
|
from hub_core.mcp.credentials import StdioTokenFile
|
||
|
|
from hub_core.runtime.cli import _run_mcp
|
||
|
|
from test_access_boundary import Owners, runtime
|
||
|
|
|
||
|
|
|
||
|
|
def test_runtime_profile_lists_only_equivalent_tools_even_when_attached():
|
||
|
|
server = HubCoreMCPServer(name='runtime', api_base='https://hub.example', backend_profile='runtime')
|
||
|
|
assert {t.name for t in asyncio.run(server.mcp.list_tools())} == RUNTIME_TOOL_NAMES
|
||
|
|
assert len(CORE_TOOL_NAMES - RUNTIME_TOOL_NAMES) == 27
|
||
|
|
server.attach_to(FastMCP('host'))
|
||
|
|
assert {t.name for t in asyncio.run(server.mcp.list_tools())} == RUNTIME_TOOL_NAMES
|
||
|
|
assert not server.trailing_slash
|
||
|
|
|
||
|
|
|
||
|
|
def test_actual_tool_calls_preserve_concurrent_caller_credentials(monkeypatch):
|
||
|
|
credential = ContextVar('caller')
|
||
|
|
server = HubCoreMCPServer(name='runtime', api_base='https://hub.example', backend_profile='runtime',
|
||
|
|
token_provider=credential.get, require_credentials=True)
|
||
|
|
seen = []
|
||
|
|
original = httpx.Client
|
||
|
|
def handle(request):
|
||
|
|
seen.append((request.headers['authorization'], request.url.path))
|
||
|
|
return httpx.Response(200, json={'ok': True})
|
||
|
|
def client(**kwargs):
|
||
|
|
assert kwargs['trust_env'] is False
|
||
|
|
assert kwargs['follow_redirects'] is False
|
||
|
|
return original(**kwargs, transport=httpx.MockTransport(handle))
|
||
|
|
monkeypatch.setattr('hub_core.mcp.server.httpx.Client', client)
|
||
|
|
async def invoke(token):
|
||
|
|
credential.set(token)
|
||
|
|
await asyncio.sleep(0)
|
||
|
|
await server.mcp.call_tool('query_workloads', {})
|
||
|
|
async def run():
|
||
|
|
await asyncio.gather(invoke('caller-a'), invoke('caller-b'))
|
||
|
|
asyncio.run(run())
|
||
|
|
assert sorted(seen) == [('Bearer caller-a','/ports/projections/workloads'),
|
||
|
|
('Bearer caller-b','/ports/projections/workloads')]
|
||
|
|
|
||
|
|
|
||
|
|
def test_rotating_stdio_credential_reaches_real_enforced_runtime(tmp_path, monkeypatch):
|
||
|
|
path = tmp_path / 'credential'
|
||
|
|
path.write_text('verified-root\n')
|
||
|
|
owners = Owners()
|
||
|
|
seen = []
|
||
|
|
server = HubCoreMCPServer(name='runtime', api_base='https://hub.example', backend_profile='runtime',
|
||
|
|
token_provider=StdioTokenFile(path), require_credentials=True)
|
||
|
|
original = httpx.Client
|
||
|
|
with TestClient(runtime(owners)) as target:
|
||
|
|
def handle(request):
|
||
|
|
response = target.get(request.url.path, headers={'Authorization':request.headers['authorization']})
|
||
|
|
seen.append(response.status_code)
|
||
|
|
return httpx.Response(response.status_code, content=response.content)
|
||
|
|
monkeypatch.setattr('hub_core.mcp.server.httpx.Client',
|
||
|
|
lambda **kw: original(**kw, transport=httpx.MockTransport(handle)))
|
||
|
|
asyncio.run(server.mcp.call_tool('query_workloads', {}))
|
||
|
|
assert owners.requests[-1].actor.subject == 'immutable-root'
|
||
|
|
# Absent workload projection returns 503 after successful authorization.
|
||
|
|
path.write_text('invalid-caller')
|
||
|
|
result = asyncio.run(server.mcp.call_tool('query_workloads', {}))
|
||
|
|
assert seen[-1] == 401
|
||
|
|
assert '401' in str(result)
|
||
|
|
path.unlink()
|
||
|
|
before = len(seen)
|
||
|
|
result = asyncio.run(server.mcp.call_tool('query_workloads', {}))
|
||
|
|
assert len(seen) == before
|
||
|
|
assert 'Request failed' in str(result)
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.parametrize('value', ['', 'two tokens', 'x'*16386, 'é', 'x'*16384+'\nextra'])
|
||
|
|
def test_credential_file_rejects_invalid_values(tmp_path, value):
|
||
|
|
path = tmp_path / 'credential'
|
||
|
|
path.write_text(value)
|
||
|
|
with pytest.raises((ValueError, UnicodeError)):
|
||
|
|
StdioTokenFile(path)()
|
||
|
|
|
||
|
|
|
||
|
|
def test_cli_enforcement_requires_private_stdio(monkeypatch, tmp_path):
|
||
|
|
monkeypatch.setenv('HUB_CORE_ENV','production')
|
||
|
|
for transport, path in [('http',None), ('http',tmp_path/'token'), ('stdio',None)]:
|
||
|
|
with pytest.raises(SystemExit):
|
||
|
|
_run_mcp('127.0.0.1',8011,transport,'https://hub.example',path)
|
||
|
|
captured = []
|
||
|
|
monkeypatch.setattr(FastMCP, 'run', lambda self, **kw: captured.append(kw))
|
||
|
|
_run_mcp('127.0.0.1',8011,'stdio','https://hub.example',tmp_path/'token')
|
||
|
|
assert captured == [{'transport':'stdio'}]
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.parametrize('url',['http://hub.example','https://user:password@hub.example','https://localhost'])
|
||
|
|
def test_credential_transport_requires_trusted_https(url):
|
||
|
|
with pytest.raises(ValueError):
|
||
|
|
HubCoreMCPServer(name='bad',api_base=url,token_provider=lambda:'token')
|
||
|
|
|
||
|
|
|
||
|
|
@pytest.mark.parametrize('segment',['..','../docs','a/b','a?b','a#b','%2e%2e','a\\b'])
|
||
|
|
def test_facet_arguments_cannot_change_route(segment):
|
||
|
|
with pytest.raises(ValueError):
|
||
|
|
HubCoreMCPServer._segment(segment)
|
||
|
|
|
||
|
|
|
||
|
|
def test_all_runtime_tools_target_admitted_runtime_get_routes(monkeypatch):
|
||
|
|
from starlette.routing import Match
|
||
|
|
from hub_core.security.boundary import iter_routes, route_key
|
||
|
|
from pathlib import Path
|
||
|
|
catalog = json.loads(Path('hub_core/security/routes.json').read_text())['routes']
|
||
|
|
app = runtime(Owners())
|
||
|
|
server = HubCoreMCPServer(name='runtime',api_base='https://hub.example',backend_profile='runtime',
|
||
|
|
token_provider=lambda:'caller',require_credentials=True)
|
||
|
|
original = httpx.Client
|
||
|
|
seen = []
|
||
|
|
def handle(request):
|
||
|
|
scope = {'type':'http','path':request.url.path,'root_path':'','method':request.method}
|
||
|
|
route = next(r for r in iter_routes(app) if r.matches(scope)[0] == Match.FULL)
|
||
|
|
assert route_key(route,request.method) in catalog
|
||
|
|
seen.append(request.url.path)
|
||
|
|
return httpx.Response(200,json={'ok':True})
|
||
|
|
monkeypatch.setattr('hub_core.mcp.server.httpx.Client',
|
||
|
|
lambda **kw: original(**kw,transport=httpx.MockTransport(handle)))
|
||
|
|
calls = {'query_repository_navigation':{},
|
||
|
|
'get_repository_navigation_facet':{'facet_kind':'category','facet_value':'tools'},
|
||
|
|
'query_workloads':{}, 'resolve_workload_reference':{'rapp_id':'rapp:test','name':'test'}}
|
||
|
|
async def run():
|
||
|
|
for name, arguments in calls.items():
|
||
|
|
await server.mcp.call_tool(name,arguments)
|
||
|
|
asyncio.run(run())
|
||
|
|
assert set(calls) == RUNTIME_TOOL_NAMES
|
||
|
|
assert len(seen) == 4
|