207 lines
7.3 KiB
Python
207 lines
7.3 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
import tempfile
|
|
import unittest
|
|
from pathlib import Path
|
|
from unittest.mock import patch
|
|
|
|
from core.executor_docker import _sandbox_env
|
|
from platform_sources.registry import (
|
|
MaterialsLibraryProvider,
|
|
MaterialsProjectProvider,
|
|
PaperServerProvider,
|
|
PlatformSourceContext,
|
|
build_platform_source_tools,
|
|
)
|
|
from platform_sources.materials_library import MaterialsLibrarySearchTool
|
|
from platform_sources.materials_project import MaterialsProjectSearchTool
|
|
from platform_sources.paper_server import PaperServerSearchTool
|
|
from tools.run_python import RunPythonTool
|
|
|
|
|
|
class _Provider:
|
|
capabilities = frozenset({"test"})
|
|
|
|
def __init__(self, source_id, *, available=True, tools=None, failure=None):
|
|
self.source_id = source_id
|
|
self._available = available
|
|
self._tools = list(tools or [])
|
|
self._failure = failure
|
|
|
|
def available(self):
|
|
if self._failure == "available":
|
|
raise RuntimeError("secret-value")
|
|
return self._available
|
|
|
|
def build_tools(self, context):
|
|
if self._failure == "build":
|
|
raise RuntimeError("secret-value")
|
|
return self._tools
|
|
|
|
|
|
class _Tool:
|
|
name = "survivor"
|
|
|
|
|
|
class PlatformSourceRegistryTests(unittest.TestCase):
|
|
def _context(self, root: Path) -> PlatformSourceContext:
|
|
return PlatformSourceContext(root, root, root)
|
|
|
|
def test_provider_failure_isolated_and_diagnostic_is_secret_free(self):
|
|
with tempfile.TemporaryDirectory() as tmp, self.assertLogs(
|
|
"platform_sources.registry", level="WARNING"
|
|
) as logs:
|
|
built = build_platform_source_tools(
|
|
self._context(Path(tmp)),
|
|
providers=(
|
|
_Provider("bad_available", failure="available"),
|
|
_Provider("bad_build", failure="build"),
|
|
_Provider("good", tools=[_Tool()]),
|
|
),
|
|
)
|
|
self.assertEqual([tool.name for tool in built], ["survivor"])
|
|
output = "\n".join(logs.output)
|
|
self.assertIn("bad_available", output)
|
|
self.assertIn("bad_build", output)
|
|
self.assertIn("RuntimeError", output)
|
|
self.assertNotIn("secret-value", output)
|
|
|
|
def test_each_env_gate_is_independent(self):
|
|
providers = (
|
|
PaperServerProvider(),
|
|
MaterialsLibraryProvider(),
|
|
MaterialsProjectProvider(),
|
|
)
|
|
empty = {
|
|
"PAPER_SERVER_API_KEY": "",
|
|
"DOCUMENT_SEARCH_API_KEY": "",
|
|
"MP_API_KEY": "",
|
|
}
|
|
with patch.dict(os.environ, empty, clear=False), patch(
|
|
"platform_sources.registry.MPRester", object()
|
|
):
|
|
self.assertEqual([provider.available() for provider in providers], [False] * 3)
|
|
for env_name, index in (
|
|
("PAPER_SERVER_API_KEY", 0),
|
|
("DOCUMENT_SEARCH_API_KEY", 1),
|
|
("MP_API_KEY", 2),
|
|
):
|
|
values = dict(empty)
|
|
values[env_name] = "configured"
|
|
with patch.dict(os.environ, values, clear=False):
|
|
available = [provider.available() for provider in providers]
|
|
self.assertTrue(available[index])
|
|
self.assertEqual(sum(available), 1)
|
|
|
|
def test_invalid_source_config_does_not_block_other_source(self):
|
|
with tempfile.TemporaryDirectory() as tmp, patch.dict(
|
|
os.environ,
|
|
{
|
|
"PAPER_SERVER_API_KEY": "paper-secret",
|
|
"PAPER_SERVER_URL": "not-a-url",
|
|
"DOCUMENT_SEARCH_API_KEY": "library-secret",
|
|
"MP_API_KEY": "",
|
|
},
|
|
clear=False,
|
|
), self.assertLogs("platform_sources.registry", level="WARNING") as logs:
|
|
tools = build_platform_source_tools(self._context(Path(tmp)))
|
|
names = {tool.name for tool in tools}
|
|
self.assertEqual(
|
|
names,
|
|
{
|
|
"materials_library_list",
|
|
"materials_library_search",
|
|
"materials_library_fetch",
|
|
},
|
|
)
|
|
output = "\n".join(logs.output)
|
|
self.assertIn("paper_server", output)
|
|
self.assertIn("ValueError", output)
|
|
self.assertNotIn("paper-secret", output)
|
|
|
|
def test_registered_tool_names_are_source_specific(self):
|
|
with tempfile.TemporaryDirectory() as tmp, patch.dict(
|
|
os.environ,
|
|
{
|
|
"PAPER_SERVER_API_KEY": "paper-secret",
|
|
"DOCUMENT_SEARCH_API_KEY": "library-secret",
|
|
"MP_API_KEY": "mp-secret",
|
|
},
|
|
clear=False,
|
|
), patch("platform_sources.registry.MPRester", object()):
|
|
tools = build_platform_source_tools(self._context(Path(tmp)))
|
|
names = {tool.schema["function"]["name"] for tool in tools}
|
|
self.assertEqual(
|
|
names,
|
|
{
|
|
"paper_server_search",
|
|
"paper_server_get",
|
|
"paper_server_fetch",
|
|
"materials_library_list",
|
|
"materials_library_search",
|
|
"materials_library_fetch",
|
|
"materials_project_search",
|
|
"materials_project_get_structure",
|
|
"materials_project_get_entries",
|
|
},
|
|
)
|
|
self.assertFalse(
|
|
names
|
|
& {
|
|
"paper_search",
|
|
"paper_get",
|
|
"paper_fetch",
|
|
"document_list_kb",
|
|
"document_search",
|
|
"document_download",
|
|
"mp_search_summary",
|
|
"mp_get_structure",
|
|
"mp_get_entries",
|
|
"platform_source_call",
|
|
}
|
|
)
|
|
|
|
def test_platform_credentials_do_not_enter_execution_environments(self):
|
|
secrets = {
|
|
"PAPER_SERVER_API_KEY": "paper-secret",
|
|
"DOCUMENT_SEARCH_API_KEY": "library-secret",
|
|
"MP_API_KEY": "mp-secret",
|
|
}
|
|
with patch.dict(os.environ, secrets, clear=False):
|
|
host_env = RunPythonTool()._filtered_env()
|
|
docker_env = _sandbox_env()
|
|
for name in secrets:
|
|
self.assertNotIn(name, host_env)
|
|
self.assertNotIn(name, docker_env)
|
|
|
|
def test_platform_credentials_are_redacted_from_tool_results(self):
|
|
secrets = {
|
|
"PAPER_SERVER_API_KEY": "paper-secret",
|
|
"DOCUMENT_SEARCH_API_KEY": "library-secret",
|
|
"MP_API_KEY": "mp-secret",
|
|
}
|
|
with patch.dict(os.environ, secrets, clear=False), patch(
|
|
"platform_sources.paper_server._config",
|
|
side_effect=RuntimeError("paper-secret"),
|
|
), patch(
|
|
"platform_sources.materials_library.client.search",
|
|
side_effect=RuntimeError("api_key=library-secret"),
|
|
), patch(
|
|
"platform_sources.materials_project._mpr",
|
|
side_effect=RuntimeError("mp-secret"),
|
|
):
|
|
results = (
|
|
PaperServerSearchTool().execute(keyword="cement"),
|
|
MaterialsLibrarySearchTool().execute(queries=["cement"]),
|
|
MaterialsProjectSearchTool().execute(formula="Ca3SiO5"),
|
|
)
|
|
joined = "\n".join(results)
|
|
for value in secrets.values():
|
|
self.assertNotIn(value, joined)
|
|
self.assertIn("[REDACTED]", joined)
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|