From 2e1efd5e10524a2e32948747c717eda303daf66b Mon Sep 17 00:00:00 2001 From: Niels Rogge Date: Wed, 16 Sep 2026 14:00:19 +0000 Subject: [PATCH 1/2] Improve MCP agent experience and catalog coverage --- mcp_server/README.md | 19 +- mcp_server/SKILL.md | 18 +- mcp_server/SPEC.md | 15 +- mcp_server/pyproject.toml | 2 +- mcp_server/src/pwc_mcp/__init__.py | 2 +- mcp_server/src/pwc_mcp/app.py | 72 ++++++- mcp_server/src/pwc_mcp/catalog.py | 280 +++++++++++++++++++++--- mcp_server/src/pwc_mcp/models.py | 139 ++++++++++-- mcp_server/src/pwc_mcp/server.py | 331 +++++++++++++++++++++++++---- mcp_server/tests/test_app.py | 47 ++-- mcp_server/tests/test_catalog.py | 54 ++++- mcp_server/tests/test_server.py | 112 +++++++++- mcp_server/tests/test_skill.py | 4 + mcp_server/uv.lock | 2 +- 14 files changed, 967 insertions(+), 130 deletions(-) diff --git a/mcp_server/README.md b/mcp_server/README.md index 41176c6..95773b3 100644 --- a/mcp_server/README.md +++ b/mcp_server/README.md @@ -3,9 +3,10 @@ Anonymous, read-only Model Context Protocol access to the public [Papers With Code](https://paperswithcode.co) catalog. -The server uses MCP `2026-07-28` over Streamable HTTP and serves legacy -2025-era clients on the same `/mcp` endpoint. It returns versioned structured -output with compact text fallbacks. +The server uses stock-client MCP `2025-11-25` over Streamable HTTP and also +serves the experimental `2026-07-28` discovery protocol on the same `/mcp` +endpoint. It returns versioned structured output with compact Markdown +fallbacks. ## Run locally @@ -27,13 +28,19 @@ curl http://127.0.0.1:7860/health - `get_paper_info` - `read_paper` - `get_related_papers` +- `get_trending_papers` +- `get_paper_evaluations` - `get_paper_lineage` - `get_task` +- `list_tasks` - `get_method` +- `list_methods` - `list_benchmarks` - `get_benchmark` -All tools are annotated read-only and idempotent. Search is deterministic; +All tools are annotated read-only and idempotent. Expected failures use typed +messages (`not_found`, `ambiguous`, `no_markdown`, and `upstream_timeout`), with +candidate IDs and slugs for ambiguous references. Search is deterministic; the caller controls keyword or semantic mode. `read_paper` fetches at most one 64 KiB catalog chunk per call and returns a signed, one-hour continuation cursor when more Markdown remains. Continuations stay pinned to the resolved paper and @@ -69,6 +76,10 @@ Tool inputs cap list results at 25, catalog calls time out after 25 seconds, and request, upstream, and serialized MCP response bodies are bounded to 2 MiB. Markdown chunks use a bounded 256-entry/16 MiB in-memory cache. +`GET /.well-known/mcp` exposes connection metadata and `GET /docs` publishes +the live input/output schema for every tool. A bare `GET /mcp` returns `405`; +MCP requests use `POST /mcp`. Rate limits return `429` with `Retry-After`. + ## Test ```bash diff --git a/mcp_server/SKILL.md b/mcp_server/SKILL.md index a3059e0..0a584c0 100644 --- a/mcp_server/SKILL.md +++ b/mcp_server/SKILL.md @@ -4,7 +4,7 @@ description: "Papers With Code MCP tools for searching and reading AI/ML papers, compatibility: "Requires an MCP client connected to https://paperswithcode.co/mcp with the Papers With Code tools available." --- -Generated for `pwc-mcp v0.1.0` and MCP protocol `2026-07-28`. +Generated for `pwc-mcp v0.2.0` and MCP protocol `2025-11-25`. The tools query the public [Papers With Code](https://paperswithcode.co) catalog anonymously and are read-only. If live tool discovery and this skill disagree, @@ -36,12 +36,16 @@ matching IDs. ## Tools - `search_papers({"query": QUERY, "limit": LIMIT, "page": PAGE, "mode": "keyword"|"semantic", "published_after": START_DATE, "published_before": END_DATE, "has_official_implementation": BOOLEAN})` — search papers. Omit optional arguments when they are not needed. -- `get_paper_info({"paper": PAPER})` — show paper metadata, abstract, tasks, methods, repositories, and project pages. +- `get_paper_info({"paper": PAPER, "include_resources": BOOLEAN, "repo_limit": LIMIT})` — show paper metadata, repository count, and official code by default; optionally add capped repositories, project pages, and Hugging Face models/datasets. - `read_paper({"paper": PAPER})` — read one stored paper Markdown chunk. If `truncated` is true, call `read_paper` again with the same `paper` and the returned `next_cursor`; repeat until `truncated` is false. Treat the cursor as opaque and use it within one hour. - `list_papers({"page": PAGE, "limit": LIMIT, "search": SEARCH, "published_after": START_DATE, "published_before": END_DATE, "task": TASK, "method": METHOD, "conference": CONFERENCE, "framework": FRAMEWORK, "organization": ORGANIZATION, "authors": [AUTHOR], "order_by": "date_published"|"citation_count"|"title", "order_direction": "asc"|"desc"})` — list and filter papers. Omit optional arguments when they are not needed. - `get_related_papers({"paper": PAPER, "limit": LIMIT})` — list related papers. +- `get_trending_papers({"limit": LIMIT, "max_age_days": DAYS, "min_velocity": VELOCITY})` — list trending papers. +- `get_paper_evaluations({"paper": PAPER, "limit": LIMIT})` — list benchmark evaluations reported by a paper. - `get_paper_lineage({"paper": PAPER})` — list explicit predecessors and successors. -- `get_task({"task": TASK})` — inspect one exact task by ID, slug, or name, including its area, parents, children, and benchmarks. +- `list_tasks({"search": SEARCH, "page": PAGE, "limit": LIMIT})` — discover task slugs and IDs. +- `get_task({"task": TASK, "benchmark_limit": LIMIT})` — inspect one exact task by ID or slug with a capped benchmark list. +- `list_methods({"search": SEARCH, "page": PAGE, "limit": LIMIT})` — discover method slugs and IDs. - `get_method({"method": METHOD})` — inspect one exact method by ID, slug, full name, or name. - `list_benchmarks({"page": PAGE, "limit": LIMIT, "search": SEARCH, "task": TASK, "include_descendants": BOOLEAN, "minimum_evaluations": MINIMUM_EVALUATIONS, "is_open": BOOLEAN})` — list and filter benchmarks. Omit optional arguments when they are not needed. - `get_benchmark({"benchmark": BENCHMARK, "limit": LIMIT, "is_open": BOOLEAN})` — inspect one exact benchmark and its leading evaluation rows. @@ -51,7 +55,7 @@ when the user asks for more results than one response contains; do not infer that a missing item does not exist until the relevant pages have been checked. The MCP server does not expose standalone CLI commands for paper editing, -authentication, skill installation, version display, taxonomy enumeration, or +authentication, skill installation, version display, or advanced benchmark metric/parameter/Pareto filtering. Do not invent equivalent tools. Use the separate `pwc` CLI only when it is available and the user needs one of those capabilities. @@ -60,8 +64,10 @@ one of those capabilities. 1. Use `list_benchmarks({"task": TASK})` to discover active benchmarks, then `get_benchmark({"benchmark": NAME})` to inspect a leaderboard. -2. Use `get_paper_info({"paper": PAPER})` to inspect promising results. Its - response includes repositories and project pages. +2. Use `get_paper_info({"paper": PAPER})` to inspect promising results. The + default response includes official repositories and a repository count; pass + `include_resources: true` for other repositories, project pages, and Hugging + Face artifacts. 3. Use exact `list_papers` `authors`, `task`, `method`, `conference`, `framework`, and `organization` arguments for known identities or catalog associations. Combine them to require every association; do not substitute diff --git a/mcp_server/SPEC.md b/mcp_server/SPEC.md index 1af6141..21e7e6f 100644 --- a/mcp_server/SPEC.md +++ b/mcp_server/SPEC.md @@ -13,31 +13,38 @@ server. the MCP contract. - Serve anonymous, read-only requests. Search is deterministic and contains no embedded language model. -- Use stateless Streamable HTTP at `/mcp`, supporting MCP `2026-07-28` and - legacy 2025 clients on the same endpoint. Expose `/health` for operations. +- Use stateless Streamable HTTP at `/mcp`, advertising stock-client MCP + `2025-11-25` while also supporting experimental `2026-07-28` discovery. + Expose `/health`, `/.well-known/mcp`, and a generated `/docs` schema. ## Public contract -Expose exactly these tools: +Expose these tools: - `search_papers` - `list_papers` - `get_paper_info` - `read_paper` - `get_related_papers` +- `get_trending_papers` +- `get_paper_evaluations` - `get_paper_lineage` - `get_task` +- `list_tasks` - `get_method` +- `list_methods` - `list_benchmarks` - `get_benchmark` -Expose these resource templates and no prompts: +Expose these resource templates: - `pwc://papers/{paper}` - `pwc://papers/{paper}/markdown` - `pwc://tasks/{task}` - `pwc://benchmarks/{benchmark}` +Expose the `find_papers`, `compare_leaderboard`, and `survey_task` prompts. + Responses use stable, MCP-specific versioned structured outputs with a text fallback. `read_paper` performs one upstream read of at most 64 KiB per call and returns a signed opaque continuation cursor when more Markdown remains. The diff --git a/mcp_server/pyproject.toml b/mcp_server/pyproject.toml index 896cd32..76875d9 100644 --- a/mcp_server/pyproject.toml +++ b/mcp_server/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "pwc-mcp" -version = "0.1.0" +version = "0.2.0" description = "Read-only Papers With Code MCP server" readme = "README.md" requires-python = ">=3.10" diff --git a/mcp_server/src/pwc_mcp/__init__.py b/mcp_server/src/pwc_mcp/__init__.py index ffd884d..cab4aac 100644 --- a/mcp_server/src/pwc_mcp/__init__.py +++ b/mcp_server/src/pwc_mcp/__init__.py @@ -1,3 +1,3 @@ """Read-only Papers With Code MCP server.""" -__version__ = "0.1.0" +__version__ = "0.2.0" diff --git a/mcp_server/src/pwc_mcp/app.py b/mcp_server/src/pwc_mcp/app.py index 1967f42..357abb4 100644 --- a/mcp_server/src/pwc_mcp/app.py +++ b/mcp_server/src/pwc_mcp/app.py @@ -29,7 +29,8 @@ # MCP SDK diagnostics can include peer-supplied tool names and resource URIs. # OperationalTelemetryMiddleware is the server's sole request log surface. logging.getLogger("mcp").setLevel(logging.CRITICAL + 1) -PROTOCOL_VERSION = "2026-07-28" +PROTOCOL_VERSION = "2025-11-25" +EXPERIMENTAL_PROTOCOL_VERSION = "2026-07-28" MAX_REQUEST_BODY_SIZE = 2 * 1024 * 1024 MAX_RESPONSE_BODY_SIZE = 2 * 1024 * 1024 KNOWN_TOOLS = { @@ -38,21 +39,80 @@ "get_paper_info", "read_paper", "get_related_papers", + "get_trending_papers", + "get_paper_evaluations", "get_paper_lineage", "get_task", + "list_tasks", "get_method", + "list_methods", "list_benchmarks", "get_benchmark", } KNOWN_PROTOCOLS = { PROTOCOL_VERSION, - "2025-11-25", + EXPERIMENTAL_PROTOCOL_VERSION, "2025-06-18", "2025-03-26", "2024-11-05", } +async def well_known_mcp(request: Request) -> JSONResponse: + return JSONResponse( + { + "name": "Papers With Code", + "description": "Anonymous read-only AI research catalog", + "transport": {"type": "streamable-http", "url": "/mcp"}, + "protocol_version": PROTOCOL_VERSION, + "supported_protocol_versions": sorted(KNOWN_PROTOCOLS, reverse=True), + "documentation_url": "/docs", + } + ) + + +def _docs_schema(server) -> dict: + tools = [] + for tool in server._tool_manager.list_tools(): + tools.append( + { + "name": tool.name, + "description": tool.description, + "inputSchema": tool.parameters, + "outputSchema": tool.output_schema, + } + ) + return { + "name": "Papers With Code MCP", + "version": __version__, + "protocol_version": PROTOCOL_VERSION, + "endpoint": "/mcp", + "tools": tools, + } + + +class MCPMethodMiddleware: + """Avoid opening an SSE response for unsupported bare GET requests.""" + + def __init__(self, app: ASGIApp): + self.app = app + + async def __call__(self, scope: Scope, receive: Receive, send: Send) -> None: + if ( + scope["type"] == "http" + and scope.get("path") == "/mcp" + and scope.get("method") == "GET" + ): + response = JSONResponse( + {"error": "method_not_allowed", "allowed": ["POST"]}, + status_code=405, + headers={"Allow": "POST"}, + ) + await response(scope, receive, send) + return + await self.app(scope, receive, send) + + def _csv_env(name: str, default: list[str]) -> list[str]: value = os.environ.get(name) if value is None: @@ -426,6 +486,13 @@ def create_app( ), ) app.routes.insert(0, Route("/health", health, methods=["GET"])) + docs_schema = _docs_schema(server) + + async def docs(_request: Request) -> JSONResponse: + return JSONResponse(docs_schema) + + app.routes.insert(1, Route("/docs", docs, methods=["GET"])) + app.routes.insert(2, Route("/.well-known/mcp", well_known_mcp, methods=["GET"])) app.state.pwc_catalog_readiness = CatalogReadiness( catalog_client, initially_ready=catalog is not None, @@ -438,6 +505,7 @@ def create_app( global_concurrency_limit=global_concurrency_limit, trust_proxy_headers=trust_proxy_headers, ) + wrapped = MCPMethodMiddleware(wrapped) wrapped = ResponseSizeLimitMiddleware(wrapped) wrapped = OperationalTelemetryMiddleware(wrapped) return CORSMiddleware( diff --git a/mcp_server/src/pwc_mcp/catalog.py b/mcp_server/src/pwc_mcp/catalog.py index c5cacc6..f561aff 100644 --- a/mcp_server/src/pwc_mcp/catalog.py +++ b/mcp_server/src/pwc_mcp/catalog.py @@ -8,7 +8,13 @@ from typing import Any, Protocol from urllib.parse import quote, urlparse -from pwc_cli.transport import Client, HTTPStatusError, Response, ResponseError +from pwc_cli.transport import ( + Client, + HTTPStatusError, + Response, + ResponseError, + TransportError, +) PAPER_ID = re.compile(r"(?:\d{4}\.\d{4,5}|[a-z][a-z0-9.-]*/\d{7}|\d+)", re.IGNORECASE) ARXIV_VERSION = re.compile(r"v\d+$", re.IGNORECASE) @@ -36,6 +42,32 @@ class PaperVersionMismatchError(ResponseError): pass +class CatalogError(ResponseError): + code = "catalog_error" + + +class NotFoundError(CatalogError): + code = "not_found" + + +class AmbiguousError(CatalogError): + code = "ambiguous" + + def __init__(self, label: str, reference: str, candidates: list[dict[str, str]]): + super().__init__(f"{label} is ambiguous: {reference}") + self.label = label + self.reference = reference + self.candidates = candidates + + +class NoMarkdownError(CatalogError): + code = "no_markdown" + + +class UpstreamTimeoutError(CatalogError): + code = "upstream_timeout" + + class Transport(Protocol): def get( self, path: str, params: dict[str, object | None] | None = None @@ -79,9 +111,9 @@ def put(self, key: object, value: object, ttl_seconds: int) -> None: class _MarkdownChunkCache: def __init__(self): - self._values: OrderedDict[ - object, tuple[float, int, PaperMarkdownChunk] - ] = OrderedDict() + self._values: OrderedDict[object, tuple[float, int, PaperMarkdownChunk]] = ( + OrderedDict() + ) self._bytes = 0 self._lock = threading.Lock() @@ -135,12 +167,27 @@ def _json( params: dict[str, object | None] | None = None, *, ttl: int, + accept_list: bool = False, ) -> dict[str, Any]: key = ("json", path, _freeze(params or {})) cached = self.cache.get(key) if isinstance(cached, dict): return cached - payload = self.transport.get(path, params).json() + try: + payload = self.transport.get(path, params).json() + except HTTPStatusError as error: + if error.status == 404: + raise NotFoundError("Catalog entity not found") from error + raise + except (TimeoutError, TransportError) as error: + if ( + "timed out" in str(error).casefold() + or "timeout" in str(error).casefold() + ): + raise UpstreamTimeoutError("Catalog request timed out") from error + raise + if accept_list and isinstance(payload, list): + payload = {"results": payload, "count": len(payload)} if not isinstance(payload, dict): raise ResponseError("API returned an unexpected response shape") self.cache.put(key, payload, ttl) @@ -172,7 +219,9 @@ def _text(self, path: str, *, ttl: int) -> str: @staticmethod def _rows(payload: dict[str, Any]) -> list[dict[str, Any]]: - values = payload.get("results") or payload.get("items") + values = payload.get("results") + if values is None: + values = payload.get("items") if not isinstance(values, list): raise ResponseError("API response did not contain a result list") return [item for item in values if isinstance(item, dict)] @@ -249,16 +298,24 @@ def _resolve_paper(self, reference: str) -> str: if not isinstance(next_page, int) or next_page <= page: break page = next_page - else: - raise ResponseError("Too many results to resolve paper title safely") if len(exact) == 1: return next(iter(exact)) if exact: - choices = "; ".join( - f"{item.get('title')} ({paper})" for paper, item in exact.items() + raise AmbiguousError( + "Paper title", + candidate, + [ + { + "id": paper, + "name": str(item.get("title") or candidate), + "reference": str( + item.get("arxiv_id") or item.get("id") or paper + ), + } + for paper, item in exact.items() + ], ) - raise ResponseError(f"Paper title is ambiguous: {candidate}; {choices}") - raise ResponseError(f"Paper title not found: {candidate}") + raise NotFoundError(f"Paper title not found: {candidate}") def resolve_paper(self, paper: str) -> str: return self._resolve_paper(paper) @@ -272,10 +329,29 @@ def _exact( ) -> dict[str, Any]: target = reference.strip().casefold() for field in fields: - for item in items: - if str(item.get(field) or "").strip().casefold() == target: - return item - raise ResponseError(f"{label} not found: {reference}") + matches = [ + item + for item in items + if str(item.get(field) or "").strip().casefold() == target + ] + if len(matches) == 1: + return matches[0] + if len(matches) > 1: + raise AmbiguousError( + label, + reference, + [ + { + "id": str(item.get("id") or ""), + "name": str( + item.get("full_name") or item.get("name") or "Unknown" + ), + "slug": str(item.get("slug") or ""), + } + for item in matches[:10] + ], + ) + raise NotFoundError(f"{label} not found: {reference}") def search_papers( self, @@ -352,6 +428,17 @@ def read_paper_chunk( raise PaperVersionMismatchError( "Paper Markdown changed; restart reading from the beginning" ) from error + if error.status == 404: + raise NoMarkdownError( + "No stored Markdown is available for this paper" + ) from error + raise + except (TimeoutError, TransportError) as error: + if ( + "timed out" in str(error).casefold() + or "timeout" in str(error).casefold() + ): + raise UpstreamTimeoutError("Catalog request timed out") from error raise returned_version = response.headers.get("x-pwc-content-version", "") @@ -361,7 +448,9 @@ def read_paper_chunk( markdown = response.body.decode("utf-8") next_offset = int(next_text) if next_text is not None else None except (UnicodeDecodeError, ValueError) as error: - raise ResponseError("Papers API returned an invalid Markdown chunk") from error + raise ResponseError( + "Papers API returned an invalid Markdown chunk" + ) from error if ( CONTENT_VERSION.fullmatch(returned_version) is None or truncated not in {"0", "1"} @@ -369,10 +458,7 @@ def read_paper_chunk( or (truncated == "1" and next_offset is None) or (truncated == "0" and next_offset is not None) or (next_offset is not None and next_offset <= offset) - or ( - next_offset is not None - and next_offset - offset != len(response.body) - ) + or (next_offset is not None and next_offset - offset != len(response.body)) ): detail = ( "Papers API Markdown offset did not advance" @@ -387,9 +473,7 @@ def read_paper_chunk( content_version=returned_version, next_offset=next_offset, ) - self.markdown_cache.put( - (reference, returned_version, offset, limit), result - ) + self.markdown_cache.put((reference, returned_version, offset, limit), result) return result def list_papers( @@ -450,6 +534,37 @@ def get_related_papers(self, paper: str, *, limit: int) -> dict[str, Any]: f"papers/{quote(reference, safe='.')}/related", {"limit": limit}, ttl=300, + accept_list=True, + ) + + def get_trending_papers( + self, *, limit: int, max_age_days: int, min_velocity: float | None + ) -> dict[str, Any]: + return self._json( + "papers/trending", + { + "limit": limit, + "max_age_days": max_age_days, + "min_velocity": min_velocity, + }, + ttl=60, + accept_list=True, + ) + + def get_paper_evaluations(self, paper: str, *, limit: int) -> dict[str, Any]: + detail = self.get_paper_info(paper, include_resources=False) + paper_id = detail.get("id") + if not paper_id: + raise NotFoundError(f"Paper not found: {paper}") + return self._json( + "evaluations/", + { + "page": 1, + "page_size": limit, + "paper_id": paper_id, + "ordering": "-benchmark_popularity", + }, + ttl=300, ) def get_paper_lineage(self, paper: str) -> dict[str, Any]: @@ -469,9 +584,81 @@ def get_task(self, task: str) -> dict[str, Any]: ttl=600, ) ) - task_id = str(self._exact(task, candidates, "Task").get("id")) + if not candidates: + for token in re.findall(r"[a-z0-9]+", task.casefold()): + if len(token) < 3: + continue + candidates.extend( + self._rows( + self._json( + "tasks/", + {"q": token, "page": 1, "page_size": 100}, + ttl=600, + ) + ) + ) + try: + matched = self._exact(task, candidates, "Task") + except NotFoundError as error: + if not candidates: + raise + unique = {str(item.get("id")): item for item in candidates} + tokens = re.findall(r"[a-z0-9]+", task.casefold()) + noun = tokens[-1] if tokens else "" + ranked = sorted( + unique.values(), + key=lambda item: ( + noun + not in str( + item.get("name") or item.get("slug") or "" + ).casefold(), + -sum( + token + in str( + item.get("name") or item.get("slug") or "" + ).casefold() + for token in tokens + ), + -int(item.get("paper_count") or 0), + ), + ) + raise AmbiguousError( + "Task name", + task, + [ + { + "id": str(item.get("id") or ""), + "name": str(item.get("name") or "Unknown"), + "slug": str(item.get("slug") or ""), + } + for item in ranked[:10] + ], + ) from error + task_id = str(matched.get("id")) return self._json(f"tasks/{quote(task_id, safe='')}/page", ttl=600) + def list_tasks( + self, *, search: str | None, page: int, limit: int + ) -> dict[str, Any]: + payload = self._json( + "tasks/", + {"q": search, "page": page, "page_size": limit}, + ttl=600, + ) + if search and not self._rows(payload): + tokens = [ + token + for token in re.findall(r"[a-z0-9]+", search.casefold()) + if len(token) >= 3 + ] + if tokens: + payload = self._json( + "tasks/", + {"q": tokens[-1], "page": page, "page_size": limit}, + ttl=600, + ) + return payload + def get_method(self, method: str) -> dict[str, Any]: if method.strip().isdigit(): method_id = method.strip() @@ -492,6 +679,15 @@ def get_method(self, method: str) -> dict[str, Any]: method_id = str(matched.get("id") or matched.get("slug")) return self._json(f"methods/{quote(method_id, safe='')}", ttl=600) + def list_methods( + self, *, search: str | None, page: int, limit: int + ) -> dict[str, Any]: + return self._json( + "methods/", + {"q": search, "page": page, "page_size": limit}, + ttl=600, + ) + def list_benchmarks( self, *, @@ -528,12 +724,34 @@ def get_benchmark( ttl=300, ) ) - matched = self._exact( - benchmark, - candidates, - "Benchmark", - fields=("name", "full_name", "slug", "id"), - ) + try: + matched = self._exact( + benchmark, + candidates, + "Benchmark", + fields=("name", "full_name", "slug", "id"), + ) + except NotFoundError as error: + if not candidates: + raise + ranked = sorted( + candidates, + key=lambda item: -int(item.get("paper_count") or 0), + ) + raise AmbiguousError( + "Benchmark name", + benchmark, + [ + { + "id": str(item.get("id") or ""), + "name": str( + item.get("full_name") or item.get("name") or "Unknown" + ), + "slug": str(item.get("slug") or ""), + } + for item in ranked[:10] + ], + ) from error benchmark_id = str(matched.get("id")) evaluations = self._json( f"datasets/{quote(benchmark_id, safe='')}/evaluations/", diff --git a/mcp_server/src/pwc_mcp/models.py b/mcp_server/src/pwc_mcp/models.py index 71bffbd..f5cc6a8 100644 --- a/mcp_server/src/pwc_mcp/models.py +++ b/mcp_server/src/pwc_mcp/models.py @@ -5,6 +5,15 @@ from pydantic import BaseModel, ConfigDict +def _absolute_url(value: object) -> str | None: + if not value: + return None + url = str(value) + if url.startswith("/"): + return f"https://paperswithcode.co{url}" + return url + + class OutputModel(BaseModel): model_config = ConfigDict(extra="forbid") @@ -33,6 +42,10 @@ class CatalogReference(OutputModel): slug: str | None = None +class TaxonomyReference(CatalogReference): + url: str + + class RepositoryReference(OutputModel): url: str is_official: bool @@ -48,10 +61,13 @@ class PaperDetail(OutputModel): citation_count: int | None = None url: str | None = None pdf_url: str | None = None + code_repository_count: int tasks: list[CatalogReference] methods: list[CatalogReference] repositories: list[RepositoryReference] project_pages: list[str] + hf_models: list[str] + hf_datasets: list[str] class PaperInfoResult(OutputModel): @@ -89,6 +105,7 @@ class BenchmarkSummary(OutputModel): id: str name: str slug: str | None = None + url: str | None = None full_name: str | None = None description: str | None = None hf_url: str | None = None @@ -99,11 +116,13 @@ class TaskDetail(OutputModel): id: str name: str slug: str + url: str description: str | None = None paper_count: int area: AreaReference | None = None parents: list[CatalogReference] children: list[CatalogReference] + benchmark_count: int benchmarks: list[BenchmarkSummary] @@ -116,6 +135,7 @@ class MethodDetail(OutputModel): id: str name: str slug: str + url: str full_name: str | None = None description: str | None = None introduced_year: int | None = None @@ -130,6 +150,12 @@ class MethodResult(OutputModel): method: MethodDetail +class TaxonomyPage(OutputModel): + schema_version: Literal["v1"] = "v1" + items: list[TaxonomyReference] + next_page: int | None = None + + class BenchmarkPage(OutputModel): schema_version: Literal["v1"] = "v1" items: list[BenchmarkSummary] @@ -155,6 +181,13 @@ class BenchmarkResult(OutputModel): evaluations: list[Evaluation] +class EvaluationPage(OutputModel): + schema_version: Literal["v1"] = "v1" + paper: str + evaluation_count: int + evaluations: list[Evaluation] + + def paper_summary(item: dict[str, Any]) -> PaperSummary: return PaperSummary( id=str(item.get("id") or ""), @@ -167,7 +200,7 @@ def paper_summary(item: dict[str, Any]) -> PaperSummary: if item.get("citation_count") is not None else None ), - url=str(item.get("url_abs") or item.get("source_url") or "") or None, + url=_absolute_url(item.get("url_abs") or item.get("source_url")), has_official_implementation=item.get("has_official_implementation") is True, code_repository_count=int(item.get("code_repository_count") or 0), ) @@ -181,16 +214,29 @@ def catalog_reference(item: dict[str, Any]) -> CatalogReference: ) -def paper_detail(item: dict[str, Any]) -> PaperDetail: +def taxonomy_reference(item: dict[str, Any], *, kind: str) -> TaxonomyReference: + reference = catalog_reference(item) + slug = reference.slug or reference.id + return TaxonomyReference( + **reference.model_dump(), + url=f"https://paperswithcode.co/{kind}/{slug}", + ) + + +def paper_detail( + item: dict[str, Any], *, include_resources: bool, repo_limit: int +) -> PaperDetail: repositories = [] for repository in item.get("repositories") or []: if isinstance(repository, dict) and repository.get("url"): - repositories.append( - RepositoryReference( - url=str(repository["url"]), - is_official=repository.get("is_official") is True, - ) + repository = RepositoryReference( + url=str(repository["url"]), + is_official=repository.get("is_official") is True, ) + if include_resources or repository.is_official: + repositories.append(repository) + repositories.sort(key=lambda repository: not repository.is_official) + repositories = repositories[:repo_limit] project_pages = [] for page in item.get("project_pages") or []: url = page.get("url") if isinstance(page, dict) else page @@ -208,8 +254,9 @@ def paper_detail(item: dict[str, Any]) -> PaperDetail: if item.get("citation_count") is not None else None ), - url=str(item.get("url_abs") or item.get("source_url") or "") or None, - pdf_url=str(item["url_pdf"]) if item.get("url_pdf") else None, + url=_absolute_url(item.get("url_abs") or item.get("source_url")), + pdf_url=_absolute_url(item.get("url_pdf")), + code_repository_count=int(item.get("code_repository_count") or 0), tasks=[ catalog_reference(task) for task in item.get("tasks") or [] @@ -221,7 +268,13 @@ def paper_detail(item: dict[str, Any]) -> PaperDetail: if isinstance(method, dict) ], repositories=repositories, - project_pages=project_pages, + project_pages=project_pages[:repo_limit] if include_resources else [], + hf_models=[str(value) for value in item.get("hf_models") or []][:repo_limit] + if include_resources + else [], + hf_datasets=[str(value) for value in item.get("hf_datasets") or []][:repo_limit] + if include_resources + else [], ) @@ -238,10 +291,18 @@ def paper_reference(item: dict[str, Any]) -> PaperReference: def benchmark_summary(item: dict[str, Any]) -> BenchmarkSummary: + slug = str(item["slug"]) if item.get("slug") else None return BenchmarkSummary( id=str(item.get("id") or ""), name=str(item.get("name") or item.get("slug") or "Unknown benchmark"), - slug=str(item["slug"]) if item.get("slug") else None, + slug=slug, + url=( + str(item.get("url_abs")) + if item.get("url_abs") + else f"https://paperswithcode.co/dataset/{slug}" + if slug + else None + ), full_name=str(item["full_name"]) if item.get("full_name") else None, description=str(item["description"]) if item.get("description") else None, hf_url=str(item["hf_url"]) if item.get("hf_url") else None, @@ -254,7 +315,7 @@ def evaluation(item: dict[str, Any]) -> Evaluation: return Evaluation( id=str(item.get("id") or ""), model_name=str(item.get("model_name") or "Unknown model"), - metrics={str(key): value for key, value in metrics.items()} + metrics={str(key): _metric_value(value) for key, value in metrics.items()} if isinstance(metrics, dict) else {}, best_rank=int(item["best_rank"]) if item.get("best_rank") is not None else None, @@ -271,3 +332,57 @@ def evaluation(item: dict[str, Any]) -> Evaluation: else None ), ) + + +def _metric_value(value: Any) -> float | int | str | None: + if value is None or isinstance(value, (int, float)) and not isinstance(value, bool): + return value + if isinstance(value, str): + try: + return float(value) + except ValueError: + return value + return str(value) + + +def merged_evaluations(items: list[dict[str, Any]]) -> list[Evaluation]: + """Return one row per paper/model/setup with all reported metrics combined.""" + merged: dict[tuple[str, ...], dict[str, Any]] = {} + parameter_counts: dict[tuple[str, ...], set[int | None]] = {} + for item in items: + key = tuple( + str(item.get(field) or "") + for field in ("paper_id", "task_id", "dataset_id", "model_name", "harness") + ) + count = item.get("num_parameters") + valid_count = ( + count + if isinstance(count, int) and not isinstance(count, bool) and count > 0 + else None + ) + parameter_counts.setdefault(key, set()).add(valid_count) + current = merged.get(key) + if current is None: + merged[key] = {**item, "metrics": dict(item.get("metrics") or {})} + continue + current["metrics"].update(item.get("metrics") or {}) + ranks = [ + rank + for rank in (current.get("best_rank"), item.get("best_rank")) + if isinstance(rank, int) + ] + current["best_rank"] = min(ranks) if ranks else None + for key, counts in parameter_counts.items(): + merged[key]["num_parameters"] = next(iter(counts)) if len(counts) == 1 else None + return [ + evaluation(item) + for item in sorted( + merged.values(), + key=lambda item: ( + item.get("best_rank") + if isinstance(item.get("best_rank"), int) + else 10**9, + str(item.get("model_name") or "").casefold(), + ), + ) + ] diff --git a/mcp_server/src/pwc_mcp/server.py b/mcp_server/src/pwc_mcp/server.py index ea30566..184fcba 100644 --- a/mcp_server/src/pwc_mcp/server.py +++ b/mcp_server/src/pwc_mcp/server.py @@ -3,17 +3,22 @@ import os import time from datetime import date -from typing import Annotated, Any, Literal, Protocol +from typing import Annotated, Any, Literal, Protocol, TypeVar from mcp.server.caching import CacheHint from mcp.server.mcpserver import MCPServer from mcp.server.mcpserver.exceptions import ToolError -from mcp.types import ToolAnnotations +from mcp.types import CallToolResult, TextContent, ToolAnnotations from pwc_cli.transport import ResponseError, TransportError -from pydantic import Field +from pydantic import BaseModel, Field from pwc_mcp import __version__ -from pwc_mcp.catalog import PaperMarkdownChunk, PaperVersionMismatchError +from pwc_mcp.catalog import ( + AmbiguousError, + CatalogError, + PaperMarkdownChunk, + PaperVersionMismatchError, +) from pwc_mcp.cursors import ( CURSOR_LIFETIME_SECONDS, MAX_CHUNK_BYTES, @@ -24,6 +29,7 @@ AreaReference, BenchmarkPage, BenchmarkResult, + EvaluationPage, MethodDetail, MethodResult, PaperInfoResult, @@ -32,12 +38,14 @@ PaperReadResult, TaskDetail, TaskResult, + TaxonomyPage, benchmark_summary, catalog_reference, - evaluation, + merged_evaluations, paper_detail, paper_reference, paper_summary, + taxonomy_reference, ) READ_ONLY = ToolAnnotations( @@ -48,6 +56,7 @@ ) Page = Annotated[int, Field(ge=1, le=100)] Limit = Annotated[int, Field(ge=1, le=25)] +ResourceLimit = Annotated[int, Field(ge=1, le=10)] Reference = Annotated[str, Field(min_length=1, max_length=500)] Query = Annotated[str, Field(min_length=1, max_length=500)] AuthorList = Annotated[list[str], Field(max_length=10)] @@ -63,12 +72,44 @@ def _validate_date_range(start: str | None, end: str | None) -> None: raise ToolError("published_after must be on or before published_before") +def _catalog_error(error: Exception) -> ToolError: + if isinstance(error, AmbiguousError): + choices = ", ".join( + " / ".join(value for value in candidate.values() if value) + for candidate in error.candidates + ) + return ToolError(f"ambiguous: {error}. Candidates: {choices}") + if isinstance(error, CatalogError): + return ToolError(f"{error.code}: {error}") + if isinstance(error, TransportError): + message = str(error).casefold() + if "timed out" in message or "timeout" in message: + return ToolError("upstream_timeout: the Papers With Code catalog timed out") + return ToolError("upstream_error: the Papers With Code catalog request failed") + + def _catalog_call(function: Any, *args: Any, **kwargs: Any) -> Any: - """Turn upstream failures into deliberately generic, non-content-bearing errors.""" try: return function(*args, **kwargs) except (ResponseError, TransportError) as error: - raise ToolError("the Papers With Code catalog request failed") from error + raise _catalog_error(error) from error + + +Output = TypeVar("Output", bound=BaseModel) + + +def _tool_result(value: Output, markdown: str) -> CallToolResult: + structured = value.model_dump(mode="json") + return CallToolResult( + content=[TextContent(type="text", text=markdown)], + structured_content=structured, + ) + + +def _structured_model(result: Any, model: Any) -> Any: + if isinstance(result, CallToolResult): + return model.model_validate(result.structured_content) + return result class Catalog(Protocol): @@ -120,12 +161,26 @@ def list_papers( def get_related_papers(self, paper: str, *, limit: int) -> dict[str, Any]: ... + def get_trending_papers( + self, *, limit: int, max_age_days: int, min_velocity: float | None + ) -> dict[str, Any]: ... + + def get_paper_evaluations(self, paper: str, *, limit: int) -> dict[str, Any]: ... + def get_paper_lineage(self, paper: str) -> dict[str, Any]: ... def get_task(self, task: str) -> dict[str, Any]: ... + def list_tasks( + self, *, search: str | None, page: int, limit: int + ) -> dict[str, Any]: ... + def get_method(self, method: str) -> dict[str, Any]: ... + def list_methods( + self, *, search: str | None, page: int, limit: int + ) -> dict[str, Any]: ... + def list_benchmarks( self, *, @@ -157,6 +212,13 @@ def build_server( "pwc", title="Papers With Code", description="Read-only access to papers, tasks, methods, and benchmarks.", + instructions=( + "Use search_papers for relevance queries and list_papers for deterministic " + "filters. Pass slugs or numeric IDs returned by list tools to exact task, " + "method, and benchmark tools. Dates use YYYY-MM-DD. When a tool returns " + "ambiguous, retry with a listed slug or ID. Structured data is in " + "structuredContent; text is only a compact Markdown summary." + ), version=__version__, website_url="https://paperswithcode.co", cache_hints={ @@ -177,7 +239,7 @@ def search_papers( published_before: str | None = None, has_official_implementation: bool = False, ) -> PaperPage: - """Search papers by title, topic, author, or arXiv ID.""" + """Search by relevance across paper title, topic, author, or arXiv ID. Use keyword for exact terms and semantic for concepts. Dates are YYYY-MM-DD. has_official_implementation means at least one author-claimed official code repository.""" _validate_date_range(published_after, published_before) payload = _catalog_call( catalog.search_papers, @@ -189,7 +251,7 @@ def search_papers( published_before=published_before, has_official_implementation=has_official_implementation, ) - return PaperPage( + result = PaperPage( items=[ paper_summary(item) for item in payload.get("results") or [] @@ -201,16 +263,41 @@ def search_papers( else None ), ) + return _tool_result( + result, + f"Found {len(result.items)} papers." + + (f" Next page: {result.next_page}." if result.next_page else ""), + ) @server.tool(annotations=READ_ONLY, structured_output=True) - def get_paper_info(paper: Reference) -> PaperInfoResult: - """Get metadata for an arXiv ID, PwC ID, URL, or exact paper title.""" + def get_paper_info( + paper: Reference, + include_resources: bool = False, + repo_limit: ResourceLimit = 3, + ) -> PaperInfoResult: + """Get a paper by arXiv ID (for example 1706.03762), numeric PwC ID, supported URL, or exact title. By default returns official repositories only; include_resources adds other repositories, project pages, and Hugging Face models/datasets up to repo_limit repositories.""" payload = _catalog_call(catalog.get_paper_info, paper, include_resources=True) - return PaperInfoResult(paper=paper_detail(payload)) + result = PaperInfoResult( + paper=paper_detail( + payload, + include_resources=include_resources, + repo_limit=repo_limit, + ) + ) + official = next( + (repo.url for repo in result.paper.repositories if repo.is_official), None + ) + markdown = f"## {result.paper.title}\n\n" + if result.paper.url: + markdown += f"[Paper]({result.paper.url})" + if official: + markdown += f" · [Official code]({official})" + markdown += f" · {result.paper.code_repository_count} code repositories" + return _tool_result(result, markdown) @server.tool(annotations=READ_ONLY, structured_output=True) def read_paper(paper: Reference, cursor: str | None = None) -> PaperReadResult: - """Read stored paper Markdown, continuing oversized documents with a cursor.""" + """Read stored paper Markdown by arXiv/PwC ID, supported URL, or exact title. Pass next_cursor unchanged to continue a document; no_markdown means the catalog has no stored text.""" reference = paper.strip() try: state = codec.decode(cursor, reference=reference) if cursor else None @@ -239,9 +326,11 @@ def read_paper(paper: Reference, cursor: str | None = None) -> PaperReadResult: resolved=True, ) except PaperVersionMismatchError as error: - raise ToolError("paper changed; restart reading from the beginning") from error + raise ToolError( + "paper changed; restart reading from the beginning" + ) from error except (ResponseError, TransportError) as error: - raise ToolError("the Papers With Code catalog request failed") from error + raise _catalog_error(error) from error if chunk.paper != canonical or chunk.source != source: raise ToolError("the Papers With Code catalog request failed") next_cursor = None @@ -257,12 +346,13 @@ def read_paper(paper: Reference, cursor: str | None = None) -> PaperReadResult: expires_at=expires_at, ) ) - return PaperReadResult( + result = PaperReadResult( paper=reference, markdown=chunk.markdown, truncated=chunk.truncated, next_cursor=next_cursor, ) + return _tool_result(result, chunk.markdown) @server.tool(annotations=READ_ONLY, structured_output=True) def list_papers( @@ -282,7 +372,7 @@ def list_papers( page: Page = 1, limit: Limit = 10, ) -> PaperPage: - """List and filter papers in a deterministic catalog order.""" + """List papers in deterministic catalog order. Unlike search_papers this is for filters and pagination. Task/method filters should be slugs returned by list_tasks/list_methods; dates use YYYY-MM-DD.""" _validate_date_range(published_after, published_before) payload = _catalog_call( catalog.list_papers, @@ -300,7 +390,7 @@ def list_papers( page=page, limit=limit, ) - return PaperPage( + result = PaperPage( items=[ paper_summary(item) for item in payload.get("results") or [] @@ -312,12 +402,13 @@ def list_papers( else None ), ) + return _tool_result(result, f"Listed {len(result.items)} papers.") @server.tool(annotations=READ_ONLY, structured_output=True) def get_related_papers(paper: Reference, limit: Limit = 10) -> PaperPage: - """Find catalog papers related to one paper.""" + """Find papers related to one arXiv/PwC ID, supported URL, or exact title.""" payload = _catalog_call(catalog.get_related_papers, paper, limit=limit) - return PaperPage( + result = PaperPage( items=[ paper_summary(item) for item in payload.get("results") or [] @@ -325,6 +416,47 @@ def get_related_papers(paper: Reference, limit: Limit = 10) -> PaperPage: ], next_page=None, ) + return _tool_result(result, f"Found {len(result.items)} related papers.") + + @server.tool(annotations=READ_ONLY, structured_output=True) + def get_trending_papers( + limit: Limit = 10, + max_age_days: Annotated[int, Field(ge=1, le=3650)] = 30, + min_velocity: Annotated[float, Field(ge=0)] | None = None, + ) -> PaperPage: + """List currently trending papers. max_age_days bounds paper age; min_velocity optionally filters the catalog's citation-velocity score.""" + payload = _catalog_call( + catalog.get_trending_papers, + limit=limit, + max_age_days=max_age_days, + min_velocity=min_velocity, + ) + result = PaperPage( + items=[ + paper_summary(item) + for item in payload.get("results") or payload.get("items") or [] + if isinstance(item, dict) + ], + next_page=None, + ) + return _tool_result(result, f"Found {len(result.items)} trending papers.") + + @server.tool(annotations=READ_ONLY, structured_output=True) + def get_paper_evaluations(paper: Reference, limit: Limit = 10) -> EvaluationPage: + """Get benchmark evaluation rows reported for one paper. paper accepts an arXiv/PwC ID, supported URL, or exact title.""" + payload = _catalog_call(catalog.get_paper_evaluations, paper, limit=limit) + rows = payload.get("results") or payload.get("items") or [] + result = EvaluationPage( + paper=paper, + evaluation_count=int(payload.get("count") or len(rows)), + evaluations=merged_evaluations( + [item for item in rows if isinstance(item, dict)] + ), + ) + return _tool_result( + result, + f"Found {result.evaluation_count} evaluation rows; returned {len(result.evaluations)}.", + ) @server.tool(annotations=READ_ONLY, structured_output=True) def get_paper_lineage(paper: Reference) -> PaperLineageResult: @@ -333,7 +465,7 @@ def get_paper_lineage(paper: Reference) -> PaperLineageResult: current = payload.get("paper") if not isinstance(current, dict): raise TypeError("lineage response did not contain a paper") - return PaperLineageResult( + result = PaperLineageResult( paper=paper_reference(current), predecessors=[ paper_reference(item) @@ -346,10 +478,14 @@ def get_paper_lineage(paper: Reference) -> PaperLineageResult: if isinstance(item, dict) ], ) + return _tool_result( + result, + f"{result.paper.title}: {len(result.predecessors)} predecessors, {len(result.successors)} successors.", + ) @server.tool(annotations=READ_ONLY, structured_output=True) - def get_task(task: Reference) -> TaskResult: - """Get an exact task by ID, slug, or name, including its benchmarks.""" + def get_task(task: Reference, benchmark_limit: ResourceLimit = 10) -> TaskResult: + """Get an exact task by numeric ID or slug from list_tasks. Display names are accepted only when unambiguous. Benchmarks are capped by benchmark_limit.""" payload = _catalog_call(catalog.get_task, task) item = payload.get("task") if not isinstance(item, dict): @@ -363,11 +499,20 @@ def get_task(task: Reference) -> TaskResult: if isinstance(area_item, dict) else None ) - return TaskResult( + benchmarks = sorted( + [ + benchmark_summary(value) + for value in payload.get("benchmarks") or [] + if isinstance(value, dict) + ], + key=lambda benchmark: (-benchmark.paper_count, benchmark.name.casefold()), + ) + result = TaskResult( task=TaskDetail( id=str(item.get("id") or ""), name=str(item.get("name") or "Unknown task"), slug=str(item.get("slug") or item.get("id") or ""), + url=f"https://paperswithcode.co/tasks/{item.get('slug') or item.get('id')}", description=( str(item["description"]) if item.get("description") else None ), @@ -383,23 +528,49 @@ def get_task(task: Reference) -> TaskResult: for value in payload.get("children") or [] if isinstance(value, dict) ], - benchmarks=[ - benchmark_summary(value) - for value in payload.get("benchmarks") or [] - if isinstance(value, dict) - ], + benchmark_count=len(benchmarks), + benchmarks=benchmarks[:benchmark_limit], ) ) + return _tool_result( + result, + f"## {result.task.name}\n\n{result.task.paper_count} papers · {result.task.benchmark_count} benchmarks; returned {len(result.task.benchmarks)}.", + ) + + @server.tool(annotations=READ_ONLY, structured_output=True) + def list_tasks( + search: str | None = None, + page: Page = 1, + limit: Limit = 10, + ) -> TaxonomyPage: + """List or search research tasks. Use this before get_task; pass the returned slug or numeric ID to avoid ambiguous display names.""" + payload = _catalog_call( + catalog.list_tasks, search=search, page=page, limit=limit + ) + result = TaxonomyPage( + items=[ + taxonomy_reference(item, kind="tasks") + for item in payload.get("results") or payload.get("items") or [] + if isinstance(item, dict) + ], + next_page=( + int(payload["next_page"]) + if payload.get("next_page") is not None + else None + ), + ) + return _tool_result(result, f"Listed {len(result.items)} tasks.") @server.tool(annotations=READ_ONLY, structured_output=True) def get_method(method: Reference) -> MethodResult: - """Get an exact method by ID, slug, full name, or name.""" + """Get an exact method by numeric ID or slug from list_methods. Full/display names are accepted only when unambiguous.""" item = _catalog_call(catalog.get_method, method) - return MethodResult( + result = MethodResult( method=MethodDetail( id=str(item.get("id") or ""), name=str(item.get("name") or "Unknown method"), slug=str(item.get("slug") or item.get("id") or ""), + url=f"https://paperswithcode.co/methods/{item.get('slug') or item.get('id')}", full_name=str(item["full_name"]) if item.get("full_name") else None, description=( str(item["description"]) if item.get("description") else None @@ -414,13 +585,46 @@ def get_method(method: Reference) -> MethodResult: if item.get("source_paper_id") else None ), - source_url=str(item["source_url"]) if item.get("source_url") else None, + source_url=( + f"https://paperswithcode.co{item['source_url']}" + if str(item.get("source_url") or "").startswith("/") + else str(item["source_url"]) + if item.get("source_url") + else None + ), source_title=( str(item["source_title"]) if item.get("source_title") else None ), paper_count=int(item.get("paper_count") or 0), ) ) + return _tool_result( + result, f"## {result.method.name}\n\n{result.method.paper_count} papers." + ) + + @server.tool(annotations=READ_ONLY, structured_output=True) + def list_methods( + search: str | None = None, + page: Page = 1, + limit: Limit = 10, + ) -> TaxonomyPage: + """List or search methods. Use this before get_method; pass the returned slug or numeric ID to avoid ambiguous display names such as Mamba.""" + payload = _catalog_call( + catalog.list_methods, search=search, page=page, limit=limit + ) + result = TaxonomyPage( + items=[ + taxonomy_reference(item, kind="methods") + for item in payload.get("results") or payload.get("items") or [] + if isinstance(item, dict) + ], + next_page=( + int(payload["next_page"]) + if payload.get("next_page") is not None + else None + ), + ) + return _tool_result(result, f"Listed {len(result.items)} methods.") @server.tool(annotations=READ_ONLY, structured_output=True) def list_benchmarks( @@ -432,7 +636,7 @@ def list_benchmarks( page: Page = 1, limit: Limit = 10, ) -> BenchmarkPage: - """List benchmark datasets with optional task and availability filters.""" + """List benchmark datasets, ordered by coverage, with optional task and availability filters. task should be a slug from list_tasks. Use search for a name query; minimum_evaluations filters coverage.""" payload = _catalog_call( catalog.list_benchmarks, search=search, @@ -443,7 +647,7 @@ def list_benchmarks( page=page, limit=limit, ) - return BenchmarkPage( + result = BenchmarkPage( items=[ benchmark_summary(item) for item in payload.get("results") or [] @@ -455,6 +659,7 @@ def list_benchmarks( else None ), ) + return _tool_result(result, f"Listed {len(result.items)} benchmarks.") @server.tool(annotations=READ_ONLY, structured_output=True) def get_benchmark( @@ -462,21 +667,54 @@ def get_benchmark( limit: Limit = 10, is_open: bool | None = None, ) -> BenchmarkResult: - """Get an exact benchmark and its top evaluation rows.""" + """Get one exact benchmark by numeric ID or slug from list_benchmarks and its top model rows. Display names are accepted only when unambiguous; is_open filters reproducible/open implementations.""" payload = _catalog_call( catalog.get_benchmark, benchmark, limit=limit, is_open=is_open ) item = payload.get("benchmark") if not isinstance(item, dict): raise TypeError("benchmark response did not contain a benchmark") - return BenchmarkResult( + result = BenchmarkResult( benchmark=benchmark_summary(item), evaluation_count=int(payload.get("count") or 0), - evaluations=[ - evaluation(value) - for value in payload.get("results") or [] - if isinstance(value, dict) - ], + evaluations=merged_evaluations( + [ + value + for value in payload.get("results") or [] + if isinstance(value, dict) + ] + ), + ) + return _tool_result( + result, + f"## {result.benchmark.name}\n\n{result.evaluation_count} evaluation rows; returned {len(result.evaluations)} models.", + ) + + @server.prompt(name="find_papers", title="Find papers") + def find_papers_prompt(topic: str) -> str: + """Find relevant papers and their official implementations for a topic.""" + return ( + f"Search Papers With Code for papers about {topic}. Start with " + "search_papers, summarize the strongest matches, then call get_paper_info " + "for official repositories. Cite the absolute paper and repository URLs." + ) + + @server.prompt(name="compare_leaderboard", title="Compare a leaderboard") + def compare_leaderboard_prompt(benchmark: str) -> str: + """Compare leading models on a named benchmark.""" + return ( + f"Use list_benchmarks to resolve {benchmark} without guessing, then call " + "get_benchmark with its slug. Compare models only within that returned " + "benchmark and explain metric direction when known." + ) + + @server.prompt(name="survey_task", title="Survey a research task") + def survey_task_prompt(task: str) -> str: + """Survey papers, methods, and benchmarks for a research task.""" + return ( + f"Resolve {task} with list_tasks, inspect it with get_task, then use its " + "slug with list_papers and list_benchmarks. Summarize representative " + "papers, official code, and high-coverage benchmarks." ) @server.resource( @@ -487,7 +725,8 @@ def get_benchmark( mime_type="application/json", ) def paper_info_resource(paper: str) -> str: - return get_paper_info(paper).model_dump_json() + result = _structured_model(get_paper_info(paper), PaperInfoResult) + return result.model_dump_json() @server.resource( "pwc://papers/{paper}/markdown", @@ -497,7 +736,7 @@ def paper_info_resource(paper: str) -> str: mime_type="text/markdown", ) def paper_markdown_resource(paper: str) -> str: - result = read_paper(paper) + result = _structured_model(read_paper(paper), PaperReadResult) if result.truncated: raise ValueError( "paper is too large for one resource response; use read_paper with its continuation cursor" @@ -512,7 +751,8 @@ def paper_markdown_resource(paper: str) -> str: mime_type="application/json", ) def task_resource(task: str) -> str: - return get_task(task).model_dump_json() + result = _structured_model(get_task(task), TaskResult) + return result.model_dump_json() @server.resource( "pwc://benchmarks/{benchmark}", @@ -522,6 +762,7 @@ def task_resource(task: str) -> str: mime_type="application/json", ) def benchmark_resource(benchmark: str) -> str: - return get_benchmark(benchmark).model_dump_json() + result = _structured_model(get_benchmark(benchmark), BenchmarkResult) + return result.model_dump_json() return server diff --git a/mcp_server/tests/test_app.py b/mcp_server/tests/test_app.py index b446b01..14a4908 100644 --- a/mcp_server/tests/test_app.py +++ b/mcp_server/tests/test_app.py @@ -43,14 +43,36 @@ def test_health_and_browser_origin_policy_are_explicit(): assert health.json() == { "status": "ok", "service": "pwc-mcp", - "version": "0.1.0", - "protocol": "2026-07-28", + "version": "0.2.0", + "protocol": "2025-11-25", } assert rejected.status_code == 403 assert preflight.status_code == 200 assert preflight.headers["access-control-allow-origin"] == "https://chatgpt.com" +def test_discovery_docs_and_bare_get_are_client_friendly(): + app = create_app(StubCatalog(), allowed_hosts=["testserver"]) + + with TestClient(app) as client: + discovery = client.get("/.well-known/mcp") + docs = client.get("/docs") + bare_get = client.get("/mcp") + + assert discovery.status_code == 200 + assert discovery.json()["transport"]["url"] == "/mcp" + assert discovery.json()["protocol_version"] == "2025-11-25" + assert docs.status_code == 200 + assert {tool["name"] for tool in docs.json()["tools"]} >= { + "search_papers", + "get_paper_info", + "get_benchmark", + } + assert docs.json()["tools"][0]["inputSchema"]["type"] == "object" + assert bare_get.status_code == 405 + assert bare_get.headers["allow"] == "POST" + + def test_wildcard_browser_origin_is_rejected_at_startup(): try: create_app(StubCatalog(), allowed_origins=["*"]) @@ -206,18 +228,15 @@ def test_one_http_endpoint_serves_modern_and_legacy_protocol_eras(): def test_proxy_identity_trusts_only_an_exact_loopback_peer(): headers = Headers({"x-forwarded-for": "203.0.113.9"}) - assert _client_address( - {"client": ("127.0.0.1", 1234)}, headers, True - ) == "203.0.113.9" - assert _client_address( - {"client": ("::1", 1234)}, headers, True - ) == "203.0.113.9" - assert _client_address( - {"client": ("10.0.0.2", 1234)}, headers, True - ) == "10.0.0.2" - assert _client_address( - {"client": ("192.168.1.2", 1234)}, headers, True - ) == "192.168.1.2" + assert ( + _client_address({"client": ("127.0.0.1", 1234)}, headers, True) == "203.0.113.9" + ) + assert _client_address({"client": ("::1", 1234)}, headers, True) == "203.0.113.9" + assert _client_address({"client": ("10.0.0.2", 1234)}, headers, True) == "10.0.0.2" + assert ( + _client_address({"client": ("192.168.1.2", 1234)}, headers, True) + == "192.168.1.2" + ) def test_serialized_mcp_response_limit_fails_closed(): diff --git a/mcp_server/tests/test_catalog.py b/mcp_server/tests/test_catalog.py index 6c137ba..eaa7f15 100644 --- a/mcp_server/tests/test_catalog.py +++ b/mcp_server/tests/test_catalog.py @@ -4,7 +4,12 @@ import pytest from pwc_cli.transport import Response, ResponseError -from pwc_mcp.catalog import CatalogClient +from pwc_mcp.catalog import ( + AmbiguousError, + CatalogClient, + NotFoundError, + UpstreamTimeoutError, +) class StubTransport: @@ -173,8 +178,53 @@ def test_catalog_resolves_exact_titles_and_rejects_ambiguous_titles(): ) catalog = CatalogClient(transport=transport) - with pytest.raises(ResponseError, match="ambiguous"): + with pytest.raises(AmbiguousError, match="ambiguous") as captured: catalog.get_paper_lineage("Same Title") + assert [candidate["reference"] for candidate in captured.value.candidates] == [ + "1111.11111", + "2222.22222", + ] + + +def test_catalog_reports_a_typed_missing_title(): + catalog = CatalogClient(transport=StubTransport({"papers/search": {"results": []}})) + + with pytest.raises(NotFoundError, match="Paper title not found"): + catalog.get_related_papers("A Paper That Does Not Exist", limit=5) + + +def test_catalog_reports_a_typed_upstream_timeout(): + class TimeoutTransport: + def get(self, _path, _params=None): + raise TimeoutError("timed out") + + catalog = CatalogClient(transport=TimeoutTransport()) + + with pytest.raises(UpstreamTimeoutError): + catalog.search_papers(query="transformers") + + +def test_catalog_rejects_ambiguous_taxonomy_names_with_candidates(): + catalog = CatalogClient( + transport=StubTransport( + { + "datasets/": { + "results": [ + {"id": "1", "name": "ImageNet", "slug": "imagenet-a"}, + {"id": "2", "name": "ImageNet", "slug": "imagenet-b"}, + ] + } + } + ) + ) + + with pytest.raises(AmbiguousError) as captured: + catalog.get_benchmark("ImageNet", limit=10, is_open=None) + + assert [candidate["slug"] for candidate in captured.value.candidates] == [ + "imagenet-a", + "imagenet-b", + ] def test_catalog_resolves_pwc_urls_and_dotted_legacy_arxiv_ids(): diff --git a/mcp_server/tests/test_server.py b/mcp_server/tests/test_server.py index 07b0ee2..d2d4033 100644 --- a/mcp_server/tests/test_server.py +++ b/mcp_server/tests/test_server.py @@ -5,7 +5,7 @@ from mcp.client import Client from pwc_cli.transport import ResponseError -from pwc_mcp.catalog import PaperMarkdownChunk +from pwc_mcp.catalog import NoMarkdownError, PaperMarkdownChunk from pwc_mcp.server import build_server @@ -45,6 +45,9 @@ def get_paper_info(self, paper: str, *, include_resources: bool): "citation_count": 190_373, "url_abs": "https://arxiv.org/abs/1706.03762v7", "url_pdf": "https://arxiv.org/pdf/1706.03762v7.pdf", + "code_repository_count": 595, + "hf_models": ["huggingface/transformers"], + "hf_datasets": ["huggingface/example"], "tasks": [ { "id": "6", @@ -83,7 +86,9 @@ def read_paper_chunk( self.read_calls.append((offset, content_version, limit)) raw = b"abcdefgh" markdown = raw[offset : offset + limit].decode() - next_offset = offset + len(markdown) if offset + len(markdown) < len(raw) else None + next_offset = ( + offset + len(markdown) if offset + len(markdown) < len(raw) else None + ) return PaperMarkdownChunk( paper=paper, source="arxiv", @@ -100,6 +105,17 @@ def get_related_papers(self, paper: str, *, limit: int): assert limit == 2 return self.search_papers() + def get_trending_papers(self, *, limit: int, max_age_days: int, min_velocity): + assert limit == 2 + assert max_age_days == 30 + assert min_velocity is None + return self.search_papers() + + def get_paper_evaluations(self, paper: str, *, limit: int): + assert paper == "1706.03762" + assert limit == 5 + return self.get_benchmark("imagenet-1k", limit=5, is_open=None) + def get_paper_lineage(self, paper: str): assert paper == "1706.03762" return { @@ -137,6 +153,18 @@ def get_task(self, task: str): ], } + def list_tasks(self, **_kwargs): + return { + "results": [ + { + "id": "1", + "name": "Image Classification", + "slug": "image-classification", + } + ], + "next_page": None, + } + def get_method(self, method: str): assert method == "transformer" return { @@ -152,6 +180,12 @@ def get_method(self, method: str): "paper_count": 13505, } + def list_methods(self, **_kwargs): + return { + "results": [{"id": "2", "name": "Transformer", "slug": "transformer"}], + "next_page": None, + } + def list_benchmarks(self, **_kwargs): return { "next_page": None, @@ -171,19 +205,30 @@ def get_benchmark(self, benchmark: str, *, limit: int, is_open: bool | None): assert is_open in {True, None} return { "benchmark": {"id": "72", "name": "ImageNet-1k", "slug": "imagenet-1k"}, - "count": 1, + "count": 2, "results": [ { "id": "10", "model_name": "ExampleNet", - "metrics": {"Accuracy": 90.1}, + "metrics": {"Accuracy": "90.1"}, "best_rank": 1, "paper_id": "755", "paper_title": "Attention Is All You Need", "paper_arxiv_id": "1706.03762", "is_open": True, "num_parameters": 1000, - } + }, + { + "id": "11", + "model_name": "ExampleNet", + "metrics": {"F1": 88}, + "best_rank": 2, + "paper_id": "755", + "paper_title": "Attention Is All You Need", + "paper_arxiv_id": "1706.03762", + "is_open": True, + "num_parameters": 1000, + }, ], } @@ -260,7 +305,7 @@ async def exercise(): assert result.is_error is True assert result.content[0].text == ( "Error executing tool search_papers: " - "the Papers With Code catalog request failed" + "upstream_error: the Papers With Code catalog request failed" ) assert secret_query not in caplog.text @@ -288,6 +333,11 @@ async def exercise(): {"id": "6", "name": "Machine Translation", "slug": "machine-translation"} ] assert info.structured_content["paper"]["repositories"][0]["is_official"] is True + assert info.structured_content["paper"]["code_repository_count"] == 595 + assert info.structured_content["paper"]["project_pages"] == [] + assert info.structured_content["paper"]["hf_models"] == [] + assert info.content[0].text.startswith("## Attention Is All You Need") + assert '"schema_version"' not in info.content[0].text assert first.structured_content["markdown"] == "abcde" assert first.structured_content["truncated"] is True assert first.structured_content["next_cursor"] @@ -302,6 +352,33 @@ async def exercise(): assert catalog.read_calls == [(0, None, 5), (5, "a" * 64, 5)] +def test_discovery_tools_and_prompts_cover_common_agent_flows(): + async def exercise(): + async with Client(build_server(StubCatalog())) as client: + trending = await client.call_tool("get_trending_papers", {"limit": 2}) + evaluations = await client.call_tool( + "get_paper_evaluations", {"paper": "1706.03762", "limit": 5} + ) + tasks = await client.call_tool("list_tasks", {"search": "image"}) + methods = await client.call_tool("list_methods", {"search": "transformer"}) + prompts = {prompt.name for prompt in (await client.list_prompts()).prompts} + return trending, evaluations, tasks, methods, prompts + + trending, evaluations, tasks, methods, prompts = asyncio.run(exercise()) + + assert trending.structured_content["items"][0]["id"] == "755" + assert evaluations.structured_content["evaluations"][0]["metrics"] == { + "Accuracy": 90.1, + "F1": 88, + } + assert tasks.structured_content["items"][0]["slug"] == "image-classification" + assert tasks.structured_content["items"][0]["url"] == ( + "https://paperswithcode.co/tasks/image-classification" + ) + assert methods.structured_content["items"][0]["slug"] == "transformer" + assert prompts == {"find_papers", "compare_leaderboard", "survey_task"} + + def test_read_paper_rejects_invalid_continuation_as_an_expected_error(): async def exercise(): async with Client(build_server(StubCatalog())) as client: @@ -317,6 +394,21 @@ async def exercise(): ) +def test_read_paper_reports_typed_no_markdown_error(): + class MissingMarkdownCatalog(StubCatalog): + def read_paper_chunk(self, *_args, **_kwargs): + raise NoMarkdownError("No stored Markdown is available for this paper") + + async def exercise(): + async with Client(build_server(MissingMarkdownCatalog())) as client: + return await client.call_tool("read_paper", {"paper": "1706.03762"}) + + result = asyncio.run(exercise()) + + assert result.is_error is True + assert "no_markdown:" in result.content[0].text + + def test_paper_listing_related_work_and_lineage_are_composable(): async def exercise(): async with Client(build_server(StubCatalog())) as client: @@ -364,9 +456,13 @@ async def exercise(): "get_paper_info", "read_paper", "get_related_papers", + "get_trending_papers", + "get_paper_evaluations", "get_paper_lineage", "get_task", + "list_tasks", "get_method", + "list_methods", "list_benchmarks", "get_benchmark", } @@ -374,8 +470,10 @@ async def exercise(): assert method.structured_content["method"]["introduced_year"] == 2017 assert benchmarks.structured_content["items"][0]["slug"] == "imagenet-1k" assert benchmark.structured_content["evaluations"][0]["metrics"] == { - "Accuracy": 90.1 + "Accuracy": 90.1, + "F1": 88, } + assert len(benchmark.structured_content["evaluations"]) == 1 def test_resources_expose_canonical_papers_tasks_and_benchmarks(): diff --git a/mcp_server/tests/test_skill.py b/mcp_server/tests/test_skill.py index 0ccb341..338b9a0 100644 --- a/mcp_server/tests/test_skill.py +++ b/mcp_server/tests/test_skill.py @@ -10,9 +10,13 @@ "get_paper_info", "read_paper", "get_related_papers", + "get_trending_papers", + "get_paper_evaluations", "get_paper_lineage", "get_task", + "list_tasks", "get_method", + "list_methods", "list_benchmarks", "get_benchmark", } diff --git a/mcp_server/uv.lock b/mcp_server/uv.lock index f9c4aca..d4fee0d 100644 --- a/mcp_server/uv.lock +++ b/mcp_server/uv.lock @@ -443,7 +443,7 @@ source = { editable = "../standalone_cli" } [[package]] name = "pwc-mcp" -version = "0.1.0" +version = "0.2.0" source = { editable = "." } dependencies = [ { name = "mcp" }, From 13f726f3d278ed894382e11f559d881523a78689 Mon Sep 17 00:00:00 2001 From: Niels Rogge Date: Thu, 17 Sep 2026 13:44:42 +0000 Subject: [PATCH 2/2] Expose interpretable paginated MCP evaluations --- mcp_server/README.md | 11 +- mcp_server/SKILL.md | 14 +- mcp_server/SPEC.md | 6 + mcp_server/pyproject.toml | 2 +- mcp_server/src/pwc_mcp/__init__.py | 2 +- mcp_server/src/pwc_mcp/app.py | 3 +- mcp_server/src/pwc_mcp/catalog.py | 16 +- mcp_server/src/pwc_mcp/models.py | 245 ++++++++++++++++++++++++++--- mcp_server/src/pwc_mcp/server.py | 65 +++++--- mcp_server/tests/test_app.py | 6 +- mcp_server/tests/test_server.py | 68 +++++++- mcp_server/uv.lock | 2 +- 12 files changed, 384 insertions(+), 56 deletions(-) diff --git a/mcp_server/README.md b/mcp_server/README.md index 95773b3..61f2ed6 100644 --- a/mcp_server/README.md +++ b/mcp_server/README.md @@ -45,6 +45,10 @@ the caller controls keyword or semantic mode. `read_paper` fetches at most one 64 KiB catalog chunk per call and returns a signed, one-hour continuation cursor when more Markdown remains. Continuations stay pinned to the resolved paper and content version, so a changed paper fails with an explicit restart response. +Benchmark and paper evaluations are paginated. Equivalent duplicate rows are +merged while preserving every task-specific rank scope; rows also expose the +reported protocol, split, shot count when stated, source, update timestamp, +openness, and metric direction. ## Resources @@ -77,8 +81,11 @@ request, upstream, and serialized MCP response bodies are bounded to 2 MiB. Markdown chunks use a bounded 256-entry/16 MiB in-memory cache. `GET /.well-known/mcp` exposes connection metadata and `GET /docs` publishes -the live input/output schema for every tool. A bare `GET /mcp` returns `405`; -MCP requests use `POST /mcp`. Rate limits return `429` with `Retry-After`. +the live input/output schema for every tool. On the canonical host these are +also available as `/.well-known/mcp` and `/mcp/schema`, while a human GET of +`/mcp` opens the setup guide. Direct package servers return `405` for a bare +`GET /mcp`; MCP requests use `POST /mcp`. Rate limits return `429` with +`Retry-After`. ## Test diff --git a/mcp_server/SKILL.md b/mcp_server/SKILL.md index 0a584c0..7a6dca3 100644 --- a/mcp_server/SKILL.md +++ b/mcp_server/SKILL.md @@ -4,7 +4,7 @@ description: "Papers With Code MCP tools for searching and reading AI/ML papers, compatibility: "Requires an MCP client connected to https://paperswithcode.co/mcp with the Papers With Code tools available." --- -Generated for `pwc-mcp v0.2.0` and MCP protocol `2025-11-25`. +Generated for `pwc-mcp v0.2.1` and MCP protocol `2025-11-25`. The tools query the public [Papers With Code](https://paperswithcode.co) catalog anonymously and are read-only. If live tool discovery and this skill disagree, @@ -41,14 +41,14 @@ matching IDs. - `list_papers({"page": PAGE, "limit": LIMIT, "search": SEARCH, "published_after": START_DATE, "published_before": END_DATE, "task": TASK, "method": METHOD, "conference": CONFERENCE, "framework": FRAMEWORK, "organization": ORGANIZATION, "authors": [AUTHOR], "order_by": "date_published"|"citation_count"|"title", "order_direction": "asc"|"desc"})` — list and filter papers. Omit optional arguments when they are not needed. - `get_related_papers({"paper": PAPER, "limit": LIMIT})` — list related papers. - `get_trending_papers({"limit": LIMIT, "max_age_days": DAYS, "min_velocity": VELOCITY})` — list trending papers. -- `get_paper_evaluations({"paper": PAPER, "limit": LIMIT})` — list benchmark evaluations reported by a paper. -- `get_paper_lineage({"paper": PAPER})` — list explicit predecessors and successors. +- `get_paper_evaluations({"paper": PAPER, "page": PAGE, "limit": LIMIT})` — list paginated benchmark evaluations reported by a paper. +- `get_paper_lineage({"paper": PAPER})` — list explicit catalog predecessors and successors; empty results do not prove that none exist. - `list_tasks({"search": SEARCH, "page": PAGE, "limit": LIMIT})` — discover task slugs and IDs. - `get_task({"task": TASK, "benchmark_limit": LIMIT})` — inspect one exact task by ID or slug with a capped benchmark list. - `list_methods({"search": SEARCH, "page": PAGE, "limit": LIMIT})` — discover method slugs and IDs. - `get_method({"method": METHOD})` — inspect one exact method by ID, slug, full name, or name. - `list_benchmarks({"page": PAGE, "limit": LIMIT, "search": SEARCH, "task": TASK, "include_descendants": BOOLEAN, "minimum_evaluations": MINIMUM_EVALUATIONS, "is_open": BOOLEAN})` — list and filter benchmarks. Omit optional arguments when they are not needed. -- `get_benchmark({"benchmark": BENCHMARK, "limit": LIMIT, "is_open": BOOLEAN})` — inspect one exact benchmark and its leading evaluation rows. +- `get_benchmark({"benchmark": BENCHMARK, "page": PAGE, "limit": LIMIT, "is_open": BOOLEAN})` — inspect one exact benchmark and one page of evaluation rows. All page numbers start at 1. `limit` is between 1 and 25. Follow `next_page` when the user asks for more results than one response contains; do not infer @@ -83,6 +83,12 @@ one of those capabilities. - Tool results use stable, versioned structured output with compact text fallbacks. Prefer structured fields over parsing the text fallback. +- Evaluation ranks are task-scoped. Compare ranks only within matching + `rank_scopes`; use `metric_directions`, protocol, split, shots, and source to + decide whether scores are comparable. Follow `next_page` for more rows. +- `has_official_implementation` means the catalog marks linked code as official; + it is not an independent audit. Evaluation `is_open` is the catalog's + implementation-availability flag and is null when not recorded. - Search mode is deterministic: choose `keyword` by default and use `semantic` when conceptual similarity is more useful. The MCP server does not support the CLI's `hybrid` mode. diff --git a/mcp_server/SPEC.md b/mcp_server/SPEC.md index 21e7e6f..57e32c2 100644 --- a/mcp_server/SPEC.md +++ b/mcp_server/SPEC.md @@ -56,6 +56,12 @@ Paper references accept arXiv IDs, numeric PwC external IDs, arXiv/Hugging Face/Papers With Code URLs, and exact titles. Ambiguous exact titles fail rather than selecting one result. +Benchmark and paper evaluation results paginate with `page` and `next_page`. +Equivalent result rows merge their metrics while retaining task-scoped ranks, +evaluation protocol, split, shots when reported, source URL, openness, and +update timestamp. Metric direction is explicit and unknown directions remain +`unknown` rather than being guessed. + ## Safety and operations - Require a strict configurable browser Origin allowlist; native clients may diff --git a/mcp_server/pyproject.toml b/mcp_server/pyproject.toml index 76875d9..211434f 100644 --- a/mcp_server/pyproject.toml +++ b/mcp_server/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "pwc-mcp" -version = "0.2.0" +version = "0.2.1" description = "Read-only Papers With Code MCP server" readme = "README.md" requires-python = ">=3.10" diff --git a/mcp_server/src/pwc_mcp/__init__.py b/mcp_server/src/pwc_mcp/__init__.py index cab4aac..ba5b9ca 100644 --- a/mcp_server/src/pwc_mcp/__init__.py +++ b/mcp_server/src/pwc_mcp/__init__.py @@ -1,3 +1,3 @@ """Read-only Papers With Code MCP server.""" -__version__ = "0.2.0" +__version__ = "0.2.1" diff --git a/mcp_server/src/pwc_mcp/app.py b/mcp_server/src/pwc_mcp/app.py index 357abb4..f257586 100644 --- a/mcp_server/src/pwc_mcp/app.py +++ b/mcp_server/src/pwc_mcp/app.py @@ -66,7 +66,8 @@ async def well_known_mcp(request: Request) -> JSONResponse: "transport": {"type": "streamable-http", "url": "/mcp"}, "protocol_version": PROTOCOL_VERSION, "supported_protocol_versions": sorted(KNOWN_PROTOCOLS, reverse=True), - "documentation_url": "/docs", + "documentation_url": "https://paperswithcode.co/mcp/schema", + "setup_url": "https://paperswithcode.co/mcp", } ) diff --git a/mcp_server/src/pwc_mcp/catalog.py b/mcp_server/src/pwc_mcp/catalog.py index f561aff..9ec913d 100644 --- a/mcp_server/src/pwc_mcp/catalog.py +++ b/mcp_server/src/pwc_mcp/catalog.py @@ -551,7 +551,9 @@ def get_trending_papers( accept_list=True, ) - def get_paper_evaluations(self, paper: str, *, limit: int) -> dict[str, Any]: + def get_paper_evaluations( + self, paper: str, *, limit: int, page: int = 1 + ) -> dict[str, Any]: detail = self.get_paper_info(paper, include_resources=False) paper_id = detail.get("id") if not paper_id: @@ -559,7 +561,7 @@ def get_paper_evaluations(self, paper: str, *, limit: int) -> dict[str, Any]: return self._json( "evaluations/", { - "page": 1, + "page": page, "page_size": limit, "paper_id": paper_id, "ordering": "-benchmark_popularity", @@ -715,7 +717,12 @@ def list_benchmarks( ) def get_benchmark( - self, benchmark: str, *, limit: int, is_open: bool | None + self, + benchmark: str, + *, + limit: int, + is_open: bool | None, + page: int = 1, ) -> dict[str, Any]: candidates = self._rows( self._json( @@ -756,7 +763,7 @@ def get_benchmark( evaluations = self._json( f"datasets/{quote(benchmark_id, safe='')}/evaluations/", { - "page": 1, + "page": page, "page_size": limit, "ordering": "best_rank", "is_open": is_open, @@ -767,4 +774,5 @@ def get_benchmark( "benchmark": matched, "count": evaluations.get("count") or 0, "results": evaluations.get("results") or [], + "next_page": evaluations.get("next_page"), } diff --git a/mcp_server/src/pwc_mcp/models.py b/mcp_server/src/pwc_mcp/models.py index f5cc6a8..802baf5 100644 --- a/mcp_server/src/pwc_mcp/models.py +++ b/mcp_server/src/pwc_mcp/models.py @@ -1,8 +1,9 @@ from __future__ import annotations +import re from typing import Any, Literal -from pydantic import BaseModel, ConfigDict +from pydantic import BaseModel, ConfigDict, Field def _absolute_url(value: object) -> str | None: @@ -26,7 +27,12 @@ class PaperSummary(OutputModel): published: str | None = None citation_count: int | None = None url: str | None = None - has_official_implementation: bool + has_official_implementation: bool = Field( + description=( + "True when the catalog marks at least one linked repository as the " + "paper's official implementation; this is not an independent code audit." + ) + ) code_repository_count: int @@ -94,6 +100,11 @@ class PaperLineageResult(OutputModel): paper: PaperReference predecessors: list[PaperReference] successors: list[PaperReference] + coverage: Literal["explicit_catalog_links_only"] = "explicit_catalog_links_only" + coverage_note: str = ( + "Only explicit catalog relationships are returned; an empty list does not " + "prove that no predecessor or successor exists." + ) class AreaReference(OutputModel): @@ -108,6 +119,7 @@ class BenchmarkSummary(OutputModel): url: str | None = None full_name: str | None = None description: str | None = None + split: str | None = None hf_url: str | None = None paper_count: int @@ -162,16 +174,41 @@ class BenchmarkPage(OutputModel): next_page: int | None = None +class EvaluationRankScope(OutputModel): + task_id: str | None = None + task_name: str | None = None + task_slug: str | None = None + rank: int | None = None + + class Evaluation(OutputModel): id: str model_name: str metrics: dict[str, float | int | str | None] - best_rank: int | None = None + best_rank: int | None = Field( + default=None, + description="Best rank across the task-specific rank scopes listed in rank_scopes.", + ) + rank_scopes: list[EvaluationRankScope] paper_id: str | None = None paper_title: str | None = None paper_arxiv_id: str | None = None - is_open: bool + is_open: bool | None = Field( + description=( + "Catalog openness flag for the evaluated implementation: true=open, " + "false=closed, null=not recorded." + ) + ) num_parameters: int | None = None + split: str | None = None + shots: int | None = None + evaluation_protocol: str | None = None + harness: str | None = None + source_url: str | None = None + code_url: str | None = None + hf_model_url: str | None = None + updated_at: str | None = None + uses_additional_data: bool | None = None class BenchmarkResult(OutputModel): @@ -179,6 +216,10 @@ class BenchmarkResult(OutputModel): benchmark: BenchmarkSummary evaluation_count: int evaluations: list[Evaluation] + metric_directions: dict[str, Literal["higher", "lower", "unknown"]] + page: int + next_page: int | None = None + ranking_note: str = "Ranks are scoped by task and are not necessarily comparable across rank_scopes." class EvaluationPage(OutputModel): @@ -186,6 +227,8 @@ class EvaluationPage(OutputModel): paper: str evaluation_count: int evaluations: list[Evaluation] + page: int + next_page: int | None = None def paper_summary(item: dict[str, Any]) -> PaperSummary: @@ -305,6 +348,11 @@ def benchmark_summary(item: dict[str, Any]) -> BenchmarkSummary: ), full_name=str(item["full_name"]) if item.get("full_name") else None, description=str(item["description"]) if item.get("description") else None, + split=( + str(item.get("split_name") or item.get("split")) + if item.get("split_name") or item.get("split") + else None + ), hf_url=str(item["hf_url"]) if item.get("hf_url") else None, paper_count=int(item.get("paper_count") or 0), ) @@ -324,16 +372,71 @@ def evaluation(item: dict[str, Any]) -> Evaluation: paper_arxiv_id=( str(item["paper_arxiv_id"]) if item.get("paper_arxiv_id") else None ), - is_open=item.get("is_open") is not False, + rank_scopes=_rank_scopes(item), + is_open=( + item.get("is_open") if isinstance(item.get("is_open"), bool) else None + ), num_parameters=( int(item["num_parameters"]) if isinstance(item.get("num_parameters"), int) and not isinstance(item.get("num_parameters"), bool) else None ), + split=( + str(item.get("split") or item.get("split_name") or item.get("evaluated_on")) + if item.get("split") or item.get("split_name") or item.get("evaluated_on") + else None + ), + shots=_shot_count(item.get("methodology")), + evaluation_protocol=( + str(item["methodology"]) if item.get("methodology") else None + ), + harness=str(item["harness"]) if item.get("harness") else None, + source_url=_absolute_url( + item.get("result_url") + or item.get("source_url") + or item.get("external_source_url") + ), + code_url=_absolute_url(item.get("code_url")), + hf_model_url=_absolute_url(item.get("hf_model_url")), + updated_at=str(item["updated_at"]) if item.get("updated_at") else None, + uses_additional_data=( + item.get("uses_additional_data") + if isinstance(item.get("uses_additional_data"), bool) + else None + ), ) +def _rank_scopes(item: dict[str, Any]) -> list[EvaluationRankScope]: + existing = item.get("rank_scopes") + if isinstance(existing, list): + return [ + EvaluationRankScope.model_validate(scope) + for scope in existing + if isinstance(scope, dict) + ] + if not any( + item.get(field) is not None for field in ("task_id", "task_name", "best_rank") + ): + return [] + return [ + EvaluationRankScope( + task_id=str(item["task_id"]) if item.get("task_id") else None, + task_name=str(item["task_name"]) if item.get("task_name") else None, + task_slug=str(item["task_slug"]) if item.get("task_slug") else None, + rank=int(item["best_rank"]) if item.get("best_rank") is not None else None, + ) + ] + + +def _shot_count(methodology: object) -> int | None: + if not methodology: + return None + match = re.search(r"\b(\d+)\s*[- ]?shot\b", str(methodology), re.IGNORECASE) + return int(match.group(1)) if match else None + + def _metric_value(value: Any) -> float | int | str | None: if value is None or isinstance(value, (int, float)) and not isinstance(value, bool): return value @@ -346,38 +449,102 @@ def _metric_value(value: Any) -> float | int | str | None: def merged_evaluations(items: list[dict[str, Any]]) -> list[Evaluation]: - """Return one row per paper/model/setup with all reported metrics combined.""" - merged: dict[tuple[str, ...], dict[str, Any]] = {} - parameter_counts: dict[tuple[str, ...], set[int | None]] = {} + """Combine equivalent rows while retaining each task-specific rank scope.""" + merged: dict[tuple[str, ...], list[dict[str, Any]]] = {} for item in items: - key = tuple( - str(item.get(field) or "") - for field in ("paper_id", "task_id", "dataset_id", "model_name", "harness") + shots = _shot_count(item.get("methodology")) + base_key = ( + *( + str(item.get(field) or "") + for field in ("paper_id", "dataset_id", "model_name", "harness") + ), + str( + item.get("split") + or item.get("split_name") + or item.get("evaluated_on") + or "" + ), + str(shots) if shots is not None else "", ) + source_metrics = item.get("metrics") + metrics = { + str(name): _metric_value(value) + for name, value in ( + source_metrics.items() if isinstance(source_metrics, dict) else [] + ) + } + candidates = merged.setdefault(base_key, []) + current = next( + ( + candidate + for candidate in candidates + if all( + name not in candidate["metrics"] + or candidate["metrics"][name] == value + for name, value in metrics.items() + ) + ), + None, + ) + scope = _rank_scopes(item) count = item.get("num_parameters") valid_count = ( count if isinstance(count, int) and not isinstance(count, bool) and count > 0 else None ) - parameter_counts.setdefault(key, set()).add(valid_count) - current = merged.get(key) if current is None: - merged[key] = {**item, "metrics": dict(item.get("metrics") or {})} + candidates.append( + { + **item, + "metrics": metrics, + "rank_scopes": [value.model_dump() for value in scope], + "_parameter_counts": {valid_count}, + "_open_values": { + item.get("is_open") + if isinstance(item.get("is_open"), bool) + else None + }, + } + ) continue - current["metrics"].update(item.get("metrics") or {}) + current["metrics"].update(metrics) + known_scopes = { + (value.get("task_id"), value.get("rank")) + for value in current["rank_scopes"] + } + current["rank_scopes"].extend( + value.model_dump() + for value in scope + if (value.task_id, value.rank) not in known_scopes + ) ranks = [ rank for rank in (current.get("best_rank"), item.get("best_rank")) if isinstance(rank, int) ] current["best_rank"] = min(ranks) if ranks else None - for key, counts in parameter_counts.items(): - merged[key]["num_parameters"] = next(iter(counts)) if len(counts) == 1 else None + current["_parameter_counts"].add(valid_count) + current["_open_values"].add( + item.get("is_open") if isinstance(item.get("is_open"), bool) else None + ) + if len(str(item.get("methodology") or "")) > len( + str(current.get("methodology") or "") + ): + current["methodology"] = item["methodology"] + if str(item.get("updated_at") or "") > str(current.get("updated_at") or ""): + current["updated_at"] = item["updated_at"] + + rows = [row for candidates in merged.values() for row in candidates] + for row in rows: + counts = row.pop("_parameter_counts") + row["num_parameters"] = next(iter(counts)) if len(counts) == 1 else None + openness = row.pop("_open_values") + row["is_open"] = next(iter(openness)) if len(openness) == 1 else None return [ evaluation(item) for item in sorted( - merged.values(), + rows, key=lambda item: ( item.get("best_rank") if isinstance(item.get("best_rank"), int) @@ -386,3 +553,45 @@ def merged_evaluations(items: list[dict[str, Any]]) -> list[Evaluation]: ), ) ] + + +def metric_directions( + evaluations: list[Evaluation], +) -> dict[str, Literal["higher", "lower", "unknown"]]: + names = {name for row in evaluations for name in row.metrics} + lower_markers = ( + "error", + "loss", + "perplexity", + "latency", + "runtime", + "wer", + "cer", + "eer", + "fid", + "mae", + "rmse", + ) + higher_markers = ( + "accuracy", + "precision", + "recall", + "f1", + "bleu", + "rouge", + "map", + "auc", + "score", + "em", + ) + result: dict[str, Literal["higher", "lower", "unknown"]] = {} + for name in sorted(names): + normalized = re.sub(r"[^a-z0-9]+", " ", name.casefold()) + words = set(normalized.split()) + if any(marker in words for marker in lower_markers): + result[name] = "lower" + elif any(marker in words for marker in higher_markers): + result[name] = "higher" + else: + result[name] = "unknown" + return result diff --git a/mcp_server/src/pwc_mcp/server.py b/mcp_server/src/pwc_mcp/server.py index 184fcba..9a05e50 100644 --- a/mcp_server/src/pwc_mcp/server.py +++ b/mcp_server/src/pwc_mcp/server.py @@ -42,6 +42,7 @@ benchmark_summary, catalog_reference, merged_evaluations, + metric_directions, paper_detail, paper_reference, paper_summary, @@ -165,7 +166,9 @@ def get_trending_papers( self, *, limit: int, max_age_days: int, min_velocity: float | None ) -> dict[str, Any]: ... - def get_paper_evaluations(self, paper: str, *, limit: int) -> dict[str, Any]: ... + def get_paper_evaluations( + self, paper: str, *, page: int, limit: int + ) -> dict[str, Any]: ... def get_paper_lineage(self, paper: str) -> dict[str, Any]: ... @@ -194,7 +197,7 @@ def list_benchmarks( ) -> dict[str, Any]: ... def get_benchmark( - self, benchmark: str, *, limit: int, is_open: bool | None + self, benchmark: str, *, page: int, limit: int, is_open: bool | None ) -> dict[str, Any]: ... @@ -217,7 +220,10 @@ def build_server( "filters. Pass slugs or numeric IDs returned by list tools to exact task, " "method, and benchmark tools. Dates use YYYY-MM-DD. When a tool returns " "ambiguous, retry with a listed slug or ID. Structured data is in " - "structuredContent; text is only a compact Markdown summary." + "structuredContent; text is only a compact Markdown summary. Official " + "implementation means catalog-designated official code, not an independent " + "audit. Evaluation is_open is the catalog's implementation availability " + "flag and may be null when unknown." ), version=__version__, website_url="https://paperswithcode.co", @@ -239,7 +245,7 @@ def search_papers( published_before: str | None = None, has_official_implementation: bool = False, ) -> PaperPage: - """Search by relevance across paper title, topic, author, or arXiv ID. Use keyword for exact terms and semantic for concepts. Dates are YYYY-MM-DD. has_official_implementation means at least one author-claimed official code repository.""" + """Search by relevance across paper title, topic, author, or arXiv ID. Use keyword for exact terms and semantic for concepts. Dates are YYYY-MM-DD. has_official_implementation means at least one repository is marked official in the catalog; it is not an independent code audit.""" _validate_date_range(published_after, published_before) payload = _catalog_call( catalog.search_papers, @@ -442,9 +448,13 @@ def get_trending_papers( return _tool_result(result, f"Found {len(result.items)} trending papers.") @server.tool(annotations=READ_ONLY, structured_output=True) - def get_paper_evaluations(paper: Reference, limit: Limit = 10) -> EvaluationPage: - """Get benchmark evaluation rows reported for one paper. paper accepts an arXiv/PwC ID, supported URL, or exact title.""" - payload = _catalog_call(catalog.get_paper_evaluations, paper, limit=limit) + def get_paper_evaluations( + paper: Reference, page: Page = 1, limit: Limit = 10 + ) -> EvaluationPage: + """Get paginated benchmark evaluations reported for one paper, including protocol, source, update time, and task-scoped ranks. paper accepts an arXiv/PwC ID, supported URL, or exact title.""" + payload = _catalog_call( + catalog.get_paper_evaluations, paper, page=page, limit=limit + ) rows = payload.get("results") or payload.get("items") or [] result = EvaluationPage( paper=paper, @@ -452,15 +462,22 @@ def get_paper_evaluations(paper: Reference, limit: Limit = 10) -> EvaluationPage evaluations=merged_evaluations( [item for item in rows if isinstance(item, dict)] ), + page=page, + next_page=( + int(payload["next_page"]) + if payload.get("next_page") is not None + else None + ), ) return _tool_result( result, - f"Found {result.evaluation_count} evaluation rows; returned {len(result.evaluations)}.", + f"Found {result.evaluation_count} evaluation rows; returned {len(result.evaluations)} on page {page}." + + (f" Next page: {result.next_page}." if result.next_page else ""), ) @server.tool(annotations=READ_ONLY, structured_output=True) def get_paper_lineage(paper: Reference) -> PaperLineageResult: - """Get explicit predecessor and successor relationships for a paper.""" + """Get explicit catalog predecessor and successor links for a paper. Empty lists mean no relationships are recorded, not proof that none exist.""" payload = _catalog_call(catalog.get_paper_lineage, paper) current = payload.get("paper") if not isinstance(current, dict): @@ -480,7 +497,7 @@ def get_paper_lineage(paper: Reference) -> PaperLineageResult: ) return _tool_result( result, - f"{result.paper.title}: {len(result.predecessors)} predecessors, {len(result.successors)} successors.", + f"{result.paper.title}: {len(result.predecessors)} recorded predecessors, {len(result.successors)} recorded successors. Coverage is explicit catalog links only.", ) @server.tool(annotations=READ_ONLY, structured_output=True) @@ -664,30 +681,40 @@ def list_benchmarks( @server.tool(annotations=READ_ONLY, structured_output=True) def get_benchmark( benchmark: Reference, + page: Page = 1, limit: Limit = 10, is_open: bool | None = None, ) -> BenchmarkResult: - """Get one exact benchmark by numeric ID or slug from list_benchmarks and its top model rows. Display names are accepted only when unambiguous; is_open filters reproducible/open implementations.""" + """Get one exact benchmark and a page of evaluation rows. Rows expose split, shots/protocol, source, update time, metric direction, and task-specific rank scopes. Use next_page for continuation. is_open filters the catalog's open-implementation flag; null means unrecorded.""" payload = _catalog_call( - catalog.get_benchmark, benchmark, limit=limit, is_open=is_open + catalog.get_benchmark, + benchmark, + page=page, + limit=limit, + is_open=is_open, ) item = payload.get("benchmark") if not isinstance(item, dict): raise TypeError("benchmark response did not contain a benchmark") + evaluations = merged_evaluations( + [value for value in payload.get("results") or [] if isinstance(value, dict)] + ) result = BenchmarkResult( benchmark=benchmark_summary(item), evaluation_count=int(payload.get("count") or 0), - evaluations=merged_evaluations( - [ - value - for value in payload.get("results") or [] - if isinstance(value, dict) - ] + evaluations=evaluations, + metric_directions=metric_directions(evaluations), + page=page, + next_page=( + int(payload["next_page"]) + if payload.get("next_page") is not None + else None ), ) return _tool_result( result, - f"## {result.benchmark.name}\n\n{result.evaluation_count} evaluation rows; returned {len(result.evaluations)} models.", + f"## {result.benchmark.name}\n\n{result.evaluation_count} evaluation rows; returned {len(result.evaluations)} models on page {page}. Ranks are task-scoped." + + (f" Next page: {result.next_page}." if result.next_page else ""), ) @server.prompt(name="find_papers", title="Find papers") diff --git a/mcp_server/tests/test_app.py b/mcp_server/tests/test_app.py index 14a4908..69526f7 100644 --- a/mcp_server/tests/test_app.py +++ b/mcp_server/tests/test_app.py @@ -43,7 +43,7 @@ def test_health_and_browser_origin_policy_are_explicit(): assert health.json() == { "status": "ok", "service": "pwc-mcp", - "version": "0.2.0", + "version": "0.2.1", "protocol": "2025-11-25", } assert rejected.status_code == 403 @@ -62,6 +62,10 @@ def test_discovery_docs_and_bare_get_are_client_friendly(): assert discovery.status_code == 200 assert discovery.json()["transport"]["url"] == "/mcp" assert discovery.json()["protocol_version"] == "2025-11-25" + assert discovery.json()["documentation_url"] == ( + "https://paperswithcode.co/mcp/schema" + ) + assert discovery.json()["setup_url"] == "https://paperswithcode.co/mcp" assert docs.status_code == 200 assert {tool["name"] for tool in docs.json()["tools"]} >= { "search_papers", diff --git a/mcp_server/tests/test_server.py b/mcp_server/tests/test_server.py index d2d4033..e14771f 100644 --- a/mcp_server/tests/test_server.py +++ b/mcp_server/tests/test_server.py @@ -111,8 +111,9 @@ def get_trending_papers(self, *, limit: int, max_age_days: int, min_velocity): assert min_velocity is None return self.search_papers() - def get_paper_evaluations(self, paper: str, *, limit: int): + def get_paper_evaluations(self, paper: str, *, page: int, limit: int): assert paper == "1706.03762" + assert page == 1 assert limit == 5 return self.get_benchmark("imagenet-1k", limit=5, is_open=None) @@ -199,13 +200,17 @@ def list_benchmarks(self, **_kwargs): ], } - def get_benchmark(self, benchmark: str, *, limit: int, is_open: bool | None): + def get_benchmark( + self, benchmark: str, *, page: int = 1, limit: int, is_open: bool | None + ): assert benchmark == "imagenet-1k" + assert page in {1, 2} assert limit in {5, 10} assert is_open in {True, None} return { "benchmark": {"id": "72", "name": "ImageNet-1k", "slug": "imagenet-1k"}, - "count": 2, + "count": 4, + "next_page": 2 if page == 1 else None, "results": [ { "id": "10", @@ -217,6 +222,12 @@ def get_benchmark(self, benchmark: str, *, limit: int, is_open: bool | None): "paper_arxiv_id": "1706.03762", "is_open": True, "num_parameters": 1000, + "task_id": "1", + "task_name": "Image Classification", + "task_slug": "image-classification", + "methodology": "5-shot evaluation on the validation split.", + "result_url": "/paper/1706.03762", + "updated_at": "2026-09-16T00:00:00Z", }, { "id": "11", @@ -228,6 +239,29 @@ def get_benchmark(self, benchmark: str, *, limit: int, is_open: bool | None): "paper_arxiv_id": "1706.03762", "is_open": True, "num_parameters": 1000, + "task_id": "2", + "task_name": "Visual Recognition", + "task_slug": "visual-recognition", + "methodology": "5-shot evaluation on the validation split.", + "result_url": "/paper/1706.03762", + "updated_at": "2026-09-16T00:00:00Z", + }, + { + "id": "12", + "model_name": "ExampleNet", + "metrics": {"Accuracy": "90.1"}, + "best_rank": 1, + "paper_id": "755", + "paper_title": "Attention Is All You Need", + "paper_arxiv_id": "1706.03762", + "is_open": True, + "num_parameters": 1000, + "task_id": "1", + "task_name": "Image Classification", + "task_slug": "image-classification", + "methodology": "0-shot evaluation on the validation split.", + "result_url": "/paper/1706.03762", + "updated_at": "2026-09-16T00:00:00Z", }, ], } @@ -371,6 +405,7 @@ async def exercise(): "Accuracy": 90.1, "F1": 88, } + assert evaluations.structured_content["next_page"] == 2 assert tasks.structured_content["items"][0]["slug"] == "image-classification" assert tasks.structured_content["items"][0]["url"] == ( "https://paperswithcode.co/tasks/image-classification" @@ -431,6 +466,7 @@ async def exercise(): assert lineage.structured_content["successors"] == [ {"id": "900", "reference": "2001.00001", "title": "A Follow-up"} ] + assert lineage.structured_content["coverage"] == "explicit_catalog_links_only" def test_taxonomy_and_benchmark_tools_return_stable_catalog_entities(): @@ -473,7 +509,31 @@ async def exercise(): "Accuracy": 90.1, "F1": 88, } - assert len(benchmark.structured_content["evaluations"]) == 1 + assert len(benchmark.structured_content["evaluations"]) == 2 + assert benchmark.structured_content["next_page"] == 2 + assert benchmark.structured_content["metric_directions"] == { + "Accuracy": "higher", + "F1": "higher", + } + evaluation = benchmark.structured_content["evaluations"][0] + assert evaluation["shots"] == 5 + assert evaluation["source_url"] == "https://paperswithcode.co/paper/1706.03762" + assert evaluation["updated_at"] == "2026-09-16T00:00:00Z" + assert evaluation["rank_scopes"] == [ + { + "task_id": "1", + "task_name": "Image Classification", + "task_slug": "image-classification", + "rank": 1, + }, + { + "task_id": "2", + "task_name": "Visual Recognition", + "task_slug": "visual-recognition", + "rank": 2, + }, + ] + assert benchmark.structured_content["evaluations"][1]["shots"] == 0 def test_resources_expose_canonical_papers_tasks_and_benchmarks(): diff --git a/mcp_server/uv.lock b/mcp_server/uv.lock index d4fee0d..4c7d497 100644 --- a/mcp_server/uv.lock +++ b/mcp_server/uv.lock @@ -443,7 +443,7 @@ source = { editable = "../standalone_cli" } [[package]] name = "pwc-mcp" -version = "0.2.0" +version = "0.2.1" source = { editable = "." } dependencies = [ { name = "mcp" },