Skip to content
Open
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
44 changes: 42 additions & 2 deletions internal/confidential/frame_guard.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ import (
const (
maxEncryptedResponseFrame = 64 << 20
maxProblemResponseBody = 64 << 10
maxUnencryptedExcerpt = 512
)

type guardingTransport struct {
Expand Down Expand Up @@ -47,13 +48,52 @@ func (t *guardingTransport) RoundTrip(request *http.Request) (*http.Response, er
if len(responseNonces) == 1 && response.Body != nil {
response.Body = &frameGuardReader{source: response.Body}
}
if response.StatusCode == http.StatusUnprocessableEntity &&
isMediaType(response.Header.Get("Content-Type"), protocol.ProblemJSONMediaType) && response.Body != nil {
keyMismatch := response.StatusCode == http.StatusUnprocessableEntity &&
isMediaType(response.Header.Get("Content-Type"), protocol.ProblemJSONMediaType)
if keyMismatch && response.Body != nil {
response.Body = &boundedReadCloser{source: response.Body, remaining: maxProblemResponseBody}
}
// An encrypted request whose reply has no nonce was answered by something
// in front of the Gateway's EHBP handler (an edge, a load balancer or an
// early rejection). EHBP would only report the missing header, so describe
// what actually answered. The key-mismatch reply is the one exception: the
// EHBP client needs it to re-verify.
if len(responseNonces) == 0 && !keyMismatch && request.Header.Get(protocol.EncapsulatedKeyHeader) != "" {
return nil, unencryptedReplyError(response)
}
return response, nil
}

// unencryptedReplyError consumes and closes response. It never sees request
// or response plaintext: the reply was not encrypted to begin with.
func unencryptedReplyError(response *http.Response) error {
var excerpt []byte
if response.Body != nil {
excerpt, _ = io.ReadAll(io.LimitReader(response.Body, maxUnencryptedExcerpt))
_ = response.Body.Close()
}
details := []string{}
for _, name := range []string{"Content-Type", "Server", "Retry-After", "Cf-Ray"} {
if value := response.Header.Get(name); value != "" {
details = append(details, fmt.Sprintf("%s=%s", name, singleLine(value)))
}
}
message := fmt.Sprintf("Gateway replied without EHBP encryption: HTTP %d", response.StatusCode)
if len(details) > 0 {
message += " (" + strings.Join(details, ", ") + ")"
}
if text := singleLine(string(excerpt)); text != "" {
message += ": " + text
}
return errors.New(message)
}

func singleLine(value string) string {
return strings.Join(strings.FieldsFunc(value, func(r rune) bool {
return r < 0x20 || r == 0x7f || r == ' '
}), " ")
}

func canonicalResponseNonce(encoded string) bool {
if len(encoded) != 64 {
return false
Expand Down
108 changes: 108 additions & 0 deletions internal/confidential/frame_guard_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -111,6 +111,114 @@ func TestGuardRejectsMalformedOrDuplicateResponseNonce(t *testing.T) {
}
}

func TestGuardDescribesUnencryptedReplyToEncryptedRequest(t *testing.T) {
closed := false
response := &http.Response{
StatusCode: http.StatusTooManyRequests,
Header: http.Header{
"Content-Type": {"text/html"},
"Server": {"cloudflare"},
"Retry-After": {"5"},
"Cf-Ray": {"8f00aa-MAD"},
},
Body: &closeRecorder{
Reader: strings.NewReader("<html>\n <title>Rate limited</title>\n" + strings.Repeat("x", 2*maxUnencryptedExcerpt)),
closed: &closed,
},
}
transport := GuardEHBPResponses(&staticRoundTripper{response: response})
_, err := transport.RoundTrip(encryptedRequest(t))
if err == nil {
t.Fatal("RoundTrip() error = nil, want unencrypted reply error")
}
for _, want := range []string{
"without EHBP encryption",
"HTTP 429",
"Content-Type=text/html",
"Server=cloudflare",
"Retry-After=5",
"Cf-Ray=8f00aa-MAD",
"<html> <title>Rate limited</title>",
} {
if !strings.Contains(err.Error(), want) {
t.Errorf("error %q does not contain %q", err, want)
}
}
if strings.Contains(err.Error(), "\n") {
t.Errorf("error %q spans several lines", err)
}
if len(err.Error()) > 2*maxUnencryptedExcerpt {
t.Errorf("error is %d bytes, want the body excerpt capped", len(err.Error()))
}
if !closed {
t.Error("response body was not closed")
}
}

func TestGuardPassesKeyMismatchAndUnencryptedRequests(t *testing.T) {
tests := []struct {
name string
response *http.Response
request func(*testing.T) *http.Request
}{
{
name: "key mismatch keeps EHBP re-verification",
response: &http.Response{
StatusCode: http.StatusUnprocessableEntity,
Header: http.Header{"Content-Type": {protocol.ProblemJSONMediaType}},
Body: io.NopCloser(strings.NewReader(`{"title":"key configuration mismatch"}`)),
},
request: encryptedRequest,
},
{
name: "request that was never encrypted",
response: &http.Response{
StatusCode: http.StatusOK,
Header: http.Header{"Content-Type": {"application/json"}},
Body: io.NopCloser(strings.NewReader(`{}`)),
},
request: func(t *testing.T) *http.Request {
request, err := http.NewRequest(http.MethodGet, "https://gateway.example/v1/models", nil)
if err != nil {
t.Fatal(err)
}
return request
},
},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
transport := GuardEHBPResponses(&staticRoundTripper{response: test.response})
got, err := transport.RoundTrip(test.request(t))
if err != nil {
t.Fatalf("RoundTrip() error = %v, want the response passed through", err)
}
if got.StatusCode != test.response.StatusCode {
t.Fatalf("status = %d, want %d", got.StatusCode, test.response.StatusCode)
}
})
}
}

func encryptedRequest(t *testing.T) *http.Request {
request, err := http.NewRequest(http.MethodPost, "https://gateway.example", strings.NewReader("ciphertext"))
if err != nil {
t.Fatal(err)
}
request.Header.Set(protocol.EncapsulatedKeyHeader, strings.Repeat("ab", 32))
return request
}

type closeRecorder struct {
io.Reader
closed *bool
}

func (r *closeRecorder) Close() error {
*r.closed = true
return nil
}

func frame(payload []byte) []byte {
result := make([]byte, 4+len(payload))
binary.BigEndian.PutUint32(result, uint32(len(payload)))
Expand Down
50 changes: 50 additions & 0 deletions internal/proxy/handler_integration_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -8,6 +8,7 @@ import (
"encoding/json"
"errors"
"io"
"log"
"net/http"
"net/http/httptest"
"regexp"
Expand Down Expand Up @@ -576,6 +577,55 @@ func TestOutboundGatewayRedirectIsNeverFollowed(t *testing.T) {
}
}

func TestUnencryptedGatewayReplyIsLoggedWithItsStatus(t *testing.T) {
serverIdentity, err := identity.NewIdentity()
if err != nil {
t.Fatal(err)
}
edge := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
_, _ = io.Copy(io.Discard, r.Body)
w.Header().Set("Content-Type", "text/plain")
w.Header().Set("Retry-After", "5")
w.WriteHeader(http.StatusTooManyRequests)
_, _ = io.WriteString(w, "rate limited at the edge")
}))
defer edge.Close()

wireClient := &http.Client{
Transport: confidential.GuardEHBPResponses(edge.Client().Transport),
CheckRedirect: func(_ *http.Request, _ []*http.Request) error {
return errors.New("redirects are not allowed")
},
}
client, err := confidential.NewClient(edge.URL+attestation.ConfidentialEndpoint,
&fakeEvidenceVerifier{key: serverIdentity.MarshalPublicKey()}, wireClient)
if err != nil {
t.Fatal(err)
}
var logs bytes.Buffer
handler, err := NewHandler(client, log.New(&logs, "", 0))
if err != nil {
t.Fatal(err)
}
local := httptest.NewServer(handler)
defer local.Close()

response := postJSON(t, local.URL+LocalChatEndpoint, `{"prompt":"`+promptCanary+`"}`, "Bearer key")
defer response.Body.Close()
_, _ = io.Copy(io.Discard, response.Body)
if response.StatusCode != http.StatusBadGateway {
t.Fatalf("local response status = %d, want 502", response.StatusCode)
}
for _, want := range []string{"without EHBP encryption", "HTTP 429", "Retry-After=5", "rate limited at the edge"} {
if !strings.Contains(logs.String(), want) {
t.Errorf("proxy log %q does not contain %q", logs.String(), want)
}
}
if strings.Contains(logs.String(), promptCanary) {
t.Fatalf("proxy log leaked the prompt: %q", logs.String())
}
}

func TestKeyRotationFailsWithoutReplayAndReattestsNextCallerRequest(t *testing.T) {
activeIdentity, err := identity.NewIdentity()
if err != nil {
Expand Down