pi-codemcp 1.3.0 → 1.3.2

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.
package/package.json CHANGED
@@ -1,6 +1,6 @@
1
1
  {
2
2
  "name": "pi-codemcp",
3
- "version": "1.3.0",
3
+ "version": "1.3.2",
4
4
  "description": "Typed, sandboxed Code Mode access to configured MCP servers for Pi",
5
5
  "type": "module",
6
6
  "packageManager": "bun@1.3.10",
@@ -4,21 +4,34 @@ import hashlib
4
4
  import json
5
5
  import os
6
6
  import re
7
- from typing import TYPE_CHECKING, override
7
+ import time
8
+ from contextlib import suppress
9
+ from typing import TYPE_CHECKING, Any, override
8
10
  from urllib.parse import urlsplit
9
11
 
12
+ import httpx
10
13
  from fastmcp.client.auth import OAuth
14
+ from fastmcp.client.auth.oauth import TokenStorageAdapter
11
15
  from fastmcp.mcp_config import (
12
16
  MCPConfig,
13
17
  RemoteMCPServer,
14
18
  StdioMCPServer,
15
19
  infer_transport_type_from_url,
16
20
  )
21
+ from key_value.aio.adapters.pydantic import PydanticAdapter
17
22
  from key_value.aio.stores.filetree import FileTreeStore
18
23
  from key_value.aio.stores.filetree.store import (
19
24
  FileTreeV1CollectionSanitizationStrategy,
20
25
  FileTreeV1KeySanitizationStrategy,
21
26
  )
27
+ from mcp.client.auth.utils import (
28
+ build_oauth_authorization_server_metadata_discovery_urls,
29
+ build_protected_resource_metadata_discovery_urls,
30
+ create_oauth_metadata_request,
31
+ handle_auth_metadata_response,
32
+ handle_protected_resource_response,
33
+ )
34
+ from mcp.shared.auth import OAuthMetadata, ProtectedResourceMetadata
22
35
  from pydantic import BaseModel, ConfigDict
23
36
 
24
37
  from .json_types import JSON_VALUE_ADAPTER, JsonObject, JsonValue
@@ -27,7 +40,9 @@ from .models import NormalizedServerInfo, ServerAuth
27
40
  if TYPE_CHECKING:
28
41
  from pathlib import Path
29
42
 
30
- import httpx
43
+ from key_value.aio.protocols import AsyncKeyValue
44
+ from mcp.shared.auth import OAuthClientInformationFull
45
+ from pydantic import AnyUrl
31
46
 
32
47
  PI_ONLY_FIELDS = {"directTools", "lifecycle", "idleTimeout", "disabled", "enabled"}
33
48
  REMOTE_TRANSPORTS = {"http", "streamable-http", "sse"}
@@ -59,16 +74,165 @@ class NormalizedConfig(BaseModel):
59
74
  servers: list[NormalizedServerInfo]
60
75
 
61
76
 
77
+ class CodemcpTokenStorage(TokenStorageAdapter):
78
+ """Token storage that keeps OAuth state usable across sidecar restarts.
79
+
80
+ The upstream adapter evicts the client registration once the server-announced
81
+ client_secret_expires_at passes (Linear announces 24 hours), which silently
82
+ forces a full browser re-login after the next token expiry. A genuinely dead
83
+ secret still surfaces as invalid_client and re-registers, so persisting the
84
+ registration is strictly better. This adapter also persists the discovered
85
+ authorization-server metadata so token refresh hits the real token endpoint
86
+ instead of the SDK's "<origin>/token" fallback (a 404 for e.g. Outline).
87
+ """
88
+
89
+ def __init__(self, async_key_value: AsyncKeyValue, server_url: str) -> None:
90
+ super().__init__(async_key_value, server_url)
91
+ self._storage_oauth_metadata = PydanticAdapter[OAuthMetadata](
92
+ default_collection="mcp-oauth-metadata",
93
+ key_value=async_key_value,
94
+ pydantic_model=OAuthMetadata,
95
+ raise_on_validation_error=True,
96
+ )
97
+ self._storage_protected_resource = PydanticAdapter[ProtectedResourceMetadata](
98
+ default_collection="mcp-oauth-protected-resource",
99
+ key_value=async_key_value,
100
+ pydantic_model=ProtectedResourceMetadata,
101
+ raise_on_validation_error=True,
102
+ )
103
+
104
+ @override
105
+ async def set_client_info(self, client_info: OAuthClientInformationFull) -> None:
106
+ await self._storage_client_info.put(
107
+ key=self._get_client_info_cache_key(),
108
+ value=client_info,
109
+ )
110
+
111
+ @override
112
+ async def clear(self) -> None:
113
+ await super().clear()
114
+ await self._storage_oauth_metadata.delete(key=self._oauth_metadata_cache_key())
115
+ await self._storage_protected_resource.delete(key=self._protected_resource_cache_key())
116
+
117
+ async def get_oauth_metadata(self) -> OAuthMetadata | None:
118
+ result: OAuthMetadata | None = await self._storage_oauth_metadata.get(
119
+ key=self._oauth_metadata_cache_key()
120
+ )
121
+ return result
122
+
123
+ async def set_oauth_metadata(self, metadata: OAuthMetadata) -> None:
124
+ await self._storage_oauth_metadata.put(
125
+ key=self._oauth_metadata_cache_key(),
126
+ value=metadata,
127
+ )
128
+
129
+ async def get_protected_resource_metadata(self) -> ProtectedResourceMetadata | None:
130
+ result: ProtectedResourceMetadata | None = await self._storage_protected_resource.get(
131
+ key=self._protected_resource_cache_key()
132
+ )
133
+ return result
134
+
135
+ async def set_protected_resource_metadata(self, metadata: ProtectedResourceMetadata) -> None:
136
+ await self._storage_protected_resource.put(
137
+ key=self._protected_resource_cache_key(),
138
+ value=metadata,
139
+ )
140
+
141
+ def _oauth_metadata_cache_key(self) -> str:
142
+ return f"{self._server_url}/oauth_metadata"
143
+
144
+ def _protected_resource_cache_key(self) -> str:
145
+ return f"{self._server_url}/protected_resource"
146
+
147
+
62
148
  class PersistentCallbackOAuth(OAuth):
63
- """Reuse the callback registered with a persisted dynamic OAuth client."""
149
+ """OAuth provider hardened for long-lived shared file token storage.
150
+
151
+ On top of reusing the callback registered with a persisted dynamic client:
152
+ - discovered authorization-server metadata is persisted and restored so token
153
+ refresh works in fresh sidecar processes (the SDK otherwise falls back to
154
+ "<origin>/token", which 404s for servers like Outline and turns every
155
+ access-token expiry into a forced interactive re-login);
156
+ - a refresh response without refresh_token keeps the previous one (RFC 6749
157
+ section 6 allows omission when the refresh token does not rotate);
158
+ - a failed refresh adopts fresher tokens another sidecar process may have
159
+ stored instead of dropping straight into the browser flow.
160
+ """
161
+
162
+ def __init__(
163
+ self,
164
+ *,
165
+ mcp_url: str,
166
+ client_name: str,
167
+ token_storage: AsyncKeyValue,
168
+ additional_client_metadata: dict[str, Any] | None = None,
169
+ ) -> None:
170
+ self._codemcp_token_store = token_storage
171
+ self._persisted_oauth_metadata: OAuthMetadata | None = None
172
+ self._persisted_protected_resource: ProtectedResourceMetadata | None = None
173
+ super().__init__(
174
+ mcp_url=mcp_url,
175
+ client_name=client_name,
176
+ token_storage=token_storage,
177
+ additional_client_metadata=additional_client_metadata,
178
+ )
179
+
180
+ @override
181
+ def _bind(self, mcp_url: str) -> None:
182
+ super()._bind(mcp_url)
183
+ if isinstance(self.token_storage_adapter, CodemcpTokenStorage):
184
+ return
185
+ storage = CodemcpTokenStorage(self._codemcp_token_store, self.mcp_url)
186
+ self.token_storage_adapter = storage
187
+ self.context.storage = storage
64
188
 
65
189
  @override
66
190
  async def _initialize(self) -> None:
67
191
  await super()._initialize()
68
192
  client_info = self.context.client_info
69
- if client_info is None or not client_info.redirect_uris:
193
+ if client_info is not None and client_info.redirect_uris:
194
+ self._reuse_registered_callback(client_info.redirect_uris[0])
195
+ storage = self.token_storage_adapter
196
+ if not isinstance(storage, CodemcpTokenStorage):
70
197
  return
71
- redirect_uri = client_info.redirect_uris[0]
198
+ if self.context.oauth_metadata is None:
199
+ self.context.oauth_metadata = await storage.get_oauth_metadata()
200
+ self._persisted_oauth_metadata = self.context.oauth_metadata
201
+ if self.context.protected_resource_metadata is None:
202
+ self.context.protected_resource_metadata = (
203
+ await storage.get_protected_resource_metadata()
204
+ )
205
+ self._persisted_protected_resource = self.context.protected_resource_metadata
206
+ tokens = self.context.current_tokens
207
+ if tokens is not None and tokens.expires_in and await storage.get_token_expiry() is None:
208
+ # Without the absolute expiry record the upstream fallback re-applies the
209
+ # stale relative expires_in from now; the expired access token then looks
210
+ # valid, gets rejected with a 401, and the flow goes interactive.
211
+ self.context.token_expiry_time = time.time() - 1
212
+
213
+ @override
214
+ async def _refresh_token(self) -> httpx.Request:
215
+ if self.context.oauth_metadata is None:
216
+ await self._discover_server_metadata()
217
+ return await super()._refresh_token()
218
+
219
+ @override
220
+ async def _handle_token_response(self, response: httpx.Response) -> None:
221
+ await super()._handle_token_response(response)
222
+ await self._persist_discovered_metadata()
223
+
224
+ @override
225
+ async def _handle_refresh_response(self, response: httpx.Response) -> bool:
226
+ previous_tokens = self.context.current_tokens
227
+ previous_access_token = previous_tokens.access_token if previous_tokens else None
228
+ previous_refresh_token = previous_tokens.refresh_token if previous_tokens else None
229
+ if await super()._handle_refresh_response(response):
230
+ await self._restore_unrotated_refresh_token(previous_refresh_token)
231
+ await self._persist_discovered_metadata()
232
+ return True
233
+ return await self._adopt_tokens_refreshed_elsewhere(previous_access_token)
234
+
235
+ def _reuse_registered_callback(self, redirect_uri: AnyUrl) -> None:
72
236
  parsed = urlsplit(str(redirect_uri))
73
237
  if (
74
238
  parsed.scheme != "http"
@@ -82,6 +246,90 @@ class PersistentCallbackOAuth(OAuth):
82
246
  self._callback_host = parsed.hostname
83
247
  self.context.client_metadata.redirect_uris = [redirect_uri]
84
248
 
249
+ async def _discover_server_metadata(self) -> None:
250
+ """Best-effort OAuth discovery so refresh uses the real token endpoint."""
251
+ with suppress(httpx.HTTPError, ValueError):
252
+ async with self.httpx_client_factory() as client:
253
+ await self._discover_protected_resource(client)
254
+ await self._discover_authorization_server(client)
255
+ await self._persist_discovered_metadata()
256
+
257
+ async def _discover_protected_resource(self, client: httpx.AsyncClient) -> None:
258
+ if self.context.protected_resource_metadata is not None:
259
+ return
260
+ for url in build_protected_resource_metadata_discovery_urls(
261
+ None,
262
+ self.context.server_url,
263
+ ):
264
+ response = await client.send(create_oauth_metadata_request(url))
265
+ prm = await handle_protected_resource_response(response)
266
+ if prm is not None:
267
+ self.context.protected_resource_metadata = prm
268
+ self.context.auth_server_url = str(prm.authorization_servers[0])
269
+ return
270
+
271
+ async def _discover_authorization_server(self, client: httpx.AsyncClient) -> None:
272
+ if self.context.oauth_metadata is not None:
273
+ return
274
+ for url in build_oauth_authorization_server_metadata_discovery_urls(
275
+ self.context.auth_server_url,
276
+ self.context.server_url,
277
+ ):
278
+ response = await client.send(create_oauth_metadata_request(url))
279
+ ok, metadata = await handle_auth_metadata_response(response)
280
+ if not ok:
281
+ return
282
+ if metadata is not None:
283
+ self.context.oauth_metadata = metadata
284
+ return
285
+
286
+ async def _persist_discovered_metadata(self) -> None:
287
+ storage = self.token_storage_adapter
288
+ if not isinstance(storage, CodemcpTokenStorage):
289
+ return
290
+ metadata = self.context.oauth_metadata
291
+ if metadata is not None and metadata != self._persisted_oauth_metadata:
292
+ await storage.set_oauth_metadata(metadata)
293
+ self._persisted_oauth_metadata = metadata
294
+ resource = self.context.protected_resource_metadata
295
+ if resource is not None and resource != self._persisted_protected_resource:
296
+ await storage.set_protected_resource_metadata(resource)
297
+ self._persisted_protected_resource = resource
298
+
299
+ async def _restore_unrotated_refresh_token(self, previous_refresh_token: str | None) -> None:
300
+ tokens = self.context.current_tokens
301
+ if tokens is None or tokens.refresh_token is not None or previous_refresh_token is None:
302
+ return
303
+ # RFC 6749 section 6: the server may omit refresh_token when it does not
304
+ # rotate; the SDK overwrites the stored token set and would lose it.
305
+ tokens.refresh_token = previous_refresh_token
306
+ await self.context.storage.set_tokens(tokens)
307
+
308
+ async def _adopt_tokens_refreshed_elsewhere(self, previous_access_token: str | None) -> bool:
309
+ storage = self.token_storage_adapter
310
+ if not isinstance(storage, CodemcpTokenStorage):
311
+ return False
312
+ stored = await storage.get_tokens()
313
+ if (
314
+ stored is None
315
+ or not stored.access_token
316
+ or stored.access_token == previous_access_token
317
+ ):
318
+ return False
319
+ expiry = await storage.get_token_expiry()
320
+ if expiry is not None and time.time() > expiry:
321
+ return False
322
+ # Another sidecar process rotated the refresh token first and stored the
323
+ # result; adopt it instead of dropping into the interactive flow.
324
+ self.context.current_tokens = stored
325
+ if expiry is not None:
326
+ self.context.token_expiry_time = expiry
327
+ elif stored.expires_in is not None:
328
+ self.context.token_expiry_time = time.time() + stored.expires_in
329
+ else:
330
+ self.context.token_expiry_time = None
331
+ return True
332
+
85
333
 
86
334
  def load_mcp_json(path: Path) -> JsonObject:
87
335
  if not path.exists():
@@ -1047,42 +1047,67 @@ def _resolve_ref(ref: str, root: JsonSchema) -> JsonObject:
1047
1047
  def _merge_all_of(
1048
1048
  schema: JsonObject,
1049
1049
  root: JsonSchema,
1050
- ) -> JsonObject | None:
1051
- merged_properties: JsonObject = {}
1052
- merged_required: list[JsonValue] = []
1053
- additional_properties: JsonValue = True
1054
- has_additional_properties = False
1050
+ ) -> JsonSchema | None:
1055
1051
  raw_members = schema.get("allOf")
1056
1052
  if not isinstance(raw_members, list):
1057
1053
  return None
1054
+ # Sibling constraints (type, properties, ...) apply on top of allOf members.
1055
+ parts: list[JsonObject] = [{key: value for key, value in schema.items() if key != "allOf"}]
1058
1056
  for member in raw_members:
1059
1057
  if not isinstance(member, dict):
1060
1058
  return None
1061
1059
  ref = member.get("$ref")
1062
- resolved = _resolve_ref(ref, root) if isinstance(ref, str) else member
1063
- if resolved.get("type") not in {None, "object"}:
1064
- return None
1065
- raw_properties = resolved.get("properties")
1060
+ parts.append(_resolve_ref(ref, root) if isinstance(ref, str) else member)
1061
+
1062
+ declared_types: set[str] = set()
1063
+ for part in parts:
1064
+ part_type = part.get("type")
1065
+ if isinstance(part_type, str):
1066
+ declared_types.add(part_type)
1067
+ if len(declared_types) > 1:
1068
+ return None
1069
+ merged_type = declared_types.pop() if declared_types else None
1070
+
1071
+ if merged_type == "object" or (
1072
+ merged_type is None and any("properties" in part for part in parts)
1073
+ ):
1074
+ return _merge_object_parts(parts)
1075
+
1076
+ if merged_type is None:
1077
+ # Annotation-only allOf (descriptions, patterns, ...): no type information.
1078
+ return True
1079
+ merged: JsonObject = {"type": merged_type}
1080
+ for part in parts:
1081
+ if "items" in part and "items" not in merged:
1082
+ merged["items"] = part["items"]
1083
+ return merged
1084
+
1085
+
1086
+ def _merge_object_parts(parts: list[JsonObject]) -> JsonObject:
1087
+ merged_properties: JsonObject = {}
1088
+ merged_required: list[JsonValue] = []
1089
+ additional_properties: JsonValue = True
1090
+ has_additional_properties = False
1091
+ for part in parts:
1092
+ raw_properties = part.get("properties")
1066
1093
  if isinstance(raw_properties, dict):
1067
1094
  merged_properties.update(raw_properties)
1068
- raw_required = resolved.get("required")
1095
+ raw_required = part.get("required")
1069
1096
  if isinstance(raw_required, list):
1070
1097
  for required in raw_required:
1071
1098
  if isinstance(required, str) and required not in merged_required:
1072
1099
  merged_required.append(required)
1073
- if "additionalProperties" in resolved:
1074
- additional_properties = JSON_VALUE_ADAPTER.validate_python(
1075
- resolved["additionalProperties"]
1076
- )
1100
+ if "additionalProperties" in part:
1101
+ additional_properties = JSON_VALUE_ADAPTER.validate_python(part["additionalProperties"])
1077
1102
  has_additional_properties = True
1078
- merged: JsonObject = {
1103
+ merged_object: JsonObject = {
1079
1104
  "type": "object",
1080
1105
  "properties": merged_properties,
1081
1106
  "required": merged_required,
1082
1107
  }
1083
1108
  if has_additional_properties:
1084
- merged["additionalProperties"] = additional_properties
1085
- return merged
1109
+ merged_object["additionalProperties"] = additional_properties
1110
+ return merged_object
1086
1111
 
1087
1112
 
1088
1113
  def _dedupe(blocks: Iterable[str]) -> list[str]: