#!/usr/bin/env python3
"""Minimal OAuth 2.1 + PKCE server for Claude MCP connector."""
import json
import secrets
import hashlib
import base64
import socket
from http.server import HTTPServer, BaseHTTPRequestHandler
from urllib.parse import urlparse, parse_qs

# Store auth codes temporarily
auth_codes = {}

class OAuthHandler(BaseHTTPRequestHandler):
    def log_message(self, format, *args):
        pass
    
    def send_json(self, data, status=200):
        self.send_response(status)
        self.send_header('Content-Type', 'application/json')
        self.send_header('Access-Control-Allow-Origin', '*')
        self.end_headers()
        self.wfile.write(json.dumps(data).encode())
    
    def do_OPTIONS(self):
        self.send_response(200)
        self.send_header('Access-Control-Allow-Origin', '*')
        self.send_header('Access-Control-Allow-Methods', 'GET, POST, OPTIONS')
        self.send_header('Access-Control-Allow-Headers', 'Content-Type, Authorization')
        self.end_headers()
    
    def do_GET(self):
        parsed = urlparse(self.path)
        params = parse_qs(parsed.query)
        
        if parsed.path == '/.well-known/oauth-authorization-server':
            self.send_json({
                'issuer': 'https://robblake.cloud',
                'authorization_endpoint': 'https://robblake.cloud/oauth/authorize',
                'token_endpoint': 'https://robblake.cloud/oauth/token',
                'registration_endpoint': 'https://robblake.cloud/oauth/register',
                'response_types_supported': ['code'],
                'grant_types_supported': ['authorization_code'],
                'code_challenge_methods_supported': ['S256']
            })
        
        elif parsed.path == '/oauth/authorize':
            client_id = params.get('client_id', [''])[0]
            redirect_uri = params.get('redirect_uri', [''])[0]
            state = params.get('state', [''])[0]
            code_challenge = params.get('code_challenge', [''])[0]
            
            if not redirect_uri:
                self.send_error(400, 'Missing redirect_uri')
                return
            
            code = secrets.token_urlsafe(32)
            auth_codes[code] = {
                'client_id': client_id,
                'redirect_uri': redirect_uri,
                'code_challenge': code_challenge
            }
            
            sep = '&' if '?' in redirect_uri else '?'
            location = f"{redirect_uri}{sep}code={code}&state={state}"
            self.send_response(302)
            self.send_header('Location', location)
            self.send_header('Access-Control-Allow-Origin', '*')
            self.end_headers()
        
        else:
            self.send_error(404)
    
    def do_POST(self):
        parsed = urlparse(self.path)
        
        if parsed.path == '/oauth/register':
            content_length = int(self.headers.get('Content-Length', 0))
            body = self.rfile.read(content_length).decode() if content_length else '{}'
            try:
                data = json.loads(body)
            except:
                data = {}
            
            redirect_uris = data.get('redirect_uris', ['https://claude.ai/api/mcp/auth_callback'])
            
            self.send_json({
                'client_id': 'claude-desktop',
                'client_secret': 'static_secret_' + secrets.token_urlsafe(16),
                'client_id_issued_at': 1690000000,
                'client_secret_expires_at': 0,
                'redirect_uris': redirect_uris,
                'token_endpoint_auth_method': 'client_secret_post'
            })
        
        elif parsed.path == '/oauth/token':
            content_length = int(self.headers.get('Content-Length', 0))
            body = self.rfile.read(content_length).decode() if content_length else ''
            params = parse_qs(body)
            
            grant_type = params.get('grant_type', [''])[0]
            code = params.get('code', [''])[0]
            code_verifier = params.get('code_verifier', [''])[0]
            
            if grant_type != 'authorization_code':
                self.send_json({'error': 'unsupported_grant_type'}, 400)
                return
            
            if code not in auth_codes:
                self.send_json({'error': 'invalid_grant'}, 400)
                return
            
            stored = auth_codes[code]
            
            if stored['code_challenge']:
                if not code_verifier:
                    self.send_json({'error': 'invalid_request', 'error_description': 'code_verifier required'}, 400)
                    return
                digest = hashlib.sha256(code_verifier.encode()).digest()
                computed = base64.urlsafe_b64encode(digest).decode().rstrip('=')
                if computed != stored['code_challenge']:
                    self.send_json({'error': 'invalid_grant', 'error_description': 'PKCE verification failed'}, 400)
                    return
            
            del auth_codes[code]
            
            self.send_json({
                'access_token': '7d42cb17a34acc345e1ad3064386889fcc80bb583e7ed455',
                'token_type': 'Bearer',
                'expires_in': 3600
            })
        
        elif parsed.path == '/mcp':
            auth = self.headers.get('Authorization', '')
            if not auth.startswith('Bearer '):
                self.send_error(401, 'Unauthorized')
                return
            
            token = auth[7:]
            if token != '7d42cb17a34acc345e1ad3064386889fcc80bb583e7ed455':
                self.send_error(401, 'Invalid token')
                return
            
            # Forward to supergateway using raw socket for SSE support
            content_length = int(self.headers.get('Content-Length', 0))
            body = self.rfile.read(content_length) if content_length else b''
            
            try:
                # Connect to supergateway
                sock = socket.socket(socket.AF_INET, socket.SOCK_STREAM)
                sock.connect(('127.0.0.1', 9121))
                
                # Send request
                request = f"POST /mcp HTTP/1.1\r\nHost: 127.0.0.1:9121\r\nContent-Type: application/json\r\nAccept: application/json, text/event-stream\r\nContent-Length: {len(body)}\r\nConnection: close\r\n\r\n".encode() + body
                sock.sendall(request)
                
                # Read response headers
                response = b''
                while b'\r\n\r\n' not in response:
                    chunk = sock.recv(1)
                    if not chunk:
                        break
                    response += chunk
                
                header_end = response.find(b'\r\n\r\n')
                headers = response[:header_end].decode()
                rest = response[header_end+4:]
                
                # Parse status
                status_line = headers.split('\r\n')[0]
                status_code = int(status_line.split()[1])
                
                # Send status and headers to client
                self.send_response(status_code)
                for line in headers.split('\r\n')[1:]:
                    if ':' in line:
                        k, v = line.split(':', 1)
                        k, v = k.strip(), v.strip()
                        if k.lower() not in ('transfer-encoding', 'connection', 'content-length'):
                            self.send_header(k, v)
                self.end_headers()
                
                # Stream the body
                if rest:
                    self.wfile.write(rest)
                while True:
                    chunk = sock.recv(4096)
                    if not chunk:
                        break
                    self.wfile.write(chunk)
                    self.wfile.flush()
                
                sock.close()
            except Exception as e:
                self.send_error(502, str(e))
        
        else:
            self.send_error(404)

if __name__ == '__main__':
    server = HTTPServer(('127.0.0.1', 9122), OAuthHandler)
    print('OAuth server on 127.0.0.1:9122')
    server.serve_forever()
