package main

import (
	"context"
	"encoding/json"
	"errors"
	"fmt"
	"io"
	"log"
	"net/http"
	"net/url"
	"os"
	"regexp"
	"strings"
	"time"

	"github.com/golang-jwt/jwt/v5"
	"github.com/modelcontextprotocol/go-sdk/mcp"
)

type OrderInput struct {
	OrderID string `json:"order_id" jsonschema:"Order ID, using letters, numbers, or hyphens"`
}
type Order struct {
	ID     string `json:"id"`
	Status string `json:"status"`
}

func envOr(key, fallback string) string {
	if value := os.Getenv(key); value != "" {
		return value
	}
	return fallback
}
func main() {
	apiSecret, mcpSecret := os.Getenv("API_JWT_SECRET"), os.Getenv("MCP_JWT_SECRET")
	if len(apiSecret) < 32 || len(mcpSecret) < 32 || apiSecret == mcpSecret {
		log.Fatal("Use separate API and MCP signing secrets of at least 32 characters")
	}
	apiBase, err := url.Parse(envOr("API_BASE_URL", "http://127.0.0.1:8081"))
	if err != nil || apiBase.Host == "" {
		log.Fatal("Set a valid API_BASE_URL")
	}
	origin, err := url.Parse(envOr("PUBLIC_ORIGIN", "http://127.0.0.1:8080"))
	if err != nil || origin.Host == "" {
		log.Fatal("Set a valid PUBLIC_ORIGIN")
	}
	client := &http.Client{Timeout: 5 * time.Second, CheckRedirect: func(*http.Request, []*http.Request) error { return errors.New("redirects are not allowed") }}
	validID := regexp.MustCompile(`^[a-zA-Z0-9-]{1,64}$`)
	type principalKey struct{}
	createServer := func(subject string) *mcp.Server {
		server := mcp.NewServer(&mcp.Implementation{Name: "orders-mcp", Version: "1.0.0"}, nil)
		mcp.AddTool(server, &mcp.Tool{Name: "get_order", Description: "Read an order's shipping status from the Orders API.", Annotations: &mcp.ToolAnnotations{ReadOnlyHint: true}},
			func(ctx context.Context, _ *mcp.CallToolRequest, input OrderInput) (*mcp.CallToolResult, Order, error) {
				if !validID.MatchString(input.OrderID) {
					return nil, Order{}, errors.New("invalid order ID")
				}
				// Carry the verified caller in a new JWT intended for the API.
				now := time.Now()
				apiToken, err := jwt.NewWithClaims(jwt.SigningMethodHS256, jwt.MapClaims{
					"iss": "orders-demo", "aud": "orders-api", "sub": subject,
					"scope": "orders:read", "iat": now.Unix(), "exp": now.Add(5 * time.Minute).Unix(),
				}).SignedString([]byte(apiSecret))
				if err != nil {
					return nil, Order{}, errors.New("could not authorize API request")
				}
				endpoint := apiBase.ResolveReference(&url.URL{Path: "/orders/" + input.OrderID})
				req, err := http.NewRequestWithContext(ctx, http.MethodGet, endpoint.String(), nil)
				if err != nil {
					return nil, Order{}, errors.New("could not prepare API request")
				}
				req.Header.Set(envOr("API_AUTH_HEADER", "Authorization"), "Bearer "+apiToken)
				res, err := client.Do(req)
				if err != nil {
					return nil, Order{}, errors.New("could not read order")
				}
				defer res.Body.Close()
				var order Order
				if res.StatusCode != http.StatusOK || json.NewDecoder(io.LimitReader(res.Body, 65536)).Decode(&order) != nil || order.ID == "" || order.Status == "" {
					return nil, Order{}, errors.New("could not read order; check its ID and API access")
				}
				return nil, order, nil
			})
		return server
	}
	handler := mcp.NewStreamableHTTPHandler(func(r *http.Request) *mcp.Server {
		subject, _ := r.Context().Value(principalKey{}).(string)
		return createServer(subject)
	}, &mcp.StreamableHTTPOptions{Stateless: true, JSONResponse: true})
	mux := http.NewServeMux()
	mux.HandleFunc("GET /healthz", func(w http.ResponseWriter, _ *http.Request) {
		w.Header().Set("Content-Type", "application/json")
		fmt.Fprint(w, `{"ok":true}`)
	})
	mux.HandleFunc("/mcp", func(w http.ResponseWriter, r *http.Request) {
		if (r.Host != origin.Host && r.Host != "127.0.0.1:8080" && r.Host != "localhost:8080") || (r.Header.Get("Origin") != "" && r.Header.Get("Origin") != origin.Scheme+"://"+origin.Host) {
			http.Error(w, "Invalid host or origin", http.StatusForbidden)
			return
		}
		header := r.Header.Get(envOr("MCP_AUTH_HEADER", "Authorization"))
		if !strings.HasPrefix(header, "Bearer ") {
			http.Error(w, "Missing access token", http.StatusUnauthorized)
			return
		}
		token, err := jwt.Parse(strings.TrimPrefix(header, "Bearer "), func(*jwt.Token) (any, error) { return []byte(mcpSecret), nil },
			jwt.WithValidMethods([]string{"HS256"}), jwt.WithIssuer("orders-demo"), jwt.WithAudience("orders-mcp"), jwt.WithExpirationRequired())
		if err != nil || !token.Valid {
			http.Error(w, "Invalid or expired access token", http.StatusUnauthorized)
			return
		}
		claims, ok := token.Claims.(jwt.MapClaims)
		subject, subjectErr := claims.GetSubject()
		if !ok || subjectErr != nil || subject == "" {
			http.Error(w, "Missing subject", http.StatusUnauthorized)
			return
		}
		scopes, _ := claims["scope"].(string)
		permitted := false
		for _, scope := range strings.Fields(scopes) {
			if scope == "orders:read" {
				permitted = true
			}
		}
		if !permitted {
			http.Error(w, "orders:read permission required", http.StatusForbidden)
			return
		}
		r.Body = http.MaxBytesReader(w, r.Body, 65536)
		handler.ServeHTTP(w, r.WithContext(context.WithValue(r.Context(), principalKey{}, subject)))
	})
	httpServer := &http.Server{Addr: envOr("HOST", "127.0.0.1") + ":" + envOr("PORT", "8080"), Handler: mux, ReadHeaderTimeout: 5 * time.Second, IdleTimeout: 60 * time.Second}
	log.Fatal(httpServer.ListenAndServe())
}
