zcbot/platform_sources/registry.py

160 lines
4.8 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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