easy-mcp-kit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- easy_mcp/__init__.py +62 -0
- easy_mcp/decorators.py +175 -0
- easy_mcp/exceptions.py +112 -0
- easy_mcp/logging.py +69 -0
- easy_mcp/py.typed +0 -0
- easy_mcp/schema.py +246 -0
- easy_mcp/security/__init__.py +13 -0
- easy_mcp/security/auth.py +156 -0
- easy_mcp/security/ratelimit.py +67 -0
- easy_mcp/server.py +520 -0
- easy_mcp/transport/__init__.py +7 -0
- easy_mcp/transport/base.py +56 -0
- easy_mcp/transport/sse.py +247 -0
- easy_mcp_kit-0.1.0.dist-info/METADATA +276 -0
- easy_mcp_kit-0.1.0.dist-info/RECORD +17 -0
- easy_mcp_kit-0.1.0.dist-info/WHEEL +4 -0
- easy_mcp_kit-0.1.0.dist-info/licenses/LICENSE +21 -0
easy_mcp/__init__.py
ADDED
|
@@ -0,0 +1,62 @@
|
|
|
1
|
+
"""easy_mcp — build secure MCP (Model Context Protocol) servers from plain
|
|
2
|
+
Python functions.
|
|
3
|
+
|
|
4
|
+
Quickstart::
|
|
5
|
+
|
|
6
|
+
from easy_mcp import MCPServer
|
|
7
|
+
|
|
8
|
+
server = MCPServer(port=8000)
|
|
9
|
+
|
|
10
|
+
@server.tool
|
|
11
|
+
def add(a: int, b: int) -> int:
|
|
12
|
+
\"\"\"Add two numbers.\"\"\"
|
|
13
|
+
return a + b
|
|
14
|
+
|
|
15
|
+
server.run()
|
|
16
|
+
"""
|
|
17
|
+
|
|
18
|
+
from .decorators import ToolDefinition
|
|
19
|
+
from .exceptions import (
|
|
20
|
+
AuthenticationError,
|
|
21
|
+
AuthorizationError,
|
|
22
|
+
EasyMCPError,
|
|
23
|
+
PayloadTooLargeError,
|
|
24
|
+
ProtocolError,
|
|
25
|
+
RateLimitError,
|
|
26
|
+
SchemaError,
|
|
27
|
+
SessionLimitError,
|
|
28
|
+
ToolError,
|
|
29
|
+
ToolRegistrationError,
|
|
30
|
+
ValidationError,
|
|
31
|
+
)
|
|
32
|
+
from .security.auth import APIKeyAuth, ClientIdentity
|
|
33
|
+
from .security.ratelimit import SlidingWindowRateLimiter
|
|
34
|
+
from .server import PROTOCOL_VERSION, MCPServer
|
|
35
|
+
from .transport.base import ClientContext, Transport
|
|
36
|
+
from .transport.sse import SSETransport
|
|
37
|
+
|
|
38
|
+
__version__ = "0.1.0"
|
|
39
|
+
|
|
40
|
+
__all__ = [
|
|
41
|
+
"APIKeyAuth",
|
|
42
|
+
"AuthenticationError",
|
|
43
|
+
"AuthorizationError",
|
|
44
|
+
"ClientContext",
|
|
45
|
+
"ClientIdentity",
|
|
46
|
+
"EasyMCPError",
|
|
47
|
+
"MCPServer",
|
|
48
|
+
"PROTOCOL_VERSION",
|
|
49
|
+
"PayloadTooLargeError",
|
|
50
|
+
"ProtocolError",
|
|
51
|
+
"RateLimitError",
|
|
52
|
+
"SSETransport",
|
|
53
|
+
"SchemaError",
|
|
54
|
+
"SessionLimitError",
|
|
55
|
+
"SlidingWindowRateLimiter",
|
|
56
|
+
"ToolDefinition",
|
|
57
|
+
"ToolError",
|
|
58
|
+
"ToolRegistrationError",
|
|
59
|
+
"Transport",
|
|
60
|
+
"ValidationError",
|
|
61
|
+
"__version__",
|
|
62
|
+
]
|
easy_mcp/decorators.py
ADDED
|
@@ -0,0 +1,175 @@
|
|
|
1
|
+
"""Tool registration: the ``@server.tool`` decorator machinery and registry."""
|
|
2
|
+
|
|
3
|
+
from __future__ import annotations
|
|
4
|
+
|
|
5
|
+
import inspect
|
|
6
|
+
import logging
|
|
7
|
+
import re
|
|
8
|
+
import threading
|
|
9
|
+
from collections.abc import Callable, Iterable, Mapping
|
|
10
|
+
from dataclasses import dataclass
|
|
11
|
+
from typing import Any
|
|
12
|
+
|
|
13
|
+
from .exceptions import SchemaError, ToolRegistrationError
|
|
14
|
+
from .schema import build_input_schema, parse_docstring
|
|
15
|
+
|
|
16
|
+
logger = logging.getLogger("easy_mcp.registry")
|
|
17
|
+
|
|
18
|
+
_TOOL_NAME_RE = re.compile(r"^[A-Za-z][A-Za-z0-9_-]{0,63}$")
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
@dataclass(frozen=True, slots=True)
|
|
22
|
+
class ToolDefinition:
|
|
23
|
+
"""Everything the server knows about one registered tool."""
|
|
24
|
+
|
|
25
|
+
name: str
|
|
26
|
+
description: str
|
|
27
|
+
fn: Callable[..., Any]
|
|
28
|
+
input_schema: dict[str, Any]
|
|
29
|
+
is_async: bool
|
|
30
|
+
requires_auth: bool = False
|
|
31
|
+
scopes: frozenset[str] = frozenset()
|
|
32
|
+
tags: tuple[str, ...] = ()
|
|
33
|
+
category: str | None = None
|
|
34
|
+
examples: tuple[Mapping[str, Any], ...] = ()
|
|
35
|
+
timeout: float | None = None
|
|
36
|
+
max_calls_per_session: int | None = None
|
|
37
|
+
|
|
38
|
+
def to_mcp(self) -> dict[str, Any]:
|
|
39
|
+
"""Serialize this tool for a ``tools/list`` response."""
|
|
40
|
+
entry: dict[str, Any] = {
|
|
41
|
+
"name": self.name,
|
|
42
|
+
"description": self.description,
|
|
43
|
+
"inputSchema": self.input_schema,
|
|
44
|
+
}
|
|
45
|
+
meta: dict[str, Any] = {}
|
|
46
|
+
if self.tags:
|
|
47
|
+
meta["tags"] = list(self.tags)
|
|
48
|
+
if self.category:
|
|
49
|
+
meta["category"] = self.category
|
|
50
|
+
if self.examples:
|
|
51
|
+
meta["examples"] = [dict(example) for example in self.examples]
|
|
52
|
+
if meta:
|
|
53
|
+
# `_meta` is MCP's designated slot for implementation metadata.
|
|
54
|
+
entry["_meta"] = {"easy_mcp": meta}
|
|
55
|
+
return entry
|
|
56
|
+
|
|
57
|
+
|
|
58
|
+
def build_tool(
|
|
59
|
+
fn: Callable[..., Any],
|
|
60
|
+
*,
|
|
61
|
+
name: str | None = None,
|
|
62
|
+
description: str | None = None,
|
|
63
|
+
requires_auth: bool = False,
|
|
64
|
+
scopes: Iterable[str] = (),
|
|
65
|
+
tags: Iterable[str] = (),
|
|
66
|
+
category: str | None = None,
|
|
67
|
+
examples: Iterable[Mapping[str, Any]] = (),
|
|
68
|
+
timeout: float | None = None,
|
|
69
|
+
max_calls_per_session: int | None = None,
|
|
70
|
+
) -> ToolDefinition:
|
|
71
|
+
"""Introspect *fn* and produce a :class:`ToolDefinition`.
|
|
72
|
+
|
|
73
|
+
Args:
|
|
74
|
+
fn: The plain (or async) Python function to expose.
|
|
75
|
+
name: Override for the tool name (defaults to ``fn.__name__``).
|
|
76
|
+
description: Override for the description (defaults to the docstring
|
|
77
|
+
summary line).
|
|
78
|
+
requires_auth: Mark the tool as callable only by authenticated clients.
|
|
79
|
+
scopes: Scopes an API key must hold to call the tool. A non-empty
|
|
80
|
+
value implies ``requires_auth``.
|
|
81
|
+
tags: Free-form labels surfaced to clients in tool metadata.
|
|
82
|
+
category: Optional grouping label surfaced in tool metadata.
|
|
83
|
+
examples: Example invocations, e.g. ``({"arguments": {...}},)``.
|
|
84
|
+
timeout: Per-tool execution timeout in seconds (overrides the server
|
|
85
|
+
default).
|
|
86
|
+
max_calls_per_session: Cap on how often one session may call the tool.
|
|
87
|
+
|
|
88
|
+
Raises:
|
|
89
|
+
ToolRegistrationError: If the function cannot be exposed safely.
|
|
90
|
+
"""
|
|
91
|
+
if not callable(fn):
|
|
92
|
+
raise ToolRegistrationError(f"@tool target must be callable, got {type(fn).__name__}")
|
|
93
|
+
tool_name = name or getattr(fn, "__name__", "")
|
|
94
|
+
if not _TOOL_NAME_RE.match(tool_name or ""):
|
|
95
|
+
raise ToolRegistrationError(
|
|
96
|
+
f"invalid tool name {tool_name!r}: use 1-64 chars [A-Za-z0-9_-], "
|
|
97
|
+
"starting with a letter"
|
|
98
|
+
)
|
|
99
|
+
if timeout is not None and timeout <= 0:
|
|
100
|
+
raise ToolRegistrationError("timeout must be positive")
|
|
101
|
+
if max_calls_per_session is not None and max_calls_per_session < 1:
|
|
102
|
+
raise ToolRegistrationError("max_calls_per_session must be >= 1")
|
|
103
|
+
|
|
104
|
+
summary, param_docs = parse_docstring(inspect.getdoc(fn))
|
|
105
|
+
tool_description = (description or summary).strip()
|
|
106
|
+
if not tool_description:
|
|
107
|
+
# LLM clients pick tools by their descriptions; a tool without one is
|
|
108
|
+
# effectively invisible to them.
|
|
109
|
+
logger.warning(
|
|
110
|
+
"tool %r has no description; add a docstring or description=", tool_name
|
|
111
|
+
)
|
|
112
|
+
tool_description = tool_name
|
|
113
|
+
|
|
114
|
+
try:
|
|
115
|
+
input_schema = build_input_schema(fn, param_docs)
|
|
116
|
+
except SchemaError as exc:
|
|
117
|
+
raise ToolRegistrationError(f"cannot register tool {tool_name!r}: {exc}") from exc
|
|
118
|
+
|
|
119
|
+
scope_set = frozenset(scopes)
|
|
120
|
+
return ToolDefinition(
|
|
121
|
+
name=tool_name,
|
|
122
|
+
description=tool_description,
|
|
123
|
+
fn=fn,
|
|
124
|
+
input_schema=input_schema,
|
|
125
|
+
is_async=inspect.iscoroutinefunction(fn),
|
|
126
|
+
# A scope requirement implies the tool is protected.
|
|
127
|
+
requires_auth=bool(requires_auth or scope_set),
|
|
128
|
+
scopes=scope_set,
|
|
129
|
+
tags=tuple(tags),
|
|
130
|
+
category=category,
|
|
131
|
+
examples=tuple(dict(example) for example in examples),
|
|
132
|
+
timeout=timeout,
|
|
133
|
+
max_calls_per_session=max_calls_per_session,
|
|
134
|
+
)
|
|
135
|
+
|
|
136
|
+
|
|
137
|
+
class ToolRegistry:
|
|
138
|
+
"""Thread-safe, deterministic registry of tool definitions."""
|
|
139
|
+
|
|
140
|
+
def __init__(self) -> None:
|
|
141
|
+
self._tools: dict[str, ToolDefinition] = {}
|
|
142
|
+
self._lock = threading.Lock()
|
|
143
|
+
|
|
144
|
+
def register(self, tool: ToolDefinition, *, replace: bool = False) -> None:
|
|
145
|
+
"""Add a tool; refuses silent overwrites unless ``replace=True``."""
|
|
146
|
+
with self._lock:
|
|
147
|
+
if tool.name in self._tools and not replace:
|
|
148
|
+
raise ToolRegistrationError(f"a tool named {tool.name!r} is already registered")
|
|
149
|
+
self._tools[tool.name] = tool
|
|
150
|
+
|
|
151
|
+
def unregister(self, name: str) -> ToolDefinition:
|
|
152
|
+
"""Remove and return a tool by name."""
|
|
153
|
+
with self._lock:
|
|
154
|
+
try:
|
|
155
|
+
return self._tools.pop(name)
|
|
156
|
+
except KeyError:
|
|
157
|
+
raise ToolRegistrationError(f"no tool named {name!r} is registered") from None
|
|
158
|
+
|
|
159
|
+
def get(self, name: str) -> ToolDefinition | None:
|
|
160
|
+
"""Look up a tool by name, or ``None``."""
|
|
161
|
+
with self._lock:
|
|
162
|
+
return self._tools.get(name)
|
|
163
|
+
|
|
164
|
+
def list(self) -> list[ToolDefinition]:
|
|
165
|
+
"""Tools sorted by name, so ``tools/list`` output is reproducible."""
|
|
166
|
+
with self._lock:
|
|
167
|
+
return sorted(self._tools.values(), key=lambda tool: tool.name)
|
|
168
|
+
|
|
169
|
+
def __contains__(self, name: object) -> bool:
|
|
170
|
+
with self._lock:
|
|
171
|
+
return name in self._tools
|
|
172
|
+
|
|
173
|
+
def __len__(self) -> int:
|
|
174
|
+
with self._lock:
|
|
175
|
+
return len(self._tools)
|
easy_mcp/exceptions.py
ADDED
|
@@ -0,0 +1,112 @@
|
|
|
1
|
+
"""Exception hierarchy and JSON-RPC error codes for easy_mcp.
|
|
2
|
+
|
|
3
|
+
Two kinds of errors exist:
|
|
4
|
+
|
|
5
|
+
* :class:`ProtocolError` subclasses map directly onto JSON-RPC error
|
|
6
|
+
responses. Their messages are written to be safe to send to clients.
|
|
7
|
+
* Everything else is an *internal* error. Outside debug mode the server
|
|
8
|
+
never forwards its message or traceback to a client; it logs the full
|
|
9
|
+
detail server-side under a unique ``error_id`` instead.
|
|
10
|
+
"""
|
|
11
|
+
|
|
12
|
+
from __future__ import annotations
|
|
13
|
+
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
# --- Standard JSON-RPC 2.0 error codes --------------------------------------
|
|
17
|
+
PARSE_ERROR = -32700
|
|
18
|
+
INVALID_REQUEST = -32600
|
|
19
|
+
METHOD_NOT_FOUND = -32601
|
|
20
|
+
INVALID_PARAMS = -32602
|
|
21
|
+
INTERNAL_ERROR = -32603
|
|
22
|
+
|
|
23
|
+
# --- easy_mcp error codes (JSON-RPC reserves -32000..-32099 for servers) -----
|
|
24
|
+
AUTHENTICATION_REQUIRED = -32001
|
|
25
|
+
FORBIDDEN = -32002
|
|
26
|
+
RATE_LIMITED = -32003
|
|
27
|
+
PAYLOAD_TOO_LARGE = -32004
|
|
28
|
+
TOOL_TIMEOUT = -32005
|
|
29
|
+
SESSION_LIMIT_EXCEEDED = -32006
|
|
30
|
+
|
|
31
|
+
|
|
32
|
+
class EasyMCPError(Exception):
|
|
33
|
+
"""Base class for every easy_mcp exception."""
|
|
34
|
+
|
|
35
|
+
|
|
36
|
+
class ToolRegistrationError(EasyMCPError):
|
|
37
|
+
"""A function could not be registered as a tool."""
|
|
38
|
+
|
|
39
|
+
|
|
40
|
+
class SchemaError(ToolRegistrationError):
|
|
41
|
+
"""A type annotation could not be converted to JSON Schema."""
|
|
42
|
+
|
|
43
|
+
|
|
44
|
+
class ToolError(EasyMCPError):
|
|
45
|
+
"""Raised *inside a tool* to return an intentional, safe error message.
|
|
46
|
+
|
|
47
|
+
Unlike arbitrary exceptions (which are sanitized down to an opaque
|
|
48
|
+
``error_id``), the message of a ``ToolError`` is sent to the client
|
|
49
|
+
verbatim. Only raise it with text you would show an end user.
|
|
50
|
+
"""
|
|
51
|
+
|
|
52
|
+
|
|
53
|
+
class ProtocolError(EasyMCPError):
|
|
54
|
+
"""An error with a JSON-RPC error code, safe to serialize to clients."""
|
|
55
|
+
|
|
56
|
+
code: int = INVALID_REQUEST
|
|
57
|
+
|
|
58
|
+
def __init__(self, message: str, *, code: int | None = None, data: Any = None) -> None:
|
|
59
|
+
super().__init__(message)
|
|
60
|
+
if code is not None:
|
|
61
|
+
self.code = code
|
|
62
|
+
self.data = data
|
|
63
|
+
|
|
64
|
+
|
|
65
|
+
class ValidationError(ProtocolError):
|
|
66
|
+
"""Tool arguments failed schema validation."""
|
|
67
|
+
|
|
68
|
+
code = INVALID_PARAMS
|
|
69
|
+
|
|
70
|
+
def __init__(self, errors: list[str]) -> None:
|
|
71
|
+
self.errors = list(errors)
|
|
72
|
+
super().__init__(
|
|
73
|
+
"Invalid tool arguments: " + "; ".join(self.errors),
|
|
74
|
+
data={"errors": self.errors},
|
|
75
|
+
)
|
|
76
|
+
|
|
77
|
+
|
|
78
|
+
class AuthenticationError(ProtocolError):
|
|
79
|
+
"""The request needs a valid API key."""
|
|
80
|
+
|
|
81
|
+
code = AUTHENTICATION_REQUIRED
|
|
82
|
+
|
|
83
|
+
|
|
84
|
+
class AuthorizationError(ProtocolError):
|
|
85
|
+
"""The authenticated client lacks a scope the tool requires."""
|
|
86
|
+
|
|
87
|
+
code = FORBIDDEN
|
|
88
|
+
|
|
89
|
+
|
|
90
|
+
class RateLimitError(ProtocolError):
|
|
91
|
+
"""The client exceeded its request budget."""
|
|
92
|
+
|
|
93
|
+
code = RATE_LIMITED
|
|
94
|
+
|
|
95
|
+
def __init__(self, retry_after_seconds: float) -> None:
|
|
96
|
+
self.retry_after_seconds = retry_after_seconds
|
|
97
|
+
super().__init__(
|
|
98
|
+
f"Rate limit exceeded; retry in {retry_after_seconds:.1f}s",
|
|
99
|
+
data={"retry_after_seconds": round(retry_after_seconds, 3)},
|
|
100
|
+
)
|
|
101
|
+
|
|
102
|
+
|
|
103
|
+
class PayloadTooLargeError(ProtocolError):
|
|
104
|
+
"""The request body exceeded ``max_request_bytes``."""
|
|
105
|
+
|
|
106
|
+
code = PAYLOAD_TOO_LARGE
|
|
107
|
+
|
|
108
|
+
|
|
109
|
+
class SessionLimitError(ProtocolError):
|
|
110
|
+
"""A per-session tool usage limit was reached."""
|
|
111
|
+
|
|
112
|
+
code = SESSION_LIMIT_EXCEEDED
|
easy_mcp/logging.py
ADDED
|
@@ -0,0 +1,69 @@
|
|
|
1
|
+
"""Structured JSON logging and the security audit trail.
|
|
2
|
+
|
|
3
|
+
All server logs go through the ``easy_mcp`` logger; audit events go through
|
|
4
|
+
``easy_mcp.audit``. Audit callers must never pass secrets — pass API-key
|
|
5
|
+
*fingerprints* (see :func:`easy_mcp.security.auth.fingerprint`), never keys.
|
|
6
|
+
"""
|
|
7
|
+
|
|
8
|
+
from __future__ import annotations
|
|
9
|
+
|
|
10
|
+
import json
|
|
11
|
+
import logging
|
|
12
|
+
import sys
|
|
13
|
+
from datetime import UTC, datetime
|
|
14
|
+
from typing import Any
|
|
15
|
+
|
|
16
|
+
LOGGER_NAME = "easy_mcp"
|
|
17
|
+
AUDIT_LOGGER_NAME = "easy_mcp.audit"
|
|
18
|
+
|
|
19
|
+
|
|
20
|
+
class JSONLogFormatter(logging.Formatter):
|
|
21
|
+
"""Format log records as single-line JSON objects."""
|
|
22
|
+
|
|
23
|
+
def format(self, record: logging.LogRecord) -> str:
|
|
24
|
+
payload: dict[str, Any] = {
|
|
25
|
+
"timestamp": datetime.now(UTC)
|
|
26
|
+
.isoformat(timespec="milliseconds")
|
|
27
|
+
.replace("+00:00", "Z"),
|
|
28
|
+
"level": record.levelname,
|
|
29
|
+
"logger": record.name,
|
|
30
|
+
"message": record.getMessage(),
|
|
31
|
+
}
|
|
32
|
+
event = getattr(record, "event", None)
|
|
33
|
+
if isinstance(event, dict):
|
|
34
|
+
payload["event"] = event
|
|
35
|
+
if record.exc_info:
|
|
36
|
+
# Tracebacks are for server-side logs only; the dispatcher never
|
|
37
|
+
# sends them to clients outside debug mode.
|
|
38
|
+
payload["exception"] = self.formatException(record.exc_info)
|
|
39
|
+
return json.dumps(payload, ensure_ascii=False, default=str)
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
def configure_logging(*, debug: bool = False, json_logs: bool = True) -> logging.Logger:
|
|
43
|
+
"""Attach a stderr handler to the ``easy_mcp`` logger (idempotent)."""
|
|
44
|
+
logger = logging.getLogger(LOGGER_NAME)
|
|
45
|
+
logger.setLevel(logging.DEBUG if debug else logging.INFO)
|
|
46
|
+
if not any(getattr(handler, "_easy_mcp", False) for handler in logger.handlers):
|
|
47
|
+
handler = logging.StreamHandler(sys.stderr)
|
|
48
|
+
handler._easy_mcp = True # type: ignore[attr-defined]
|
|
49
|
+
if json_logs:
|
|
50
|
+
handler.setFormatter(JSONLogFormatter())
|
|
51
|
+
else:
|
|
52
|
+
handler.setFormatter(
|
|
53
|
+
logging.Formatter("%(asctime)s %(levelname)s %(name)s %(message)s")
|
|
54
|
+
)
|
|
55
|
+
logger.addHandler(handler)
|
|
56
|
+
logger.propagate = False
|
|
57
|
+
return logger
|
|
58
|
+
|
|
59
|
+
|
|
60
|
+
def audit(event_type: str, **fields: Any) -> None:
|
|
61
|
+
"""Emit a structured audit event.
|
|
62
|
+
|
|
63
|
+
Args:
|
|
64
|
+
event_type: Short machine-readable event name, e.g. ``tool_call``.
|
|
65
|
+
**fields: Event payload. Must not contain secrets.
|
|
66
|
+
"""
|
|
67
|
+
logging.getLogger(AUDIT_LOGGER_NAME).info(
|
|
68
|
+
event_type, extra={"event": {"type": event_type, **fields}}
|
|
69
|
+
)
|
easy_mcp/py.typed
ADDED
|
File without changes
|
easy_mcp/schema.py
ADDED
|
@@ -0,0 +1,246 @@
|
|
|
1
|
+
"""Type-hint → JSON Schema conversion, docstring parsing, and validation.
|
|
2
|
+
|
|
3
|
+
Only a deliberate subset of JSON Schema is generated and validated —
|
|
4
|
+
``type`` (with strict bool/int separation), ``items``, ``properties`` /
|
|
5
|
+
``required`` / ``additionalProperties``, ``enum`` and ``anyOf``. Keeping the
|
|
6
|
+
validator small and hand-written means there is no third-party dependency in
|
|
7
|
+
the request path and its behavior is easy to audit.
|
|
8
|
+
"""
|
|
9
|
+
|
|
10
|
+
from __future__ import annotations
|
|
11
|
+
|
|
12
|
+
import inspect
|
|
13
|
+
import json
|
|
14
|
+
import re
|
|
15
|
+
import types
|
|
16
|
+
import typing
|
|
17
|
+
from collections.abc import Callable
|
|
18
|
+
from typing import Any
|
|
19
|
+
|
|
20
|
+
from .exceptions import SchemaError, ValidationError
|
|
21
|
+
|
|
22
|
+
_SCALARS: dict[type, str] = {str: "string", int: "integer", float: "number", bool: "boolean"}
|
|
23
|
+
|
|
24
|
+
_SECTION_RE = re.compile(
|
|
25
|
+
r"^(args|arguments|parameters|returns?|raises?|yields?|examples?|notes?|attributes)\s*:\s*$",
|
|
26
|
+
re.IGNORECASE,
|
|
27
|
+
)
|
|
28
|
+
_PARAM_RE = re.compile(r"^\*{0,2}([A-Za-z_][A-Za-z0-9_]*)\s*(?:\([^)]*\))?\s*:\s*(.*)$")
|
|
29
|
+
|
|
30
|
+
|
|
31
|
+
def parse_docstring(doc: str | None) -> tuple[str, dict[str, str]]:
|
|
32
|
+
"""Split a docstring into ``(summary, {param_name: description})``.
|
|
33
|
+
|
|
34
|
+
Understands Google-style ``Args:`` sections. The summary is the first
|
|
35
|
+
paragraph joined onto a single line — LLM clients choose tools by it.
|
|
36
|
+
"""
|
|
37
|
+
if not doc:
|
|
38
|
+
return "", {}
|
|
39
|
+
lines = inspect.cleandoc(doc).splitlines()
|
|
40
|
+
|
|
41
|
+
summary_parts: list[str] = []
|
|
42
|
+
for line in lines:
|
|
43
|
+
stripped = line.strip()
|
|
44
|
+
if not stripped or _SECTION_RE.match(stripped):
|
|
45
|
+
break
|
|
46
|
+
summary_parts.append(stripped)
|
|
47
|
+
|
|
48
|
+
params: dict[str, str] = {}
|
|
49
|
+
in_args = False
|
|
50
|
+
current: str | None = None
|
|
51
|
+
for line in lines:
|
|
52
|
+
stripped = line.strip()
|
|
53
|
+
section = _SECTION_RE.match(stripped)
|
|
54
|
+
if section:
|
|
55
|
+
in_args = section.group(1).lower() in ("args", "arguments", "parameters")
|
|
56
|
+
current = None
|
|
57
|
+
continue
|
|
58
|
+
if not in_args:
|
|
59
|
+
continue
|
|
60
|
+
if not stripped:
|
|
61
|
+
current = None
|
|
62
|
+
continue
|
|
63
|
+
match = _PARAM_RE.match(stripped)
|
|
64
|
+
if match:
|
|
65
|
+
current = match.group(1)
|
|
66
|
+
params[current] = match.group(2).strip()
|
|
67
|
+
elif current is not None:
|
|
68
|
+
# Continuation line of the previous parameter's description.
|
|
69
|
+
params[current] = f"{params[current]} {stripped}".strip()
|
|
70
|
+
return " ".join(summary_parts), params
|
|
71
|
+
|
|
72
|
+
|
|
73
|
+
def annotation_to_schema(annotation: Any) -> dict[str, Any]:
|
|
74
|
+
"""Convert a Python type annotation into a JSON Schema fragment.
|
|
75
|
+
|
|
76
|
+
Supported: ``str``, ``int``, ``float``, ``bool``, ``list``/``list[T]``,
|
|
77
|
+
``dict``/``dict[str, T]``, ``Optional``/unions, ``Literal`` and ``Any``.
|
|
78
|
+
|
|
79
|
+
Raises:
|
|
80
|
+
SchemaError: For any annotation outside the supported set.
|
|
81
|
+
"""
|
|
82
|
+
if annotation is inspect.Parameter.empty:
|
|
83
|
+
raise SchemaError("parameter is missing a type annotation")
|
|
84
|
+
if annotation is None or annotation is type(None):
|
|
85
|
+
return {"type": "null"}
|
|
86
|
+
if annotation is Any:
|
|
87
|
+
return {}
|
|
88
|
+
origin = typing.get_origin(annotation)
|
|
89
|
+
if origin is None:
|
|
90
|
+
scalar = _SCALARS.get(annotation)
|
|
91
|
+
if scalar is not None:
|
|
92
|
+
return {"type": scalar}
|
|
93
|
+
if annotation is list:
|
|
94
|
+
return {"type": "array"}
|
|
95
|
+
if annotation is dict:
|
|
96
|
+
return {"type": "object"}
|
|
97
|
+
raise SchemaError(f"unsupported type annotation: {annotation!r}")
|
|
98
|
+
if origin is list:
|
|
99
|
+
args = typing.get_args(annotation)
|
|
100
|
+
if not args:
|
|
101
|
+
return {"type": "array"}
|
|
102
|
+
return {"type": "array", "items": annotation_to_schema(args[0])}
|
|
103
|
+
if origin is dict:
|
|
104
|
+
args = typing.get_args(annotation)
|
|
105
|
+
if not args:
|
|
106
|
+
return {"type": "object"}
|
|
107
|
+
key_type, value_type = args
|
|
108
|
+
if key_type is not str:
|
|
109
|
+
raise SchemaError("dict keys must be str for JSON compatibility")
|
|
110
|
+
return {"type": "object", "additionalProperties": annotation_to_schema(value_type)}
|
|
111
|
+
if origin is typing.Union or origin is types.UnionType:
|
|
112
|
+
return {"anyOf": [annotation_to_schema(arg) for arg in typing.get_args(annotation)]}
|
|
113
|
+
if origin is typing.Literal:
|
|
114
|
+
values = list(typing.get_args(annotation))
|
|
115
|
+
for value in values:
|
|
116
|
+
if not isinstance(value, (str, int, float, bool)):
|
|
117
|
+
raise SchemaError(f"Literal values must be JSON scalars, got {value!r}")
|
|
118
|
+
return {"enum": values}
|
|
119
|
+
raise SchemaError(f"unsupported type annotation: {annotation!r}")
|
|
120
|
+
|
|
121
|
+
|
|
122
|
+
def build_input_schema(
|
|
123
|
+
fn: Callable[..., Any], param_docs: dict[str, str] | None = None
|
|
124
|
+
) -> dict[str, Any]:
|
|
125
|
+
"""Build a strict JSON Schema object for *fn*'s parameters.
|
|
126
|
+
|
|
127
|
+
Every parameter must have a supported type annotation. ``*args`` /
|
|
128
|
+
``**kwargs`` and positional-only parameters are rejected because tool
|
|
129
|
+
arguments arrive as a JSON object and are passed by keyword.
|
|
130
|
+
"""
|
|
131
|
+
param_docs = param_docs or {}
|
|
132
|
+
signature = inspect.signature(fn)
|
|
133
|
+
try:
|
|
134
|
+
hints = typing.get_type_hints(fn)
|
|
135
|
+
except Exception as exc: # unresolvable forward references etc.
|
|
136
|
+
raise SchemaError(f"could not resolve type hints: {exc}") from exc
|
|
137
|
+
|
|
138
|
+
properties: dict[str, Any] = {}
|
|
139
|
+
required: list[str] = []
|
|
140
|
+
for name, param in signature.parameters.items():
|
|
141
|
+
if param.kind in (inspect.Parameter.VAR_POSITIONAL, inspect.Parameter.VAR_KEYWORD):
|
|
142
|
+
raise SchemaError("*args/**kwargs parameters are not supported")
|
|
143
|
+
if param.kind is inspect.Parameter.POSITIONAL_ONLY:
|
|
144
|
+
raise SchemaError("positional-only parameters are not supported")
|
|
145
|
+
annotation = hints.get(name, param.annotation)
|
|
146
|
+
try:
|
|
147
|
+
prop = annotation_to_schema(annotation)
|
|
148
|
+
except SchemaError as exc:
|
|
149
|
+
raise SchemaError(f"parameter '{name}': {exc}") from exc
|
|
150
|
+
if name in param_docs:
|
|
151
|
+
prop = {**prop, "description": param_docs[name]}
|
|
152
|
+
if param.default is inspect.Parameter.empty:
|
|
153
|
+
required.append(name)
|
|
154
|
+
else:
|
|
155
|
+
try:
|
|
156
|
+
json.dumps(param.default)
|
|
157
|
+
except (TypeError, ValueError):
|
|
158
|
+
pass # non-JSON default: still optional, just not advertised
|
|
159
|
+
else:
|
|
160
|
+
prop = {**prop, "default": param.default}
|
|
161
|
+
properties[name] = prop
|
|
162
|
+
|
|
163
|
+
# additionalProperties: false makes unknown fields a hard error — clients
|
|
164
|
+
# cannot smuggle unexpected arguments into a tool call.
|
|
165
|
+
return {
|
|
166
|
+
"type": "object",
|
|
167
|
+
"properties": properties,
|
|
168
|
+
"required": required,
|
|
169
|
+
"additionalProperties": False,
|
|
170
|
+
}
|
|
171
|
+
|
|
172
|
+
|
|
173
|
+
def validate_arguments(arguments: dict[str, Any], schema: dict[str, Any]) -> None:
|
|
174
|
+
"""Validate *arguments* against *schema*.
|
|
175
|
+
|
|
176
|
+
Raises:
|
|
177
|
+
ValidationError: Listing every violation found (not just the first).
|
|
178
|
+
"""
|
|
179
|
+
errors = _check(arguments, schema, "arguments")
|
|
180
|
+
if errors:
|
|
181
|
+
raise ValidationError(errors)
|
|
182
|
+
|
|
183
|
+
|
|
184
|
+
def _type_ok(value: Any, expected: str) -> bool:
|
|
185
|
+
# bool is a subclass of int in Python, but JSON treats them as distinct
|
|
186
|
+
# types — so booleans are rejected wherever numbers are expected.
|
|
187
|
+
if expected == "boolean":
|
|
188
|
+
return isinstance(value, bool)
|
|
189
|
+
if expected == "integer":
|
|
190
|
+
return isinstance(value, int) and not isinstance(value, bool)
|
|
191
|
+
if expected == "number":
|
|
192
|
+
return isinstance(value, (int, float)) and not isinstance(value, bool)
|
|
193
|
+
if expected == "string":
|
|
194
|
+
return isinstance(value, str)
|
|
195
|
+
if expected == "null":
|
|
196
|
+
return value is None
|
|
197
|
+
return False
|
|
198
|
+
|
|
199
|
+
|
|
200
|
+
def _check(value: Any, schema: dict[str, Any], path: str) -> list[str]:
|
|
201
|
+
if "enum" in schema:
|
|
202
|
+
allowed = schema["enum"]
|
|
203
|
+
if any(type(value) is type(option) and value == option for option in allowed):
|
|
204
|
+
return []
|
|
205
|
+
return [f"{path}: must be one of {allowed!r}"]
|
|
206
|
+
if "anyOf" in schema:
|
|
207
|
+
for option in schema["anyOf"]:
|
|
208
|
+
if not _check(value, option, path):
|
|
209
|
+
return []
|
|
210
|
+
return [f"{path}: does not match any allowed type"]
|
|
211
|
+
expected = schema.get("type")
|
|
212
|
+
if expected is None:
|
|
213
|
+
return [] # Any: accept everything
|
|
214
|
+
if expected == "array":
|
|
215
|
+
if not isinstance(value, list):
|
|
216
|
+
return [f"{path}: expected array, got {type(value).__name__}"]
|
|
217
|
+
items = schema.get("items")
|
|
218
|
+
if not items:
|
|
219
|
+
return []
|
|
220
|
+
errors: list[str] = []
|
|
221
|
+
for index, item in enumerate(value):
|
|
222
|
+
errors.extend(_check(item, items, f"{path}[{index}]"))
|
|
223
|
+
return errors
|
|
224
|
+
if expected == "object":
|
|
225
|
+
if not isinstance(value, dict):
|
|
226
|
+
return [f"{path}: expected object, got {type(value).__name__}"]
|
|
227
|
+
errors = []
|
|
228
|
+
properties = schema.get("properties", {})
|
|
229
|
+
additional = schema.get("additionalProperties", True)
|
|
230
|
+
for name in schema.get("required", []):
|
|
231
|
+
if name not in value:
|
|
232
|
+
errors.append(f"{path}.{name}: missing required parameter")
|
|
233
|
+
for key, item in value.items():
|
|
234
|
+
if not isinstance(key, str):
|
|
235
|
+
errors.append(f"{path}: object keys must be strings")
|
|
236
|
+
continue
|
|
237
|
+
if key in properties:
|
|
238
|
+
errors.extend(_check(item, properties[key], f"{path}.{key}"))
|
|
239
|
+
elif additional is False:
|
|
240
|
+
errors.append(f"{path}.{key}: unexpected parameter")
|
|
241
|
+
elif isinstance(additional, dict):
|
|
242
|
+
errors.extend(_check(item, additional, f"{path}.{key}"))
|
|
243
|
+
return errors
|
|
244
|
+
if not _type_ok(value, expected):
|
|
245
|
+
return [f"{path}: expected {expected}, got {type(value).__name__}"]
|
|
246
|
+
return []
|
|
@@ -0,0 +1,13 @@
|
|
|
1
|
+
"""Security primitives: authentication, authorization, and rate limiting."""
|
|
2
|
+
|
|
3
|
+
from .auth import APIKeyAuth, ClientIdentity, authorize, fingerprint, visible
|
|
4
|
+
from .ratelimit import SlidingWindowRateLimiter
|
|
5
|
+
|
|
6
|
+
__all__ = [
|
|
7
|
+
"APIKeyAuth",
|
|
8
|
+
"ClientIdentity",
|
|
9
|
+
"SlidingWindowRateLimiter",
|
|
10
|
+
"authorize",
|
|
11
|
+
"fingerprint",
|
|
12
|
+
"visible",
|
|
13
|
+
]
|