152 lines
4.9 KiB
Python
152 lines
4.9 KiB
Python
#!/usr/bin/env python3
|
|
"""Step 1: Register OAuth client and generate authorization URL."""
|
|
import sys, os, json, hashlib, base64, secrets, string, asyncio
|
|
from urllib.parse import urlencode, urlparse
|
|
|
|
sys.path.insert(0, "/usr/local/lib/hermes-agent")
|
|
|
|
SERVER_NAME = sys.argv[1] if len(sys.argv) > 1 else "cloudflare"
|
|
|
|
from hermes_cli.config import load_config
|
|
config = load_config()
|
|
servers = config.get("mcp_servers", {})
|
|
server_config = servers.get(SERVER_NAME, {})
|
|
url = server_config.get("url", "")
|
|
oauth_config = server_config.get("oauth", {})
|
|
|
|
from tools.mcp_oauth import HermesTokenStorage, _configure_callback_port
|
|
from tools.mcp_oauth_manager import get_manager
|
|
|
|
# Clear old tokens
|
|
storage = HermesTokenStorage(SERVER_NAME)
|
|
for p in [storage._tokens_path(), storage._client_info_path(), storage._meta_path()]:
|
|
if p.exists():
|
|
p.unlink()
|
|
print(f" ✓ Removed {p.name}")
|
|
get_manager().remove(SERVER_NAME)
|
|
|
|
cfg = dict(oauth_config or {})
|
|
_configure_callback_port(cfg, storage)
|
|
resolved_port = cfg.get("_resolved_port", 38255)
|
|
redirect_uri = f"http://127.0.0.1:{resolved_port}/callback"
|
|
|
|
# Generate PKCE
|
|
code_verifier = "".join(secrets.choice(string.ascii_letters + string.digits + "-._~") for _ in range(128))
|
|
digest = hashlib.sha256(code_verifier.encode()).digest()
|
|
code_challenge = base64.urlsafe_b64encode(digest).decode().rstrip("=")
|
|
state = secrets.token_urlsafe(32)
|
|
|
|
# Register client at the correct OAuth domain
|
|
parsed_url = urlparse(url)
|
|
oauth_domain = f"{parsed_url.scheme}://{parsed_url.netloc}"
|
|
|
|
import httpx
|
|
|
|
async def register():
|
|
async with httpx.AsyncClient(timeout=15.0) as client:
|
|
reg_payload = {
|
|
"client_name": "Hermes Agent",
|
|
"redirect_uris": [redirect_uri],
|
|
"grant_types": ["authorization_code", "refresh_token"],
|
|
"response_types": ["code"],
|
|
"token_endpoint_auth_method": "none",
|
|
}
|
|
resp = await client.post(
|
|
f"{oauth_domain}/register",
|
|
json=reg_payload,
|
|
headers={"Accept": "application/json"}
|
|
)
|
|
if resp.status_code not in (200, 201):
|
|
print(f" ✗ Registration failed: HTTP {resp.status_code}")
|
|
print(f" Response: {resp.text[:500]}")
|
|
return None, None
|
|
|
|
reg_data = resp.json()
|
|
client_id = reg_data.get("client_id")
|
|
|
|
# Discover token endpoint
|
|
token_endpoint = None
|
|
for disc_url in [
|
|
f"{oauth_domain}/.well-known/oauth-authorization-server",
|
|
f"{oauth_domain}/.well-known/openid-configuration",
|
|
]:
|
|
try:
|
|
resp2 = await client.get(disc_url, headers={"Accept": "application/json"})
|
|
if resp2.status_code == 200:
|
|
data = resp2.json()
|
|
token_endpoint = data.get("token_endpoint")
|
|
if token_endpoint:
|
|
break
|
|
except:
|
|
pass
|
|
|
|
if not token_endpoint:
|
|
token_endpoint = f"{oauth_domain}/token"
|
|
|
|
print(f" ✓ Registered client: {client_id}")
|
|
print(f" Token endpoint: {token_endpoint}")
|
|
return client_id, token_endpoint
|
|
|
|
client_id, token_endpoint = asyncio.run(register())
|
|
if not client_id:
|
|
sys.exit(1)
|
|
|
|
# Save state for step 2
|
|
state_data = {
|
|
"server_name": SERVER_NAME,
|
|
"client_id": client_id,
|
|
"token_endpoint": token_endpoint,
|
|
"redirect_uri": redirect_uri,
|
|
"code_verifier": code_verifier,
|
|
"state": state,
|
|
"url": url,
|
|
}
|
|
with open(f"/tmp/mcp_oauth_{SERVER_NAME}.json", "w") as f:
|
|
json.dump(state_data, f)
|
|
|
|
# Build auth URL using discovered metadata
|
|
params = {
|
|
"response_type": "code",
|
|
"client_id": client_id,
|
|
"redirect_uri": redirect_uri,
|
|
"state": state,
|
|
"code_challenge": code_challenge,
|
|
"code_challenge_method": "S256",
|
|
"resource": url,
|
|
}
|
|
|
|
# Discover the correct authorization endpoint
|
|
import httpx
|
|
|
|
async def discover_auth_endpoint():
|
|
async with httpx.AsyncClient(timeout=10.0) as client:
|
|
for disc_url in [
|
|
f"{oauth_domain}/.well-known/oauth-authorization-server",
|
|
]:
|
|
try:
|
|
resp = await client.get(disc_url, headers={"Accept": "application/json"})
|
|
if resp.status_code == 200:
|
|
data = resp.json()
|
|
auth_ep = data.get("authorization_endpoint")
|
|
if auth_ep:
|
|
return auth_ep
|
|
except:
|
|
pass
|
|
return f"{oauth_domain}/authorize"
|
|
|
|
auth_endpoint = asyncio.run(discover_auth_endpoint())
|
|
print(f" Auth endpoint: {auth_endpoint}")
|
|
|
|
auth_url = f"{auth_endpoint}?{urlencode(params)}"
|
|
|
|
print(f"\n {'='*70}")
|
|
print(f" Server: {SERVER_NAME}")
|
|
print(f" URL: {url}")
|
|
print(f" OAuth domain: {oauth_domain}")
|
|
print(f" {'='*70}")
|
|
print(f"\n OPEN THIS URL IN YOUR BROWSER:\n")
|
|
print(f" {auth_url}")
|
|
print(f"\n {'='*70}")
|
|
print(f" Then run: python3 /root/workspace/oauth-step2.py {SERVER_NAME} <callback_url>")
|
|
print(f" {'='*70}")
|