160 lines
4.8 KiB
Python
160 lines
4.8 KiB
Python
"""平台托管来源的显式注册表和统一生命周期入口。
|
||
|
||
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
|