zcbot/tests/test_platform_sources.py

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()