"""平台托管来源的显式注册表和统一生命周期入口。 Provider 只统一来源标识、能力声明、可用性和 Tool 装配;各来源的业务 API 保持专用 typed contract,不在这里抽象成 search/get/fetch 万能接口。 """ from __future__ import annotations import logging import os from collections.abc import Sequence from dataclasses import dataclass from pathlib import Path from typing import Protocol from urllib.parse import urlparse from tools.base import Tool from .materials_library import ( MaterialsLibraryFetchTool, MaterialsLibraryListTool, MaterialsLibrarySearchTool, ) from .materials_project import ( MPRester, MaterialsProjectGetEntriesTool, MaterialsProjectGetStructureTool, MaterialsProjectSearchTool, ) from .paper_server import ( PaperServerFetchTool, PaperServerGetTool, PaperServerSearchTool, ) logger = logging.getLogger(__name__) @dataclass(frozen=True) class PlatformSourceContext: """Provider 装配 Tool 所需的非敏感宿主上下文。""" base_dir: Path user_root: Path working_dir: Path class PlatformSourceProvider(Protocol): source_id: str capabilities: frozenset[str] def available(self) -> bool: ... def build_tools(self, context: PlatformSourceContext) -> list[Tool]: ... def _env_set(name: str) -> bool: return bool(os.environ.get(name, "").strip()) def _require_valid_http_url(name: str, default: str) -> None: value = os.environ.get(name, default).strip() parsed = urlparse(value) if parsed.scheme not in {"http", "https"} or not parsed.netloc: raise ValueError(f"invalid {name}") class PaperServerProvider: source_id = "paper_server" capabilities = frozenset({"search", "metadata", "download"}) def available(self) -> bool: if not _env_set("PAPER_SERVER_API_KEY"): return False _require_valid_http_url("PAPER_SERVER_URL", "http://paper.xxhhcty.xyz:8080") return True def build_tools(self, context: PlatformSourceContext) -> list[Tool]: common = {"base_dir": context.base_dir, "user_root": context.user_root} return [ PaperServerSearchTool(**common), PaperServerGetTool(**common), PaperServerFetchTool(working_dir=context.working_dir, **common), ] class MaterialsLibraryProvider: source_id = "materials_library" capabilities = frozenset({"catalog", "batch_search", "batch_download"}) def available(self) -> bool: if not _env_set("DOCUMENT_SEARCH_API_KEY"): return False _require_valid_http_url( "DOCUMENT_SEARCH_URL", "https://ai.ctc-zc.com:8100/api" ) return True def build_tools(self, context: PlatformSourceContext) -> list[Tool]: common = {"base_dir": context.base_dir, "user_root": context.user_root} return [ MaterialsLibraryListTool(**common), MaterialsLibrarySearchTool(**common), MaterialsLibraryFetchTool(working_dir=context.working_dir, **common), ] class MaterialsProjectProvider: source_id = "materials_project" capabilities = frozenset({"summary_search", "structure", "entries"}) def available(self) -> bool: if not _env_set("MP_API_KEY"): return False if MPRester is None: raise RuntimeError("Materials Project SDK unavailable") return True def build_tools(self, context: PlatformSourceContext) -> list[Tool]: common = {"base_dir": context.base_dir, "user_root": context.user_root} return [ MaterialsProjectSearchTool(working_dir=context.working_dir, **common), MaterialsProjectGetStructureTool( working_dir=context.working_dir, **common ), MaterialsProjectGetEntriesTool( working_dir=context.working_dir, **common ), ] # 显式可信列表:不扫描目录、不动态 import 部署侧 definition。 TRUSTED_PROVIDERS: tuple[PlatformSourceProvider, ...] = ( PaperServerProvider(), MaterialsLibraryProvider(), MaterialsProjectProvider(), ) def build_platform_source_tools( context: PlatformSourceContext, *, providers: Sequence[PlatformSourceProvider] = TRUSTED_PROVIDERS, ) -> list[Tool]: """逐来源隔离装配;失败只跳过该来源,日志不含配置值或异常正文。""" built: list[Tool] = [] for provider in providers: source_id = getattr(provider, "source_id", "unknown") try: if not provider.available(): continue built.extend(provider.build_tools(context)) except Exception as exc: logger.warning( "platform source %s registration failed (%s)", source_id, type(exc).__name__, ) return built