agent endpoints: build the front and registration from params

Front is an http.Handler, so fallback, identity and singleKeyFallback were read
on every request with no synchronization while WithFallback, WithIdentity and
WithSingleKeyFallback wrote them on the live object. Each returned the same
pointer, so the chain read as though it built a value. Every call site happens
to run before the handler is mounted, so nothing races today, but nothing in
the type prevents it, and WithSingleKeyFallback took no argument and could only
be turned on. NewFront now takes FrontParams and configuration is fixed at
construction.

NewWorkerRegisterer takes its EndpointSettingsFunc directly, and
HandshakeAgentWorker takes one in place of a variadic of raw closures that
existed only to reach the setter it replaced.

Registration gains NewRegistration and RegistrationParams, moving the Draining
callback off an exported mutable field and folding SetSession into
construction. IsDraining absorbs the nil check at both call sites.

Access carried a three-state ladder as two bools, with "granted implies
credentialed" documented but unenforced. It is now an ordered AccessLevel, so
callers compare a rank rather than combining flags and the invariant holds by
construction.

Registry.Register returned an error that was always nil, with a dead branch at
each call site. Both registry maps and the per-registration session are
read-heavy, so they take RWMutex.

CopyBody returned two errors to separate a source failure from a destination
one. It returns one, wrapping a source failure in *SourceError, which is what
the caller discriminates on.

NewWebTransportServer took a callback to break the handler/server init cycle;
the caller assigns wt.H3.Handler after construction instead. StartWebTransport
reads Development off the service rather than taking a bool, and returns a nil
stop function where it starts no listener.

Smaller: slices.Sort for sort.Strings, for range for an unused counter, a nil
slice for regs[:0:0], p2c generic over its slice so the call site passes a
method expression rather than allocating a closure per request, and streamCode
deduplicated into wire.StreamCode so both peers map reset codes in one place.

Comments on the touched declarations drop remote behavior, migration narration
and contrastive framing, keeping the constraints and invariants.

Co-authored-by: Claude Opus 5 (1M context) <noreply@anthropic.com>
This commit is contained in:
Paul Wells
2026-09-14 17:56:52 -07:00
co-authored by Claude Opus 5
parent 9ac10fba0e
commit d9a643243a
20 changed files with 342 additions and 291 deletions
+6 -6
View File
@@ -180,13 +180,13 @@ func (a *attempt) writeRequest(w io.Writer) error {
bufp := a.pools.getBuf()
defer a.pools.putBuf(bufp)
n, srcErr, dstErr := wire.CopyBody(bw, a.body, *bufp)
if dstErr != nil {
return dstErr
}
n, err := wire.CopyBody(bw, a.body, *bufp)
var srcErr *wire.SourceError
switch {
case srcErr != nil:
return bw.Close(wire.CompletionPeerGone, srcErr.Error())
case errors.As(err, &srcErr):
return bw.Close(wire.CompletionPeerGone, srcErr.Err.Error())
case err != nil:
return err
case cl > 0 && n < cl:
return bw.Close(wire.CompletionTruncated, fmt.Sprintf("declared %d bytes, read %d", cl, n))
}
+3 -4
View File
@@ -130,11 +130,10 @@ func rawScriptedFront(t *testing.T, script func(r io.Reader, w *workerSide)) *ht
require.NoError(t, err)
reg := NewRegistry()
r := &Registration{WorkerID: "w1", APIKey: "proj", AgentName: "a", Deployment: "d", Manifest: m}
r.SetSession(&scriptedSession{script: script})
require.NoError(t, reg.Register(r))
r := NewRegistration(RegistrationParams{WorkerID: "w1", APIKey: "proj", AgentName: "a", Deployment: "d", Manifest: m, Session: &scriptedSession{script: script}})
reg.Register(r)
ts := httptest.NewServer(NewFront(reg, grantedTo("proj"), logger.GetLogger()))
ts := httptest.NewServer(NewFront(FrontParams{Registry: reg, ResolveAccess: grantedTo("proj"), Logger: logger.GetLogger()}))
t.Cleanup(ts.Close)
return ts
}
+4 -8
View File
@@ -267,7 +267,7 @@ func (w *Worker) serve(ctx context.Context, stream *webtransport.Stream) {
// the deadline and the session ending both have to reach a blocked copy
stop := context.AfterFunc(ctx, func() {
_ = conn.Close()
stream.CancelRead(streamCode(livekit.AgentHttp_HSR_ABORT))
stream.CancelRead(wire.StreamCode(livekit.AgentHttp_HSR_ABORT))
})
defer stop()
@@ -290,17 +290,13 @@ func pipe(stream *webtransport.Stream, conn net.Conn) {
_, _ = io.Copy(stream, conn)
// the target is done answering, so nothing more of the request is wanted
stream.CancelRead(streamCode(livekit.AgentHttp_HSR_ABORT))
stream.CancelRead(wire.StreamCode(livekit.AgentHttp_HSR_ABORT))
<-reqDone
}
func streamCode(c livekit.AgentHttp_HttpStreamResetCode) webtransport.StreamErrorCode {
return webtransport.StreamErrorCode(c)
}
// resetStream reports an outcome that happened before any HTTP bytes flowed,
// the only point at which a reset can carry one without racing them.
func resetStream(stream *webtransport.Stream, c livekit.AgentHttp_HttpStreamResetCode) {
stream.CancelWrite(streamCode(c))
stream.CancelRead(streamCode(c))
stream.CancelWrite(wire.StreamCode(c))
stream.CancelRead(wire.StreamCode(c))
}
+82 -81
View File
@@ -18,6 +18,7 @@ import (
"bufio"
"context"
"errors"
"fmt"
"io"
"math/rand/v2"
"net/http"
@@ -67,16 +68,41 @@ var (
errHeadTooLarge = errors.New("endpoint: response head too large")
)
// AccessLevel is how far a request's caller is trusted. Callers compare against
// it, so a new level must be inserted at its correct rank.
type AccessLevel int
const (
// AccessNone presented no credential.
AccessNone AccessLevel = iota
// AccessCredentialed presented a valid token carrying no agent-endpoint
// grant for the addressed agent and deployment.
AccessCredentialed
// AccessGranted presented a token whose agent-endpoint grant covers the
// addressed agent and deployment.
AccessGranted
)
func (a AccessLevel) String() string {
switch a {
case AccessNone:
return "none"
case AccessCredentialed:
return "credentialed"
case AccessGranted:
return "granted"
default:
return fmt.Sprintf("%d", int(a))
}
}
// Access is what the front knows about a request's caller, for the agent and
// deployment its URL addresses. Granted implies Credentialed.
// deployment its URL addresses.
type Access struct {
// APIKey is the registry scope the request is served from; empty means the
// request cannot be placed.
APIKey string
// Credentialed selects 401 over 403 for a denied request.
Credentialed bool
// Granted opens non-public routes.
Granted bool
Level AccessLevel
}
// AccessResolver maps an inbound request, plus the agent and deployment its URL
@@ -84,39 +110,36 @@ type Access struct {
type AccessResolver func(r *http.Request, agentName, deployment string) Access
type Front struct {
registry *Registry
resolveAccess AccessResolver
logger logger.Logger
pools *bridgePools
// fallback is consulted when nothing local can serve the request (no
// candidates, no route match, or every match without capacity); a
// multi-node deployment plugs its resolve-and-relay here. nil means local
// misses are final.
fallback Fallback
// see WithSingleKeyFallback
singleKeyFallback bool
// see WithIdentity
identity Identity
params FrontParams
pools *bridgePools
}
// Identity resolves the agent and deployment a request addresses. Reporting
// false leaves them to the URL.
type Identity func(r *http.Request) (agentName, deployment string, ok bool)
// WithIdentity resolves the agent and deployment from the request itself.
func (f *Front) WithIdentity(fn Identity) *Front {
f.identity = fn
return f
// FrontParams configures a Front. Fields are read on every request once the
// Front is serving, so none may change after construction.
type FrontParams struct {
Registry *Registry
ResolveAccess AccessResolver
Logger logger.Logger
// Fallback is consulted when nothing local can serve the request: no
// candidates, no route match, or every match without capacity. nil means
// local misses are final.
Fallback Fallback
Identity Identity
// SingleKeyFallback resolves unauthenticated requests to the registry's
// single api key when the resolver yields none. The key comes from the
// registry, so this is sound only where every registration belongs to one
// tenant.
SingleKeyFallback bool
}
func NewFront(registry *Registry, resolveAccess AccessResolver, log logger.Logger) *Front {
return &Front{
registry: registry,
resolveAccess: resolveAccess,
logger: log.WithComponent("agents.endpoint"),
pools: newBridgePools(),
}
func NewFront(params FrontParams) *Front {
params.Logger = params.Logger.WithComponent("agents.endpoint")
return &Front{params: params, pools: newBridgePools()}
}
// FallbackRequest describes a request nothing local could serve. The request
@@ -132,22 +155,6 @@ type FallbackRequest struct {
// status mapping.
type Fallback func(w http.ResponseWriter, r *http.Request, req *FallbackRequest) bool
// WithFallback installs the miss handler consulted when nothing local can
// serve a request.
func (f *Front) WithFallback(fb Fallback) *Front {
f.fallback = fb
return f
}
// WithSingleKeyFallback resolves unauthenticated requests to the registry's
// single api key when the resolver yields none. Self-hosted convenience only: a
// multi-tenant front must never guess an api key from what happens to be
// registered.
func (f *Front) WithSingleKeyFallback() *Front {
f.singleKeyFallback = true
return f
}
// writeUnavailable writes a 503 with a Retry-After hint: no local worker can
// serve the request and no fallback placed it elsewhere.
func (f *Front) writeUnavailable(w http.ResponseWriter, msg string) {
@@ -166,9 +173,9 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return
}
agentName, deployment, path, escPath := ep.agentName, ep.deployment, ep.path, ep.escPath
if f.identity != nil {
if f.params.Identity != nil {
// must precede resolveAccess, which may consume its source headers
if name, dep, ok := f.identity(r); ok {
if name, dep, ok := f.params.Identity(r); ok {
agentName, deployment = name, dep
}
}
@@ -179,11 +186,11 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return
}
access := f.resolveAccess(r, agentName, deployment)
if access.APIKey == "" && f.singleKeyFallback {
access := f.params.ResolveAccess(r, agentName, deployment)
if access.APIKey == "" && f.params.SingleKeyFallback {
// unauthenticated: OSS serves public routes when the worker fleet
// belongs to a single key. A guessed api key confers no access.
access.APIKey, _ = f.registry.SingleAPIKey()
access.APIKey, _ = f.params.Registry.SingleAPIKey()
}
if access.APIKey == "" {
w.Header().Set("WWW-Authenticate", "Bearer")
@@ -191,8 +198,8 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
return
}
candidates := f.registry.Candidates(access.APIKey, agentName, deployment)
if len(candidates) == 0 && f.fallback == nil {
candidates := f.params.Registry.Candidates(access.APIKey, agentName, deployment)
if len(candidates) == 0 && f.params.Fallback == nil {
f.writeUnavailable(w, "no workers available for deployment")
return
}
@@ -204,7 +211,7 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
rt, res := reg.Manifest.Match(p, r.Method)
switch res {
case MatchFull:
if !access.Granted && !rt.Public {
if access.Level < AccessGranted && !rt.Public {
denied = true
continue
}
@@ -232,12 +239,12 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
matched, route, partial, denied = matchAll(path)
}
}
if route == nil && f.fallback != nil {
if route == nil && f.params.Fallback != nil {
// nothing local matched: hand off to the multi-node fallback (relay to a
// node holding the deployment) before the local status mapping. The
// serving node's relay listener installs no fallback of its own, so a
// relayed request is served or errored there and never re-relays.
if f.fallback(w, r, &FallbackRequest{
if f.params.Fallback(w, r, &FallbackRequest{
Access: access,
AgentName: agentName, Deployment: deployment,
}) {
@@ -252,7 +259,7 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
switch {
case denied:
// access does not vary across candidates, so one verdict covers them all
if access.Credentialed {
if access.Level >= AccessCredentialed {
http.Error(w, "forbidden", http.StatusForbidden)
} else {
w.Header().Set("WWW-Authenticate", "Bearer")
@@ -273,7 +280,7 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
target: requestTarget(escPath, r.URL.RawQuery),
route: route,
requestID: reqID,
granted: access.Granted,
granted: access.Level >= AccessGranted,
pools: f.pools,
}
a.body = &countingReader{r: r.Body, n: &bodyConsumed}
@@ -284,13 +291,13 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
http.Error(w, "request header fields too large", http.StatusRequestHeaderFieldsTooLarge)
return
}
f.logger.Debugw("agent endpoint rejected a request head", "error", err, "requestID", reqID)
f.params.Logger.Debugw("agent endpoint rejected a request head", "error", err, "requestID", reqID)
http.Error(w, "bad request", http.StatusBadRequest)
return
}
attempted := make(map[*Registration]bool)
for i := 0; i < maxAttempts; i++ {
for range maxAttempts {
reg := pickWorker(matched, attempted)
if reg == nil {
break
@@ -315,8 +322,8 @@ func (f *Front) ServeHTTP(w http.ResponseWriter, r *http.Request) {
// capacity elsewhere. Safe exactly while no request bytes were consumed -
// reaching this point implies it, since consuming attempts are never
// retryable.
if bodyConsumed.Load() == 0 && f.fallback != nil {
if f.fallback(w, r, &FallbackRequest{
if bodyConsumed.Load() == 0 && f.params.Fallback != nil {
if f.params.Fallback(w, r, &FallbackRequest{
Access: access,
AgentName: agentName, Deployment: deployment,
}) {
@@ -394,18 +401,12 @@ func splitEndpointPath(u *url.URL) (endpointPath, error) {
}
// pickWorker chooses a worker by the power of two choices: sample two eligible
// registrations at random and take the one with fewer in-flight streams (least
// outstanding requests). This approximates optimal load spreading without global
// coordination or the herding of exact least-loaded, and - unlike a
// self-reported load - the in-flight count is observed here from the worker's
// own session. Eligible = not already attempted, has a live session, not draining.
// registrations at random and take the one with fewer in-flight streams.
// Eligible = not already attempted, has a live session, not draining.
func pickWorker(regs []*Registration, ignore map[*Registration]bool) *Registration {
eligible := regs[:0:0]
var eligible []*Registration
for _, reg := range regs {
if ignore[reg] || !reg.HasSession() {
continue
}
if reg.Draining != nil && reg.Draining() {
if ignore[reg] || !reg.HasSession() || reg.IsDraining() {
continue
}
eligible = append(eligible, reg)
@@ -413,13 +414,13 @@ func pickWorker(regs []*Registration, ignore map[*Registration]bool) *Registrati
if len(eligible) == 0 {
return nil
}
return eligible[p2c(len(eligible), func(i int) int { return eligible[i].InflightStreams() })]
return eligible[p2c(eligible, (*Registration).InflightStreams)]
}
// p2c returns the index of the less-loaded of two distinct random draws from
// [0,n) (n >= 1). With n == 2 both are always sampled, so it is exact; larger n
// trades a little optimality for O(1) work and no herding.
func p2c(n int, load func(int) int) int {
// items, which must be non-empty.
func p2c[T any](items []T, load func(T) int) int {
n := len(items)
if n == 1 {
return 0
}
@@ -428,7 +429,7 @@ func p2c(n int, load func(int) int) int {
if j >= i { // fold to a distinct second draw
j++
}
if load(i) <= load(j) {
if load(items[i]) <= load(items[j]) {
return i
}
return j
@@ -482,7 +483,7 @@ func (f *Front) bridge(w http.ResponseWriter, a *attempt, reg *Registration) bri
if err != nil {
err = completionError(err)
if !a.retryable(err) {
f.logger.Warnw("agent endpoint request failed", err,
f.params.Logger.Warnw("agent endpoint request failed", err,
"workerID", reg.WorkerID, "path", a.escPath, "requestID", a.requestID)
writeGatewayError(w, err)
return bridgeDone
@@ -535,7 +536,7 @@ func (f *Front) aborted(err error, reg *Registration, a *attempt, writeErrCh <-c
f.logAborted(err, reg, a)
select {
case werr := <-writeErrCh:
f.logger.Debugw("request write result after response failure", "error", werr)
f.params.Logger.Debugw("request write result after response failure", "error", werr)
default:
}
return bridgeAbort
@@ -544,12 +545,12 @@ func (f *Front) aborted(err error, reg *Registration, a *attempt, writeErrCh <-c
func (f *Front) logAborted(err error, reg *Registration, a *attempt) {
var ce *wire.CompletionError
if errors.As(err, &ce) {
f.logger.Infow("agent endpoint response aborted",
f.params.Logger.Infow("agent endpoint response aborted",
"workerID", reg.WorkerID, "path", a.escPath, "requestID", a.requestID,
"completion", string(ce.Completion), "reason", ce.Reason)
return
}
f.logger.Infow("agent endpoint response aborted",
f.params.Logger.Infow("agent endpoint response aborted",
"workerID", reg.WorkerID, "path", a.escPath, "requestID", a.requestID, "error", err)
}
+24 -26
View File
@@ -20,7 +20,7 @@ import (
// grantedTo resolves every request to apiKey with full access.
func grantedTo(apiKey string) AccessResolver {
return func(*http.Request, string, string) Access {
return Access{APIKey: apiKey, Credentialed: true, Granted: true}
return Access{APIKey: apiKey, Level: AccessGranted}
}
}
@@ -31,15 +31,15 @@ func fallbackFront(t *testing.T, fb Fallback, withWorker bool) *Front {
{Path: "/known", Methods: []string{"GET"}, Public: true},
})
require.NoError(t, err)
r := &Registration{WorkerID: "w1", APIKey: "proj", AgentName: "a", Deployment: "d", Manifest: m}
r.SetSession(&fakeSession{})
require.NoError(t, reg.Register(r))
r := NewRegistration(RegistrationParams{WorkerID: "w1", APIKey: "proj", AgentName: "a", Deployment: "d", Manifest: m, Session: &fakeSession{}})
reg.Register(r)
}
f := NewFront(reg, grantedTo("proj"), logger.GetLogger())
if fb != nil {
f = f.WithFallback(fb)
}
return f
return NewFront(FrontParams{
Registry: reg,
ResolveAccess: grantedTo("proj"),
Logger: logger.GetLogger(),
Fallback: fb,
})
}
func serveFront(f *Front, path string) *httptest.ResponseRecorder {
@@ -62,8 +62,7 @@ func TestFrontFallbackFires(t *testing.T) {
require.Equal(t, http.StatusTeapot, w.Code)
require.NotNil(t, got)
require.Equal(t, "proj", got.APIKey)
require.True(t, got.Granted)
require.True(t, got.Credentialed)
require.Equal(t, AccessGranted, got.Level)
require.Equal(t, "a", got.AgentName)
require.Equal(t, "d", got.Deployment)
}
@@ -183,22 +182,22 @@ func accessFront(t *testing.T, a Access, fb Fallback) *Front {
{Path: "/private", Methods: []string{"GET"}, Public: false},
})
require.NoError(t, err)
r := &Registration{WorkerID: "w1", APIKey: "proj", AgentName: "a", Deployment: "d", Manifest: m}
r.SetSession(&fakeSession{})
require.NoError(t, reg.Register(r))
r := NewRegistration(RegistrationParams{WorkerID: "w1", APIKey: "proj", AgentName: "a", Deployment: "d", Manifest: m, Session: &fakeSession{}})
reg.Register(r)
f := NewFront(reg, func(*http.Request, string, string) Access { return a }, logger.GetLogger())
if fb != nil {
f = f.WithFallback(fb)
}
return f
return NewFront(FrontParams{
Registry: reg,
ResolveAccess: func(*http.Request, string, string) Access { return a },
Logger: logger.GetLogger(),
Fallback: fb,
})
}
// fakeSession opens no stream, so a request that clears authorization reaches 503.
func TestFrontPrivateRouteAccessMapping(t *testing.T) {
anonymous := Access{APIKey: "proj"}
credentialed := Access{APIKey: "proj", Credentialed: true}
granted := Access{APIKey: "proj", Credentialed: true, Granted: true}
anonymous := Access{APIKey: "proj", Level: AccessNone}
credentialed := Access{APIKey: "proj", Level: AccessCredentialed}
granted := Access{APIKey: "proj", Level: AccessGranted}
t.Run("anonymous is challenged", func(t *testing.T) {
w := serveFront(accessFront(t, anonymous, nil), "/private")
@@ -223,14 +222,14 @@ func TestFrontPrivateRouteAccessMapping(t *testing.T) {
// the slash-normalized form of a private route is still private.
func TestFrontDeniedAppliesToNormalizedPath(t *testing.T) {
f := accessFront(t, Access{APIKey: "proj", Credentialed: true}, nil)
f := accessFront(t, Access{APIKey: "proj", Level: AccessCredentialed}, nil)
require.Equal(t, http.StatusForbidden, serveFront(f, "/private/").Code)
}
// another node's worker may declare the same path public.
func TestFrontDeniedStillRelays(t *testing.T) {
var got *FallbackRequest
f := accessFront(t, Access{APIKey: "proj", Credentialed: true}, func(w http.ResponseWriter, _ *http.Request, fr *FallbackRequest) bool {
f := accessFront(t, Access{APIKey: "proj", Level: AccessCredentialed}, func(w http.ResponseWriter, _ *http.Request, fr *FallbackRequest) bool {
got = fr
w.WriteHeader(http.StatusTeapot)
return true
@@ -238,8 +237,7 @@ func TestFrontDeniedStillRelays(t *testing.T) {
require.Equal(t, http.StatusTeapot, serveFront(f, "/private").Code)
require.NotNil(t, got)
require.True(t, got.Credentialed)
require.False(t, got.Granted)
require.Equal(t, AccessCredentialed, got.Level)
}
// the split runs before decoding, so a name or route param may carry any byte
+5 -5
View File
@@ -10,18 +10,18 @@ import (
func TestP2CChoice(t *testing.T) {
// a single candidate is always chosen
require.Equal(t, 0, p2c(1, func(int) int { return 42 }))
require.Equal(t, 0, p2c([]int{42}, func(int) int { return 42 }))
// with two candidates both are always drawn, so the lower-load one wins
// deterministically regardless of the random draw
for i := 0; i < 100; i++ {
require.Equal(t, 1, p2c(2, func(i int) int { return []int{5, 2}[i] }))
require.Equal(t, 0, p2c(2, func(i int) int { return []int{2, 5}[i] }))
require.Equal(t, 1, p2c([]int{5, 2}, func(v int) int { return v }))
require.Equal(t, 0, p2c([]int{2, 5}, func(v int) int { return v }))
}
// larger n: the pick is always in range
for i := 0; i < 500; i++ {
idx := p2c(5, func(int) int { return 0 })
idx := p2c(make([]int, 5), func(int) int { return 0 })
require.GreaterOrEqual(t, idx, 0)
require.Less(t, idx, 5)
}
@@ -32,7 +32,7 @@ func TestP2CChoice(t *testing.T) {
loads := []int{0, 100, 100, 100, 100}
low := 0
for i := 0; i < 4000; i++ {
if p2c(len(loads), func(i int) int { return loads[i] }) == 0 {
if p2c(loads, func(v int) int { return v }) == 0 {
low++
}
}
+51 -29
View File
@@ -63,27 +63,50 @@ type Registration struct {
Deployment string
Manifest *Manifest
// Draining is provided by the control-plane layer that owns the worker; a
// draining worker takes no new streams. Worker selection uses live in-flight
// streams, not a reported load.
Draining func() bool
draining func() bool
mu sync.Mutex
lock sync.RWMutex
session Session
closed bool
}
// SetSession attaches the worker's live data-plane session. One session per
// worker: the WebTransport session that also carries its control stream.
func (r *Registration) SetSession(s Session) {
r.mu.Lock()
r.session = s
r.mu.Unlock()
// RegistrationParams is fixed for the life of the registration, which lasts
// exactly as long as the session.
type RegistrationParams struct {
WorkerID string
APIKey string
AgentName string
Deployment string
Manifest *Manifest
// Session is the worker's live data-plane session: the WebTransport session
// that also carries its control stream. One session per worker.
Session Session
// Draining reports that the worker is shedding; a shedding worker takes no
// new streams.
Draining func() bool
}
func NewRegistration(params RegistrationParams) *Registration {
return &Registration{
WorkerID: params.WorkerID,
APIKey: params.APIKey,
AgentName: params.AgentName,
Deployment: params.Deployment,
Manifest: params.Manifest,
draining: params.Draining,
session: params.Session,
}
}
// IsDraining is false when no drain signal was supplied.
func (r *Registration) IsDraining() bool {
return r.draining != nil && r.draining()
}
func (r *Registration) getSession() Session {
r.mu.Lock()
defer r.mu.Unlock()
r.lock.RLock()
defer r.lock.RUnlock()
if r.closed {
return nil
}
@@ -129,15 +152,15 @@ func (r *Registration) SpareStreams() int {
}
func (r *Registration) close() {
r.mu.Lock()
r.lock.Lock()
if r.closed {
r.mu.Unlock()
r.lock.Unlock()
return
}
r.closed = true
s := r.session
r.session = nil
r.mu.Unlock()
r.lock.Unlock()
if s != nil {
s.Close("registration closed")
}
@@ -146,7 +169,7 @@ func (r *Registration) close() {
// Registry tracks data-plane registrations on this node, keyed by
// (api key, agent name, deployment). The api key is the project identity in OSS.
type Registry struct {
mu sync.Mutex
lock sync.RWMutex
regs map[string]*Registration // by worker id
byKey map[regKey][]*Registration
}
@@ -168,23 +191,22 @@ func NewRegistry() *Registry {
// worker ids are stable across reconnects, and the retiring session must not be
// able to strand the new epoch (its own Deregister is a no-op once replaced).
// The superseded epoch's session is closed.
func (g *Registry) Register(r *Registration) error {
func (g *Registry) Register(r *Registration) {
key := regKey{r.APIKey, r.AgentName, normalizeDeployment(r.Deployment)}
g.mu.Lock()
g.lock.Lock()
old := g.regs[r.WorkerID]
if old != nil {
g.removeLocked(old)
}
g.regs[r.WorkerID] = r
g.byKey[key] = append(g.byKey[key], r)
g.mu.Unlock()
g.lock.Unlock()
if old != nil {
old.close()
}
return nil
}
// removeLocked unlinks a registration from all indexes. Callers hold g.mu.
// removeLocked unlinks a registration from all indexes. Callers hold g.lock.
func (g *Registry) removeLocked(r *Registration) {
delete(g.regs, r.WorkerID)
key := regKey{r.APIKey, r.AgentName, normalizeDeployment(r.Deployment)}
@@ -203,20 +225,20 @@ func (g *Registry) removeLocked(r *Registration) {
// Deregister removes exactly this registration; it is a no-op when a newer
// epoch has already superseded it.
func (g *Registry) Deregister(r *Registration) {
g.mu.Lock()
g.lock.Lock()
if g.regs[r.WorkerID] != r {
g.mu.Unlock()
g.lock.Unlock()
return
}
g.removeLocked(r)
g.mu.Unlock()
g.lock.Unlock()
r.close()
}
// Candidates returns the registrations for (api key, agent name, deployment segment).
func (g *Registry) Candidates(apiKey, agentName, deployment string) []*Registration {
g.mu.Lock()
defer g.mu.Unlock()
g.lock.RLock()
defer g.lock.RUnlock()
return slices.Clone(g.byKey[regKey{apiKey, agentName, normalizeDeployment(deployment)}])
}
@@ -224,8 +246,8 @@ func (g *Registry) Candidates(apiKey, agentName, deployment string) []*Registrat
// resolution for unauthenticated requests to public endpoints. ok is false when
// zero or multiple keys are present.
func (g *Registry) SingleAPIKey() (string, bool) {
g.mu.Lock()
defer g.mu.Unlock()
g.lock.RLock()
defer g.lock.RUnlock()
var key string
for _, r := range g.regs {
if key == "" {
+12 -14
View File
@@ -33,19 +33,18 @@ func TestRegistrySupersede(t *testing.T) {
require.NoError(t, err)
mk := func() (*Registration, *fakeSession) {
r := &Registration{
s := &fakeSession{}
return NewRegistration(RegistrationParams{
WorkerID: "AW_1", APIKey: "key",
AgentName: "agent", Deployment: "production", Manifest: manifest,
}
s := &fakeSession{}
r.SetSession(s)
return r, s
Session: s,
}), s
}
oldReg, oldSess := mk()
require.NoError(t, g.Register(oldReg))
g.Register(oldReg)
newReg, newSess := mk()
require.NoError(t, g.Register(newReg))
g.Register(newReg)
require.Equal(t, []*Registration{newReg}, g.Candidates("key", "agent", "production"))
require.True(t, oldSess.closed, "superseded epoch's session must be closed")
@@ -71,19 +70,18 @@ func TestRegistryAgentScoping(t *testing.T) {
require.NoError(t, err)
mk := func(workerID, agentName, deployment string) *Registration {
r := &Registration{
return NewRegistration(RegistrationParams{
WorkerID: workerID, APIKey: "key",
AgentName: agentName, Deployment: deployment, Manifest: manifest,
}
r.SetSession(&fakeSession{})
return r
Session: &fakeSession{},
})
}
a := mk("AW_a", "alpha", "production")
b := mk("AW_b", "beta", "production")
staging := mk("AW_c", "alpha", "staging")
require.NoError(t, g.Register(a))
require.NoError(t, g.Register(b))
require.NoError(t, g.Register(staging))
g.Register(a)
g.Register(b)
g.Register(staging)
require.Equal(t, []*Registration{a}, g.Candidates("key", "alpha", "production"))
require.Equal(t, []*Registration{b}, g.Candidates("key", "beta", "production"))
+8 -4
View File
@@ -72,10 +72,14 @@ func startFramedWorker(t *testing.T, targetAddr string, eps []*livekit.AgentHttp
reg := endpoint.NewRegistry()
base := startWTServer(t, reg)
front := endpoint.NewFront(reg, func(*http.Request, string, string) endpoint.Access {
return endpoint.Access{}
}, logger.GetLogger()).
WithSingleKeyFallback()
front := endpoint.NewFront(endpoint.FrontParams{
Registry: reg,
ResolveAccess: func(*http.Request, string, string) endpoint.Access {
return endpoint.Access{}
},
Logger: logger.GetLogger(),
SingleKeyFallback: true,
})
ts := httptest.NewUnstartedServer(front)
// raised past net/http's 1 MiB default so the front's own head bound is what
// rejects an oversized head
+2 -9
View File
@@ -89,7 +89,7 @@ func (s *wtStream) Reset(code livekit.AgentHttp_HttpStreamResetCode, _ string) {
s.mu.Lock()
s.sendClosed = true
s.mu.Unlock()
c := streamCode(code)
c := wire.StreamCode(code)
s.qs.CancelWrite(c)
s.qs.CancelRead(c)
s.release()
@@ -105,7 +105,7 @@ func (s *wtStream) Close() error {
cancelWrite := !s.sendClosed
s.sendClosed = true
s.mu.Unlock()
abort := streamCode(livekit.AgentHttp_HSR_ABORT)
abort := wire.StreamCode(livekit.AgentHttp_HSR_ABORT)
if cancelWrite {
s.qs.CancelWrite(abort)
}
@@ -126,13 +126,6 @@ func (s *wtStream) release() {
s.sess.open.Add(-1)
}
// streamCode converts a protocol reset code to the WebTransport code that
// carries it. HSR_ABORT is zero, so a teardown with nothing to say sends the
// plain cancel code.
func streamCode(c livekit.AgentHttp_HttpStreamResetCode) webtransport.StreamErrorCode {
return webtransport.StreamErrorCode(c)
}
// translateStreamError turns a peer reset into the protocol's own error. Only a
// remote reset carries meaning: cancelling this side says nothing about what the
// worker did with the request.
+27 -14
View File
@@ -86,15 +86,15 @@ func handleSession(reg *endpoint.Registry, sess *webtransport.Session) {
if err != nil {
return
}
registration := &endpoint.Registration{
registration := endpoint.NewRegistration(endpoint.RegistrationParams{
WorkerID: rw.GetInstanceId(),
APIKey: "test",
AgentName: rw.GetAgentName(),
Deployment: rw.GetDeployment(),
Manifest: manifest,
}
registration.SetSession(endpoint.NewWebTransportSession(sess, endpoint.DefaultMaxStreams))
_ = reg.Register(registration)
Session: endpoint.NewWebTransportSession(sess, endpoint.DefaultMaxStreams),
})
reg.Register(registration)
_ = wire.WriteControlMessage(control, &livekit.ServerMessage{
Message: &livekit.ServerMessage_Register{
Register: &livekit.RegisterWorkerResponse{
@@ -136,10 +136,14 @@ func TestWebTransportEndpointRoundTrip(t *testing.T) {
reg := endpoint.NewRegistry()
base := startWTServer(t, reg)
front := endpoint.NewFront(reg, func(*http.Request, string, string) endpoint.Access {
return endpoint.Access{}
}, logger.GetLogger()).
WithSingleKeyFallback()
front := endpoint.NewFront(endpoint.FrontParams{
Registry: reg,
ResolveAccess: func(*http.Request, string, string) endpoint.Access {
return endpoint.Access{}
},
Logger: logger.GetLogger(),
SingleKeyFallback: true,
})
ts := httptest.NewServer(front)
defer ts.Close()
@@ -203,14 +207,23 @@ func TestWebTransportPrivateEndpointRequiresGrant(t *testing.T) {
reg := endpoint.NewRegistry()
base := startWTServer(t, reg)
anonymous := httptest.NewServer(endpoint.NewFront(reg, func(*http.Request, string, string) endpoint.Access {
return endpoint.Access{}
}, logger.GetLogger()).WithSingleKeyFallback())
anonymous := httptest.NewServer(endpoint.NewFront(endpoint.FrontParams{
Registry: reg,
ResolveAccess: func(*http.Request, string, string) endpoint.Access {
return endpoint.Access{}
},
Logger: logger.GetLogger(),
SingleKeyFallback: true,
}))
defer anonymous.Close()
granted := httptest.NewServer(endpoint.NewFront(reg, func(*http.Request, string, string) endpoint.Access {
return endpoint.Access{APIKey: "test", Credentialed: true, Granted: true}
}, logger.GetLogger()))
granted := httptest.NewServer(endpoint.NewFront(endpoint.FrontParams{
Registry: reg,
ResolveAccess: func(*http.Request, string, string) endpoint.Access {
return endpoint.Access{APIKey: "test", Level: endpoint.AccessGranted}
},
Logger: logger.GetLogger(),
}))
defer granted.Close()
ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second)
+19 -8
View File
@@ -18,7 +18,7 @@ import (
"fmt"
"io"
"net/http"
"sort"
"slices"
"strconv"
"strings"
@@ -198,7 +198,7 @@ func BuildRequestHead(method, target, host string, h http.Header, contentLength
}
keys = append(keys, k)
}
sort.Strings(keys) // deterministic across attempts and across nodes
slices.Sort(keys) // deterministic across attempts and across nodes
for _, k := range keys {
if !httpguts.ValidHeaderFieldName(k) {
return nil, fmt.Errorf("endpoint: invalid header name %q", k)
@@ -309,23 +309,34 @@ func sanitizeReason(s string) string {
return s
}
// SourceError reports that the body's source failed mid-copy. The stream is
// still intact, so the caller can close the body with a completion.
type SourceError struct{ Err error }
func (e *SourceError) Error() string {
return fmt.Sprintf("endpoint: body source failed: %s", e.Err)
}
func (e *SourceError) Unwrap() error { return e.Err }
// CopyBody pumps src into dst one write per read, so a streaming body stays
// incremental. srcErr and dstErr are separate because only a source failure
// leaves a stream to report the outcome on.
func CopyBody(dst BodyWriter, src io.Reader, buf []byte) (n int64, srcErr, dstErr error) {
// incremental. A src failure is wrapped in *SourceError; any other error is
// from dst.
func CopyBody(dst BodyWriter, src io.Reader, buf []byte) (int64, error) {
var n int64
for {
nr, rerr := src.Read(buf)
if nr > 0 {
if _, werr := dst.Write(buf[:nr]); werr != nil {
return n, nil, werr
return n, werr
}
n += int64(nr)
}
if rerr == io.EOF {
return n, nil, nil
return n, nil
}
if rerr != nil {
return n, rerr, nil
return n, &SourceError{Err: rerr}
}
}
}
+12 -11
View File
@@ -17,6 +17,7 @@ package wire
import (
"bufio"
"bytes"
"errors"
"io"
"net/http"
"strings"
@@ -233,9 +234,8 @@ func TestCopyBodySplitReads(t *testing.T) {
src := bytes.Repeat([]byte("abcdefgh"), 1024)
var buf bytes.Buffer
bw := NewIdentityBody(&buf)
n, srcErr, dstErr := CopyBody(bw, oneByteReader{bytes.NewReader(src)}, make([]byte, 512))
require.NoError(t, srcErr)
require.NoError(t, dstErr)
n, err := CopyBody(bw, oneByteReader{bytes.NewReader(src)}, make([]byte, 512))
require.NoError(t, err)
require.EqualValues(t, len(src), n)
require.Equal(t, src, buf.Bytes())
}
@@ -244,16 +244,17 @@ type failingWriter struct{}
func (failingWriter) Write([]byte) (int, error) { return 0, io.ErrClosedPipe }
// A stream failure and a source failure are different outcomes: one can still be
// reported to the peer, the other cannot.
// Only a source failure leaves the stream able to report the outcome.
func TestCopyBodySeparatesSourceAndSinkFailures(t *testing.T) {
_, srcErr, dstErr := CopyBody(NewIdentityBody(failingWriter{}), strings.NewReader("xxxx"), make([]byte, 2))
require.NoError(t, srcErr)
require.Error(t, dstErr)
// a destination failure is returned unwrapped
_, err := CopyBody(NewIdentityBody(failingWriter{}), strings.NewReader("xxxx"), make([]byte, 2))
require.Error(t, err)
var srcErr *SourceError
require.False(t, errors.As(err, &srcErr))
_, srcErr, dstErr = CopyBody(NewIdentityBody(io.Discard), iotestErrReader{}, make([]byte, 2))
require.Error(t, srcErr)
require.NoError(t, dstErr)
// a source failure is wrapped, so the caller can still close the body
_, err = CopyBody(NewIdentityBody(io.Discard), iotestErrReader{}, make([]byte, 2))
require.ErrorAs(t, err, &srcErr)
}
type iotestErrReader struct{}
+28
View File
@@ -0,0 +1,28 @@
// Copyright 2026 LiveKit, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package wire
import (
"github.com/quic-go/webtransport-go"
"github.com/livekit/protocol/livekit"
)
// StreamCode converts a protocol reset code to the WebTransport code that
// carries it. HSR_ABORT is zero, so a teardown with nothing to say sends the
// plain cancel code.
func StreamCode(c livekit.AgentHttp_HttpStreamResetCode) webtransport.StreamErrorCode {
return webtransport.StreamErrorCode(c)
}
+2 -8
View File
@@ -195,22 +195,16 @@ type WorkerRegisterer struct {
registered bool
}
func NewWorkerRegisterer(conn SignalConn, serverInfo *livekit.ServerInfo, base WorkerRegistration) *WorkerRegisterer {
func NewWorkerRegisterer(conn SignalConn, serverInfo *livekit.ServerInfo, base WorkerRegistration, endpointSettings EndpointSettingsFunc) *WorkerRegisterer {
return &WorkerRegisterer{
WorkerPingHandler: WorkerPingHandler{conn: conn},
serverInfo: serverInfo,
registration: base,
deadline: time.Now().Add(RegisterTimeout),
endpointSettings: endpointSettings,
}
}
// WithEndpointSettings enables the HTTP endpoints data plane for registrations that
// declare endpoints.
func (h *WorkerRegisterer) WithEndpointSettings(f EndpointSettingsFunc) *WorkerRegisterer {
h.endpointSettings = f
return h
}
func (h *WorkerRegisterer) Deadline() time.Time {
return h.deadline
}
+3 -3
View File
@@ -46,8 +46,8 @@ func TestHandleRegisterEndpointAgentNames(t *testing.T) {
for _, c := range cases {
t.Run(c.name, func(t *testing.T) {
h := agent.NewWorkerRegisterer(nopSignalConn{}, &livekit.ServerInfo{}, agent.WorkerRegistration{}).
WithEndpointSettings(func(*livekit.RegisterWorkerRequest) (*livekit.AgentHttp_AgentEndpointSettings, error) {
h := agent.NewWorkerRegisterer(nopSignalConn{}, &livekit.ServerInfo{}, agent.WorkerRegistration{},
func(*livekit.RegisterWorkerRequest) (*livekit.AgentHttp_AgentEndpointSettings, error) {
return &livekit.AgentHttp_AgentEndpointSettings{}, nil
})
@@ -68,7 +68,7 @@ func TestHandleRegisterEndpointAgentNames(t *testing.T) {
// the reserved names constrain only workers that declare endpoints.
func TestHandleRegisterWithoutEndpointsIgnoresReservedNames(t *testing.T) {
for _, name := range []string{"", "_", ".", ".."} {
h := agent.NewWorkerRegisterer(nopSignalConn{}, &livekit.ServerInfo{}, agent.WorkerRegistration{})
h := agent.NewWorkerRegisterer(nopSignalConn{}, &livekit.ServerInfo{}, agent.WorkerRegistration{}, nil)
require.NoError(t, h.HandleRegister(&livekit.RegisterWorkerRequest{
Type: livekit.JobType_JT_ROOM,
AgentName: name,
+26 -28
View File
@@ -115,11 +115,8 @@ func DispatchAgentWorkerSignal(c agent.SignalConn, h agent.WorkerSignalHandler,
return true
}
func HandshakeAgentWorker(c agent.SignalConn, serverInfo *livekit.ServerInfo, registration agent.WorkerRegistration, l logger.Logger, opts ...func(*agent.WorkerRegisterer)) (r agent.WorkerRegistration, ok bool) {
wr := agent.NewWorkerRegisterer(c, serverInfo, registration)
for _, opt := range opts {
opt(wr)
}
func HandshakeAgentWorker(c agent.SignalConn, serverInfo *livekit.ServerInfo, registration agent.WorkerRegistration, l logger.Logger, endpointSettings agent.EndpointSettingsFunc) (r agent.WorkerRegistration, ok bool) {
wr := agent.NewWorkerRegisterer(c, serverInfo, registration, endpointSettings)
if err := c.SetReadDeadline(wr.Deadline()); err != nil {
return
}
@@ -138,6 +135,9 @@ type AgentService struct {
upgrader AgentSocketUpgrader
signalMessageSizeLimit int64
// developmentMode allows the WebTransport listener to run without a
// configured certificate.
developmentMode bool
*AgentHandler
}
@@ -187,6 +187,7 @@ func NewAgentService(
) (*AgentService, error) {
s := &AgentService{
signalMessageSizeLimit: conf.Limit.AgentSignalMessageSizeLimit,
developmentMode: conf.Development,
}
serverInfo := &livekit.ServerInfo{
@@ -227,20 +228,22 @@ func NewAgentService(
// token is present; a non-public route additionally requires an agent-endpoint
// grant scoped to this agent and deployment.
func (s *AgentService) EndpointFront() http.Handler {
front := endpoint.NewFront(s.endpointRegistry, func(r *http.Request, agentName, deployment string) endpoint.Access {
if claims := GetGrants(r.Context()); claims != nil {
return endpoint.Access{
APIKey: GetAPIKey(r.Context()),
Credentialed: true,
Granted: claims.AgentEndpoint.Allows(agentName, deployment),
return endpoint.NewFront(endpoint.FrontParams{
Registry: s.endpointRegistry,
ResolveAccess: func(r *http.Request, agentName, deployment string) endpoint.Access {
if claims := GetGrants(r.Context()); claims != nil {
level := endpoint.AccessCredentialed
if claims.AgentEndpoint.Allows(agentName, deployment) {
level = endpoint.AccessGranted
}
return endpoint.Access{APIKey: GetAPIKey(r.Context()), Level: level}
}
}
// unauthenticated: with a single configured key the api key is
// unambiguous even when this node holds no registrations (multi-node)
return endpoint.Access{APIKey: s.singleAPIKey}
}, s.logger)
front.WithSingleKeyFallback()
return front
// unauthenticated: one configured key makes the api key unambiguous
return endpoint.Access{APIKey: s.singleAPIKey}
},
Logger: s.logger,
SingleKeyFallback: true,
})
}
func (s *AgentService) ServeHTTP(w http.ResponseWriter, r *http.Request) {
@@ -321,9 +324,7 @@ func (h *AgentHandler) HandleConnection(ctx context.Context, conn agent.SignalCo
// exchanges share it); nil for a WebSocket control connection, which serves no
// endpoints.
func (h *AgentHandler) handleConnection(ctx context.Context, conn agent.SignalConn, registration agent.WorkerRegistration, sess endpoint.Session) {
registration, ok := HandshakeAgentWorker(conn, h.serverInfo, registration, h.logger, func(wr *agent.WorkerRegisterer) {
wr.WithEndpointSettings(h.endpointSettings)
})
registration, ok := HandshakeAgentWorker(conn, h.serverInfo, registration, h.logger, h.endpointSettings)
if !ok {
return
}
@@ -363,19 +364,16 @@ func (h *AgentHandler) registerEndpoints(w *agent.Worker, sess endpoint.Session)
w.Logger().Errorw("endpoint manifest failed to re-parse", err)
return nil
}
reg := &endpoint.Registration{
reg := endpoint.NewRegistration(endpoint.RegistrationParams{
WorkerID: w.ID,
APIKey: w.APIKey(),
AgentName: w.AgentName,
Deployment: w.Deployment,
Manifest: manifest,
Session: sess,
Draining: w.Draining,
}
reg.SetSession(sess)
if err := h.endpointRegistry.Register(reg); err != nil {
w.Logger().Errorw("failed to register endpoints", err)
return nil
}
})
h.endpointRegistry.Register(reg)
w.Logger().Infow("endpoints registered",
"deployment", w.Deployment, "routes", len(w.Endpoints))
return reg
+24 -29
View File
@@ -44,22 +44,21 @@ import (
// StartWebTransport starts the worker WebTransport listener (control + data) on
// the configured UDP port, if endpoints are enabled and a port is set. It
// returns a stop function (always non-nil). QUIC has no plaintext mode, so a
// TLS certificate is required: from tls_cert_file/tls_key_file, or a generated
// self-signed cert in dev mode.
func (s *AgentService) StartWebTransport(dev bool) (func(), error) {
noop := func() {}
// returns the stop function, nil when no listener was started. QUIC has no
// plaintext mode, so a TLS certificate is required: from
// tls_cert_file/tls_key_file, or a generated self-signed cert in dev mode.
func (s *AgentService) StartWebTransport() (func(), error) {
cfg := s.endpointsConfig
if cfg.Disabled || cfg.WebTransportPort == 0 {
return noop, nil
return nil, nil
}
tlsConf, err := WebTransportTLS(cfg.TLSCertFile, cfg.TLSKeyFile, dev)
tlsConf, err := WebTransportTLS(cfg.TLSCertFile, cfg.TLSKeyFile, s.developmentMode)
if err != nil {
return noop, err
return nil, err
}
udp, err := net.ListenUDP("udp", &net.UDPAddr{Port: int(cfg.WebTransportPort)})
if err != nil {
return noop, err
return nil, err
}
wt := NewAgentWebTransportServer(s, s.keyProvider, tlsConf)
go func() {
@@ -72,7 +71,7 @@ func (s *AgentService) StartWebTransport(dev bool) (func(), error) {
}
// WebTransportTLS builds the listener's TLS config from cert files, or a
// generated self-signed cert in dev mode. Shared by the OSS and cloud servers.
// generated self-signed cert in dev mode.
func WebTransportTLS(certFile, keyFile string, dev bool) (*tls.Config, error) {
if certFile != "" && keyFile != "" {
cert, err := tls.LoadX509KeyPair(certFile, keyFile)
@@ -107,19 +106,15 @@ func WebTransportTLS(certFile, keyFile string, dev bool) (*tls.Config, error) {
}, nil
}
// NewWebTransportServer wraps an HTTP/3 WebTransport server around a handler.
// register is called with the constructed server so the caller can mount routes
// that Upgrade on it (Upgrade needs the *webtransport.Server); it returns the
// HTTP/3 handler. tlsConf must be usable for HTTP/3 (the h3 ALPN is set here if
// absent). Shared by the OSS and cloud agent servers.
func NewWebTransportServer(tlsConf *tls.Config, register func(*webtransport.Server) http.Handler) *webtransport.Server {
// NewWebTransportServer wraps an HTTP/3 WebTransport server around tlsConf (the
// h3 ALPN is set here if absent). The caller must assign wt.H3.Handler: routes
// that Upgrade need the *webtransport.Server itself.
func NewWebTransportServer(tlsConf *tls.Config) *webtransport.Server {
tlsConf = tlsConf.Clone()
if len(tlsConf.NextProtos) == 0 {
tlsConf.NextProtos = []string{http3.NextProtoH3}
}
wt := &webtransport.Server{H3: &http3.Server{TLSConfig: tlsConf}}
wt.H3.Handler = register(wt)
return wt
return &webtransport.Server{H3: &http3.Server{TLSConfig: tlsConf}}
}
// NewAgentWebTransportServer builds the WebTransport server that terminates a
@@ -129,16 +124,16 @@ func NewWebTransportServer(tlsConf *tls.Config, register func(*webtransport.Serv
// handler runs behind the same api-key auth middleware as the rest of the
// agent surface, so the agent grant is enforced identically.
func NewAgentWebTransportServer(svc *AgentService, keyProvider auth.KeyProvider, tlsConf *tls.Config) *webtransport.Server {
return NewWebTransportServer(tlsConf, func(wt *webtransport.Server) http.Handler {
authMW := NewAPIKeyAuthMiddleware(keyProvider)
mux := http.NewServeMux()
mux.HandleFunc("/agent", func(w http.ResponseWriter, r *http.Request) {
svc.ServeWebTransport(wt, w, r)
})
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
authMW.ServeHTTP(w, r, mux.ServeHTTP)
})
wt := NewWebTransportServer(tlsConf)
authMW := NewAPIKeyAuthMiddleware(keyProvider)
mux := http.NewServeMux()
mux.HandleFunc("/agent", func(w http.ResponseWriter, r *http.Request) {
svc.ServeWebTransport(wt, w, r)
})
wt.H3.Handler = http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
authMW.ServeHTTP(w, r, mux.ServeHTTP)
})
return wt
}
// ServeWebTransport verifies the agent grant, upgrades the request to a
@@ -203,7 +198,7 @@ type wtSignalConn struct {
}
// NewWTSignalConn adapts a WebTransport session's control stream to
// agent.SignalConn. Shared by the OSS and cloud agent servers.
// agent.SignalConn.
func NewWTSignalConn(sess *webtransport.Session, control *webtransport.Stream) agent.SignalConn {
return &wtSignalConn{sess: sess, control: control}
}
+3 -3
View File
@@ -76,9 +76,9 @@ func (m *APIKeyAuthMiddleware) ServeHTTP(w http.ResponseWriter, r *http.Request,
authToken = authHeader[len(bearerPrefix):]
} else {
// attempt to find from the query string. FormValue would also parse
// url-encoded POST bodies, consuming the body of any request that gets
// proxied further (agent HTTP endpoints)
// the body must survive for requests proxied further (agent HTTP
// endpoints), so the token comes from the query string alone. URL is nil
// on hand-built requests.
if r.URL != nil {
authToken = r.URL.Query().Get(accessTokenParam)
}
+1 -1
View File
@@ -219,7 +219,7 @@ func (s *LivekitServer) Start() error {
}
if s.agentService != nil {
stop, err := s.agentService.StartWebTransport(s.config.Development)
stop, err := s.agentService.StartWebTransport()
if err != nil {
return err
}