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