Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
67 changes: 63 additions & 4 deletions cli/internal/loopback/loopback.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,9 @@ package loopback

import (
"context"
"crypto/rand"
"crypto/subtle"
"encoding/base64"
"errors"
"fmt"
"net"
Expand All @@ -30,6 +33,41 @@ const (
// tab doesn't sit on a blank page — signet itself has already moved on.
const callbackResponse = "Signet: link received. You can close this tab.\n"

// rejectedResponse is served to a callback whose state does not match. It is
// deliberately not an invitation to try again — the only party that should be
// hitting this port already knows the state.
const rejectedResponse = "Signet: this request was not expected. Ignoring it.\n"

// Accept decides whether a callback is the one being waited for. Returning
// false rejects that request *without* ending the wait.
type Accept func(url.Values) bool

// NewState returns a cryptographically random value to thread through the
// approval link and back in the callback.
func NewState() (string, error) {
buf := make([]byte, 32)
if _, err := rand.Read(buf); err != nil {
return "", fmt.Errorf("loopback: generating state: %w", err)
}
return base64.RawURLEncoding.EncodeToString(buf), nil
}

// MatchState is the Accept that RFC 8252's loopback-redirect hardening calls
// for: the callback must carry exactly the state that went out in the link.
//
// While `signet link` is listening, *any* page the developer's browser visits
// can issue requests to this port. Without this check a hostile page could
// hand the CLI its own payload — linking an attacker-chosen account — or
// simply hit the port to abort the developer's link. Compared in constant time
// out of habit rather than need: the state is not a secret an attacker gets to
// guess a byte at a time, but a timing-safe compare costs nothing here.
func MatchState(expected string) Accept {
return func(values url.Values) bool {
got := values.Get("state")
return subtle.ConstantTimeCompare([]byte(got), []byte(expected)) == 1
}
}

// Server is a one-shot loopback HTTP server. Create with New, read Port to
// build the callback URL, then call Wait exactly once to block until that
// callback arrives (or ctx is cancelled, or the timeout elapses).
Expand Down Expand Up @@ -79,10 +117,23 @@ func (s *Server) Close() error {
//
// The returned url.Values are the callback request's query parameters.
func (s *Server) Wait(ctx context.Context, timeout time.Duration) (url.Values, error) {
return s.WaitFor(ctx, timeout, func(url.Values) bool { return true })
}

// WaitFor is Wait, but only a callback that `accept` approves ends it.
//
// A rejected callback is answered and discarded, and the server keeps
// listening — the wait must survive a hostile or stray request rather than
// being terminated by one, which is the whole point of #256. It still serves
// at most one *accepted* callback.
func (s *Server) WaitFor(ctx context.Context, timeout time.Duration, accept Accept) (url.Values, error) {
if s.waited {
return nil, ErrAlreadyWaiting
}
s.waited = true
if accept == nil {
accept = func(url.Values) bool { return true }
}

type result struct {
values url.Values
Expand All @@ -92,15 +143,23 @@ func (s *Server) Wait(ctx context.Context, timeout time.Duration) (url.Values, e

mux := http.NewServeMux()
mux.HandleFunc(s.path, func(w http.ResponseWriter, r *http.Request) {
values := r.URL.Query()
w.Header().Set("content-type", "text/plain; charset=utf-8")

if !accept(values) {
w.WriteHeader(http.StatusBadRequest)
_, _ = fmt.Fprint(w, rejectedResponse)
return
}

w.WriteHeader(http.StatusOK)
_, _ = fmt.Fprint(w, callbackResponse)
select {
case done <- result{values: r.URL.Query()}:
case done <- result{values: values}:
default:
// A second request raced the first here — first one wins, this
// one's result is simply dropped; the shutdown below tears down
// the listener regardless.
// A second accepted request raced the first here — first one
// wins, this one's result is simply dropped; the shutdown below
// tears down the listener regardless.
}
})

Expand Down
117 changes: 117 additions & 0 deletions cli/internal/loopback/loopback_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,8 +2,10 @@ package loopback

import (
"context"
"errors"
"net"
"net/http"
"net/url"
"strings"
"testing"
"time"
Expand Down Expand Up @@ -151,3 +153,118 @@ func TestWait_SecondCallIsAnError(t *testing.T) {
t.Fatalf("got %v, want ErrAlreadyWaiting", err)
}
}

// ── state verification (RFC 8252 loopback-redirect hardening) ────────────

func TestNewState_IsRandomAndUrlSafe(t *testing.T) {
a, err := NewState()
if err != nil {
t.Fatalf("NewState: %v", err)
}
b, err := NewState()
if err != nil {
t.Fatalf("NewState: %v", err)
}
if a == b {
t.Fatal("NewState returned the same value twice")
}
if len(a) < 32 {
t.Fatalf("state %q is too short to be unguessable", a)
}
if strings.ContainsAny(a, "+/=&?#") {
t.Fatalf("state %q is not safe to put in a URL unescaped", a)
}
}

func TestMatchState(t *testing.T) {
accept := MatchState("expected")
if !accept(url.Values{"state": {"expected"}}) {
t.Fatal("rejected the matching state")
}
for _, wrong := range []url.Values{
{"state": {"other"}},
{"state": {"expected "}},
{"state": {"EXPECTED"}},
{"state": {""}},
{},
} {
if accept(wrong) {
t.Fatalf("accepted %v", wrong)
}
}
}

func TestWaitFor_RejectsMismatchedStateWithoutEndingTheWait(t *testing.T) {
s, err := New("/callback")
if err != nil {
t.Fatalf("New: %v", err)
}
target := s.URL()

go func() {
// A hostile page hits the port first with its own state. The wait
// must survive this — the whole attack is aborting or hijacking the
// developer's link.
time.Sleep(20 * time.Millisecond)
_ = getAndDiscard(http.DefaultClient, target+"?state=attacker&code=evil")
time.Sleep(20 * time.Millisecond)
_ = getAndDiscard(http.DefaultClient, target+"?state=expected&code=real")
}()

values, err := s.WaitFor(context.Background(), 5*time.Second, MatchState("expected"))
if err != nil {
t.Fatalf("WaitFor: %v", err)
}
if values.Get("code") != "real" {
t.Fatalf("returned the attacker's callback: %v", values)
}
}

func TestWaitFor_MismatchedStateGetsA400(t *testing.T) {
s, err := New("/callback")
if err != nil {
t.Fatalf("New: %v", err)
}
target := s.URL()

codes := make(chan int, 1)
go func() {
time.Sleep(20 * time.Millisecond)
resp, err := http.Get(target + "?state=wrong")
if err == nil {
codes <- resp.StatusCode
_ = resp.Body.Close()
} else {
codes <- 0
}
time.Sleep(20 * time.Millisecond)
_ = getAndDiscard(http.DefaultClient, target+"?state=right")
}()

if _, err := s.WaitFor(context.Background(), 5*time.Second, MatchState("right")); err != nil {
t.Fatalf("WaitFor: %v", err)
}
if got := <-codes; got != http.StatusBadRequest {
t.Fatalf("mismatched callback got %d, want 400", got)
}
}

func TestWaitFor_StillTimesOutIfOnlyMismatchesArrive(t *testing.T) {
s, err := New("/callback")
if err != nil {
t.Fatalf("New: %v", err)
}
target := s.URL()

go func() {
for i := 0; i < 3; i++ {
time.Sleep(10 * time.Millisecond)
_ = getAndDiscard(http.DefaultClient, target+"?state=wrong")
}
}()

_, err = s.WaitFor(context.Background(), 150*time.Millisecond, MatchState("right"))
if !errors.Is(err, context.DeadlineExceeded) {
t.Fatalf("err = %v, want a deadline error", err)
}
}
Loading