diff --git a/internal/confidential/frame_guard.go b/internal/confidential/frame_guard.go
index d007ccc..68bc7e6 100644
--- a/internal/confidential/frame_guard.go
+++ b/internal/confidential/frame_guard.go
@@ -16,6 +16,7 @@ import (
const (
maxEncryptedResponseFrame = 64 << 20
maxProblemResponseBody = 64 << 10
+ maxUnencryptedExcerpt = 512
)
type guardingTransport struct {
@@ -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
diff --git a/internal/confidential/frame_guard_test.go b/internal/confidential/frame_guard_test.go
index 317d3bb..08e686b 100644
--- a/internal/confidential/frame_guard_test.go
+++ b/internal/confidential/frame_guard_test.go
@@ -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("\n
Rate limited\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",
+ " Rate limited",
+ } {
+ 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)))
diff --git a/internal/proxy/handler_integration_test.go b/internal/proxy/handler_integration_test.go
index aae3b85..0ab70e7 100644
--- a/internal/proxy/handler_integration_test.go
+++ b/internal/proxy/handler_integration_test.go
@@ -8,6 +8,7 @@ import (
"encoding/json"
"errors"
"io"
+ "log"
"net/http"
"net/http/httptest"
"regexp"
@@ -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 {