363 lines
13 KiB
Python
363 lines
13 KiB
Python
"""Host-side paper_server tools.
|
|
|
|
PAPER_SERVER_API_KEY stays in the host control plane. The model receives only
|
|
typed business arguments and trimmed JSON / task-local file paths.
|
|
"""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import os
|
|
import re
|
|
from pathlib import Path
|
|
from typing import Any, Optional
|
|
from urllib.parse import urljoin, urlparse
|
|
|
|
import httpx
|
|
|
|
from .base import Tool
|
|
|
|
_DEFAULT_BASE_URL = "http://paper.xxhhcty.xyz:8080"
|
|
_TIMEOUT = 30.0
|
|
_DOWNLOAD_TIMEOUT = 60.0
|
|
_MAX_DOWNLOAD_BYTES = 100 * 1024 * 1024
|
|
_TYPE_RE = re.compile(r"^[A-Za-z0-9_-]{1,64}$")
|
|
_AUTH_ERR_CODES = {"not_authenticated", "authentication_failed", "permission_denied"}
|
|
_LIST_FIELDS = (
|
|
"id",
|
|
"doi",
|
|
"title",
|
|
"first_author",
|
|
"first_author_institution",
|
|
"publication_year",
|
|
"publication_date",
|
|
"publication_name",
|
|
"has_fulltext_pdf",
|
|
"has_fulltext_xml",
|
|
"has_abstract",
|
|
"is_oa",
|
|
"type",
|
|
"abstract",
|
|
"pdf_url",
|
|
"xml_url",
|
|
)
|
|
|
|
|
|
def _config() -> tuple[str, str, str]:
|
|
base_url = os.environ.get("PAPER_SERVER_URL", _DEFAULT_BASE_URL).strip().rstrip("/")
|
|
api_key = os.environ.get("PAPER_SERVER_API_KEY", "").strip()
|
|
if not api_key:
|
|
raise RuntimeError("PAPER_SERVER_API_KEY env 未设置,无法查询 paper_server")
|
|
return base_url, f"{base_url}/api/resm/paper", api_key
|
|
|
|
|
|
def _params(api_key: str, values: Optional[dict[str, Any]] = None) -> dict[str, Any]:
|
|
result = dict(values or {})
|
|
result["api_key"] = api_key
|
|
return result
|
|
|
|
|
|
def _raise_response_error(response: httpx.Response) -> None:
|
|
err_code = ""
|
|
if response.status_code in (400, 401, 403):
|
|
try:
|
|
response.read()
|
|
body = response.json()
|
|
if isinstance(body, dict):
|
|
err_code = str(body.get("err_code") or "")
|
|
except Exception:
|
|
pass
|
|
if response.status_code in (401, 403) or err_code in _AUTH_ERR_CODES:
|
|
raise RuntimeError(
|
|
f"paper_server auth failed (HTTP {response.status_code}, "
|
|
f"{err_code or 'no err_code'}):请管理员检查平台 PAPER_SERVER_API_KEY"
|
|
)
|
|
if response.status_code >= 400:
|
|
# 不调用 raise_for_status():HTTPStatusError 会携带含 api_key 的完整请求 URL。
|
|
raise RuntimeError(f"paper_server request failed (HTTP {response.status_code})")
|
|
|
|
|
|
def _get_json(
|
|
url: str, *, api_key: str, params: Optional[dict[str, Any]] = None
|
|
) -> Any:
|
|
try:
|
|
response = httpx.get(
|
|
url,
|
|
params=_params(api_key, params),
|
|
timeout=_TIMEOUT,
|
|
)
|
|
except httpx.RequestError as exc:
|
|
raise RuntimeError(
|
|
f"paper_server connection failed: {type(exc).__name__}"
|
|
) from None
|
|
_raise_response_error(response)
|
|
try:
|
|
return response.json()
|
|
except ValueError:
|
|
raise RuntimeError("paper_server returned invalid JSON") from None
|
|
|
|
|
|
def _results(data: Any) -> list[dict[str, Any]]:
|
|
values = (
|
|
data.get("results") if isinstance(data, dict) and "results" in data else data
|
|
)
|
|
if not isinstance(values, list):
|
|
raise RuntimeError("paper_server returned an unexpected result shape")
|
|
return [item for item in values if isinstance(item, dict)]
|
|
|
|
|
|
def _is_doi(value: str) -> bool:
|
|
return "/" in value and value.lstrip().startswith("10.")
|
|
|
|
|
|
def _resolve_to_id(id_or_doi: str, api_url: str, api_key: str) -> str:
|
|
value = str(id_or_doi or "").strip()
|
|
if not value:
|
|
raise ValueError("id_or_doi 不能为空")
|
|
if not _is_doi(value):
|
|
return value
|
|
matches = _results(_get_json(api_url + "/", api_key=api_key, params={"doi": value}))
|
|
if not matches:
|
|
raise ValueError(f"doi 未命中:{value}")
|
|
if len(matches) > 1:
|
|
raise ValueError(f"doi 命中多条({len(matches)}):{value}")
|
|
paper_id = str(matches[0].get("id") or "").strip()
|
|
if not paper_id:
|
|
raise RuntimeError("paper_server DOI 查询结果缺少 id")
|
|
return paper_id
|
|
|
|
|
|
def _get_paper(id_or_doi: str) -> dict[str, Any]:
|
|
_base_url, api_url, api_key = _config()
|
|
paper_id = _resolve_to_id(id_or_doi, api_url, api_key)
|
|
data = _get_json(f"{api_url}/{paper_id}/", api_key=api_key)
|
|
if not isinstance(data, dict):
|
|
raise RuntimeError("paper_server returned an unexpected paper shape")
|
|
return data
|
|
|
|
|
|
def _safe_stem(value: str) -> str:
|
|
stem = re.sub(r"[^A-Za-z0-9._-]+", "_", value.strip()).strip("._")
|
|
return stem[:180] or "paper"
|
|
|
|
|
|
def _media_url(raw_url: str, base_url: str) -> str:
|
|
url = urljoin(base_url + "/", raw_url)
|
|
parsed = urlparse(url)
|
|
expected = urlparse(base_url)
|
|
if parsed.scheme not in ("http", "https") or (
|
|
parsed.scheme.lower(),
|
|
parsed.netloc.lower(),
|
|
) != (expected.scheme.lower(), expected.netloc.lower()):
|
|
raise RuntimeError("paper_server returned an out-of-origin media URL")
|
|
return url
|
|
|
|
|
|
class PaperSearchTool(Tool):
|
|
name = "paper_search"
|
|
description = (
|
|
"Search the platform paper_server metadata collection. "
|
|
"Use publication_type for server-side source-type filtering, e.g. 'book', "
|
|
"'book-chapter' or 'article'. publication_type uses paper_server/OpenAlex raw "
|
|
"values; normalized literature type book_chapter is represented here as raw "
|
|
"'book-chapter'. Prefer English keywords."
|
|
)
|
|
parameters = {
|
|
"type": "object",
|
|
"properties": {
|
|
"keyword": {
|
|
"type": "string",
|
|
"description": "Fuzzy title/author/institution query.",
|
|
},
|
|
"publication_type": {
|
|
"type": "string",
|
|
"description": "Exact raw source type, e.g. book or book-chapter.",
|
|
},
|
|
"year": {"type": "integer", "description": "Exact publication year."},
|
|
"year_gte": {"type": "integer", "description": "Minimum publication year."},
|
|
"year_lte": {"type": "integer", "description": "Maximum publication year."},
|
|
"doi": {"type": "string", "description": "Exact DOI."},
|
|
"first_author": {
|
|
"type": "string",
|
|
"description": "Exact first-author name.",
|
|
},
|
|
"publication_name": {
|
|
"type": "string",
|
|
"description": "Exact journal/container name.",
|
|
},
|
|
"has_pdf": {
|
|
"type": "boolean",
|
|
"description": "Filter by PDF availability.",
|
|
},
|
|
"is_oa": {
|
|
"type": "boolean",
|
|
"description": "Filter by open-access status.",
|
|
},
|
|
"limit": {
|
|
"type": "integer",
|
|
"default": 10,
|
|
"description": "Maximum records, 1-50.",
|
|
},
|
|
},
|
|
}
|
|
|
|
def execute(
|
|
self,
|
|
keyword: str = "",
|
|
publication_type: str = "",
|
|
year: Optional[int] = None,
|
|
year_gte: Optional[int] = None,
|
|
year_lte: Optional[int] = None,
|
|
doi: str = "",
|
|
first_author: str = "",
|
|
publication_name: str = "",
|
|
has_pdf: Optional[bool] = None,
|
|
is_oa: Optional[bool] = None,
|
|
limit: int = 10,
|
|
) -> str:
|
|
raw_type = str(publication_type or "").strip()
|
|
if raw_type and not _TYPE_RE.fullmatch(raw_type):
|
|
return (
|
|
"[Error] publication_type 只能包含字母、数字、下划线或连字符,"
|
|
"最长 64 字符"
|
|
)
|
|
limit = min(max(int(limit), 1), 50)
|
|
params: dict[str, Any] = {"page_size": limit}
|
|
optional = {
|
|
"search": str(keyword or "").strip(),
|
|
"type": raw_type,
|
|
"doi": str(doi or "").strip(),
|
|
"first_author": str(first_author or "").strip(),
|
|
"publication_name": str(publication_name or "").strip(),
|
|
}
|
|
params.update({key: value for key, value in optional.items() if value})
|
|
if year is not None:
|
|
params["publication_year"] = int(year)
|
|
if year_gte is not None:
|
|
params["publication_year_gte"] = int(year_gte)
|
|
if year_lte is not None:
|
|
params["publication_year_lte"] = int(year_lte)
|
|
if has_pdf is not None:
|
|
params["has_fulltext_pdf"] = "true" if has_pdf else "false"
|
|
if is_oa is not None:
|
|
params["is_oa"] = "true" if is_oa else "false"
|
|
try:
|
|
_base_url, api_url, api_key = _config()
|
|
papers = _results(_get_json(api_url + "/", api_key=api_key, params=params))
|
|
except Exception as exc:
|
|
return f"[Error] paper_search failed:{type(exc).__name__}:{exc}"
|
|
trimmed = [
|
|
{key: paper.get(key) for key in _LIST_FIELDS} for paper in papers[:limit]
|
|
]
|
|
return json.dumps(trimmed, ensure_ascii=False, indent=2)
|
|
|
|
|
|
class PaperGetTool(Tool):
|
|
name = "paper_get"
|
|
description = (
|
|
"Get one complete paper_server metadata record by internal id or exact DOI."
|
|
)
|
|
parameters = {
|
|
"type": "object",
|
|
"properties": {
|
|
"id_or_doi": {"type": "string", "description": "paper_server id or DOI."},
|
|
},
|
|
"required": ["id_or_doi"],
|
|
}
|
|
|
|
def execute(self, id_or_doi: str) -> str:
|
|
try:
|
|
paper = _get_paper(id_or_doi)
|
|
except Exception as exc:
|
|
return f"[Error] paper_get failed:{type(exc).__name__}:{exc}"
|
|
return json.dumps(paper, ensure_ascii=False, indent=2)
|
|
|
|
|
|
class PaperFetchTool(Tool):
|
|
name = "paper_fetch"
|
|
description = (
|
|
"Download an available PDF or XML from paper_server into the current task's "
|
|
"papers/ directory. Use the format actually reported by paper_search/paper_get."
|
|
)
|
|
parameters = {
|
|
"type": "object",
|
|
"properties": {
|
|
"id_or_doi": {"type": "string", "description": "paper_server id or DOI."},
|
|
"format": {
|
|
"type": "string",
|
|
"enum": ["pdf", "xml"],
|
|
"description": "File format to download.",
|
|
},
|
|
},
|
|
"required": ["id_or_doi", "format"],
|
|
}
|
|
|
|
def __init__(
|
|
self,
|
|
*,
|
|
working_dir: Path,
|
|
base_dir: Optional[Path] = None,
|
|
user_root: Optional[Path] = None,
|
|
) -> None:
|
|
super().__init__(base_dir=base_dir, user_root=user_root)
|
|
self.working_dir = Path(working_dir)
|
|
|
|
def execute(self, id_or_doi: str, format: str) -> str: # noqa: A002 - JSON tool contract
|
|
file_format = str(format or "").lower().strip()
|
|
if file_format not in {"pdf", "xml"}:
|
|
return "[Error] format 必须是 pdf 或 xml"
|
|
try:
|
|
paper = _get_paper(id_or_doi)
|
|
available_key = f"has_fulltext_{file_format}"
|
|
if not paper.get(available_key):
|
|
reason = (
|
|
paper.get("fail_reason") or f"no {file_format.upper()} on server"
|
|
)
|
|
raise RuntimeError(f"paper has no {file_format.upper()}:{reason}")
|
|
raw_url = str(paper.get(f"{file_format}_url") or "").strip()
|
|
if not raw_url:
|
|
raise RuntimeError(f"paper {file_format}_url unavailable")
|
|
base_url, _api_url, api_key = _config()
|
|
media_url = _media_url(raw_url, base_url)
|
|
identity = str(paper.get("doi") or paper.get("id") or id_or_doi)
|
|
destination = (
|
|
self.working_dir / "papers" / f"{_safe_stem(identity)}.{file_format}"
|
|
)
|
|
if destination.exists() and destination.stat().st_size > 0:
|
|
return f"saved:{self._display(destination)} (existing)"
|
|
destination.parent.mkdir(parents=True, exist_ok=True)
|
|
partial = destination.with_suffix(destination.suffix + ".part")
|
|
total = 0
|
|
try:
|
|
with httpx.stream(
|
|
"GET",
|
|
media_url,
|
|
params=_params(api_key),
|
|
timeout=_DOWNLOAD_TIMEOUT,
|
|
) as response:
|
|
_raise_response_error(response)
|
|
with partial.open("wb") as handle:
|
|
for chunk in response.iter_bytes(chunk_size=64 * 1024):
|
|
total += len(chunk)
|
|
if total > _MAX_DOWNLOAD_BYTES:
|
|
message = (
|
|
"paper_server file exceeds "
|
|
f"{_MAX_DOWNLOAD_BYTES} bytes"
|
|
)
|
|
raise RuntimeError(message)
|
|
handle.write(chunk)
|
|
partial.replace(destination)
|
|
except httpx.RequestError as exc:
|
|
raise RuntimeError(
|
|
f"paper_server download connection failed:{type(exc).__name__}"
|
|
) from None
|
|
finally:
|
|
if partial.exists():
|
|
try:
|
|
partial.unlink()
|
|
except OSError:
|
|
pass
|
|
except Exception as exc:
|
|
return f"[Error] paper_fetch failed:{type(exc).__name__}:{exc}"
|
|
return f"saved:{self._display(destination)}"
|