Skip to content
Draft
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
1 change: 1 addition & 0 deletions agent/api.go
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@ func (a *agent) apiHandler() http.Handler {
r.Use(
httpmw.Recover(a.logger),
tracing.StatusWriterMiddleware,
tracing.SessionIDMiddleware,
loggermw.Logger(a.logger, nil),
agentchat.Middleware,
)
Expand Down
14 changes: 14 additions & 0 deletions coderd/tracing/httpmw.go
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,20 @@ func sessionIDFromHeaders(h http.Header) string {
return id
}

// SessionIDMiddleware reads the session_id baggage member from the request and
// adds it to the log context so downstream request logs can be correlated by
// session. Unlike Middleware, it does not create spans, emit telemetry, or gate
// on route patterns. It is intended for the agent, per the connection-log RFC,
// which for now only requires the session ID on the log context.
func SessionIDMiddleware(next http.Handler) http.Handler {
return http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
if sessionID := sessionIDFromHeaders(r.Header); sessionID != "" {
r = r.WithContext(slog.With(r.Context(), slog.F("session_id", sessionID)))
}
next.ServeHTTP(rw, r)
})
}

// validSessionID reports whether s is a 32-character lowercase hexadecimal
// string (a 16-byte value), the encoding the RFC mandates for the session ID.
// Only lowercase is accepted so that case-sensitive searches correlate
Expand Down
108 changes: 108 additions & 0 deletions coderd/tracing/httpmw_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@ import (
"go.opentelemetry.io/otel/trace/noop"

"cdr.dev/slog/v3"
"github.com/coder/coder/v2/coderd/httpmw/loggermw"
"github.com/coder/coder/v2/coderd/tracing"
"github.com/coder/coder/v2/testutil"
)
Expand Down Expand Up @@ -227,6 +228,113 @@ func Test_Middleware_SessionID(t *testing.T) {
})
}

func Test_SessionIDMiddleware(t *testing.T) {
t.Parallel()

// downstreamFields runs a request through SessionIDMiddleware and returns
// the fields a downstream handler logs using the request context.
downstreamFields := func(t *testing.T, header string) []slog.Field {
t.Helper()

sink := testutil.NewFakeSink(t)
logger := sink.Logger()

handler := tracing.SessionIDMiddleware(http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) {
logger.Info(r.Context(), "downstream handler invoked")
rw.WriteHeader(http.StatusNoContent)
}))

r := httptest.NewRequest(http.MethodGet, "/api/v0/foo", nil)
if header != "" {
r.Header.Set("baggage", header)
}
handler.ServeHTTP(httptest.NewRecorder(), r)

entries := sink.Entries(func(e slog.SinkEntry) bool {
return e.Message == "downstream handler invoked"
})
require.Len(t, entries, 1)
return entries[0].Fields
}

fieldValue := func(fields []slog.Field, name string) (any, bool) {
for _, f := range fields {
if f.Name == name {
return f.Value, true
}
}
return nil, false
}

t.Run("ValidBaggage", func(t *testing.T) {
t.Parallel()

val, ok := fieldValue(downstreamFields(t, tracing.SessionIDBaggageKey+"="+testSessionID), "session_id")
require.True(t, ok, "session_id should be on the log context")
require.Equal(t, testSessionID, val)
})

t.Run("NoBaggage", func(t *testing.T) {
t.Parallel()

_, ok := fieldValue(downstreamFields(t, ""), "session_id")
require.False(t, ok, "session_id should be absent when no baggage is sent")
})

t.Run("MalformedBaggage", func(t *testing.T) {
t.Parallel()

_, ok := fieldValue(downstreamFields(t, tracing.SessionIDBaggageKey+"=not-a-valid-session-id"), "session_id")
require.False(t, ok, "malformed session_id should be ignored")
})

t.Run("UppercaseRejected", func(t *testing.T) {
t.Parallel()

upper := strings.ToUpper(testSessionID)
_, ok := fieldValue(downstreamFields(t, tracing.SessionIDBaggageKey+"="+upper), "session_id")
require.False(t, ok, "uppercase session_id should be rejected")
})
}

// Test_SessionIDMiddleware_AccessLog verifies that, wired in the same order as
// the agent middleware stack, the session_id lands on loggermw's request
// completion log line, not just on downstream handler logs.
func Test_SessionIDMiddleware_AccessLog(t *testing.T) {
t.Parallel()

sink := testutil.NewFakeSink(t)

// StatusWriterMiddleware is required by loggermw; SessionIDMiddleware runs
// before loggermw so the field is on the request context when the access
// log is emitted. This mirrors agent/api.go.
handler := tracing.StatusWriterMiddleware(
tracing.SessionIDMiddleware(
loggermw.Logger(sink.Logger(), nil)(
http.HandlerFunc(func(rw http.ResponseWriter, _ *http.Request) {
rw.WriteHeader(http.StatusNoContent)
}),
),
),
)

r := httptest.NewRequest(http.MethodGet, "/api/v0/foo", nil)
r.Header.Set("baggage", tracing.SessionIDBaggageKey+"="+testSessionID)
handler.ServeHTTP(httptest.NewRecorder(), r)

entries := sink.Entries()
require.Len(t, entries, 1)

var found bool
for _, f := range entries[0].Fields {
if f.Name == "session_id" {
found = true
require.Equal(t, testSessionID, f.Value)
}
}
require.True(t, found, "session_id should be on the request access log")
}

func Test_Middleware(t *testing.T) {
t.Parallel()

Expand Down
Loading