import os
import re
import time
from urllib.parse import urlparse

import httpx
import jwt
import uvicorn
from mcp.server.fastmcp import Context, FastMCP
from mcp.server.transport_security import TransportSecuritySettings
from mcp.types import ToolAnnotations
from starlette.responses import JSONResponse

api_base = os.getenv("API_BASE_URL", "http://127.0.0.1:8081").rstrip("/")
api_secret = os.environ["API_JWT_SECRET"]
mcp_secret = os.environ["MCP_JWT_SECRET"]
if min(len(api_secret), len(mcp_secret)) < 32 or api_secret == mcp_secret:
    raise ValueError("Use separate API and MCP signing secrets of at least 32 characters")
origin = os.getenv("PUBLIC_ORIGIN", "http://127.0.0.1:8080").rstrip("/")
mcp = FastMCP(
    "orders-mcp", stateless_http=True, json_response=True,
    transport_security=TransportSecuritySettings(
        enable_dns_rebinding_protection=True,
        allowed_hosts=[urlparse(origin).netloc, "127.0.0.1:8080", "localhost:8080"],
        allowed_origins=[origin],
    ),
)


@mcp.tool(annotations=ToolAnnotations(readOnlyHint=True))
async def get_order(order_id: str, ctx: Context) -> dict[str, str]:
    """Read an order's shipping status from the Orders API."""
    if not re.fullmatch(r"[a-zA-Z0-9-]{1,64}", order_id):
        raise ValueError("Use an order ID containing letters, numbers, or hyphens")
    # Identity comes from verified HTTP state, not a user-supplied tool argument.
    subject = ctx.request_context.request.state.principal
    now = int(time.time())
    api_token = jwt.encode(
        {"iss": "orders-demo", "aud": "orders-api", "sub": subject,
         "scope": "orders:read", "iat": now, "exp": now + 300},
        api_secret, algorithm="HS256",
    )
    try:
        async with httpx.AsyncClient(timeout=5.0, follow_redirects=False) as client:
            response = await client.get(
                f"{api_base}/orders/{order_id}",
                headers={os.getenv("API_AUTH_HEADER", "Authorization"): f"Bearer {api_token}"},
            )
            response.raise_for_status()
            data = response.json()
            if not isinstance(data.get("id"), str) or not isinstance(data.get("status"), str):
                raise ValueError("Unexpected API response")
            return {"id": data["id"], "status": data["status"]}
    except (httpx.HTTPError, ValueError, AttributeError):
        raise ValueError("Could not read this order. Check its ID and your API access.") from None


@mcp.custom_route("/healthz", methods=["GET"])
async def health(_request):
    return JSONResponse({"ok": True})


class JWTGate:
    def __init__(self, application):
        self.application = application

    async def __call__(self, scope, receive, send):
        if scope["type"] == "http" and scope["path"] == "/mcp":
            headers = dict(scope["headers"])
            try:
                match = re.fullmatch(rb"Bearer (\S+)", headers.get(os.getenv("MCP_AUTH_HEADER", "Authorization").lower().encode(), b""))
                if not match:
                    raise ValueError("Missing token")
                claims = jwt.decode(
                    match[1], mcp_secret, algorithms=["HS256"],
                    issuer="orders-demo", audience="orders-mcp",
                    options={"require": ["sub", "exp"]},
                )
                if not isinstance(claims["sub"], str) or not claims["sub"]:
                    raise ValueError("Missing subject")
                if not isinstance(claims.get("scope"), str) or "orders:read" not in claims["scope"].split():
                    await JSONResponse({"error": "orders:read permission required"}, status_code=403)(scope, receive, send)
                    return
                scope.setdefault("state", {})["principal"] = claims["sub"]
            except (jwt.PyJWTError, ValueError):
                await JSONResponse(
                    {"error": "Invalid or expired access token"}, status_code=401,
                    headers={"WWW-Authenticate": 'Bearer error="invalid_token"'},
                )(scope, receive, send)
                return
        await self.application(scope, receive, send)


app = JWTGate(mcp.streamable_http_app())
if __name__ == "__main__":
    uvicorn.run(app, host=os.getenv("HOST", "127.0.0.1"), port=int(os.getenv("PORT", "8080")))
