diff --git a/pay/cmd/paywatch/main.go b/pay/cmd/paywatch/main.go index 845f0ec..cdab577 100644 --- a/pay/cmd/paywatch/main.go +++ b/pay/cmd/paywatch/main.go @@ -114,7 +114,7 @@ func main() { defer stop() go w.Loop(ctx, time.Duration(pollSec)*time.Second) - srv := &http.Server{Addr: addr, Handler: httpapi.New(svc), ReadHeaderTimeout: 10 * time.Second} + srv := &http.Server{Addr: addr, Handler: httpapi.New(svc, env("PAY_CORS_ORIGINS", "https://pangolin.yanmeiai.com,https://pangolin-site.pages.dev")), ReadHeaderTimeout: 10 * time.Second} go func() { log.Info("pangolin-pay listening", "addr", addr, "poll_seconds", pollSec) if err := srv.ListenAndServe(); err != nil && err != http.ErrServerClosed { diff --git a/pay/internal/httpapi/cors.go b/pay/internal/httpapi/cors.go new file mode 100644 index 0000000..2da07ca --- /dev/null +++ b/pay/internal/httpapi/cors.go @@ -0,0 +1,40 @@ +package httpapi + +import ( + "net/http" + "strings" +) + +// withCORS lets the browser storefront (pangolin website) call the order API +// cross-origin. Auth is server-to-server / bearer-less here (no cookies), so +// Allow-Credentials is not needed. Only whitelisted origins get CORS headers; +// OPTIONS preflight is answered directly. +func withCORS(next http.Handler, allowed map[string]bool) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + origin := r.Header.Get("Origin") + if origin != "" && allowed[origin] { + h := w.Header() + h.Set("Access-Control-Allow-Origin", origin) + h.Add("Vary", "Origin") + h.Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") + h.Set("Access-Control-Allow-Headers", "Content-Type") + h.Set("Access-Control-Max-Age", "600") + } + if r.Method == http.MethodOptions { + w.WriteHeader(http.StatusNoContent) + return + } + next.ServeHTTP(w, r) + }) +} + +// parseOrigins builds the allowed-origin set from a comma-separated list. +func parseOrigins(csv string) map[string]bool { + m := map[string]bool{} + for _, o := range strings.Split(csv, ",") { + if o = strings.TrimSpace(o); o != "" { + m[o] = true + } + } + return m +} diff --git a/pay/internal/httpapi/handler.go b/pay/internal/httpapi/handler.go index eb18947..d8c2809 100644 --- a/pay/internal/httpapi/handler.go +++ b/pay/internal/httpapi/handler.go @@ -14,8 +14,9 @@ import ( type Handler struct{ svc *pay.Service } -// New wires the routes (Go 1.22 method+wildcard patterns). -func New(svc *pay.Service) http.Handler { +// New wires the routes (Go 1.22 method+wildcard patterns) and applies CORS for +// the whitelisted browser origins (comma-separated; empty = no CORS headers). +func New(svc *pay.Service, corsOrigins string) http.Handler { h := &Handler{svc: svc} mux := http.NewServeMux() mux.HandleFunc("POST /order", h.createOrder) @@ -24,7 +25,7 @@ func New(svc *pay.Service) http.Handler { w.WriteHeader(http.StatusOK) _, _ = w.Write([]byte("ok")) }) - return mux + return withCORS(mux, parseOrigins(corsOrigins)) } type createReq struct { diff --git a/pay/internal/httpapi/handler_test.go b/pay/internal/httpapi/handler_test.go index c593ea7..7b22a57 100644 --- a/pay/internal/httpapi/handler_test.go +++ b/pay/internal/httpapi/handler_test.go @@ -16,7 +16,7 @@ const recvAddr = "TRecv00000000000000000000000000000A" func TestCreateGetAndConflict(t *testing.T) { st, _ := store.Open(":memory:") t.Cleanup(func() { _ = st.Close() }) - srv := httptest.NewServer(New(pay.New(st, pay.Config{ReceiveAddress: recvAddr}))) + srv := httptest.NewServer(New(pay.New(st, pay.Config{ReceiveAddress: recvAddr}), "https://pangolin.yanmeiai.com")) t.Cleanup(srv.Close) body, _ := json.Marshal(map[string]any{"user_ref": "u1", "sku": "pro-year", "amount": 5_000000})