Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 23 additions & 1 deletion transports/bifrost-http/handlers/inference.go
Original file line number Diff line number Diff line change
Expand Up @@ -2631,6 +2631,25 @@ func (h *CompletionHandler) videoRemix(ctx *fasthttp.RequestCtx) {
SendJSON(ctx, resp)
}

// resolveBatchProvider resolves the provider (and optional model) for a batch
// create request. Per the OpenAI spec, model is optional on POST /v1/batches —
// it lives inside each JSONL request body. When model is present it is parsed
// via resolveModelAndProvider; when absent the provider is taken from the
// ?provider= query param or x-model-provider header (same as fileUpload).
func resolveBatchProvider(ctx *fasthttp.RequestCtx, config *lib.Config, model string) (schemas.ModelProvider, string, error) {
if model != "" {
return resolveModelAndProvider(ctx, config, model)
}
p := string(ctx.QueryArgs().Peek("provider"))
if p == "" {
p = string(ctx.Request.Header.Peek("x-model-provider"))
}
if p == "" {
return "", "", fmt.Errorf("provider query parameter or x-model-provider header is required when model is not specified")
}
return schemas.ModelProvider(p), "", nil
}

// batchCreate handles POST /v1/batches - Create a new batch job
func (h *CompletionHandler) batchCreate(ctx *fasthttp.RequestCtx) {
var req BatchCreateRequest
Expand All @@ -2639,7 +2658,10 @@ func (h *CompletionHandler) batchCreate(ctx *fasthttp.RequestCtx) {
return
}

provider, modelName, err := resolveModelAndProvider(ctx, h.config, req.Model)
// model is optional on POST /v1/batches per the OpenAI spec — the model lives
// inside each JSONL request body. When omitted, resolve the provider from the
// x-model-provider header or ?provider= query param (same as fileUpload).
provider, modelName, err := resolveBatchProvider(ctx, h.config, req.Model)
if err != nil {
SendError(ctx, fasthttp.StatusBadRequest, err.Error())
return
Expand Down
81 changes: 81 additions & 0 deletions transports/bifrost-http/handlers/inference_batch_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
package handlers

import (
"strings"
"testing"

"github.com/maximhq/bifrost/transports/bifrost-http/lib"
"github.com/valyala/fasthttp"
)

// TestResolveBatchProvider covers the three resolution paths introduced to make
// model optional on POST /v1/batches (OpenAI spec: model lives in the JSONL body).
func TestResolveBatchProvider(t *testing.T) {
config := &lib.Config{}

cases := []struct {
name string
model string
header string // x-model-provider; empty = unset
query string // ?provider=; empty = unset
wantProvider string
wantModel string
wantErrMsg string // non-empty = error expected, substring match
}{
{
name: "model field: provider+model parsed",
model: "openai/gpt-4o-mini",
wantProvider: "openai",
wantModel: "gpt-4o-mini",
},
{
name: "no model, x-model-provider header",
header: "openai",
wantProvider: "openai",
wantModel: "",
},
{
name: "no model, ?provider= query param",
query: "anthropic",
wantProvider: "anthropic",
wantModel: "",
},
{
name: "no model, no provider → error",
wantErrMsg: "provider query parameter or x-model-provider header is required",
},
}

for _, tc := range cases {
t.Run(tc.name, func(t *testing.T) {
ctx := &fasthttp.RequestCtx{}
if tc.header != "" {
ctx.Request.Header.Set("x-model-provider", tc.header)
}
if tc.query != "" {
ctx.QueryArgs().Set("provider", tc.query)
}

provider, modelName, err := resolveBatchProvider(ctx, config, tc.model)

if tc.wantErrMsg != "" {
if err == nil {
t.Fatalf("expected error containing %q, got nil", tc.wantErrMsg)
}
if !strings.Contains(err.Error(), tc.wantErrMsg) {
t.Fatalf("error %q does not contain %q", err.Error(), tc.wantErrMsg)
}
return
}
if err != nil {
t.Fatalf("unexpected error: %v", err)
}
if string(provider) != tc.wantProvider {
t.Fatalf("provider = %q, want %q", provider, tc.wantProvider)
}
if modelName != tc.wantModel {
t.Fatalf("modelName = %q, want %q", modelName, tc.wantModel)
}
})
}
}