From 860aee48ee626f1e999ed0b387f06bf238bfe46c Mon Sep 17 00:00:00 2001 From: Johannes du Plessis Date: Tue, 23 Jun 2026 12:07:13 -0700 Subject: [PATCH] feat: add user-scoped Notion MCP OAuth (#1593) * feat: add user-scoped Notion MCP OAuth Co-authored-by: open-swe[bot] * fix: refresh Notion token per tool call Co-authored-by: open-swe[bot] * fix: normalize Notion wrapper response format Co-authored-by: open-swe[bot] --------- Co-authored-by: open-swe[bot] --- agent/dashboard/notion_oauth.py | 324 ++++++++++++++++++++++++++++ agent/dashboard/routes.py | 113 ++++++++++ agent/dashboard/user_credentials.py | 256 +++++++++++++++++++++- agent/integrations/notion_mcp.py | 101 +++++++++ agent/server.py | 8 + tests/test_notion_oauth.py | 111 ++++++++++ tests/test_observability_tools.py | 101 ++++++++- tests/test_user_credentials.py | 107 +++++++++ ui/src/lib/api.ts | 16 ++ ui/src/routes/my-settings.tsx | 84 +++++++- 10 files changed, 1212 insertions(+), 9 deletions(-) create mode 100644 agent/dashboard/notion_oauth.py create mode 100644 agent/integrations/notion_mcp.py create mode 100644 tests/test_notion_oauth.py diff --git a/agent/dashboard/notion_oauth.py b/agent/dashboard/notion_oauth.py new file mode 100644 index 00000000..80369d27 --- /dev/null +++ b/agent/dashboard/notion_oauth.py @@ -0,0 +1,324 @@ +"""Notion MCP OAuth helpers.""" + +from __future__ import annotations + +import base64 +import hashlib +import os +import secrets +from datetime import UTC, datetime +from typing import Any +from urllib.parse import urlencode, urlparse + +import httpx +from langgraph_sdk import get_client + +from ..encryption import decrypt_token, encrypt_token + +NOTION_MCP_URL = "https://mcp.notion.com/mcp" +NOTION_STATE_COOKIE_NAME = "osw_notion_oauth_state" +NOTION_OAUTH_FLOW_NAMESPACE: list[str] = ["notion_oauth_flows"] + +_NOTION_HOST = "mcp.notion.com" +_PROTECTED_RESOURCE_METADATA_URL = "https://mcp.notion.com/.well-known/oauth-protected-resource" +_AUTHORIZATION_SERVER_METADATA_PATH = "/.well-known/oauth-authorization-server" +_HTTP_TIMEOUT = httpx.Timeout(15.0, connect=5.0) + + +class NotionOAuthError(Exception): + """Notion OAuth endpoint error.""" + + def __init__(self, status_code: int, detail: str, *, error_code: str | None = None) -> None: + super().__init__(detail) + self.status_code = status_code + self.detail = detail + self.error_code = error_code + + +def _client(): + return get_client() + + +def _is_notion_https_url(url: str) -> bool: + parsed = urlparse(url) + return parsed.scheme == "https" and parsed.hostname == _NOTION_HOST + + +def _require_notion_https_url(url: str, label: str) -> str: + if not _is_notion_https_url(url): + raise NotionOAuthError(502, f"invalid Notion OAuth {label}") + return url + + +def _metadata_url(auth_server_url: str) -> str: + parsed = urlparse(_require_notion_https_url(auth_server_url, "authorization server")) + return f"{parsed.scheme}://{parsed.netloc}{_AUTHORIZATION_SERVER_METADATA_PATH}" + + +def _base64url(data: bytes) -> str: + return base64.urlsafe_b64encode(data).decode().rstrip("=") + + +def generate_code_verifier() -> str: + return _base64url(secrets.token_bytes(32)) + + +def code_challenge_for_verifier(verifier: str) -> str: + return _base64url(hashlib.sha256(verifier.encode()).digest()) + + +def build_notion_authorize_url( + *, + authorization_endpoint: str, + client_id: str, + redirect_uri: str, + code_challenge: str, + state: str, +) -> str: + _require_notion_https_url(authorization_endpoint, "authorization endpoint") + params = { + "response_type": "code", + "client_id": client_id, + "redirect_uri": redirect_uri, + "state": state, + "code_challenge": code_challenge, + "code_challenge_method": "S256", + "prompt": "consent", + } + return f"{authorization_endpoint}?{urlencode(params)}" + + +async def discover_notion_oauth_metadata() -> dict[str, Any]: + """Discover Notion MCP OAuth endpoints.""" + try: + async with httpx.AsyncClient(timeout=_HTTP_TIMEOUT) as client: + protected_resource = await client.get(_PROTECTED_RESOURCE_METADATA_URL) + if not protected_resource.is_success: + raise _oauth_error_from_response( + protected_resource, + "Notion OAuth protected resource discovery failed", + ) + resource_metadata = protected_resource.json() + auth_servers = resource_metadata.get("authorization_servers") + if not isinstance(auth_servers, list) or not auth_servers: + raise NotionOAuthError( + 502, "Notion OAuth discovery returned no authorization server" + ) + auth_server = auth_servers[0] + if not isinstance(auth_server, str): + raise NotionOAuthError( + 502, + "Notion OAuth discovery returned invalid authorization server", + ) + + metadata_response = await client.get(_metadata_url(auth_server)) + if not metadata_response.is_success: + raise _oauth_error_from_response( + metadata_response, + "Notion OAuth authorization server discovery failed", + ) + metadata = metadata_response.json() + except httpx.RequestError as exc: + raise NotionOAuthError(503, "Notion OAuth discovery failed") from exc + + if not isinstance(metadata, dict): + raise NotionOAuthError(502, "Notion OAuth discovery returned invalid metadata") + for key in ("authorization_endpoint", "token_endpoint", "registration_endpoint"): + value = metadata.get(key) + if not isinstance(value, str) or not value: + raise NotionOAuthError(502, f"Notion OAuth discovery missing {key}") + _require_notion_https_url(value, key) + return metadata + + +async def register_notion_oauth_client( + metadata: dict[str, Any], + *, + redirect_uri: str, +) -> dict[str, Any]: + """Register this deployment as a Notion MCP OAuth client.""" + registration_endpoint = metadata.get("registration_endpoint") + if not isinstance(registration_endpoint, str): + raise NotionOAuthError(502, "Notion OAuth metadata missing registration endpoint") + _require_notion_https_url(registration_endpoint, "registration endpoint") + body: dict[str, Any] = { + "client_name": os.environ.get("NOTION_MCP_CLIENT_NAME", "Open SWE"), + "redirect_uris": [redirect_uri], + "grant_types": ["authorization_code", "refresh_token"], + "response_types": ["code"], + "token_endpoint_auth_method": "none", + } + client_uri = os.environ.get("DASHBOARD_BASE_URL", "").strip() + if client_uri: + body["client_uri"] = client_uri + + try: + async with httpx.AsyncClient(timeout=_HTTP_TIMEOUT) as client: + response = await client.post( + registration_endpoint, + headers={"Accept": "application/json", "Content-Type": "application/json"}, + json=body, + ) + except httpx.RequestError as exc: + raise NotionOAuthError(503, "Notion OAuth client registration failed") from exc + if not response.is_success: + raise _oauth_error_from_response(response, "Notion OAuth client registration failed") + data = response.json() + if not isinstance(data, dict) or not isinstance(data.get("client_id"), str): + raise NotionOAuthError(502, "Notion OAuth client registration missing client_id") + return data + + +async def store_notion_oauth_flow( + login: str, + nonce_hash: str, + *, + redirect_uri: str, + state: str, +) -> str: + """Create and store a short-lived Notion OAuth flow.""" + metadata = await discover_notion_oauth_metadata() + client_info = await register_notion_oauth_client(metadata, redirect_uri=redirect_uri) + verifier = generate_code_verifier() + authorization_endpoint = str(metadata["authorization_endpoint"]) + client_id = str(client_info["client_id"]) + client_secret = ( + client_info.get("client_secret") + if isinstance(client_info.get("client_secret"), str) + else None + ) + value = { + "login": login, + "encrypted_code_verifier": encrypt_token(verifier), + "client_id": client_id, + "encrypted_client_secret": encrypt_token(client_secret) if client_secret else None, + "token_endpoint": str(metadata["token_endpoint"]), + "redirect_uri": redirect_uri, + "created_at": datetime.now(UTC).isoformat(), + } + await _client().store.put_item([*NOTION_OAUTH_FLOW_NAMESPACE, login], nonce_hash, value) + return build_notion_authorize_url( + authorization_endpoint=authorization_endpoint, + client_id=client_id, + redirect_uri=redirect_uri, + code_challenge=code_challenge_for_verifier(verifier), + state=state, + ) + + +async def pop_notion_oauth_flow(login: str, nonce_hash: str) -> dict[str, Any] | None: + """Read and delete a pending Notion OAuth flow.""" + namespace = [*NOTION_OAUTH_FLOW_NAMESPACE, login] + try: + item = await _client().store.get_item(namespace, nonce_hash) + except httpx.HTTPStatusError as exc: + if exc.response.status_code == 404: + return None + raise + try: + await _client().store.delete_item(namespace, nonce_hash) + except httpx.HTTPStatusError as exc: + if exc.response.status_code != 404: + raise + if item is None: + return None + value = item.get("value") if isinstance(item, dict) else getattr(item, "value", None) + if not isinstance(value, dict): + return None + encrypted_code_verifier = value.pop("encrypted_code_verifier", "") + if encrypted_code_verifier: + value["code_verifier"] = decrypt_token(encrypted_code_verifier) + encrypted_client_secret = value.pop("encrypted_client_secret", "") + if encrypted_client_secret: + value["client_secret"] = decrypt_token(encrypted_client_secret) or None + return value + + +def _oauth_error_from_response(response: httpx.Response, fallback: str) -> NotionOAuthError: + error_code = None + detail = fallback + try: + data = response.json() + except ValueError: + data = None + if isinstance(data, dict): + raw_error = data.get("error") + if isinstance(raw_error, str): + error_code = raw_error + raw_description = data.get("error_description") + description = raw_description if isinstance(raw_description, str) else raw_error + detail = f"{fallback}: {description}" + elif response.text: + detail = f"{fallback}: {response.text[:200]}" + return NotionOAuthError(response.status_code, detail, error_code=error_code) + + +async def exchange_notion_code(code: str, flow: dict[str, Any]) -> dict[str, Any]: + """Exchange a Notion OAuth code for tokens.""" + token_endpoint = flow.get("token_endpoint") + client_id = flow.get("client_id") + redirect_uri = flow.get("redirect_uri") + code_verifier = flow.get("code_verifier") + if not all( + isinstance(v, str) and v for v in (token_endpoint, client_id, redirect_uri, code_verifier) + ): + raise NotionOAuthError(400, "stored Notion OAuth flow is incomplete") + _require_notion_https_url(token_endpoint, "token endpoint") + body = { + "grant_type": "authorization_code", + "code": code, + "client_id": client_id, + "redirect_uri": redirect_uri, + "code_verifier": code_verifier, + } + client_secret = flow.get("client_secret") + if isinstance(client_secret, str) and client_secret: + body["client_secret"] = client_secret + return await _request_token(token_endpoint, body, fallback="Notion OAuth token exchange failed") + + +async def refresh_notion_access_token( + *, + refresh_token: str, + token_endpoint: str, + client_id: str, + client_secret: str | None = None, +) -> dict[str, Any]: + """Refresh a Notion OAuth access token.""" + _require_notion_https_url(token_endpoint, "token endpoint") + body = { + "grant_type": "refresh_token", + "refresh_token": refresh_token, + "client_id": client_id, + } + if client_secret: + body["client_secret"] = client_secret + return await _request_token(token_endpoint, body, fallback="Notion OAuth token refresh failed") + + +async def _request_token( + token_endpoint: str, body: dict[str, str], *, fallback: str +) -> dict[str, Any]: + try: + async with httpx.AsyncClient(timeout=_HTTP_TIMEOUT) as client: + response = await client.post( + token_endpoint, + headers={ + "Accept": "application/json", + "Content-Type": "application/x-www-form-urlencoded", + "User-Agent": "OpenSWE-Notion-MCP/1.0", + }, + data=body, + ) + except httpx.RequestError as exc: + raise NotionOAuthError(503, f"{fallback}: network error") from exc + if not response.is_success: + raise _oauth_error_from_response(response, fallback) + data = response.json() + if not isinstance(data, dict) or not isinstance(data.get("access_token"), str): + raise NotionOAuthError(502, f"{fallback}: missing access_token") + return data + + +def is_reauth_required_error(exc: BaseException) -> bool: + return isinstance(exc, NotionOAuthError) and exc.error_code == "invalid_grant" diff --git a/agent/dashboard/routes.py b/agent/dashboard/routes.py index ffbd81b9..9a4084e8 100644 --- a/agent/dashboard/routes.py +++ b/agent/dashboard/routes.py @@ -35,6 +35,13 @@ from .enabled_repos import ( from .eval_jobs import ( get_reviewer_eval_status, ) +from .notion_oauth import ( + NOTION_STATE_COOKIE_NAME, + NotionOAuthError, + exchange_notion_code, + pop_notion_oauth_flow, + store_notion_oauth_flow, +) from .oauth import ( COOKIE_NAME, SESSION_TTL_SECONDS, @@ -146,8 +153,11 @@ from .thread_api import ( from .user_credentials import ( CurrentsCredentialsUpdate, connect_currents, + connect_notion, disconnect_currents, + disconnect_notion, get_currents_status, + get_notion_status, ) from .user_mappings import ( delete_mapping, @@ -290,6 +300,26 @@ def _clear_slack_state_cookie(response: Response) -> None: ) +def _set_notion_state_cookie(response: Response, nonce: str) -> None: + secure, _ = _cookie_security() + response.set_cookie( + key=NOTION_STATE_COOKIE_NAME, + value=nonce, + max_age=STATE_TTL_SECONDS, + httponly=True, + secure=secure, + samesite="lax", + path="/dashboard/api/notion", + ) + + +def _clear_notion_state_cookie(response: Response) -> None: + secure, _ = _cookie_security() + response.delete_cookie( + NOTION_STATE_COOKIE_NAME, path="/dashboard/api/notion", samesite="lax", secure=secure + ) + + @router.get("/auth/login") async def auth_login( request: Request, @@ -436,6 +466,89 @@ async def disconnect_my_currents( return status.get("currents", {"connected": False}) +@router.get("/my-credentials/notion") +async def get_my_notion_status( + session: dict[str, Any] = _SESSION_DEP, +) -> dict[str, Any]: + status = await get_notion_status(session["sub"]) + return status.get("notion", {"connected": False}) + + +@router.delete("/my-credentials/notion") +async def disconnect_my_notion( + session: dict[str, Any] = _SESSION_DEP, +) -> dict[str, Any]: + status = await disconnect_notion(session["sub"]) + return status.get("notion", {"connected": False}) + + +@router.get("/notion/login") +async def notion_login( + session: dict[str, Any] = _SESSION_DEP, +) -> RedirectResponse: + redirect_uri = f"{_api_base_url()}/dashboard/api/notion/callback" + nonce = new_state_nonce() + nonce_hash = hash_state_nonce(nonce) + state = issue_state( + redirect_to=f"{_frontend_base_url()}/my-settings", + nonce_hash=nonce_hash, + ) + try: + url = await store_notion_oauth_flow( + session["sub"], + nonce_hash, + redirect_uri=redirect_uri, + state=state, + ) + except NotionOAuthError as exc: + raise HTTPException(exc.status_code, exc.detail) from exc + response = RedirectResponse(url, status_code=302) + _set_notion_state_cookie(response, nonce) + return response + + +@router.get("/notion/callback") +async def notion_callback( + request: Request, + state: str, + code: str | None = None, + error: str | None = None, + error_description: str | None = None, + session: dict[str, Any] = _SESSION_DEP, +) -> RedirectResponse: + state_payload = decode_state(state) + nonce_hash = state_payload.get("nonce_hash") + cookie_nonce = request.cookies.get(NOTION_STATE_COOKIE_NAME) + if ( + not isinstance(nonce_hash, str) + or not cookie_nonce + or not hmac.compare_digest(hash_state_nonce(cookie_nonce), nonce_hash) + ): + raise HTTPException(400, "oauth state mismatch — please retry") + + flow = await pop_notion_oauth_flow(session["sub"], nonce_hash) + if flow is None: + raise HTTPException(400, "oauth flow expired — please retry") + if error: + detail = error_description or error + raise HTTPException(400, f"Notion OAuth failed: {detail}") + if not code: + raise HTTPException(400, "Notion OAuth callback missing code") + + try: + token_data = await exchange_notion_code(code, flow) + await connect_notion(session["sub"], token_data, flow) + except NotionOAuthError as exc: + raise HTTPException(exc.status_code, exc.detail) from exc + except ValueError as exc: + raise HTTPException(400, str(exc)) from exc + + redirect_to = sanitize_redirect_to(state_payload.get("redirect_to")) or _frontend_base_url() + response = RedirectResponse(redirect_to, status_code=302) + _clear_notion_state_cookie(response) + return response + + @router.get("/slack/login") async def slack_login( _session: dict[str, Any] = _SESSION_DEP, diff --git a/agent/dashboard/user_credentials.py b/agent/dashboard/user_credentials.py index e4940839..5c457925 100644 --- a/agent/dashboard/user_credentials.py +++ b/agent/dashboard/user_credentials.py @@ -1,27 +1,27 @@ -"""Per-user third-party service credentials (Currents.dev). - -Credentials are encrypted at rest with :mod:`agent.encryption` and stored in a -dedicated LangGraph Store namespace, keyed by the user's GitHub login. The -sandbox never holds these keys — they feed server-side read-only tools. -""" +"""Per-user third-party service credentials.""" from __future__ import annotations +import asyncio import logging -from datetime import UTC, datetime +from dataclasses import dataclass +from datetime import UTC, datetime, timedelta from typing import Any from langgraph_sdk import get_client from pydantic import BaseModel, field_validator from ..encryption import decrypt_token, encrypt_token +from .notion_oauth import is_reauth_required_error, refresh_notion_access_token logger = logging.getLogger(__name__) USER_CREDENTIALS_NAMESPACE: list[str] = ["user_credentials"] CURRENTS_KEY = "currents" +NOTION_KEY = "notion" CURRENTS_API_BASE = "https://api.currents.dev/v1" +_NOTION_TOKEN_EXPIRY_SKEW_SECONDS = 300 def _client(): @@ -68,6 +68,248 @@ class CurrentsCredentialsUpdate(BaseModel): return v.strip() +@dataclass(frozen=True) +class NotionCredentials: + access_token: str + refresh_token: str | None + token_endpoint: str + client_id: str + client_secret: str | None = None + + +def _expires_at_from_response(data: dict[str, Any], *, field: str = "expires_in") -> str | None: + raw = data.get(field) + if not isinstance(raw, int | float) or raw <= 0: + return None + return (datetime.now(UTC) + timedelta(seconds=int(raw))).isoformat() + + +def _token_expired( + expires_at: str | None, *, skew_seconds: int = _NOTION_TOKEN_EXPIRY_SKEW_SECONDS +) -> bool: + if not isinstance(expires_at, str) or not expires_at: + return False + try: + exp = datetime.fromisoformat(expires_at.replace("Z", "+00:00")) + if exp.tzinfo is None: + exp = exp.replace(tzinfo=UTC) + except ValueError: + return False + return datetime.now(UTC) + timedelta(seconds=skew_seconds) >= exp + + +async def get_notion_status(login: str) -> dict[str, Any]: + """Return a redacted view of the user's Notion MCP connection.""" + notion = await _get_provider(login, NOTION_KEY) + return { + "notion": { + "connected": True, + "token_expires_at": notion.get("token_expires_at"), + "updated_at": notion.get("updated_at"), + } + if notion + else {"connected": False}, + } + + +def _notion_record_from_response( + data: dict[str, Any], + *, + existing: dict[str, Any] | None = None, + client_id: str | None = None, + client_secret: str | None = None, + token_endpoint: str | None = None, +) -> dict[str, Any]: + existing = existing or {} + access_token = data.get("access_token") + if not isinstance(access_token, str) or not access_token: + raise ValueError("Notion OAuth response missing access_token") + refresh_token = data.get("refresh_token") + record: dict[str, Any] = { + "encrypted_access_token": encrypt_token(access_token), + "client_id": client_id or existing.get("client_id", ""), + "token_endpoint": token_endpoint or existing.get("token_endpoint", ""), + "updated_at": datetime.now(UTC).isoformat(), + } + token_type = data.get("token_type") + if isinstance(token_type, str): + record["token_type"] = token_type + scope = data.get("scope") + if isinstance(scope, str): + record["scope"] = scope + token_expires_at = _expires_at_from_response(data) + if token_expires_at: + record["token_expires_at"] = token_expires_at + elif existing.get("token_expires_at"): + record["token_expires_at"] = existing["token_expires_at"] + refresh_token_expires_at = _expires_at_from_response(data, field="refresh_token_expires_in") + if refresh_token_expires_at: + record["refresh_token_expires_at"] = refresh_token_expires_at + elif existing.get("refresh_token_expires_at"): + record["refresh_token_expires_at"] = existing["refresh_token_expires_at"] + if isinstance(refresh_token, str) and refresh_token: + record["encrypted_refresh_token"] = encrypt_token(refresh_token) + elif existing.get("encrypted_refresh_token"): + record["encrypted_refresh_token"] = existing["encrypted_refresh_token"] + if client_secret: + record["encrypted_client_secret"] = encrypt_token(client_secret) + elif existing.get("encrypted_client_secret"): + record["encrypted_client_secret"] = existing["encrypted_client_secret"] + return record + + +async def connect_notion(login: str, data: dict[str, Any], flow: dict[str, Any]) -> dict[str, Any]: + client_id = flow.get("client_id") + token_endpoint = flow.get("token_endpoint") + if not isinstance(client_id, str) or not isinstance(token_endpoint, str): + raise ValueError("stored Notion OAuth flow is incomplete") + client_secret = ( + flow.get("client_secret") if isinstance(flow.get("client_secret"), str) else None + ) + await _put_provider( + login, + NOTION_KEY, + _notion_record_from_response( + data, + client_id=client_id, + client_secret=client_secret, + token_endpoint=token_endpoint, + ), + ) + return await get_notion_status(login) + + +async def disconnect_notion(login: str) -> dict[str, Any]: + await _delete_provider(login, NOTION_KEY) + return await get_notion_status(login) + + +def _decrypt_notion_access_token(record: dict[str, Any]) -> str | None: + token = decrypt_token(record.get("encrypted_access_token", "")) + return token or None + + +def _decrypt_notion_refresh_token(record: dict[str, Any]) -> str | None: + token = decrypt_token(record.get("encrypted_refresh_token", "")) + return token or None + + +def _decrypt_notion_client_secret(record: dict[str, Any]) -> str | None: + token = decrypt_token(record.get("encrypted_client_secret", "")) + return token or None + + +_notion_refresh_locks: dict[str, asyncio.Lock] = {} + + +def _notion_refresh_lock(login: str) -> asyncio.Lock: + lock = _notion_refresh_locks.get(login) + if lock is None: + lock = asyncio.Lock() + _notion_refresh_locks[login] = lock + return lock + + +async def _refresh_stored_notion_token( + login: str, + record: dict[str, Any], +) -> tuple[str | None, bool]: + refresh_token = _decrypt_notion_refresh_token(record) + token_endpoint = record.get("token_endpoint") + client_id = record.get("client_id") + if not refresh_token or not isinstance(token_endpoint, str) or not isinstance(client_id, str): + return None, False + try: + data = await refresh_notion_access_token( + refresh_token=refresh_token, + token_endpoint=token_endpoint, + client_id=client_id, + client_secret=_decrypt_notion_client_secret(record), + ) + except Exception as exc: # noqa: BLE001 + logger.warning("Notion token refresh failed for %s", login, exc_info=True) + return None, is_reauth_required_error(exc) + await _put_provider(login, NOTION_KEY, _notion_record_from_response(data, existing=record)) + access_token = data.get("access_token") + return (access_token if isinstance(access_token, str) else None), False + + +async def get_notion_credentials( + login: str, *, force_refresh: bool = False +) -> NotionCredentials | None: + """Return a valid Notion MCP credential set for a user.""" + record = await _get_provider(login, NOTION_KEY) + if not record: + return None + access_token = _decrypt_notion_access_token(record) + if not access_token: + return None + if not force_refresh and not _token_expired(record.get("token_expires_at")): + return NotionCredentials( + access_token=access_token, + refresh_token=_decrypt_notion_refresh_token(record), + token_endpoint=record.get("token_endpoint", ""), + client_id=record.get("client_id", ""), + client_secret=_decrypt_notion_client_secret(record), + ) + if not _decrypt_notion_refresh_token(record): + return None + async with _notion_refresh_lock(login): + record = await _get_provider(login, NOTION_KEY) + if not record: + return None + access_token = _decrypt_notion_access_token(record) + if not access_token: + return None + if not force_refresh and not _token_expired(record.get("token_expires_at")): + return NotionCredentials( + access_token=access_token, + refresh_token=_decrypt_notion_refresh_token(record), + token_endpoint=record.get("token_endpoint", ""), + client_id=record.get("client_id", ""), + client_secret=_decrypt_notion_client_secret(record), + ) + refreshed, refresh_token_dead = await _refresh_stored_notion_token(login, record) + if refreshed: + refreshed_record = await _get_provider(login, NOTION_KEY) or record + return NotionCredentials( + access_token=refreshed, + refresh_token=_decrypt_notion_refresh_token(refreshed_record), + token_endpoint=refreshed_record.get("token_endpoint", ""), + client_id=refreshed_record.get("client_id", ""), + client_secret=_decrypt_notion_client_secret(refreshed_record), + ) + if refresh_token_dead: + latest = await _get_provider(login, NOTION_KEY) + if latest and latest.get("encrypted_refresh_token") != record.get( + "encrypted_refresh_token" + ): + latest_access_token = _decrypt_notion_access_token(latest) + if latest_access_token: + return NotionCredentials( + access_token=latest_access_token, + refresh_token=_decrypt_notion_refresh_token(latest), + token_endpoint=latest.get("token_endpoint", ""), + client_id=latest.get("client_id", ""), + client_secret=_decrypt_notion_client_secret(latest), + ) + logger.info("Dropping dead Notion authorization for %s; reconnect required", login) + await disconnect_notion(login) + return None + return NotionCredentials( + access_token=access_token, + refresh_token=_decrypt_notion_refresh_token(record), + token_endpoint=record.get("token_endpoint", ""), + client_id=record.get("client_id", ""), + client_secret=_decrypt_notion_client_secret(record), + ) + + +async def get_notion_access_token(login: str) -> str | None: + credentials = await get_notion_credentials(login) + return credentials.access_token if credentials else None + + async def get_currents_status(login: str) -> dict[str, Any]: """Return a redacted, dashboard-safe view of the user's Currents key.""" currents = await _get_provider(login, CURRENTS_KEY) diff --git a/agent/integrations/notion_mcp.py b/agent/integrations/notion_mcp.py new file mode 100644 index 00000000..ec150055 --- /dev/null +++ b/agent/integrations/notion_mcp.py @@ -0,0 +1,101 @@ +"""Server-side Notion tools backed by Notion's hosted MCP server.""" + +from __future__ import annotations + +import asyncio +import logging +from datetime import timedelta +from typing import Any + +from langchain_core.tools import BaseTool + +from ..dashboard.notion_oauth import NOTION_MCP_URL +from ..dashboard.user_credentials import get_notion_access_token + +logger = logging.getLogger(__name__) + +_MCP_TIMEOUT_SECONDS = 30.0 + + +async def _build_mcp_tools(access_token: str) -> list[BaseTool]: + from langchain_mcp_adapters.client import MultiServerMCPClient + + client = MultiServerMCPClient( + { + "notion": { + "transport": "streamable_http", + "url": NOTION_MCP_URL, + "headers": { + "Authorization": f"Bearer {access_token}", + }, + "timeout": timedelta(seconds=_MCP_TIMEOUT_SECONDS), + } + } + ) + return await client.get_tools() + + +async def _fresh_mcp_tool(login: str, tool_name: str) -> BaseTool: + access_token = await get_notion_access_token(login) + if not access_token: + raise RuntimeError( + "Notion MCP authorization unavailable; reconnect Notion in Profile Settings" + ) + tools = await _build_mcp_tools(access_token) + for tool in tools: + if tool.name == tool_name: + return tool + raise RuntimeError(f"Notion MCP tool {tool_name!r} is no longer available") + + +def _tool_input(args: tuple[Any, ...], kwargs: dict[str, Any]) -> str | dict[str, Any]: + if args and kwargs: + raise TypeError("Notion MCP tool received both positional and keyword input") + if not args: + return kwargs + if len(args) == 1 and isinstance(args[0], str): + return args[0] + if len(args) == 1 and isinstance(args[0], dict): + return args[0] + raise TypeError("Notion MCP tool received invalid positional input") + + +class _RefreshingNotionMCPTool(BaseTool): + login: str + mcp_tool_name: str + + def _run(self, *args: Any, **kwargs: Any) -> Any: + try: + asyncio.get_running_loop() + except RuntimeError: + return asyncio.run(self._arun(*args, **kwargs)) + raise RuntimeError("Notion MCP tools must be called asynchronously") + + async def _arun(self, *args: Any, **kwargs: Any) -> Any: + tool = await _fresh_mcp_tool(self.login, self.mcp_tool_name) + return await tool.ainvoke(_tool_input(args, kwargs)) + + +def _refreshing_tool(login: str, tool: BaseTool) -> BaseTool: + return _RefreshingNotionMCPTool( + name=tool.name, + description=tool.description, + args_schema=tool.args_schema, + response_format="content", + login=login, + mcp_tool_name=tool.name, + ) + + +async def load_notion_tools(login: str) -> list[BaseTool]: + """Return Notion MCP tools for a connected user.""" + access_token = await get_notion_access_token(login) + if not access_token: + return [] + try: + tools = await _build_mcp_tools(access_token) + except Exception: # noqa: BLE001 + logger.warning("Failed to load Notion MCP tools", exc_info=True) + return [] + logger.info("Loaded %d Notion MCP tool(s) for %s", len(tools), login) + return [_refreshing_tool(login, tool) for tool in tools] diff --git a/agent/server.py b/agent/server.py index 757a9466..291b8931 100644 --- a/agent/server.py +++ b/agent/server.py @@ -51,6 +51,7 @@ from .integrations.currents_tools import load_currents_tools from .integrations.datadog_mcp import load_datadog_tools from .integrations.langsmith import _configure_github_proxy from .integrations.langsmith_tools import load_langsmith_tools +from .integrations.notion_mcp import load_notion_tools from .middleware import ( ModelFallbackMiddleware, PlanModeMiddleware, @@ -739,12 +740,18 @@ async def get_agent(config: RunnableConfig) -> Pregel: corridor_tools = await _load_corridor_mcp_tools() currents_tools: list[Any] = [] + notion_tools: list[Any] = [] if profile_login: try: currents_tools = await load_currents_tools(profile_login) except Exception: logger.warning("Failed to load Currents tools", exc_info=True) currents_tools = [] + try: + notion_tools = await load_notion_tools(profile_login) + except Exception: + logger.warning("Failed to load Notion tools", exc_info=True) + notion_tools = [] logger.info("Returning agent with sandbox for thread %s", thread_id) main_model = make_model(model_id, **model_kwargs) @@ -785,6 +792,7 @@ async def get_agent(config: RunnableConfig) -> Pregel: *corridor_tools, *observability_tools, *currents_tools, + *notion_tools, ], subagents=[_general_purpose_subagent(subagent_model)], backend=backend_factory, diff --git a/tests/test_notion_oauth.py b/tests/test_notion_oauth.py new file mode 100644 index 00000000..90ed827a --- /dev/null +++ b/tests/test_notion_oauth.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +from typing import Any +from unittest.mock import AsyncMock +from urllib.parse import parse_qs, urlparse + +import pytest +from cryptography.fernet import Fernet + +from agent.dashboard import notion_oauth as no + + +class _FakeStore: + def __init__(self) -> None: + self.items: dict[tuple[tuple[str, ...], str], dict[str, Any]] = {} + + async def get_item(self, namespace: list[str], key: str): + value = self.items.get((tuple(namespace), key)) + return {"value": value} if value is not None else None + + async def put_item(self, namespace: list[str], key: str, value: dict[str, Any]) -> None: + self.items[(tuple(namespace), key)] = value + + async def delete_item(self, namespace: list[str], key: str) -> None: + self.items.pop((tuple(namespace), key), None) + + +class _FakeClient: + def __init__(self, store: _FakeStore) -> None: + self.store = store + + +@pytest.fixture() +def fake_store(monkeypatch: pytest.MonkeyPatch) -> _FakeStore: + store = _FakeStore() + monkeypatch.setattr(no, "_client", lambda: _FakeClient(store)) + monkeypatch.setenv("TOKEN_ENCRYPTION_KEY", Fernet.generate_key().decode()) + return store + + +def test_code_challenge_matches_rfc7636_vector() -> None: + verifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk" + assert no.code_challenge_for_verifier(verifier) == "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM" + + +def test_build_notion_authorize_url() -> None: + url = no.build_notion_authorize_url( + authorization_endpoint="https://mcp.notion.com/authorize", + client_id="cid", + redirect_uri="https://example.com/dashboard/api/notion/callback", + code_challenge="challenge", + state="state-token", + ) + parsed = urlparse(url) + query = parse_qs(parsed.query) + assert parsed.netloc == "mcp.notion.com" + assert parsed.path == "/authorize" + assert query["response_type"] == ["code"] + assert query["client_id"] == ["cid"] + assert query["code_challenge"] == ["challenge"] + assert query["code_challenge_method"] == ["S256"] + assert query["prompt"] == ["consent"] + + +def test_build_notion_authorize_url_rejects_other_hosts() -> None: + with pytest.raises(no.NotionOAuthError): + no.build_notion_authorize_url( + authorization_endpoint="https://example.com/authorize", + client_id="cid", + redirect_uri="https://example.com/callback", + code_challenge="challenge", + state="state-token", + ) + + +@pytest.mark.asyncio +async def test_store_and_pop_notion_oauth_flow( + fake_store: _FakeStore, + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr( + no, + "discover_notion_oauth_metadata", + AsyncMock( + return_value={ + "authorization_endpoint": "https://mcp.notion.com/authorize", + "token_endpoint": "https://mcp.notion.com/token", + "registration_endpoint": "https://mcp.notion.com/register", + } + ), + ) + monkeypatch.setattr( + no, + "register_notion_oauth_client", + AsyncMock(return_value={"client_id": "cid", "client_secret": "secret"}), + ) + monkeypatch.setattr(no, "generate_code_verifier", lambda: "verifier") + + url = await no.store_notion_oauth_flow( + "alice", + "nonce-hash", + redirect_uri="https://example.com/dashboard/api/notion/callback", + state="state-token", + ) + assert parse_qs(urlparse(url).query)["client_id"] == ["cid"] + + flow = await no.pop_notion_oauth_flow("alice", "nonce-hash") + assert flow is not None + assert flow["code_verifier"] == "verifier" + assert flow["client_secret"] == "secret" + assert await no.pop_notion_oauth_flow("alice", "nonce-hash") is None diff --git a/tests/test_observability_tools.py b/tests/test_observability_tools.py index e0f724e2..8dce9f70 100644 --- a/tests/test_observability_tools.py +++ b/tests/test_observability_tools.py @@ -1,12 +1,14 @@ from __future__ import annotations +from typing import Literal from unittest.mock import AsyncMock, patch import pytest +from langchain_core.tools import StructuredTool from agent import server from agent.dashboard.team_credentials import DatadogCredentials, LangSmithCredentials -from agent.integrations import datadog_mcp, langsmith_tools +from agent.integrations import datadog_mcp, langsmith_tools, notion_mcp @pytest.mark.asyncio @@ -36,6 +38,103 @@ async def test_load_datadog_tools_returns_tools() -> None: assert await datadog_mcp.load_datadog_tools() == sentinel +@pytest.mark.asyncio +async def test_load_notion_tools_empty_when_not_connected() -> None: + with patch.object(notion_mcp, "get_notion_access_token", AsyncMock(return_value=None)): + assert await notion_mcp.load_notion_tools("alice") == [] + + +@pytest.mark.asyncio +async def test_load_notion_tools_degrades_on_error() -> None: + with ( + patch.object(notion_mcp, "get_notion_access_token", AsyncMock(return_value="tok")), + patch.object(notion_mcp, "_build_mcp_tools", AsyncMock(side_effect=RuntimeError("boom"))), + ): + assert await notion_mcp.load_notion_tools("alice") == [] + + +def _notion_tool_for_token( + token: str, + response_format: Literal["content", "content_and_artifact"] = "content", +) -> StructuredTool: + async def notion_search(query: str): + """Search Notion.""" + content = {"query": query, "token": token} + if response_format == "content_and_artifact": + return content, {"artifact_token": token} + return content + + return StructuredTool.from_function( + coroutine=notion_search, + name="notion_search", + description="Search Notion", + response_format=response_format, + ) + + +@pytest.mark.asyncio +async def test_load_notion_tools_returns_wrappers() -> None: + discovered = _notion_tool_for_token("initial-token") + with ( + patch.object(notion_mcp, "get_notion_access_token", AsyncMock(return_value="tok")), + patch.object(notion_mcp, "_build_mcp_tools", AsyncMock(return_value=[discovered])), + ): + tools = await notion_mcp.load_notion_tools("alice") + assert len(tools) == 1 + assert tools[0].name == "notion_search" + assert tools[0].description == discovered.description + assert tools[0].args_schema == discovered.args_schema + assert tools[0].response_format == "content" + + +@pytest.mark.asyncio +async def test_notion_wrapper_normalizes_content_and_artifact_tools() -> None: + get_token = AsyncMock(side_effect=["initial-token", "fresh-token"]) + build_tools = AsyncMock( + side_effect=lambda token: [_notion_tool_for_token(token, "content_and_artifact")] + ) + with ( + patch.object(notion_mcp, "get_notion_access_token", get_token), + patch.object(notion_mcp, "_build_mcp_tools", build_tools), + ): + tools = await notion_mcp.load_notion_tools("alice") + assert tools[0].response_format == "content" + result = await tools[0].ainvoke({"query": "roadmap"}) + assert result == {"query": "roadmap", "token": "fresh-token"} + + +@pytest.mark.asyncio +async def test_notion_wrapper_refreshes_token_at_call_time() -> None: + get_token = AsyncMock(side_effect=["initial-token", "fresh-token"]) + build_tools = AsyncMock(side_effect=lambda token: [_notion_tool_for_token(token)]) + with ( + patch.object(notion_mcp, "get_notion_access_token", get_token), + patch.object(notion_mcp, "_build_mcp_tools", build_tools), + ): + tools = await notion_mcp.load_notion_tools("alice") + result = await tools[0].ainvoke({"query": "roadmap"}) + assert result == {"query": "roadmap", "token": "fresh-token"} + assert get_token.await_count == 2 + assert [call.args[0] for call in build_tools.await_args_list] == [ + "initial-token", + "fresh-token", + ] + + +@pytest.mark.asyncio +async def test_notion_wrapper_fails_when_token_missing_at_call_time() -> None: + get_token = AsyncMock(side_effect=["initial-token", None]) + build_tools = AsyncMock(return_value=[_notion_tool_for_token("initial-token")]) + with ( + patch.object(notion_mcp, "get_notion_access_token", get_token), + patch.object(notion_mcp, "_build_mcp_tools", build_tools), + ): + tools = await notion_mcp.load_notion_tools("alice") + with pytest.raises(RuntimeError, match="Notion MCP authorization unavailable"): + await tools[0].ainvoke({"query": "roadmap"}) + assert build_tools.await_count == 1 + + @pytest.mark.asyncio async def test_load_langsmith_tools_empty_when_not_connected() -> None: with patch.object(langsmith_tools, "get_langsmith_credentials", AsyncMock(return_value=None)): diff --git a/tests/test_user_credentials.py b/tests/test_user_credentials.py index 55b9cf32..5fa3dce2 100644 --- a/tests/test_user_credentials.py +++ b/tests/test_user_credentials.py @@ -1,12 +1,15 @@ from __future__ import annotations +from datetime import UTC, datetime, timedelta from typing import Any +from unittest.mock import AsyncMock, patch import pytest from cryptography.fernet import Fernet from pydantic import ValidationError from agent.dashboard import user_credentials as uc +from agent.dashboard.notion_oauth import NotionOAuthError from agent.dashboard.user_credentials import CurrentsCredentialsUpdate @@ -93,3 +96,107 @@ async def test_currents_status_when_not_connected(fake_store: _FakeStore) -> Non @pytest.mark.asyncio async def test_get_currents_api_key_none_when_not_connected(fake_store: _FakeStore) -> None: assert await uc.get_currents_api_key("nobody") is None + + +@pytest.mark.asyncio +async def test_notion_roundtrip_and_redaction(fake_store: _FakeStore) -> None: + status = await uc.connect_notion( + "alice", + { + "access_token": "notion-access-1234", + "refresh_token": "notion-refresh", + "expires_in": 3600, + "token_type": "Bearer", + }, + { + "client_id": "client-id", + "client_secret": "client-secret", + "token_endpoint": "https://mcp.notion.com/token", + }, + ) + assert status["notion"]["connected"] is True + + record = fake_store.items[(("user_credentials", "alice"), "notion")] + assert record["encrypted_access_token"] != "notion-access-1234" + assert record["encrypted_refresh_token"] != "notion-refresh" + assert record["encrypted_client_secret"] != "client-secret" + + creds = await uc.get_notion_credentials("alice") + assert creds is not None + assert creds.access_token == "notion-access-1234" + assert creds.refresh_token == "notion-refresh" + assert creds.client_id == "client-id" + assert creds.client_secret == "client-secret" + + after = await uc.disconnect_notion("alice") + assert after["notion"]["connected"] is False + assert await uc.get_notion_credentials("alice") is None + + +@pytest.mark.asyncio +async def test_notion_refresh_rotates_tokens(fake_store: _FakeStore) -> None: + await uc.connect_notion( + "alice", + { + "access_token": "old-access", + "refresh_token": "old-refresh", + "expires_in": 3600, + }, + { + "client_id": "client-id", + "client_secret": "client-secret", + "token_endpoint": "https://mcp.notion.com/token", + }, + ) + record = fake_store.items[(("user_credentials", "alice"), "notion")] + record["token_expires_at"] = (datetime.now(UTC) - timedelta(minutes=1)).isoformat() + + with patch.object( + uc, + "refresh_notion_access_token", + new_callable=AsyncMock, + return_value={ + "access_token": "new-access", + "refresh_token": "new-refresh", + "expires_in": 3600, + }, + ) as refresh: + creds = await uc.get_notion_credentials("alice") + + assert creds is not None + assert creds.access_token == "new-access" + assert creds.refresh_token == "new-refresh" + refresh.assert_awaited_once_with( + refresh_token="old-refresh", + token_endpoint="https://mcp.notion.com/token", + client_id="client-id", + client_secret="client-secret", + ) + + +@pytest.mark.asyncio +async def test_notion_invalid_grant_disconnects(fake_store: _FakeStore) -> None: + await uc.connect_notion( + "alice", + { + "access_token": "old-access", + "refresh_token": "old-refresh", + "expires_in": 3600, + }, + { + "client_id": "client-id", + "token_endpoint": "https://mcp.notion.com/token", + }, + ) + record = fake_store.items[(("user_credentials", "alice"), "notion")] + record["token_expires_at"] = (datetime.now(UTC) - timedelta(minutes=1)).isoformat() + + with patch.object( + uc, + "refresh_notion_access_token", + new_callable=AsyncMock, + side_effect=NotionOAuthError(400, "dead", error_code="invalid_grant"), + ): + assert await uc.get_notion_credentials("alice") is None + + assert await uc.get_notion_status("alice") == {"notion": {"connected": False}} diff --git a/ui/src/lib/api.ts b/ui/src/lib/api.ts index 8fbfd865..b357f91d 100644 --- a/ui/src/lib/api.ts +++ b/ui/src/lib/api.ts @@ -202,6 +202,12 @@ export interface CurrentsConnectBody { api_key: string } +export interface NotionCredentialStatus { + connected: boolean + token_expires_at?: string | null + updated_at?: string | null +} + export interface UserMapping { github_login: string work_email: string @@ -612,6 +618,12 @@ export const api = { request("/my-credentials/currents", { method: "DELETE", }), + getMyNotionStatus: () => + request("/my-credentials/notion"), + disconnectNotion: () => + request("/my-credentials/notion", { + method: "DELETE", + }), listEnabledReviewRepos: () => request<{ repos: Array }>("/enabled-review-repos"), setEnabledReviewRepo: (full_name: string, enabled: boolean) => @@ -685,3 +697,7 @@ export function loginUrl(redirectTo?: string): string { export function slackConnectUrl(): string { return `${API_BASE}/dashboard/api/slack/login` } + +export function notionConnectUrl(): string { + return `${API_BASE}/dashboard/api/notion/login` +} diff --git a/ui/src/routes/my-settings.tsx b/ui/src/routes/my-settings.tsx index 836eec65..2298421f 100644 --- a/ui/src/routes/my-settings.tsx +++ b/ui/src/routes/my-settings.tsx @@ -2,6 +2,7 @@ import { Navigate, createFileRoute, useNavigate } from "@tanstack/react-router" import { useMutation, useQuery, useQueryClient } from "@tanstack/react-query" import { useState } from "react" import { IoLogoSlack } from "react-icons/io5" +import { SiNotion } from "react-icons/si" import type { CurrentsConnectBody, SessionUser } from "@/lib/api" import { AppShell, SettingsRow, SettingsSection } from "@/components/AppShell" @@ -16,7 +17,7 @@ import { } from "@/components/ui/select" import { Skeleton } from "@/components/ui/skeleton" import { Switch } from "@/components/ui/switch" -import { api, slackConnectUrl } from "@/lib/api" +import { api, notionConnectUrl, slackConnectUrl } from "@/lib/api" import { buildProfileUpdate, useOptions, @@ -265,6 +266,85 @@ function CurrentsCredentialsSection() { ) } +function NotionCredentialsSection() { + const qc = useQueryClient() + const creds = useQuery({ + queryKey: ["myNotion"], + queryFn: api.getMyNotionStatus, + }) + const [connecting, setConnecting] = useState(false) + const [error, setError] = useState(null) + + const onSuccess = () => { + qc.invalidateQueries({ queryKey: ["myNotion"] }) + setError(null) + } + const onError = (e: Error) => setError(e.message) + + const disconnect = useMutation({ + mutationFn: () => api.disconnectNotion(), + onSuccess, + onError, + }) + + const connected = creds.data?.connected + const connect = () => { + setConnecting(true) + void qc.invalidateQueries({ queryKey: ["myNotion"] }) + window.location.assign(notionConnectUrl()) + } + + return ( + + + + {connected ? "Connected" : "Not connected"} + + {connected ? ( + + ) : ( + + )} + + } + /> + {error &&

{error}

} +
+ ) +} + function MySettingsPage() { const session = useSession() const qc = useQueryClient() @@ -368,6 +448,8 @@ function MySettingsPage() { + +