diff --git a/apps/edge/internal/openai/chat_handler.go b/apps/edge/internal/openai/chat_handler.go
index 106174f..32756c1 100644
--- a/apps/edge/internal/openai/chat_handler.go
+++ b/apps/edge/internal/openai/chat_handler.go
@@ -247,6 +247,14 @@ func (s *Server) handleChatCompletionsProviderPool(
return rewriteChatCompletionModel(rawBody, target, req)
}
+ // Normalized branch: chat completions use lenient decode so unknown
+ // fields survive; PrepareRun passes through the already-built Run
+ // without strict validation (SDD S02/S03). The tunnel branch is
+ // unaffected.
+ poolReq.PrepareRun = func(runReq edgeservice.SubmitRunRequest) (edgeservice.SubmitRunRequest, error) {
+ return runReq, nil
+ }
+
result, err := s.service.SubmitProviderPool(r.Context(), poolReq)
if err != nil {
metricLabels := s.usageLabelsFor(r.Context(), strings.TrimSpace(req.Model), usageEndpointChatCompletions, responseModePassthrough)
@@ -285,6 +293,7 @@ func (s *Server) handleChatCompletionsProviderPool(
Run: req,
Tunnel: poolReq.Tunnel,
PrepareTunnel: poolReq.PrepareTunnel,
+ PrepareRun: poolReq.PrepareRun,
})
}
diff --git a/apps/edge/internal/openai/responses_handler.go b/apps/edge/internal/openai/responses_handler.go
index 71857f7..8990d55 100644
--- a/apps/edge/internal/openai/responses_handler.go
+++ b/apps/edge/internal/openai/responses_handler.go
@@ -3,6 +3,7 @@ package openai
import (
"bytes"
"encoding/json"
+ "errors"
"fmt"
"io"
"net/http"
@@ -45,6 +46,19 @@ func (s *Server) handleResponses(w http.ResponseWriter, r *http.Request) {
writeError(w, http.StatusBadRequest, "invalid_request_error", "model is required")
return
}
+
+ if dispatch.ProviderPool {
+ // Provider-pool path: the one-shot SubmitProviderPool selects a candidate
+ // and dispatches either tunnel passthrough (OpenAI-compatible) or
+ // normalized RunEvent (Ollama/CLI/native) based on the candidate's
+ // executionPath. Raw body is preserved for the tunnel branch; the
+ // normalized branch performs strict decode + prompt build below.
+ s.handleResponsesProviderPool(r, rawBody, env, dispatch, w)
+ return
+ }
+
+ // Non-provider routes: tunnel or normalized. Tunnel routes relay the raw
+ // body verbatim; normalized routes require strict decode + prompt build.
if routeUsesProviderTunnel(dispatch) {
runMeta, workspace, err := parseOpenAIMetadata(env.Metadata)
if err != nil {
@@ -69,8 +83,8 @@ func (s *Server) handleResponses(w http.ResponseWriter, r *http.Request) {
return
}
- // Non-provider routes keep the normalized RunEvent path: strict decode,
- // stream/background rejection, prompt build, and SubmitRun.
+ // Normalized RunEvent path: strict decode, stream/background rejection,
+ // prompt build, and SubmitRun.
var req responsesRequest
if err := decodeResponsesRequest(json.NewDecoder(bytes.NewReader(rawBody)), &req); err != nil {
writeError(w, http.StatusBadRequest, "invalid_request_error", err.Error())
@@ -190,6 +204,269 @@ func (s *Server) handleResponses(w http.ResponseWriter, r *http.Request) {
s.completeResponse(w, r, req, handle, outputPolicy)
}
+// handleResponsesProviderPool dispatches a provider-pool /v1/responses request
+// through the one-shot SubmitProviderPool surface. After a single queue
+// admission, the selected candidate's executionPath determines whether the
+// response is delivered via tunnel passthrough (OpenAI-compatible providers)
+// or the normalized RunEvent path (Ollama/CLI/native). Provider auth is
+// validated via providerTunnelAuthHeaders for tunnel-path only (SDD D01).
+// Strict-output policy only applies to normalized dispatch, not to raw tunnel
+// passthrough (SDD D02).
+//
+// Selection-first flow: the provider is selected by SubmitProviderPool before
+// any path-specific decode, validation, or stream gating. The tunnel branch
+// only rewrites the model field and relays the raw body (SDD S04). The
+// normalized branch performs strict decode + stream/background rejection +
+// prompt build, but only after the candidate is known. This prevents
+// tunnel-preferring unknown or streaming requests from bypassing the selected
+// provider's executionPath rule (SDD S02/S03).
+func (s *Server) handleResponsesProviderPool(
+ r *http.Request,
+ rawBody []byte,
+ env responsesEnvelope,
+ dispatch routeDispatch,
+ w http.ResponseWriter,
+) {
+ // Build a populated base Run for the normalized branch and a Tunnel for
+ // the tunnel branch. Both are sent to SubmitProviderPool so the one-shot
+ // admission selects the provider before any path-specific behavior.
+ runMeta, workspace, err := parseOpenAIMetadata(env.Metadata)
+ if err != nil {
+ writeError(w, http.StatusBadRequest, "invalid_request_error", err.Error())
+ return
+ }
+ for k, v := range principalMetadata(r.Context()) {
+ runMeta[k] = v
+ }
+ if err := validateWorkspaceForRoute(dispatch, workspace); err != nil {
+ writeError(w, http.StatusBadRequest, "invalid_request_error", err.Error())
+ return
+ }
+
+ // Compute estimate and context class from the raw body so both tunnel and
+ // normalized branches carry the same observability / long-context
+ // information the direct tunnel path carries (SDD S02/S03).
+ estimate := estimateInputTokens(string(rawBody), runMeta, nil, nil)
+ contextClass := classifyContext(estimate, s.longContextThreshold())
+ runMeta["openai_model"] = env.Model
+ runMeta["openai_stream"] = strconv.FormatBool(env.Stream)
+ runMeta["estimated_input_tokens"] = strconv.Itoa(estimate)
+ runMeta["context_class"] = contextClass
+
+ // Base Run carries the model-group key that SubmitProviderPool requires
+ // for candidate resolution. Metadata, estimate, and context class are
+ // preserved by PrepareRun on the normalized branch; the tunnel branch
+ // inherits them from the base so tunnel dispatch keeps observability.
+ baseRun := edgeservice.SubmitRunRequest{
+ NodeRef: dispatch.NodeRef,
+ ModelGroupKey: strings.TrimSpace(env.Model),
+ SessionID: dispatch.SessionID,
+ Workspace: workspace,
+ Metadata: runMeta,
+ EstimatedInputTokens: estimate,
+ ContextClass: contextClass,
+ ProviderPool: true,
+ }
+
+ baseTunnel := edgeservice.SubmitProviderTunnelRequest{
+ Metadata: runMeta,
+ EstimatedInputTokens: estimate,
+ ContextClass: contextClass,
+ ModelGroupKey: strings.TrimSpace(env.Model),
+ SessionID: dispatch.SessionID,
+ Method: http.MethodPost,
+ Path: "/v1/responses",
+ Stream: env.Stream,
+ TimeoutSec: dispatch.TimeoutSec,
+ MaxQueue: dispatch.MaxQueue,
+ QueueTimeoutMS: dispatch.QueueTimeoutMS,
+ ProviderPool: true,
+ }
+
+ poolReq := edgeservice.ProviderPoolDispatchRequest{
+ Run: baseRun,
+ Tunnel: baseTunnel,
+ }
+
+ // Pre-dispatch provider auth header injection. Runs inside SubmitProviderPool
+ // BEFORE buildProviderTunnelRequest and the Node Send step, so the auth
+ // header lands on the actual wire request. On failure the slot is released
+ // and no tunnel request is sent (SDD S03).
+ poolReq.PrepareTunnel = func(tunnelReq edgeservice.SubmitProviderTunnelRequest) (edgeservice.SubmitProviderTunnelRequest, error) {
+ headers, err := s.providerTunnelAuthHeaders(r)
+ if err != nil {
+ return tunnelReq, err
+ }
+ tunnelReq.Headers = headers
+ return tunnelReq, nil
+ }
+
+ // Tunnel branch: rewrite only the model field, preserve all other fields
+ // (tools, max_output_tokens, custom fields). Stream/background gating is
+ // skipped: the provider itself enforces those constraints (SDD S04).
+ poolReq.Tunnel.BuildBody = func(target string) ([]byte, error) {
+ return rewriteResponsesModel(rawBody, target)
+ }
+
+ // Prepared values captured from the PrepareRun hook on the normalized path.
+ // PrepareRun runs strict decode + prompt build + output policy; the
+ // normalized branch must relay these exact values to completeResponse so
+ // strict-output XML wrapping survives the provider-pool path.
+ var (
+ preparedResponsesReq responsesRequest
+ preparedOutputPolicy strictOutputPolicy
+ prepareRunSucceeded bool
+ )
+
+ // Normalized branch: strict decode, stream/background rejection, prompt
+ // build, and generation policy. Only runs AFTER the candidate is selected
+ // by SubmitProviderPool (i.e. only when executionPath is normalized).
+ // Validation errors are wrapped with errProviderRequestValidation so
+ // SubmitProviderPool can relay them as HTTP 400.
+ poolReq.PrepareRun = func(runReq edgeservice.SubmitRunRequest) (edgeservice.SubmitRunRequest, error) {
+ req := responsesRequest{Model: env.Model}
+ if err := decodeResponsesRequest(json.NewDecoder(bytes.NewReader(rawBody)), &req); err != nil {
+ return edgeservice.SubmitRunRequest{}, fmt.Errorf("%w: invalid /v1/responses request format: %s", errProviderRequestValidation, err.Error())
+ }
+ if req.Stream {
+ return edgeservice.SubmitRunRequest{}, fmt.Errorf("%w: streaming is not supported for /v1/responses", errProviderRequestValidation)
+ }
+ if req.Background {
+ return edgeservice.SubmitRunRequest{}, fmt.Errorf("%w: background is not supported for /v1/responses", errProviderRequestValidation)
+ }
+
+ inputStr, parseInputErr := parseResponsesInput(req.Input)
+ if parseInputErr != nil {
+ return edgeservice.SubmitRunRequest{}, fmt.Errorf("%w: %s", errProviderRequestValidation, parseInputErr.Error())
+ }
+
+ prompt := buildResponsesPrompt(req.Instructions, inputStr)
+ outputPolicy := s.resolveOutputPolicy(prompt)
+ if instruction := strictOutputContractInstruction(outputPolicy); instruction != "" {
+ prompt = instruction + "\n" + prompt
+ }
+
+ preparedMeta, _, metaErr := parseOpenAIMetadata(req.Metadata)
+ if metaErr != nil {
+ return edgeservice.SubmitRunRequest{}, fmt.Errorf("%w: %s", errProviderRequestValidation, metaErr.Error())
+ }
+ for k, v := range principalMetadata(r.Context()) {
+ preparedMeta[k] = v
+ }
+
+ if workspaceErr := validateWorkspaceForRoute(dispatch, workspace); workspaceErr != nil {
+ return edgeservice.SubmitRunRequest{}, fmt.Errorf("%w: %s", errProviderRequestValidation, workspaceErr.Error())
+ }
+
+ var defaultThinkingTokenBudget int
+ if catalogEntry := s.findProviderPoolEntry(req.Model); catalogEntry != nil {
+ applyModelCatalogGenerationPolicyToResponses(&req, *catalogEntry)
+ if catalogEntry.DefaultThinkingTokenBudget > 0 {
+ defaultThinkingTokenBudget = catalogEntry.DefaultThinkingTokenBudget
+ }
+ }
+
+ preparedMeta["openai_model"] = req.Model
+ preparedMeta["openai_stream"] = fmt.Sprintf("%t", req.Stream)
+ preparedMeta["strict_output"] = fmt.Sprintf("%t", outputPolicy.Strict)
+
+ input := map[string]any{"prompt": prompt}
+ if defaultThinkingTokenBudget > 0 {
+ input["think"] = true
+ input["thinking_token_budget"] = defaultThinkingTokenBudget
+ } else if outputPolicy.Strict {
+ input["think"] = false
+ }
+ if options := req.providerOptions(); len(options) > 0 {
+ input["options"] = options
+ }
+
+ estimate := estimateInputTokens(prompt, preparedMeta, nil, nil)
+ contextClass := classifyContext(estimate, s.longContextThreshold())
+ preparedMeta["estimated_input_tokens"] = strconv.Itoa(estimate)
+ preparedMeta["context_class"] = contextClass
+
+ runReq.Prompt = prompt
+ runReq.Input = input
+ runReq.Metadata = preparedMeta
+ runReq.EstimatedInputTokens = estimate
+ runReq.ContextClass = contextClass
+ runReq.TimeoutSec = dispatch.TimeoutSec
+ runReq.MaxQueue = dispatch.MaxQueue
+ runReq.QueueTimeoutMS = dispatch.QueueTimeoutMS
+
+ preparedResponsesReq = req
+ preparedOutputPolicy = outputPolicy
+ prepareRunSucceeded = true
+ return runReq, nil
+ }
+
+ result, err := s.service.SubmitProviderPool(r.Context(), poolReq)
+ if err != nil {
+ metricLabels := s.usageLabelsFor(r.Context(), strings.TrimSpace(env.Model), usageEndpointResponses, responseModePassthrough)
+ emitUsageMetrics(metricLabels, usageStatusForError(err), usageObservation{})
+ // Provider auth failure is a client request error (400), not a
+ // backend dispatch error. The auth check runs inside
+ // SubmitProviderPool (PrepareTunnel) before any tunnel request is sent.
+ // Normalized validation failures (e.g. strict decode, stream) are
+ // also client errors: map to 400.
+ if errors.Is(err, errProviderAuthRequired) {
+ writeError(w, http.StatusBadRequest, "invalid_request_error", "provider auth token is required")
+ return
+ }
+ if isValidationError(err) {
+ writeError(w, http.StatusBadRequest, "invalid_request_error", err.Error())
+ return
+ }
+ writeError(w, http.StatusBadGateway, "node_dispatch_error", err.Error())
+ return
+ }
+
+ s.logger.Info("openai responses provider-pool dispatch",
+ zap.String("run_id", result.DispatchInfo.RunID),
+ zap.String("node_id", result.DispatchInfo.NodeID),
+ zap.String("provider_id", result.DispatchInfo.ProviderID),
+ zap.String("provider_type", result.DispatchInfo.ProviderType),
+ zap.String("execution_path", result.DispatchInfo.ExecutionPath),
+ zap.String("model_group", result.DispatchInfo.ModelGroupKey),
+ zap.String("adapter", result.DispatchInfo.Adapter),
+ zap.String("target", result.DispatchInfo.Target),
+ zap.String("path", string(result.Path)),
+ zap.Int("estimated_input_tokens", result.DispatchInfo.EstimatedInputTokens),
+ zap.String("context_class", result.DispatchInfo.ContextClass),
+ zap.String("queue_reason", result.DispatchInfo.QueueReason),
+ )
+
+ switch result.Path {
+ case edgeservice.ProviderPoolPathTunnel:
+ // Tunnel path: provider auth was already validated and injected via
+ // PrepareTunnel before dispatch; on failure SubmitProviderPool returns
+ // an error and no tunnel handle exists. Provider bytes are relayed as
+ // pure passthrough; caller metadata never selects a sideband surface.
+ metricLabels := s.usageLabelsFor(r.Context(), strings.TrimSpace(env.Model), usageEndpointResponses, responseModePassthrough)
+ s.writeProviderTunnelResponse(w, r, result.Tunnel, env.Stream, env.Model, metricLabels)
+
+ case edgeservice.ProviderPoolPathNormalized:
+ // Normalized path: no auth required, collect from RunEvent stream.
+ handle := result.Run
+ if handle == nil {
+ writeError(w, http.StatusInternalServerError, "run_error", "provider-pool selection returned normalized path but no run result")
+ return
+ }
+ // PrepareRun is only invoked on the normalized path. A normalized
+ // result without a successful PrepareRun is an internal mismatch and
+ // must surface as run_error (SDD D02).
+ if !prepareRunSucceeded {
+ writeError(w, http.StatusInternalServerError, "run_error", "provider-pool normalized result received without a successful PrepareRun")
+ return
+ }
+ // Relay the prepared request and output policy from PrepareRun to
+ // completeResponse so strict-output XML wrapping survives the
+ // normalized provider-pool path.
+ s.completeResponse(w, r, preparedResponsesReq, handle, preparedOutputPolicy)
+ }
+}
+
func decodeResponsesRequest(dec *json.Decoder, req *responsesRequest) error {
var raw map[string]json.RawMessage
if err := dec.Decode(&raw); err != nil {
@@ -441,6 +718,22 @@ func parseOpenAIMetadata(raw json.RawMessage) (map[string]string, string, error)
return flat, workspace, nil
}
+// errProviderRequestValidation is a sentinel for client request validation
+// failures that occur inside the normalized provider-pool PrepareRun hook
+// (strict decode, stream/background rejection, input parsing). These errors
+// are reported as HTTP 400 invalid_request_error, distinct from backend
+// dispatch errors that return 502.
+var errProviderRequestValidation = errors.New("provider_request_validation")
+
+// isValidationError reports whether err originated from a PrepareRun validation
+// failure in the normalized provider-pool path.
+func isValidationError(err error) bool {
+ return errors.Is(err, errProviderRequestValidation)
+}
+
+// metadataStringValue extracts a single string value from a json.RawMessage
+// for the metadata field. It rejects non-string values and enforces the
+// 512-character length cap.
func metadataStringValue(key string, raw json.RawMessage) (string, error) {
var value string
if err := json.Unmarshal(raw, &value); err != nil {
diff --git a/apps/edge/internal/openai/server_test.go b/apps/edge/internal/openai/server_test.go
index 4cbf9df..9bab1ba 100644
--- a/apps/edge/internal/openai/server_test.go
+++ b/apps/edge/internal/openai/server_test.go
@@ -71,6 +71,18 @@ type fakeRunService struct {
// calls. When non-empty, results are consumed in order; when empty, the
// poolDispatchPath / poolSubmitErr fields are used instead.
poolSubmitResults []edgeservice.ProviderPoolDispatchResult
+ // poolSubmitCount tracks how many times SubmitProviderPool was called
+ // (including from the chat handler). Used by tests to assert dispatch
+ // surface fidelity.
+ poolSubmitCount int
+ // poolLastRun records the last SubmitRunRequest passed to SubmitProviderPool
+ // (pre-PrepareRun). Tests assert ModelGroupKey is non-empty to catch the
+ // "empty Run" regression.
+ poolLastRun edgeservice.SubmitRunRequest
+ // poolPrepareRunCalled tracks whether the optional PrepareRun hook was
+ // invoked by SubmitProviderPool on the normalized path. Only meaningful
+ // when poolDispatchPath == "normalized".
+ poolPrepareRunCalled bool
}
type fakeTunnelHandle struct {
@@ -158,7 +170,15 @@ func (s *fakeRunService) SubmitProviderTunnel(_ context.Context, req edgeservice
// SubmitProviderPool is a minimal test double for the provider-pool one-shot
// dispatch. It records the selected path (tunnel vs normalized) and returns a
// stub result so the OpenAI handler can exercise the selection-first flow.
+// It also tracks whether the optional PrepareRun hook was invoked, so tests
+// can assert the handler only calls PrepareRun on the normalized path.
func (s *fakeRunService) SubmitProviderPool(_ context.Context, req edgeservice.ProviderPoolDispatchRequest) (*edgeservice.ProviderPoolDispatchResult, error) {
+ s.submitMu.Lock()
+ s.poolSubmitCount++
+ s.poolLastRun = req.Run
+ s.poolPrepareRunCalled = false
+ s.submitMu.Unlock()
+
// Custom results for successive calls (e.g. tool-validation retry tests).
s.submitMu.Lock()
if len(s.poolSubmitResults) > 0 {
@@ -254,6 +274,20 @@ func (s *fakeRunService) SubmitProviderPool(_ context.Context, req edgeservice.P
frames = relayed
} else if frames == nil {
frames = staticProviderTunnelFrames(`{"ok":true}`)
+ // Record the built tunnel body even when using a static frame source.
+ body := req.Tunnel.Body
+ if req.Tunnel.BuildBody != nil {
+ target := req.Tunnel.Target
+ if s.tunnelServedTarget != "" {
+ target = s.tunnelServedTarget
+ }
+ if built, buildErr := req.Tunnel.BuildBody(target); buildErr == nil {
+ body = built
+ }
+ }
+ s.submitMu.Lock()
+ s.tunnelBodies = append(s.tunnelBodies, body)
+ s.submitMu.Unlock()
} else {
body := req.Tunnel.Body
if req.Tunnel.BuildBody != nil {
@@ -295,9 +329,22 @@ func (s *fakeRunService) SubmitProviderPool(_ context.Context, req edgeservice.P
}
// Normalized path (mimics Ollama/CLI behavior).
+ // Invoke the optional PrepareRun hook so tests can assert it only runs
+ // on the normalized path (mirrors real SubmitProviderPool behavior).
+ if req.PrepareRun != nil {
+ s.submitMu.Lock()
+ s.poolPrepareRunCalled = true
+ s.submitMu.Unlock()
+ prepared, prepErr := req.PrepareRun(req.Run)
+ if prepErr != nil {
+ return nil, prepErr
+ }
+ req.Run = prepared
+ }
s.submitMu.Lock()
s.reqs = append(s.reqs, req.Run)
s.submitMu.Unlock()
+
runChan := s.poolRunFrames
if runChan == nil {
runChan = make(chan *iop.RunEvent, 1)
@@ -376,6 +423,18 @@ func (s *fakeRunService) buildPoolResult(req edgeservice.ProviderPoolDispatchReq
if disp.NodeID == "" {
disp.NodeID = "node-pool"
}
+ // Invoke the optional PrepareRun hook on the normalized path (mirrors
+ // real SubmitProviderPool behavior).
+ if req.PrepareRun != nil {
+ prepared, prepErr := req.PrepareRun(req.Run)
+ if prepErr != nil {
+ return nil, prepErr
+ }
+ req.Run = prepared
+ }
+ s.submitMu.Lock()
+ s.poolPrepareRunCalled = true
+ s.submitMu.Unlock()
s.submitMu.Lock()
s.reqs = append(s.reqs, req.Run)
s.submitMu.Unlock()
@@ -423,6 +482,32 @@ func (s *fakeRunService) tunnelHandleSnapshot() tunnelHandleSnapshot {
}
}
+// poolSubmitCountSnapshot returns the number of times SubmitProviderPool was
+// called. Tests assert exactly-once dispatch from the handler.
+func (s *fakeRunService) poolSubmitCountSnapshot() int {
+ s.submitMu.Lock()
+ defer s.submitMu.Unlock()
+ return s.poolSubmitCount
+}
+
+// poolLastRunSnapshot returns a copy of the last SubmitRunRequest passed to
+// SubmitProviderPool. Tests assert ModelGroupKey is non-empty to catch the
+// "empty Run" regression.
+func (s *fakeRunService) poolLastRunSnapshot() edgeservice.SubmitRunRequest {
+ s.submitMu.Lock()
+ defer s.submitMu.Unlock()
+ return s.poolLastRun
+}
+
+// poolPrepareRunCalledSnapshot returns whether the PrepareRun hook was
+// invoked by the last SubmitProviderPool call. Only meaningful when the last
+// call was on the normalized path.
+func (s *fakeRunService) poolPrepareRunCalledSnapshot() bool {
+ s.submitMu.Lock()
+ defer s.submitMu.Unlock()
+ return s.poolPrepareRunCalled
+}
+
func (s *fakeRunService) tunnelReqsSnapshot() []edgeservice.SubmitProviderTunnelRequest {
s.submitMu.Lock()
reqs := make([]edgeservice.SubmitProviderTunnelRequest, len(s.tunnelReqs))
@@ -5676,6 +5761,206 @@ func responsesProviderTunnelServer(frames chan *iop.ProviderTunnelFrame, servedT
return srv, fake
}
+// TestResponsesProviderPoolTunnelSelectionUsesPassthrough verifies that when a
+// provider-pool catalog match is resolved and the selected provider is an
+// OpenAI-compatible type (tunnel executionPath), the Responses handler
+// dispatches via tunnel passthrough with the raw body preserved (model rewrite
+// only). This mirrors TestChatCompletionsProviderPoolTunnelSelectionUsesPassthrough
+// for /v1/responses.
+func TestResponsesProviderPoolTunnelSelectionUsesPassthrough(t *testing.T) {
+ frames := make(chan *iop.ProviderTunnelFrame, 3)
+ frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200}
+ frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(`{"ok":true}`)}
+ frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true}
+ close(frames)
+
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
+ tunnelFrames: frames,
+ }
+ catalog := []config.ModelCatalogEntry{
+ {ID: "vllm-resp-model", Providers: map[string]string{"prov-vllm": "Qwen3-35B"}},
+ }
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog(catalog)
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"vllm-resp-model",
+ "input":"hello world"
+ }`))
+ w := httptest.NewRecorder()
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
+ }
+ // Tunnel passthrough must use SubmitProviderPool tunnel path.
+ tunnelReqs := fake.tunnelReqsSnapshot()
+ reqs := fake.reqsSnapshot()
+ if len(tunnelReqs) != 1 {
+ t.Fatalf("expected 1 tunnel dispatch for vLLM path, got %d", len(tunnelReqs))
+ }
+ if len(reqs) != 0 {
+ t.Fatalf("expected 0 normalized SubmitRun for vLLM tunnel path, got %d", len(reqs))
+ }
+ if !tunnelReqs[0].ProviderPool {
+ t.Error("tunnel dispatch must have ProviderPool=true")
+ }
+ if tunnelReqs[0].Path != "/v1/responses" {
+ t.Fatalf("tunnel path: got %q want /v1/responses", tunnelReqs[0].Path)
+ }
+}
+
+// TestResponsesProviderPoolOllamaSelectionUsesNormalizedRun verifies that when a
+// provider-pool catalog match is resolved and the selected provider is an
+// Ollama type (normalized executionPath), the Responses handler dispatches via
+// normalized RunEvent path with the decoded prompt/input, not via tunnel. This
+// mirrors TestChatCompletionsProviderPoolOllamaSelectionUsesNormalizedRun for
+// /v1/responses.
+func TestResponsesProviderPoolOllamaSelectionUsesNormalizedRun(t *testing.T) {
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathNormalized),
+ }
+ catalog := []config.ModelCatalogEntry{
+ {ID: "ollama-resp-model", Providers: map[string]string{"prov-ollama": "qwen3.6:35b"}},
+ }
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog(catalog)
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"ollama-resp-model",
+ "input":"hello world"
+ }`))
+ w := httptest.NewRecorder()
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
+ }
+ // Normalized Ollama path must use SubmitProviderPool normalized path.
+ reqs := fake.reqsSnapshot()
+ tunnelReqs := fake.tunnelReqsSnapshot()
+ if len(reqs) != 1 {
+ t.Fatalf("expected 1 normalized SubmitRun for Ollama path, got %d", len(reqs))
+ }
+ if len(tunnelReqs) != 0 {
+ t.Fatalf("expected 0 tunnel dispatch for Ollama normalized path, got %d", len(tunnelReqs))
+ }
+ if !reqs[0].ProviderPool {
+ t.Error("normalized dispatch must have ProviderPool=true")
+ }
+ if reqs[0].ModelGroupKey != "ollama-resp-model" {
+ t.Fatalf("expected model group key %q, got %q", "ollama-resp-model", reqs[0].ModelGroupKey)
+ }
+}
+
+// TestResponsesMixedProviderPoolTunnelSelection verifies S05: a mixed
+// provider-pool model group that selects an OpenAI-compatible (vLLM) provider
+// dispatches via tunnel passthrough for /v1/responses.
+func TestResponsesMixedProviderPoolTunnelSelection(t *testing.T) {
+ tunnelFrames := make(chan *iop.ProviderTunnelFrame, 3)
+ tunnelFrames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200}
+ tunnelFrames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(`{"id":"resp-mixed","object":"response"}`)}
+ tunnelFrames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true}
+ close(tunnelFrames)
+
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
+ tunnelFrames: tunnelFrames,
+ }
+ catalog := []config.ModelCatalogEntry{
+ {
+ ID: "mixed-resp-model",
+ Providers: map[string]string{
+ "prov-vllm": "served-vllm",
+ "prov-ollama": "served-ollama",
+ },
+ },
+ }
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog(catalog)
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"mixed-resp-model",
+ "input":"hello mixed"
+ }`))
+ w := httptest.NewRecorder()
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
+ }
+
+ tunnelReqs := fake.tunnelReqsSnapshot()
+ if len(tunnelReqs) != 1 {
+ t.Fatalf("expected 1 tunnel dispatch, got %d", len(tunnelReqs))
+ }
+ if !tunnelReqs[0].ProviderPool {
+ t.Error("tunnel request must have ProviderPool=true")
+ }
+ if tunnelReqs[0].Path != "/v1/responses" {
+ t.Fatalf("tunnel path: got %q want /v1/responses", tunnelReqs[0].Path)
+ }
+
+ runReqs := fake.reqsSnapshot()
+ if len(runReqs) != 0 {
+ t.Errorf("expected 0 normalized SubmitRun calls, got %d", len(runReqs))
+ }
+}
+
+// TestResponsesMixedProviderPoolOllamaSelection verifies S05: when the mixed
+// model group selects an Ollama provider, the request is dispatched over the
+// normalized path (SubmitRun), and no tunnel dispatch is issued for /v1/responses.
+func TestResponsesMixedProviderPoolOllamaSelection(t *testing.T) {
+ events := bufferedRunEvents(
+ &iop.RunEvent{Type: "delta", Delta: "ollama-resp"},
+ &iop.RunEvent{Type: "complete"},
+ )
+
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathNormalized),
+ events: events,
+ }
+ catalog := []config.ModelCatalogEntry{
+ {
+ ID: "mixed-resp-model",
+ Providers: map[string]string{
+ "prov-vllm": "served-vllm",
+ "prov-ollama": "served-ollama",
+ },
+ },
+ }
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog(catalog)
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"mixed-resp-model",
+ "input":"hello mixed"
+ }`))
+ w := httptest.NewRecorder()
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
+ }
+
+ tunnelReqs := fake.tunnelReqsSnapshot()
+ if len(tunnelReqs) != 0 {
+ t.Fatalf("expected 0 tunnel dispatches for Ollama selection, got %d", len(tunnelReqs))
+ }
+
+ runReqs := fake.reqsSnapshot()
+ if len(runReqs) != 1 {
+ t.Fatalf("expected 1 normalized SubmitRun, got %d", len(runReqs))
+ }
+ if runReqs[0].ModelGroupKey != "mixed-resp-model" {
+ t.Errorf("normalized ModelGroupKey: got %q, want mixed-resp-model", runReqs[0].ModelGroupKey)
+ }
+ if !runReqs[0].ProviderPool {
+ t.Error("normalized run must have ProviderPool=true")
+ }
+}
+
// TestResponsesProviderPoolDispatch verifies that /v1/responses sends
// provider-pool models through the raw provider tunnel (POST /v1/responses)
// instead of the normalized RunEvent path.
@@ -7440,3 +7725,321 @@ func TestProviderPoolStandardResponseNoExtensionFields(t *testing.T) {
t.Error("expected provider tunnel dispatch")
}
}
+
+// TestResponsesProviderPoolTunnelUnknownFieldsUsesPoolSelection asserts that
+// when the selected provider is tunnel-passthrough (OpenAI-compatible),
+// unknown fields (tools, parallel_tool_calls, store, custom field) are
+// preserved verbatim in the tunnel body, and that the Run sent to
+// SubmitProviderPool carries a non-empty ModelGroupKey.
+func TestResponsesProviderPoolTunnelUnknownFieldsUsesPoolSelection(t *testing.T) {
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
+ }
+
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog([]config.ModelCatalogEntry{
+ {ID: "pool-model", Providers: map[string]string{"prov-tunnel": "served-tunnel"}},
+ })
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"pool-model",
+ "input":"hello world",
+ "tools":[{"type":"function","function":{"name":"lookup"}}],
+ "parallel_tool_calls":true,
+ "store":false,
+ "custom_field":"keep-me"
+ }`))
+ w := httptest.NewRecorder()
+
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
+ }
+
+ // SubmitProviderPool must have been called with a populated Run.
+ lastRun := fake.poolLastRunSnapshot()
+ if lastRun.ModelGroupKey == "" {
+ t.Error("SubmitProviderPool was called with empty Run.ModelGroupKey")
+ }
+
+ // The tunnel body must preserve unknown fields verbatim.
+ if len(fake.tunnelBodiesSnapshot()) != 1 {
+ t.Fatalf("expected exactly one tunnel body, got %d", len(fake.tunnelBodiesSnapshot()))
+ }
+ body := fake.tunnelBodiesSnapshot()[0]
+ if !strings.Contains(string(body), "\"tools\"") {
+ t.Error("tunnel body must preserve tools field")
+ }
+ if !strings.Contains(string(body), "\"parallel_tool_calls\"") {
+ t.Error("tunnel body must preserve parallel_tool_calls field")
+ }
+ if !strings.Contains(string(body), "\"store\"") {
+ t.Error("tunnel body must preserve store field")
+ }
+ if !strings.Contains(string(body), "\"custom_field\"") {
+ t.Error("tunnel body must preserve custom_field")
+ }
+}
+
+// TestResponsesProviderPoolOllamaSelectionRejectsUnknownFieldsWithoutTunnel
+// asserts that when the selected provider is normalized (Ollama/CLI), an
+// unknown-field request returns HTTP 400 and no tunnel or run dispatch occurs.
+func TestResponsesProviderPoolOllamaSelectionRejectsUnknownFieldsWithoutTunnel(t *testing.T) {
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathNormalized),
+ }
+
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog([]config.ModelCatalogEntry{
+ {ID: "ollama-model", Providers: map[string]string{"prov-ollama": "served-ollama"}},
+ })
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"ollama-model",
+ "input":"hello",
+ "unknown_field":"reject-me"
+ }`))
+ w := httptest.NewRecorder()
+
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusBadRequest {
+ t.Fatalf("status: got %d body=%s, want 400", w.Code, w.Body.String())
+ }
+
+ // No tunnel or run dispatch should have occurred.
+ if len(fake.tunnelReqsSnapshot()) != 0 {
+ t.Error("normalized path must not dispatch tunnel for unknown fields")
+ }
+ if len(fake.reqsSnapshot()) != 0 {
+ t.Error("normalized path must not dispatch run for unknown fields")
+ }
+
+ // Verify pool dispatch was made.
+ if fake.poolSubmitCountSnapshot() != 1 {
+ t.Fatalf("pool submit count: got %d, want 1", fake.poolSubmitCountSnapshot())
+ }
+
+ // Verify PrepareRun was called on the normalized path.
+ prepareRunCalled := fake.poolPrepareRunCalledSnapshot()
+ if !prepareRunCalled {
+ t.Errorf("PrepareRun must be called on normalized path (poolSubmitCount=%d, poolLastRun.ModelGroupKey=%q)", fake.poolSubmitCountSnapshot(), fake.poolLastRunSnapshot().ModelGroupKey)
+ }
+}
+
+// TestResponsesProviderPoolOllamaSelectionRejectsStreamingWithoutTunnel asserts
+// that when the selected provider is normalized, stream:true returns HTTP 400
+// and no tunnel dispatch occurs.
+func TestResponsesProviderPoolOllamaSelectionRejectsStreamingWithoutTunnel(t *testing.T) {
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathNormalized),
+ }
+
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog([]config.ModelCatalogEntry{
+ {ID: "ollama-model", Providers: map[string]string{"prov-ollama": "served-ollama"}},
+ })
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"ollama-model",
+ "input":"hello",
+ "stream":true
+ }`))
+ w := httptest.NewRecorder()
+
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusBadRequest {
+ t.Fatalf("status: got %d body=%s, want 400", w.Code, w.Body.String())
+ }
+
+ // No tunnel dispatch for normalized path with stream:true.
+ if len(fake.tunnelReqsSnapshot()) != 0 {
+ t.Error("normalized path must not dispatch tunnel for stream:true")
+ }
+ if len(fake.reqsSnapshot()) != 0 {
+ t.Error("normalized path must not dispatch run for stream:true")
+ }
+
+ if !fake.poolPrepareRunCalledSnapshot() {
+ t.Error("PrepareRun must be called on normalized path")
+ }
+}
+
+// TestResponsesProviderPoolTunnelStreamingUsesPoolSelection asserts that when
+// the selected provider is tunnel-passthrough, stream:true is relayed through
+// SubmitProviderPool with the raw SSE body preserved.
+func TestResponsesProviderPoolTunnelStreamingUsesPoolSelection(t *testing.T) {
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
+ tunnelFrames: staticProviderTunnelFrames(`{"id":"resp-1","output":[{"type":"message","role":"assistant","content":[{"type":"output_text","text":"hello"}]}],"stream":true}\n`),
+ }
+
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog([]config.ModelCatalogEntry{
+ {ID: "pool-model", Providers: map[string]string{"prov-tunnel": "served-tunnel"}},
+ })
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"pool-model",
+ "input":"hello",
+ "stream":true
+ }`))
+ w := httptest.NewRecorder()
+
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
+ }
+
+ // SubmitProviderPool must have been called.
+ if fake.poolSubmitCountSnapshot() != 1 {
+ t.Fatalf("SubmitProviderPool call count: got %d, want 1", fake.poolSubmitCountSnapshot())
+ }
+
+ // The tunnel body must preserve the raw SSE payload.
+ if len(fake.tunnelBodiesSnapshot()) != 1 {
+ t.Fatalf("expected one tunnel body, got %d", len(fake.tunnelBodiesSnapshot()))
+ }
+ body := fake.tunnelBodiesSnapshot()[0]
+ if !strings.Contains(string(body), "\"stream\":true") {
+ t.Error("tunnel body must preserve stream:true")
+ }
+
+ // PrepareRun must NOT be called on the tunnel path.
+ if fake.poolPrepareRunCalledSnapshot() {
+ t.Error("PrepareRun must NOT be called on tunnel path")
+ }
+}
+
+// TestResponsesProviderPoolStrictOutputNormalizesAgentResponse verifies that the
+// provider-pool normalized path preserves the strict-output contract from
+// PrepareRun into completeResponse. A prompt containing an XML completion
+// contract must have its response wrapped in the expected block (SDD D02).
+func TestResponsesProviderPoolStrictOutputNormalizesAgentResponse(t *testing.T) {
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathNormalized),
+ }
+ fake.poolRunFrames = make(chan *iop.RunEvent, 3)
+ fake.poolRunFrames <- &iop.RunEvent{Type: "reasoning_delta", Delta: "hidden"}
+ fake.poolRunFrames <- &iop.RunEvent{Type: "delta", Delta: "plain answer"}
+ fake.poolRunFrames <- &iop.RunEvent{Type: "complete", RunId: "run-pool-normalized"}
+ close(fake.poolRunFrames)
+
+ catalog := []config.ModelCatalogEntry{
+ {ID: "ollama-strict-model", Providers: map[string]string{"prov-ollama": "served-ollama"}},
+ }
+ srv := NewServer(config.EdgeOpenAIConf{StrictOutput: true}, fake, nil)
+ srv.SetModelCatalog(catalog)
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"ollama-strict-model",
+ "input":"Once you've completed the user's task, you must use the attempt_completion tool to present the result.\n\n\ndone\n\n\nfinish"
+ }`))
+ w := httptest.NewRecorder()
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
+ }
+
+ // PrepareRun must have been invoked (normalized path).
+ if !fake.poolPrepareRunCalledSnapshot() {
+ t.Fatal("PrepareRun was not called on normalized path")
+ }
+
+ // The prepared request must carry strict_output and the injected contract
+ // instruction.
+ if len(fake.reqsSnapshot()) != 1 {
+ t.Fatalf("expected one prepared req, got %d", len(fake.reqsSnapshot()))
+ }
+ reqPrep := fake.reqsSnapshot()[0]
+ if reqPrep.Metadata["strict_output"] != "true" {
+ t.Fatalf("strict_output metadata lost in prepared request: %+v", reqPrep.Metadata)
+ }
+ if !strings.Contains(reqPrep.Prompt, "output exactly one block") {
+ t.Fatalf("strict contract instruction not injected into prepared prompt:\n%s", reqPrep.Prompt)
+ }
+
+ var resp responsesResponse
+ if err := json.Unmarshal(w.Body.Bytes(), &resp); err != nil {
+ t.Fatalf("decode: %v", err)
+ }
+ want := "\nplain answer\n"
+ if resp.OutputText != want {
+ t.Fatalf("output_text:\ngot %q\nwant %q", resp.OutputText, want)
+ }
+}
+
+// TestResponsesProviderPoolTunnelMetadataAndContextPropagation verifies that the
+// provider-pool tunnel path preserves metadata, estimated_input_tokens, and
+// context_class from the base request into the tunnel dispatch (SDD S02/S03).
+func TestResponsesProviderPoolTunnelMetadataAndContextPropagation(t *testing.T) {
+ frames := make(chan *iop.ProviderTunnelFrame, 3)
+ frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_RESPONSE_START, StatusCode: 200}
+ frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_BODY, Body: []byte(`{"ok":true}`)}
+ frames <- &iop.ProviderTunnelFrame{Kind: iop.ProviderTunnelFrameKind_PROVIDER_TUNNEL_FRAME_KIND_END, End: true}
+ close(frames)
+
+ fake := &fakeRunService{
+ poolDispatchPath: string(edgeservice.ProviderPoolPathTunnel),
+ tunnelFrames: frames,
+ tunnelServedTarget: "served-model",
+ }
+
+ catalog := []config.ModelCatalogEntry{
+ {ID: "pool-model", Providers: map[string]string{"prov-1": "served-model"}},
+ }
+ srv := NewServer(config.EdgeOpenAIConf{}, fake, nil)
+ srv.SetModelCatalog(catalog)
+
+ req := httptest.NewRequest(http.MethodPost, "/v1/responses", strings.NewReader(`{
+ "model":"pool-model",
+ "input":"hello",
+ "metadata":{"experiment":"pool-test"}
+ }`))
+ w := httptest.NewRecorder()
+ srv.handleResponses(w, req)
+
+ if w.Code != http.StatusOK {
+ t.Fatalf("status: got %d body=%s", w.Code, w.Body.String())
+ }
+
+ tunnelReqs := fake.tunnelReqsSnapshot()
+ if len(tunnelReqs) != 1 {
+ t.Fatalf("expected one tunnel req, got %d", len(tunnelReqs))
+ }
+ tunnelReq := tunnelReqs[0]
+
+ // Metadata must contain the principal and pool-echoed fields.
+ if tunnelReq.Metadata == nil {
+ t.Fatalf("tunnel metadata must not be nil")
+ }
+ if tunnelReq.Metadata["openai_model"] == "" {
+ t.Error("openai_model missing from tunnel metadata")
+ }
+ if tunnelReq.Metadata["openai_stream"] == "" {
+ t.Error("openai_stream missing from tunnel metadata")
+ }
+ if tunnelReq.Metadata["estimated_input_tokens"] == "" {
+ t.Error("estimated_input_tokens missing from tunnel metadata")
+ }
+ if tunnelReq.Metadata["context_class"] == "" {
+ t.Error("context_class missing from tunnel metadata")
+ }
+ if tunnelReq.Metadata["experiment"] != "pool-test" {
+ t.Errorf("caller metadata lost: got %v", tunnelReq.Metadata)
+ }
+
+ // EstimatedInputTokens and ContextClass must be non-zero/non-empty on the
+ // tunnel request itself, so Node observability preserves the same values
+ // as the direct tunnel path.
+ if tunnelReq.EstimatedInputTokens <= 0 {
+ t.Errorf("EstimatedInputTokens must be > 0: got %d", tunnelReq.EstimatedInputTokens)
+ }
+ if tunnelReq.ContextClass == "" {
+ t.Error("ContextClass must be non-empty on tunnel request")
+ }
+}
diff --git a/apps/edge/internal/service/run_dispatch.go b/apps/edge/internal/service/run_dispatch.go
index ab079b9..cba4dbb 100644
--- a/apps/edge/internal/service/run_dispatch.go
+++ b/apps/edge/internal/service/run_dispatch.go
@@ -921,6 +921,15 @@ const (
// after wire dispatch for provider-pool tunnel paths.
type prepareTunnelFunc func(req SubmitProviderTunnelRequest) (SubmitProviderTunnelRequest, error)
+// PrepareRun is an optional pre-dispatch hook that lets the caller prepare
+// or validate the Run request for the normalized execution path. It runs
+// AFTER a candidate is selected and ONLY on the normalized path (not on
+// tunnel/passthrough). If PrepareRun returns an error, the slot is released
+// and no normalized run is sent. This lets the caller enforce that a complete
+// SubmitRunRequest (including ModelGroupKey and other required fields) is
+// present before dispatch, while leaving the tunnel branch untouched.
+type prepareRunFunc func(req SubmitRunRequest) (SubmitRunRequest, error)
+
// ProviderPoolDispatchRequest bundles the Run and Tunnel surface values for
// a single one-shot provider-pool dispatch. SubmitProviderPool uses exactly
// one queue admission to select a candidate, then dispatches only the
@@ -929,6 +938,7 @@ type ProviderPoolDispatchRequest struct {
Run SubmitRunRequest
Tunnel SubmitProviderTunnelRequest
PrepareTunnel prepareTunnelFunc
+ PrepareRun prepareRunFunc
}
// ProviderPoolDispatchResult describes which execution path was selected and
@@ -987,6 +997,7 @@ func (s *Service) SubmitProviderPool(ctx context.Context, req ProviderPoolDispat
tunnelReq.SessionID = req.Run.SessionID
tunnelReq.MaxQueue = req.Run.MaxQueue
tunnelReq.QueueTimeoutMS = req.Run.QueueTimeoutMS
+ tunnelReq.Metadata = req.Run.Metadata
tunnelReq.EstimatedInputTokens = req.Run.EstimatedInputTokens
tunnelReq.ContextClass = req.Run.ContextClass
@@ -1036,7 +1047,15 @@ func (s *Service) SubmitProviderPool(ctx context.Context, req ProviderPoolDispat
}, nil
case providerExecutionPathNormalized:
- return s.dispatchProviderPoolRun(ctx, req.Run, adapter, target, selected, queueReason, longReserved)
+ runReq := req.Run
+ if req.PrepareRun != nil {
+ runReq, err = req.PrepareRun(runReq)
+ if err != nil {
+ s.queue.releaseSlotWithLong(req.Run.ModelGroupKey, selected.entry.NodeID, selected.providerID, longReserved)
+ return nil, err
+ }
+ }
+ return s.dispatchProviderPoolRun(ctx, runReq, adapter, target, selected, queueReason, longReserved)
default:
s.queue.releaseSlotWithLong(req.Run.ModelGroupKey, selected.entry.NodeID, selected.providerID, longReserved)