diff --git a/api/api.go b/api/api.go index 3372fb47..59d27cf4 100644 --- a/api/api.go +++ b/api/api.go @@ -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 @@ -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 diff --git a/api/internal_api_test.go b/api/internal_api_test.go index ada99d37..e8c4bf35 100644 --- a/api/internal_api_test.go +++ b/api/internal_api_test.go @@ -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) @@ -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", }, @@ -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 { @@ -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()) @@ -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", }, ) @@ -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 { @@ -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 +} diff --git a/pkg/asyncworker/worker.go b/pkg/asyncworker/worker.go index c8d40b28..c1388ed5 100644 --- a/pkg/asyncworker/worker.go +++ b/pkg/asyncworker/worker.go @@ -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 != "" { @@ -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 { @@ -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 @@ -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 +} diff --git a/pkg/asyncworker/worker_test.go b/pkg/asyncworker/worker_test.go index cbb0093a..63a07cc1 100644 --- a/pkg/asyncworker/worker_test.go +++ b/pkg/asyncworker/worker_test.go @@ -154,7 +154,7 @@ func TestSheddedRequest(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}, + Payload: testPayload(map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -191,7 +191,7 @@ func TestSuccessfulRequest(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}, + Payload: testPayload(map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -229,7 +229,7 @@ func TestSuccessfulRequest_PreservesActualStatusCode(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -332,7 +332,7 @@ func TestWorker_CancelledRequestSkipsInference(t *testing.T) { ID: msgID, Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -383,7 +383,7 @@ func TestWorker_CancellationCheckErrorRequeuesRequest(t *testing.T) { ID: msgID, Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -430,7 +430,7 @@ func TestWorker_FastPathChecksCancellationOnce(t *testing.T) { ID: msgID, Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -476,7 +476,7 @@ func TestWorker_PoolGateActionWaitThrottlesCancellationChecks(t *testing.T) { ID: "gate-wait-throttled", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -524,7 +524,7 @@ func TestWorker_PoolGateWaitTimeoutReenqueues(t *testing.T) { ID: "gate-wait-timeout", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) msg.QueueID = "gate-wait-timeout-q" msg.RequestQueueName = "gate-wait-timeout-queue" @@ -568,7 +568,7 @@ func TestWorker_PoolGatePublicDeadlineDoesNotCountWaitTimeout(t *testing.T) { ID: "gate-wait-public-deadline", Created: time.Now().Unix(), Deadline: time.Now().Unix() + 1, - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) msg.QueueID = "gate-wait-public-deadline-q" msg.RequestQueueName = "gate-wait-public-deadline-queue" @@ -611,7 +611,7 @@ func TestWorker_PoolGateWaitDoesNotConsumeInferenceTimeout(t *testing.T) { ID: "gate-wait-independent-timeout", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -663,7 +663,7 @@ func TestFatalError_NoRetry(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}, + Payload: testPayload(map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -717,7 +717,7 @@ func TestRateLimitRequest(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}, + Payload: testPayload(map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -753,7 +753,7 @@ func TestRequestTimeout(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -1041,7 +1041,7 @@ func TestRateLimitRequest_WithRetryAfterHeader(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -1141,7 +1141,7 @@ func TestWorker_cancelledCtxExitsPromptly(t *testing.T) { ID: "worker-cancel-test", Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) // Give the worker a moment to pick up the message and attempt the send, @@ -1179,7 +1179,7 @@ func TestClientError_NoRetry(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}, + Payload: testPayload(map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -1224,7 +1224,7 @@ func TestWorker_RetriesOnShutdown(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) <-reqStarted @@ -1299,7 +1299,7 @@ func TestWorker_DrainsBufferedMessagesOnShutdown(t *testing.T) { ID: ids[i], Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) } @@ -1424,7 +1424,7 @@ func TestMetrics_QueueDepthAndInflightBalance(t *testing.T) { ID: "depth-msg", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) // Wait for the terminal outcome so handling has completed. @@ -1479,7 +1479,7 @@ func TestMetrics_QueueDepthDecrementsOnDrain(t *testing.T) { ID: fmt.Sprintf("drain-depth-%d", i), Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) } consumeCancel() @@ -1539,7 +1539,7 @@ func TestMetrics_SuccessfulRequest(t *testing.T) { ID: "m-success", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) // Stamp ingestion time so the worker records queue residence time, mirroring // what the broker producers do when a message enters the in-process buffer. @@ -1600,7 +1600,7 @@ func TestMetrics_SuccessfulRequestRecordsTokens(t *testing.T) { ID: "m-tokens", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -1645,7 +1645,7 @@ func TestMetrics_PromptOnlyUsageRegistersBothDirections(t *testing.T) { ID: "m-prompt", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -1690,7 +1690,7 @@ func TestMetrics_NoUsageRecordsNoTokens(t *testing.T) { ID: "m-nousage", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -1732,7 +1732,7 @@ func TestMetrics_NonOpenAIEndpointNotTokenized(t *testing.T) { ID: "m-otherurl", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/embeddings", nil) select { @@ -1774,7 +1774,7 @@ func TestMetrics_RedirectNotTokenized(t *testing.T) { ID: "m-3xx", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -1815,7 +1815,7 @@ func TestMetrics_RateLimited(t *testing.T) { ID: "m-shedded", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -1855,7 +1855,7 @@ func TestMetrics_FatalError(t *testing.T) { ID: "m-fatal", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -1895,7 +1895,7 @@ func TestMetrics_DeadlineExceeded(t *testing.T) { ID: "m-deadline", Created: time.Now().Unix(), Deadline: time.Now().Add(-10 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -1931,7 +1931,7 @@ func TestMetrics_LabelsIsolated(t *testing.T) { QueueID: queueA, RequestQueueName: queueA, }, asyncapi.RequestMessage{ ID: "iso-a", Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -1944,7 +1944,7 @@ func TestMetrics_LabelsIsolated(t *testing.T) { QueueID: queueB, RequestQueueName: queueB, }, asyncapi.RequestMessage{ ID: "iso-b", Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -2048,7 +2048,7 @@ func TestWorker_SpanOnSuccess(t *testing.T) { requestChannel <- newEmb(asyncapi.RequestMessage{ ID: "span-success", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), Metadata: map[string]string{"model": "metadata-model"}, }, "http://localhost:30800/v1/completions", nil) @@ -2082,12 +2082,12 @@ func TestWorker_SpanOnSuccess(t *testing.T) { func TestWorker_SpanOmitsUnavailableModel(t *testing.T) { for _, tc := range []struct { name string - payload map[string]any + payload json.RawMessage }{ - {name: "missing", payload: map[string]any{"prompt": "hi"}}, - {name: "empty", payload: map[string]any{"model": ""}}, - {name: "non-string", payload: map[string]any{"model": 42}}, - {name: "null", payload: map[string]any{"model": nil}}, + {name: "missing", payload: testPayload(map[string]any{"prompt": "hi"})}, + {name: "empty", payload: testPayload(map[string]any{"model": ""})}, + {name: "non-string", payload: testPayload(map[string]any{"model": 42})}, + {name: "null", payload: testPayload(map[string]any{"model": nil})}, {name: "nil-payload"}, } { t.Run(tc.name, func(t *testing.T) { @@ -2146,7 +2146,7 @@ func TestWorker_SpanOnFatalError(t *testing.T) { requestChannel <- newEmb(asyncapi.RequestMessage{ ID: "span-fatal", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -2188,7 +2188,7 @@ func TestWorker_SpanOnRetryableError(t *testing.T) { requestChannel <- newEmb(asyncapi.RequestMessage{ ID: "span-429", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -2227,7 +2227,7 @@ func TestWorker_SpanOnServerError(t *testing.T) { requestChannel <- newEmb(asyncapi.RequestMessage{ ID: "span-500", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -2272,7 +2272,7 @@ func TestWorker_TraceContextExtraction(t *testing.T) { requestChannel <- newEmb(asyncapi.RequestMessage{ ID: "span-ctx", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), Metadata: metadata, }, "http://localhost:30800/v1/completions", nil) @@ -2315,7 +2315,7 @@ func TestWorker_SpanOnShutdownReenqueue(t *testing.T) { requestChannel <- newEmb(asyncapi.RequestMessage{ ID: "span-shutdown", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) <-reqStarted @@ -2377,7 +2377,7 @@ func TestWorker_SpanIncludesQueueName(t *testing.T) { asyncapi.InternalRouting{QueueID: "my-test-qid", RequestQueueName: "my-test-queue", RetryCount: 3}, asyncapi.RequestMessage{ ID: "span-queue", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -2433,7 +2433,7 @@ func TestWorker_InFlightCompletesOnConsumeCancel(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) <-reqStarted @@ -2495,7 +2495,7 @@ func TestWorker_DrainTimeoutCancelsInFlight(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) <-reqStarted @@ -2566,7 +2566,7 @@ func TestWorker_DrainWithCancelledRequestCtx(t *testing.T) { ID: ids[i], Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) } @@ -2703,7 +2703,7 @@ func TestWorker_PoolGateShutdownReenqueues(t *testing.T) { ID: "gate-shutdown-test", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) msg.QueueID = "gate-shutdown-q" msg.RequestQueueName = "gate-shutdown-queue" @@ -2759,7 +2759,7 @@ func TestWorker_PoolGateActionWaitShutdownReenqueues(t *testing.T) { ID: "gate-wait-shutdown-test", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) msg.QueueID = "gate-wait-shutdown-q" msg.RequestQueueName = "gate-wait-shutdown-queue" @@ -2816,7 +2816,7 @@ func TestWorker_PoolGateActionWaitHonorsCancellation(t *testing.T) { ID: msgID, Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -2874,7 +2874,7 @@ func TestWorker_RechecksCancellationAfterGateContinue(t *testing.T) { ID: msgID, Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -2933,7 +2933,7 @@ func TestWorker_RechecksCancellationImmediatelyBeforeSend(t *testing.T) { ID: msgID, Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", nil) select { @@ -3009,7 +3009,7 @@ func TestWorker_QueueGateAndPoolGateRace(t *testing.T) { ID: "req-race-test", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) var queueReleases []pipeline.GateReleaseFunc @@ -3101,7 +3101,7 @@ func TestWorker_PoolGateDecisionsMetrics(t *testing.T) { ID: "req-dropped", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, "http://localhost:30800/v1/completions", nil) msg.WorkerPoolID = poolID requestChannel <- msg @@ -3139,7 +3139,7 @@ func TestWorker_PoolGateDecisionsMetrics(t *testing.T) { ID: "req-closed", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, "http://localhost:30800/v1/completions", nil) msg.WorkerPoolID = poolID requestChannel <- msg @@ -3179,7 +3179,7 @@ func TestWorker_PoolGateDecisionsMetrics(t *testing.T) { ID: "req-quota", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) ir.SetClassification(asyncapi.ClassificationOverflow) @@ -3223,7 +3223,7 @@ func TestWorker_PoolGateDecisionsMetrics(t *testing.T) { ID: "req-error", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, "http://localhost:30800/v1/completions", nil) msg.WorkerPoolID = poolID requestChannel <- msg @@ -3261,7 +3261,7 @@ func TestWorker_PoolGateDecisionsMetrics(t *testing.T) { ID: "req-wait", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, "http://localhost:30800/v1/completions", nil) msg.WorkerPoolID = poolID requestChannel <- msg @@ -3301,7 +3301,7 @@ func TestDeadlineAbortedSendClassifiedAsDeadlineExceeded(t *testing.T) { ID: msgId, Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Second).Unix(), - Payload: map[string]any{"model": "m", "prompt": "hi"}, + Payload: testPayload(map[string]any{"model": "m", "prompt": "hi"}), }, "http://localhost:30800/v1/completions", map[string]string{}) select { @@ -3315,3 +3315,95 @@ func TestDeadlineAbortedSendClassifiedAsDeadlineExceeded(t *testing.T) { t.Fatal("timed out waiting for result") } } + +func TestValidateAndMarshal_ForwardsPayloadBytes(t *testing.T) { + for _, tc := range []struct { + name string + payload json.RawMessage + want string + }{ + {name: "verbatim", payload: json.RawMessage(`{"z":1, "a":[2, 3]}`), want: `{"z":1, "a":[2, 3]}`}, + {name: "missing payload", payload: nil, want: `null`}, + } { + t.Run(tc.name, func(t *testing.T) { + resultChannel := make(chan asyncapi.ResultMessage, 1) + msg := newEmb(asyncapi.RequestMessage{ + ID: "fwd", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), + Payload: tc.payload, + }, "http://localhost/v1/completions", nil) + got := validateAndMarshal(context.Background(), resultChannel, msg, nil) + if string(got) != tc.want { + t.Fatalf("body = %s, want %s", got, tc.want) + } + select { + case r := <-resultChannel: + t.Fatalf("unexpected result %+v", r) + default: + } + }) + } +} + +func TestWorker_SpanPrefersRequestModel(t *testing.T) { + exporter := setupTestTracer(t) + httpclient := NewTestClient(func(req *http.Request) (*http.Response, error) { + return &http.Response{StatusCode: http.StatusOK, Body: io.NopCloser(bytes.NewReader(nil)), Header: make(http.Header)}, nil + }) + inferenceClient := NewHTTPInferenceClient(httpclient) + requestChannel := make(chan pipeline.EmbelishedRequestMessage, 1) + retryChannel := make(chan pipeline.RetryMessage, 1) + resultChannel := make(chan asyncapi.ResultMessage, 1) + + ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) + defer cancel() + go Worker(ctx, ctx, pipeline.Characteristics{}, inferenceClient, requestChannel, retryChannel, resultChannel, defaultRequestTimeout, nil) + + requestChannel <- newEmb(asyncapi.RequestMessage{ + ID: "span-model-field", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), + Payload: testPayload(map[string]any{"model": "from-payload"}), + Model: "from-request", + }, "http://localhost:30800/v1/completions", nil) + + select { + case <-resultChannel: + case <-time.After(2 * time.Second): + t.Fatal("timeout waiting for result") + } + + s := findSpan(getSpansEventually(t, exporter, 1), "process-request") + if s == nil { + t.Fatal("expected 'process-request' span") + } + assertSpanAttributes(t, s, attribute.String("gen_ai.request.model", "from-request")) +} + +func TestFallbackModel(t *testing.T) { + withModel := &asyncapi.RequestMessage{Model: "set", Payload: json.RawMessage(`{"model":"payload"}`)} + withoutModel := &asyncapi.RequestMessage{Payload: json.RawMessage(`{"model":"payload"}`)} + unparseable := &asyncapi.RequestMessage{Payload: json.RawMessage(`{not json`)} + + for _, tc := range []struct { + name string + req asyncapi.Request + recording bool + want string + }{ + {"model set, sampled", withModel, true, ""}, + {"model set, not sampled", withModel, false, ""}, + {"no model, not sampled", withoutModel, false, ""}, + {"no model, sampled", withoutModel, true, "payload"}, + {"no model, sampled, unparseable payload", unparseable, true, ""}, + } { + if got := fallbackModel(tc.req, tc.recording); got != tc.want { + t.Errorf("%s: fallbackModel = %q, want %q", tc.name, got, tc.want) + } + } +} + +func testPayload(m map[string]any) json.RawMessage { + b, err := json.Marshal(m) + if err != nil { + panic(err) + } + return b +} diff --git a/pkg/asyncworker/worker_transform_test.go b/pkg/asyncworker/worker_transform_test.go index aa53fcf9..d538f542 100644 --- a/pkg/asyncworker/worker_transform_test.go +++ b/pkg/asyncworker/worker_transform_test.go @@ -78,7 +78,7 @@ func TestWorker_TransformRewritesOutgoingRequest(t *testing.T) { ID: "m1", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "whisper", "gcs_uri": "https://storage.example/audio.mp3?signature=test"}, + Payload: testPayload(map[string]any{"model": "whisper", "gcs_uri": "https://storage.example/audio.mp3?signature=test"}), Metadata: map[string]string{"provider": "whisper"}, }, "http://localhost/v1/audio/transcriptions", map[string]string{}) @@ -125,7 +125,7 @@ func TestWorker_SpanOnTransformError(t *testing.T) { requestChannel <- newEmb(asyncapi.RequestMessage{ ID: "span-transform", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "whisper"}, + Payload: testPayload(map[string]any{"model": "whisper"}), }, "http://localhost/v1/audio/transcriptions", nil) select { @@ -185,7 +185,7 @@ func TestWorker_TransformValidateFatal(t *testing.T) { ID: "m2", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Second).Unix(), - Payload: map[string]any{"model": "whisper"}, + Payload: testPayload(map[string]any{"model": "whisper"}), Metadata: map[string]string{"provider": "whisper"}, }, "http://localhost/v1/audio/transcriptions", map[string]string{}) diff --git a/pkg/redis/claim_test.go b/pkg/redis/claim_test.go index 52b3869e..aec0e0b4 100644 --- a/pkg/redis/claim_test.go +++ b/pkg/redis/claim_test.go @@ -41,7 +41,7 @@ func claimEnvelope(t *testing.T, id string, deadline int64) (*api.InternalReques ID: id, Created: time.Now().Unix(), Deadline: deadline, - Payload: map[string]any{"model": "m", "prompt": "p"}, + Payload: testPayload(map[string]any{"model": "m", "prompt": "p"}), }) b, err := json.Marshal(ir) if err != nil { diff --git a/pkg/redis/sortedset_impl_test.go b/pkg/redis/sortedset_impl_test.go index 36f27eb2..db3616a2 100644 --- a/pkg/redis/sortedset_impl_test.go +++ b/pkg/redis/sortedset_impl_test.go @@ -66,7 +66,7 @@ func registerTestClaim(ctx context.Context, flow *RedisSortedSetFlow, queueName, ID: reqID, Created: time.Now().Unix(), Deadline: time.Now().Add(time.Hour).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) payloadBytes, _ := json.Marshal(ir) @@ -203,7 +203,7 @@ func TestSortedSetFlow_MessageProcessing(t *testing.T) { ID: "msg-1", Created: time.Now().Unix(), Deadline: 9999999999, - Payload: map[string]any{"test": "data"}, + Payload: testPayload(map[string]any{"test": "data"}), } rdb.ZAdd(ctx, queue, redis.Z{Score: float64(time.Now().Unix()), Member: envelopeJSON(msg)}) @@ -1194,7 +1194,7 @@ func TestSortedSetFlow_ZeroBudget(t *testing.T) { ID: "test-zero-budget", Created: time.Now().Unix(), Deadline: 9999999999, - Payload: map[string]any{"test": "data"}, + Payload: testPayload(map[string]any{"test": "data"}), } rdb.ZAdd(ctx, queue, redis.Z{Score: float64(time.Now().Unix()), Member: envelopeJSON(msg)}) @@ -1258,7 +1258,7 @@ func TestSortedSetFlow_ClosedGateRecordsGateClosedDecision(t *testing.T) { ID: "gate-closed-1", Created: time.Now().Unix(), Deadline: 9999999999, - Payload: map[string]any{"test": "data"}, + Payload: testPayload(map[string]any{"test": "data"}), } rdb.ZAdd(ctx, queue, redis.Z{Score: float64(time.Now().Unix()), Member: envelopeJSON(msg)}) @@ -1673,7 +1673,7 @@ func TestSortedSetFlow_RequestWorkerRequeuesOnShutdown(t *testing.T) { ID: "requeue-1", Created: time.Now().Unix(), Deadline: 9999999999, - Payload: map[string]any{"key": "value"}, + Payload: testPayload(map[string]any{"key": "value"}), }) msgBytes, _ := json.Marshal(ir) score := float64(time.Now().Unix()) @@ -2359,3 +2359,11 @@ func TestSortedSetFlow_QueueLabelsSetOnDequeue(t *testing.T) { t.Fatal("Timeout waiting for message") } } + +func testPayload(m map[string]any) json.RawMessage { + b, err := json.Marshal(m) + if err != nil { + panic(err) + } + return b +} diff --git a/producer/redis_sortedset_producer.go b/producer/redis_sortedset_producer.go index deeec8d8..0816bb15 100644 --- a/producer/redis_sortedset_producer.go +++ b/producer/redis_sortedset_producer.go @@ -189,6 +189,7 @@ func toInternalRequest(req api.Request) *api.InternalRequest { Metadata: req.ReqMetadata(), Headers: req.ReqHeaders(), Endpoint: req.ReqEndpoint(), + Model: req.ReqModel(), } return ir } diff --git a/producer/redis_sortedset_producer_test.go b/producer/redis_sortedset_producer_test.go index d627f1aa..3c99aed3 100644 --- a/producer/redis_sortedset_producer_test.go +++ b/producer/redis_sortedset_producer_test.go @@ -17,8 +17,9 @@ import ( // default branch of toInternalRequest. type customRequest struct { id, endpoint string + model string created, deadline int64 - payload map[string]any + payload json.RawMessage metadata map[string]string headers map[string]string } @@ -26,10 +27,11 @@ type customRequest struct { func (r *customRequest) ReqID() string { return r.id } func (r *customRequest) ReqCreated() int64 { return r.created } func (r *customRequest) ReqDeadline() int64 { return r.deadline } -func (r *customRequest) ReqPayload() map[string]any { return r.payload } +func (r *customRequest) ReqPayload() json.RawMessage { return r.payload } func (r *customRequest) ReqMetadata() map[string]string { return r.metadata } func (r *customRequest) ReqHeaders() map[string]string { return r.headers } func (r *customRequest) ReqEndpoint() string { return r.endpoint } +func (r *customRequest) ReqModel() string { return r.model } func setupTestProducer(t *testing.T) (*RedisSortedSetProducer, *miniredis.Miniredis) { t.Helper() @@ -62,10 +64,10 @@ func TestSubmitRequest(t *testing.T) { ID: "test-123", Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]interface{}{ + Payload: testPayload(map[string]interface{}{ "model": "gpt-3.5-turbo", "prompt": "Hello, world!", - }, + }), Metadata: map[string]string{ "user": "test-user", }, @@ -102,10 +104,11 @@ func TestToInternalRequest_PubSubIDCopiesToInternalRouting(t *testing.T) { func TestToInternalRequest_CustomRequestPreservesHeadersAndEndpoint(t *testing.T) { req := &customRequest{ id: "custom-1", created: 1, deadline: 2, - payload: map[string]any{"k": "v"}, + payload: testPayload(map[string]any{"k": "v"}), metadata: map[string]string{"m": "d"}, headers: map[string]string{"Authorization": "Bearer tok"}, endpoint: "/v1/chat/completions", + model: "m1", } ir := toInternalRequest(req) rm, ok := ir.PublicRequest.(*api.RequestMessage) @@ -114,6 +117,7 @@ func TestToInternalRequest_CustomRequestPreservesHeadersAndEndpoint(t *testing.T assert.Equal(t, map[string]string{"Authorization": "Bearer tok"}, rm.Headers) assert.Equal(t, "/v1/chat/completions", rm.Endpoint) assert.Equal(t, map[string]string{"m": "d"}, rm.Metadata) + assert.Equal(t, "m1", rm.Model) } func TestToInternalRequest_RedisQueueFieldsCopyToInternalRouting(t *testing.T) { @@ -176,7 +180,7 @@ func TestSubmitRequest_Validation(t *testing.T) { req: &api.RequestMessage{ Created: time.Now().Unix(), Deadline: time.Now().Unix(), - Payload: map[string]interface{}{}, + Payload: testPayload(map[string]interface{}{}), }, wantErr: "request ID is required", }, @@ -186,7 +190,7 @@ func TestSubmitRequest_Validation(t *testing.T) { ID: "test", Created: time.Now().Unix(), Deadline: 0, - Payload: map[string]interface{}{}, + Payload: testPayload(map[string]interface{}{}), }, wantErr: "deadline is required", }, @@ -196,7 +200,7 @@ func TestSubmitRequest_Validation(t *testing.T) { ID: "test", Created: time.Now().Unix(), Deadline: 0, - Payload: map[string]interface{}{}, + Payload: testPayload(map[string]interface{}{}), }, wantErr: "deadline is required", }, @@ -206,7 +210,7 @@ func TestSubmitRequest_Validation(t *testing.T) { ID: "test", Created: time.Now().Unix(), Deadline: time.Now().Add(-1 * time.Minute).Unix(), - Payload: map[string]interface{}{}, + Payload: testPayload(map[string]interface{}{}), }, wantErr: "deadline has already expired", }, @@ -261,7 +265,7 @@ func TestSubmitRequest_ClearsStaleCancellationMarker(t *testing.T) { ID: requestID, Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }) require.NoError(t, err) assert.False(t, mr.Exists(api.RequestCancellationKey(requestID))) @@ -425,7 +429,7 @@ func TestMultipleTenantsIsolation(t *testing.T) { ID: "alpha-request", Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]interface{}{"tenant": "alpha"}, + Payload: testPayload(map[string]interface{}{"tenant": "alpha"}), } err = tenant1Producer.SubmitRequest(ctx, req1) require.NoError(t, err) @@ -434,7 +438,7 @@ func TestMultipleTenantsIsolation(t *testing.T) { ID: "beta-request", Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]interface{}{"tenant": "beta"}, + Payload: testPayload(map[string]interface{}{"tenant": "beta"}), } err = tenant2Producer.SubmitRequest(ctx, req2) require.NoError(t, err) @@ -521,7 +525,7 @@ func TestProducerAuth(t *testing.T) { ID: "auth-test", Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]interface{}{"test": true}, + Payload: testPayload(map[string]interface{}{"test": true}), } err = producer.SubmitRequest(ctx, req) assert.NoError(t, err) @@ -609,7 +613,7 @@ func TestWithRedisClient(t *testing.T) { ID: "inject-test", Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]interface{}{"test": true}, + Payload: testPayload(map[string]interface{}{"test": true}), } err = p.SubmitRequest(ctx, req) assert.NoError(t, err) @@ -651,7 +655,7 @@ func TestCloseOwnership(t *testing.T) { ID: "post-close", Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]interface{}{}, + Payload: testPayload(map[string]interface{}{}), }) assert.Error(t, err) }) @@ -742,7 +746,7 @@ func TestResultQueueNameNoNamespacing(t *testing.T) { ID: "short-name-request", Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]interface{}{"job": "batch"}, + Payload: testPayload(map[string]interface{}{"job": "batch"}), } require.NoError(t, producer.SubmitRequest(ctx, req)) @@ -795,7 +799,7 @@ func TestComplexResultQueueKeyShape(t *testing.T) { ID: "complex-key-request", Created: time.Now().Unix(), Deadline: time.Now().Add(1 * time.Hour).Unix(), - Payload: map[string]interface{}{"pool": "a"}, + Payload: testPayload(map[string]interface{}{"pool": "a"}), } require.NoError(t, producer.SubmitRequest(ctx, req)) @@ -820,3 +824,11 @@ func TestComplexResultQueueKeyShape(t *testing.T) { require.NoError(t, err) assert.Equal(t, "complex-key-request", result.ID) } + +func testPayload(m map[string]any) json.RawMessage { + b, err := json.Marshal(m) + if err != nil { + panic(err) + } + return b +} diff --git a/release-notes.d/unreleased/456.md b/release-notes.d/unreleased/456.md new file mode 100644 index 00000000..1f2beadb --- /dev/null +++ b/release-notes.d/unreleased/456.md @@ -0,0 +1,7 @@ +--- +pr: 456 +url: https://github.com/llm-d/llm-d-async/pull/456 +author: wseaton +date: 2026-09-17 +--- +Adds `RequestMessage.Model`. Producers copy it, and the worker labels trace spans from it instead of parsing the payload. diff --git a/test/e2e/e2e_benchmark_test.go b/test/e2e/e2e_benchmark_test.go index f9a9b377..b4f4773c 100644 --- a/test/e2e/e2e_benchmark_test.go +++ b/test/e2e/e2e_benchmark_test.go @@ -76,11 +76,11 @@ var _ = ginkgo.Describe("Async Processor Performance Benchmark E2E", ginkgo.Orde ID: fmt.Sprintf("bench-msg-%d", i), Created: now, Deadline: now + 600, // 10 minutes deadline - Payload: map[string]any{ + Payload: testPayload(map[string]any{ "model": "test-model", "prompt": prompt, "max_tokens": 500, - }, + }), } } @@ -162,11 +162,11 @@ var _ = ginkgo.Describe("Async Processor Performance Benchmark E2E", ginkgo.Orde ID: fmt.Sprintf("bench-pool-msg-%d", i), Created: now, Deadline: now + 600, // 10 minutes deadline - Payload: map[string]any{ + Payload: testPayload(map[string]any{ "model": "test-model", "prompt": prompt, "max_tokens": 500, - }, + }), } } @@ -248,11 +248,11 @@ var _ = ginkgo.Describe("Async Processor Performance Benchmark E2E", ginkgo.Orde ID: fmt.Sprintf("bench-pubsub-msg-%d", i), Created: now, Deadline: now + 600, // 10 minutes deadline - Payload: map[string]any{ + Payload: testPayload(map[string]any{ "model": "test-model", "prompt": prompt, "max_tokens": 500, - }, + }), } } diff --git a/test/e2e/e2e_multitenant_merge_test.go b/test/e2e/e2e_multitenant_merge_test.go index c547f1c4..d2313c1e 100644 --- a/test/e2e/e2e_multitenant_merge_test.go +++ b/test/e2e/e2e_multitenant_merge_test.go @@ -33,7 +33,7 @@ func makeTeamMessages(team string, n int) []api.RequestMessage { for i := 0; i < n; i++ { m := makeRequestMessage(fmt.Sprintf("%s-%d", team, i), 5*time.Minute) m.Metadata = map[string]string{"team": team} - m.Payload = map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 64} + m.Payload = testPayload(map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 64}) msgs = append(msgs, m) } return msgs diff --git a/test/e2e/e2e_otel_test.go b/test/e2e/e2e_otel_test.go index a62d0670..481091ef 100644 --- a/test/e2e/e2e_otel_test.go +++ b/test/e2e/e2e_otel_test.go @@ -33,7 +33,7 @@ var _ = ginkgo.Describe("OpenTelemetry tracing", ginkgo.Ordered, func() { ID: "otel-propagation-test", Created: time.Now().Unix(), Deadline: time.Now().Add(2 * time.Minute).Unix(), - Payload: map[string]any{"model": "otel-propagation-test", "prompt": "test"}, + Payload: testPayload(map[string]any{"model": "otel-propagation-test", "prompt": "test"}), Metadata: map[string]string{ "traceparent": traceparent, }, diff --git a/test/e2e/e2e_test.go b/test/e2e/e2e_test.go index 0d44ea3c..3a2584c6 100644 --- a/test/e2e/e2e_test.go +++ b/test/e2e/e2e_test.go @@ -419,7 +419,7 @@ var _ = ginkgo.Describe("Redis Dispatch Gate E2E", func() { ID: requestID, Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test-model", "prompt": "cancel me"}, + Payload: testPayload(map[string]any{"model": "test-model", "prompt": "cancel me"}), })).To(gomega.Succeed()) gomega.Eventually(func() int64 { diff --git a/test/e2e/utils_test.go b/test/e2e/utils_test.go index 85d059b5..c13740a6 100644 --- a/test/e2e/utils_test.go +++ b/test/e2e/utils_test.go @@ -208,7 +208,7 @@ func makeRequestMessage(id string, deadlineOffset time.Duration) api.RequestMess ID: id, Created: time.Now().Unix(), Deadline: deadline.Unix(), - Payload: map[string]any{"model": id, "prompt": "test"}, + Payload: testPayload(map[string]any{"model": id, "prompt": "test"}), } } @@ -403,3 +403,11 @@ func setDispatchGateBudget(ctx context.Context, rdb *redis.Client, budget string func clearDispatchGateBudget(ctx context.Context, rdb *redis.Client) { rdb.Del(ctx, dispatchGateBudgetKey) //nolint:errcheck } + +func testPayload(m map[string]any) json.RawMessage { + b, err := json.Marshal(m) + if err != nil { + panic(err) + } + return b +} diff --git a/test/integration/claim_expiry_loss_test.go b/test/integration/claim_expiry_loss_test.go index d84f3f81..5117043b 100644 --- a/test/integration/claim_expiry_loss_test.go +++ b/test/integration/claim_expiry_loss_test.go @@ -86,7 +86,7 @@ func enqueueShutdownLossRequests(t *testing.T, rdb *goredis.Client, queue string ID: id, Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hello"}), }, ) member, err := ir.MarshalJSON() diff --git a/test/integration/durable_result_delivery_test.go b/test/integration/durable_result_delivery_test.go index c2e4fc41..c6fe6e9c 100644 --- a/test/integration/durable_result_delivery_test.go +++ b/test/integration/durable_result_delivery_test.go @@ -71,7 +71,7 @@ func TestDurableResultDelivery_RedeliversAfterConsumerLoss(t *testing.T) { ID: "durable-result", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hello"}), })) receiveCtx, receiveCancel := context.WithTimeout(context.Background(), 5*time.Second) diff --git a/test/integration/merge_policy_test.go b/test/integration/merge_policy_test.go index dac9433e..52223ed3 100644 --- a/test/integration/merge_policy_test.go +++ b/test/integration/merge_policy_test.go @@ -56,7 +56,7 @@ func TestRandomRobinPolicy_ConcurrentProducers(t *testing.T) { ID: msgID(idx, m), Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) channels[idx].Channel <- ir diff --git a/test/integration/pool_gate_test.go b/test/integration/pool_gate_test.go index fb5953fe..31122ff7 100644 --- a/test/integration/pool_gate_test.go +++ b/test/integration/pool_gate_test.go @@ -73,7 +73,7 @@ func TestPoolGating_Blocking(t *testing.T) { ID: id, Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hello"}), }, ) requestChannel <- pipeline.EmbelishedRequestMessage{ @@ -159,7 +159,7 @@ func TestPoolGating_Timeout(t *testing.T) { ID: "req-slow", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) requestChannel <- pipeline.EmbelishedRequestMessage{ @@ -174,7 +174,7 @@ func TestPoolGating_Timeout(t *testing.T) { ID: "req-timeout", Created: time.Now().Unix(), Deadline: time.Now().Add(100 * time.Millisecond).Unix(), // 100ms deadline - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) requestChannel <- pipeline.EmbelishedRequestMessage{ @@ -260,7 +260,7 @@ func TestPoolGating_ActionWait(t *testing.T) { ID: "req-wait", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) requestChannel <- pipeline.EmbelishedRequestMessage{ @@ -324,7 +324,7 @@ func TestPoolGating_ActionRefuse(t *testing.T) { ID: "req-refuse", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) requestChannel <- pipeline.EmbelishedRequestMessage{ @@ -414,7 +414,7 @@ func TestPoolGating_RedisLeasedRateWaitsUntilLeasePermits(t *testing.T) { requestChannel <- pipeline.EmbelishedRequestMessage{ InternalRequest: asyncapi.NewInternalRequest(asyncapi.InternalRouting{RequestQueueName: "batch-queue"}, &asyncapi.RequestMessage{ - ID: "leased-wait", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), Payload: map[string]any{"model": "test"}, + ID: "leased-wait", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), Payload: testPayload(map[string]any{"model": "test"}), }), RequestURL: server.URL, WorkerPoolID: "batch-pool", @@ -461,7 +461,7 @@ func TestPoolGating_RedisLeasedRateWaitTimeoutRequeues(t *testing.T) { requestChannel <- pipeline.EmbelishedRequestMessage{ InternalRequest: asyncapi.NewInternalRequest(asyncapi.InternalRouting{RequestQueueName: "batch-queue"}, &asyncapi.RequestMessage{ - ID: "leased-timeout", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), Payload: map[string]any{"model": "test"}, + ID: "leased-timeout", Created: time.Now().Unix(), Deadline: time.Now().Add(30 * time.Second).Unix(), Payload: testPayload(map[string]any{"model": "test"}), }), RequestURL: "http://unused.invalid", WorkerPoolID: "batch-pool", diff --git a/test/integration/redis_pubsub_test.go b/test/integration/redis_pubsub_test.go index c794a0cf..bb6882c5 100644 --- a/test/integration/redis_pubsub_test.go +++ b/test/integration/redis_pubsub_test.go @@ -43,7 +43,7 @@ func TestRedisPubSub_PublishSubscribeResultDelivery(t *testing.T) { ID: "pubsub-test-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test-model", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test-model", "prompt": "hello"}), }, ) irBytes, err := json.Marshal(ir) @@ -99,7 +99,7 @@ func TestRedisSortedSet_EnqueueDequeueRetryRoundTrip(t *testing.T) { ID: "retry-roundtrip-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) irBytes, err := json.Marshal(ir) diff --git a/test/integration/redisimpl_test.go b/test/integration/redisimpl_test.go index 442f2b2a..98128805 100644 --- a/test/integration/redisimpl_test.go +++ b/test/integration/redisimpl_test.go @@ -54,7 +54,7 @@ func TestRedisImpl(t *testing.T) { ID: "test-id", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}, + Payload: testPayload(map[string]any{"model": "food-review", "prompt": "hi", "max_tokens": 10, "temperature": 0}), }, ), RequestURL: "http://localhost:30800/v1/completions", @@ -151,7 +151,7 @@ func TestRedisImplWithAuth(t *testing.T) { ID: "test-auth-id", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) member, err := ir.MarshalJSON() diff --git a/test/integration/sortedset_deadline_views_test.go b/test/integration/sortedset_deadline_views_test.go index c127027c..b7886f73 100644 --- a/test/integration/sortedset_deadline_views_test.go +++ b/test/integration/sortedset_deadline_views_test.go @@ -62,7 +62,7 @@ func TestSortedSetDeadlineViews_PollToMetrics(t *testing.T) { ID: fmt.Sprintf("int-msg-%d", i), Created: now, Deadline: deadline, - Payload: map[string]any{"prompt": "hi"}, + Payload: testPayload(map[string]any{"prompt": "hi"}), }, ) irBytes, err := json.Marshal(ir) diff --git a/test/integration/sortedset_quota_gate_test.go b/test/integration/sortedset_quota_gate_test.go index dca36cd7..806f9344 100644 --- a/test/integration/sortedset_quota_gate_test.go +++ b/test/integration/sortedset_quota_gate_test.go @@ -39,7 +39,7 @@ func TestSortedSetQuotaGate_AcquireDequeueRelease(t *testing.T) { ID: id, Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), Metadata: map[string]string{"userid": "user-a"}, }, ) @@ -125,7 +125,7 @@ func TestSortedSetQuotaGate_RateLimitRequeue(t *testing.T) { &api.RequestMessage{ ID: "rl-msg-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), Metadata: map[string]string{"userid": "user-b"}, }, ) @@ -142,7 +142,7 @@ func TestSortedSetQuotaGate_RateLimitRequeue(t *testing.T) { &api.RequestMessage{ ID: "rl-msg-2", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), Metadata: map[string]string{"userid": "user-b"}, }, ) diff --git a/test/integration/tier_priority_integration_test.go b/test/integration/tier_priority_integration_test.go index 9401a381..052f0ab6 100644 --- a/test/integration/tier_priority_integration_test.go +++ b/test/integration/tier_priority_integration_test.go @@ -74,7 +74,7 @@ func TestTierPriorityGate_Integration(t *testing.T) { ID: "req-drop", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) ir.SetClassification(asyncapi.ClassificationOverflow) @@ -109,7 +109,7 @@ func TestTierPriorityGate_Integration(t *testing.T) { ID: "req-refuse", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) ir.SetClassification(asyncapi.ClassificationOverflow) @@ -143,7 +143,7 @@ func TestTierPriorityGate_Integration(t *testing.T) { ID: "req-refuse-batch", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) ir.SetClassification(asyncapi.ClassificationOverflow) diff --git a/test/integration/worker_dispatch_test.go b/test/integration/worker_dispatch_test.go index 2959e617..c76390a4 100644 --- a/test/integration/worker_dispatch_test.go +++ b/test/integration/worker_dispatch_test.go @@ -63,7 +63,7 @@ func TestWorkerDispatch_DrainsBufferedOnShutdown(t *testing.T) { ID: id, Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hello"}), }, ) requestChannel <- pipeline.EmbelishedRequestMessage{ @@ -132,7 +132,7 @@ func TestWorkerDispatch_MockIGW(t *testing.T) { ID: "dispatch-test-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test-model", "prompt": "hello world"}, + Payload: testPayload(map[string]any{"model": "test-model", "prompt": "hello world"}), }, ) @@ -201,7 +201,7 @@ func TestWorkerDispatch_EndpointOverride(t *testing.T) { ID: "endpoint-override-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) @@ -248,7 +248,7 @@ func TestWorkerDispatch_ServerErrorTriggersRetry(t *testing.T) { ID: "retry-test-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), }, ) @@ -299,7 +299,7 @@ func TestWorkerDispatch_ResultCallback(t *testing.T) { ID: "callback-test-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test"}, + Payload: testPayload(map[string]any{"model": "test"}), Metadata: map[string]string{"trace_id": "abc-123"}, }) @@ -358,7 +358,7 @@ func TestWorkerDispatch_RequeuesOnShutdown(t *testing.T) { ID: "shutdown-requeue-1", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hello"}), }, ) @@ -442,7 +442,7 @@ func TestWorkerDispatch_PoolIsolation(t *testing.T) { ID: "msg-blocked", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{}, + Payload: testPayload(map[string]any{}), }, ) reqChanBlocked <- pipeline.EmbelishedRequestMessage{ @@ -465,7 +465,7 @@ func TestWorkerDispatch_PoolIsolation(t *testing.T) { ID: "msg-active", Created: time.Now().Unix(), Deadline: time.Now().Add(5 * time.Minute).Unix(), - Payload: map[string]any{}, + Payload: testPayload(map[string]any{}), }, ) reqChanActive <- pipeline.EmbelishedRequestMessage{ @@ -500,3 +500,11 @@ func TestWorkerDispatch_PoolIsolation(t *testing.T) { t.Fatal("Blocked request did not complete after release") } } + +func testPayload(m map[string]any) json.RawMessage { + b, err := json.Marshal(m) + if err != nil { + panic(err) + } + return b +} diff --git a/test/integration/worker_otel_test.go b/test/integration/worker_otel_test.go index 41d3fcdd..55f14040 100644 --- a/test/integration/worker_otel_test.go +++ b/test/integration/worker_otel_test.go @@ -76,7 +76,7 @@ func TestWorkerDispatch_TraceparentInjected(t *testing.T) { ID: "traceparent-inject-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hello"}), }, ) @@ -134,7 +134,7 @@ func TestWorkerDispatch_SpanHierarchy(t *testing.T) { ID: "hierarchy-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hello"}), }, ) @@ -218,7 +218,7 @@ func TestWorkerDispatch_MetadataTraceContextPropagation(t *testing.T) { ID: "e2e-propagation-1", Created: time.Now().Unix(), Deadline: time.Now().Add(time.Minute).Unix(), - Payload: map[string]any{"model": "test", "prompt": "hello"}, + Payload: testPayload(map[string]any{"model": "test", "prompt": "hello"}), Metadata: metadata, }, )