Skip to content
Closed
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
11 changes: 8 additions & 3 deletions api/api.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,10 +12,11 @@ type Request interface {
ReqID() string
ReqCreated() int64
ReqDeadline() int64
ReqPayload() map[string]any
ReqPayload() json.RawMessage
ReqMetadata() map[string]string
ReqHeaders() map[string]string
ReqEndpoint() string
ReqModel() string
}

// FairnessIDHeader is the request header llm-d-router's flow control reads to
Expand All @@ -41,19 +42,23 @@ type RequestMessage struct {
ID string `json:"id"`
Created int64 `json:"created"` // Unix seconds
Deadline int64 `json:"deadline"` // Unix seconds
Payload map[string]any `json:"payload"`
Payload json.RawMessage `json:"payload"`
Metadata map[string]string `json:"metadata,omitempty"`
Headers map[string]string `json:"headers,omitempty"`
Endpoint string `json:"endpoint,omitempty"`
// Model names the model the payload targets, so consumers can label a
// request without parsing its payload. Empty when the producer does not set it.
Model string `json:"model,omitempty"`
}

func (r *RequestMessage) ReqID() string { return r.ID }
func (r *RequestMessage) ReqCreated() int64 { return r.Created }
func (r *RequestMessage) ReqDeadline() int64 { return r.Deadline }
func (r *RequestMessage) ReqPayload() map[string]any { return r.Payload }
func (r *RequestMessage) ReqPayload() json.RawMessage { return r.Payload }
func (r *RequestMessage) ReqMetadata() map[string]string { return r.Metadata }
func (r *RequestMessage) ReqHeaders() map[string]string { return r.Headers }
func (r *RequestMessage) ReqEndpoint() string { return r.Endpoint }
func (r *RequestMessage) ReqModel() string { return r.Model }

// RedisRequest is the concrete Request implementation for Redis-based flows.
// Per-message queue fields here override producer defaults; producers merge them
Expand Down
41 changes: 32 additions & 9 deletions api/internal_api_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -64,8 +64,9 @@ func TestRoundTrip_PlainRequestMessage(t *testing.T) {
InternalRouting{RetryCount: 2, RequestQueueName: "rq", ResultQueueName: "resq", ResultTTLSeconds: 60, ResultRoutingResolved: true},
&RequestMessage{
ID: "plain-1", Created: 1000, Deadline: 2000,
Payload: map[string]any{"model": "m1"},
Payload: testPayload(map[string]any{"model": "m1"}),
Metadata: map[string]string{"k": "v"},
Model: "m1",
},
)
b, err := json.Marshal(ir)
Expand All @@ -86,19 +87,33 @@ func TestRoundTrip_PlainRequestMessage(t *testing.T) {
if rm.ID != "plain-1" || rm.Created != 1000 || rm.Deadline != 2000 {
t.Errorf("field mismatch: %+v", rm)
}
if rm.Payload["model"] != "m1" {
t.Errorf("payload mismatch: %v", rm.Payload)
if string(rm.Payload) != `{"model":"m1"}` {
t.Errorf("payload mismatch: %s", rm.Payload)
}
if rm.Metadata["k"] != "v" {
t.Errorf("metadata mismatch: %v", rm.Metadata)
}
if rm.Model != "m1" {
t.Errorf("model mismatch: %q", rm.Model)
}
}

func TestUnmarshal_EnvelopeWithoutModel(t *testing.T) {
b := []byte(`{"internal":{},"request_kind":"plain","data":{"id":"a","created":1,"deadline":2,"payload":{"model":"m"}}}`)
var got InternalRequest
if err := json.Unmarshal(b, &got); err != nil {
t.Fatalf("unmarshal: %v", err)
}
if m := got.PublicRequest.ReqModel(); m != "" {
t.Errorf("ReqModel() = %q, want empty for an envelope written without the field", m)
}
}

func TestRoundTrip_RedisRequest(t *testing.T) {
ir := NewInternalRequest(
InternalRouting{RetryCount: 1, RequestQueueName: "rq", ResultQueueName: "resq", TransportCorrelationID: "tc"},
&RedisRequest{
RequestMessage: RequestMessage{ID: "redis-1", Created: 100, Deadline: 200, Payload: map[string]any{"p": 1}},
RequestMessage: RequestMessage{ID: "redis-1", Created: 100, Deadline: 200, Payload: testPayload(map[string]any{"p": 1})},
RequestQueueName: "per-msg-rq",
ResultQueueName: "per-msg-resq",
},
Expand Down Expand Up @@ -221,7 +236,7 @@ func TestUnmarshal_EmptyData(t *testing.T) {
func TestRoundTrip_PublicRequestInterface(t *testing.T) {
ir := NewInternalRequest(
InternalRouting{},
&RequestMessage{ID: "iface-test", Created: 1, Deadline: 2, Payload: map[string]any{"k": "v"}, Metadata: map[string]string{"m": "d"}},
&RequestMessage{ID: "iface-test", Created: 1, Deadline: 2, Payload: testPayload(map[string]any{"k": "v"}), Metadata: map[string]string{"m": "d"}},
)
b, err := json.Marshal(ir)
if err != nil {
Expand All @@ -245,8 +260,8 @@ func TestRoundTrip_PublicRequestInterface(t *testing.T) {
if r.ReqDeadline() != 2 {
t.Errorf("ReqDeadlineUnixSec() = %d", r.ReqDeadline())
}
if r.ReqPayload()["k"] != "v" {
t.Errorf("ReqPayload() = %v", r.ReqPayload())
if string(r.ReqPayload()) != `{"k":"v"}` {
t.Errorf("ReqPayload() = %s", r.ReqPayload())
}
if r.ReqMetadata()["m"] != "d" {
t.Errorf("ReqMetadata() = %v", r.ReqMetadata())
Expand All @@ -258,7 +273,7 @@ func TestRoundTrip_EndpointField(t *testing.T) {
InternalRouting{RequestQueueName: "rq"},
&RequestMessage{
ID: "ep-test", Created: 1, Deadline: 2,
Payload: map[string]any{"model": "m"},
Payload: testPayload(map[string]any{"model": "m"}),
Endpoint: "/v1/custom",
},
)
Expand Down Expand Up @@ -287,7 +302,7 @@ func TestRoundTrip_EndpointField(t *testing.T) {
func TestRoundTrip_EndpointOmittedWhenEmpty(t *testing.T) {
ir := NewInternalRequest(
InternalRouting{},
&RequestMessage{ID: "no-ep", Created: 1, Deadline: 2, Payload: map[string]any{}},
&RequestMessage{ID: "no-ep", Created: 1, Deadline: 2, Payload: testPayload(map[string]any{})},
)
b, err := json.Marshal(ir)
if err != nil {
Expand Down Expand Up @@ -332,3 +347,11 @@ func assertRouting(t *testing.T, got, want InternalRouting) {
t.Errorf("TransportCorrelationID = %q, want %q", got.TransportCorrelationID, want.TransportCorrelationID)
}
}

func testPayload(m map[string]any) json.RawMessage {
b, err := json.Marshal(m)
if err != nil {
panic(err)
}
return b
}
36 changes: 27 additions & 9 deletions pkg/asyncworker/worker.go
Original file line number Diff line number Diff line change
Expand Up @@ -272,7 +272,7 @@ func WorkerWithGateTimeout(consumeCtx, requestCtx context.Context, characteristi
attribute.Int(uotel.AttrRetryCount, msg.RetryCount),
attribute.Int(uotel.LegacyAttrRetryCount, msg.RetryCount),
}
if model, ok := msg.PublicRequest.ReqPayload()["model"].(string); ok && model != "" {
if model := msg.PublicRequest.ReqModel(); model != "" {
spanAttrs = append(spanAttrs, attribute.String(uotel.AttrRequestModel, model))
}
if queueID != "" {
Expand All @@ -291,6 +291,9 @@ func WorkerWithGateTimeout(consumeCtx, requestCtx context.Context, characteristi
trace.WithAttributes(spanAttrs...),
)
defer span.End()
if model := fallbackModel(msg.PublicRequest, span.IsRecording()); model != "" {
span.SetAttributes(attribute.String(uotel.AttrRequestModel, model))
}

reqDeadline := time.Now().Add(requestTimeout)
if dline := msg.PublicRequest.ReqDeadline(); dline > 0 {
Expand Down Expand Up @@ -476,14 +479,9 @@ func validateAndMarshal(ctx context.Context, resultChannel chan asyncapi.ResultM
return nil
}

payloadBytes, err := json.Marshal(r.ReqPayload())
if err != nil {
metrics.RecordFailedReq(queueID, queueName, msg.WorkerPoolID)
select {
case resultChannel <- asyncapi.NewErrorResult(r, msg.InternalRouting, asyncapi.ErrCodeInvalidRequest, fmt.Sprintf("Failed to marshal message's payload: %s", err.Error())):
case <-ctx.Done():
}
return nil
payloadBytes := []byte(r.ReqPayload())
if len(payloadBytes) == 0 {
payloadBytes = []byte("null")
}

// Pre-dispatch transform validation (e.g. signed object URL expiry). A
Expand Down Expand Up @@ -647,3 +645,23 @@ func expBackoffDuration(retryCount int, secondsToDeadline int) float64 {
half := temp / 2
return half + rand.Float64()*half // #nosec G404 -- non-security jitter, crypto/rand unnecessary
}

// fallbackModel reads the model from the payload for producers that leave
// Model unset, and only for a sampled span.
func fallbackModel(r asyncapi.Request, recording bool) string {
if !recording || r.ReqModel() != "" {
return ""
}
return payloadModel(r.ReqPayload())
}

// payloadModel returns the payload's model field, or "" if it has none.
func payloadModel(payload json.RawMessage) string {
var p struct {
Model string `json:"model"`
}
if err := json.Unmarshal(payload, &p); err != nil {
return ""
}
return p.Model
}
Loading
Loading