Honor explicit model URL overrides for NVIDIA endpoints (#18152)

Some NVIDIA hosted models (e.g. meta/llama-3.2-11b-vision-instruct )
expose a full
endpoint URL per model that does not follow the normal base_url +
url_suffix
assembly. Previously the Go driver always called {base}/chat/completions
, so chat
requests for these vision models hit the wrong endpoint and failed.

This PR adds an optional per-model url field in conf/models/nvidia.json
. When
present, every NVIDIA driver request (chat, streaming chat, embedding,
rerank, model
listing) uses it directly; otherwise the standard assembly is unchanged.
This commit is contained in:
mkaaad
2026-08-12 14:17:13 +08:00
committed by GitHub
parent 7b2d052f8a
commit 1b02abd487
4 changed files with 189 additions and 13 deletions

View File

@@ -72,6 +72,7 @@
"name": "meta/llama-3.2-11b-vision-instruct",
"content_length": 131072,
"max_output": 8192,
"url":"https://integrate.api.nvidia.com/v1/meta/llama-3.2-11b-vision-instruct",
"model_types": [
"chat",
"vision"
@@ -97,6 +98,7 @@
"name": "meta/llama-3.2-90b-vision-instruct",
"content_length": 131072,
"max_output": 8192,
"url":"https://integrate.api.nvidia.com/v1/meta/llama-3.2-90b-vision-instruct",
"model_types": [
"chat",
"vision"

View File

@@ -166,6 +166,7 @@ type Model struct {
Thinking *ModelThinking `json:"thinking"`
Tools *ModelTools `json:"tools"`
Class *string `json:"class"`
URL string `json:"url"`
MaxDimension *int `json:"max_dimension"` // used by embedding models
MaxBatchSize *int `json:"max_batch_size"` // used by embedding models
Dimensions []int `json:"dimensions"`

View File

@@ -64,6 +64,30 @@ func (n *NvidiaModel) Name() string {
return "nvidia"
}
// resolveEndpoint returns the endpoint URL for modelName. When the
// model's preset config in conf/models/nvidia.json carries an explicit
// url, that full endpoint is used as-is because it does not follow the
// standard baseURL + suffix assembly; otherwise the usual assembly
// applies. Every endpoint (chat, embedding, rerank, models list) goes
// through this path so the override is honored uniformly.
func (n *NvidiaModel) resolveEndpoint(modelName string, apiConfig *APIConfig, suffix string) (string, error) {
if modelName != "" {
if pm := GetProviderManager(); pm != nil {
if provider := pm.FindProvider("NVIDIA"); provider != nil {
if model := pm.FindModel(provider, modelName); model != nil && model.URL != "" {
return strings.TrimSuffix(model.URL, "/"), nil
}
}
}
}
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
if err != nil {
return "", err
}
return fmt.Sprintf("%s/%s", strings.TrimRight(resolvedBaseURL, "/"), strings.TrimLeft(suffix, "/")), nil
}
func (n *NvidiaModel) ChatWithMessages(ctx context.Context, modelName string, messages []Message, apiConfig *APIConfig, chatModelConfig *ChatConfig, modelUsage *common.ModelUsage) (*ChatResponse, error) {
if err := n.baseModel.APIConfigCheck(apiConfig); err != nil {
return nil, err
@@ -73,11 +97,10 @@ func (n *NvidiaModel) ChatWithMessages(ctx context.Context, modelName string, me
return nil, fmt.Errorf("messages is empty")
}
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
baseURL, err := n.resolveEndpoint(modelName, apiConfig, n.baseModel.URLSuffix.Chat)
if err != nil {
return nil, err
}
baseURL := fmt.Sprintf("%s/%s", resolvedBaseURL, n.baseModel.URLSuffix.Chat)
reqBody := buildRequestBody(chatModelConfig, modelName, messages, false)
if chatModelConfig != nil {
@@ -108,11 +131,10 @@ func (n *NvidiaModel) ChatStreamlyWithSender(ctx context.Context, modelName stri
return fmt.Errorf("messages is empty")
}
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
baseURL, err := n.resolveEndpoint(modelName, apiConfig, n.baseModel.URLSuffix.Chat)
if err != nil {
return err
}
baseURL := fmt.Sprintf("%s/%s", resolvedBaseURL, n.baseModel.URLSuffix.Chat)
reqBody := buildRequestBody(modelConfig, modelName, messages, true)
if modelConfig != nil {
@@ -152,13 +174,11 @@ func (n *NvidiaModel) Embed(ctx context.Context, modelName *string, texts []stri
return nil, fmt.Errorf("model name is required")
}
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
baseURL, err := n.resolveEndpoint(*modelName, apiConfig, n.baseModel.URLSuffix.Embedding)
if err != nil {
return nil, err
}
baseURL := fmt.Sprintf("%s/%s", strings.TrimSuffix(resolvedBaseURL, "/"), n.baseModel.URLSuffix.Embedding)
reqBody := map[string]interface{}{
"model": *modelName,
"input": texts,
@@ -264,13 +284,11 @@ func (n *NvidiaModel) Rerank(ctx context.Context, modelName *string, query strin
return nil, fmt.Errorf("model name is required")
}
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
baseURL, err := n.resolveEndpoint(*modelName, apiConfig, n.baseModel.URLSuffix.Rerank)
if err != nil {
return nil, err
}
baseURL := fmt.Sprintf("%s/%s", strings.TrimSuffix(resolvedBaseURL, "/"), n.baseModel.URLSuffix.Rerank)
topN := len(documents)
if rerankConfig != nil && rerankConfig.TopN > 0 && rerankConfig.TopN < topN {
topN = rerankConfig.TopN
@@ -376,11 +394,10 @@ func (n *NvidiaModel) ListModels(ctx context.Context, apiConfig *APIConfig) ([]L
return nil, err
}
resolvedBaseURL, err := n.baseModel.GetBaseURL(apiConfig)
baseURL, err := n.resolveEndpoint("", apiConfig, n.baseModel.URLSuffix.Models)
if err != nil {
return nil, err
}
baseURL := fmt.Sprintf("%s/%s", strings.TrimRight(resolvedBaseURL, "/"), strings.TrimLeft(n.baseModel.URLSuffix.Models, "/"))
modelListCtx, cancel := context.WithTimeout(ctx, nonStreamCallTimeout)
defer cancel()
@@ -422,7 +439,7 @@ func (n *NvidiaModel) ListModels(ctx context.Context, apiConfig *APIConfig) ([]L
provider = pm.FindProvider("NVIDIA")
}
models := parseNvidiaModelList(modelList, provider)
if n.usesHostedCatalog(resolvedBaseURL) {
if n.usesHostedCatalog(baseURL) {
catalogCtx, catalogCancel := context.WithTimeout(ctx, nonStreamCallTimeout)
catalog, catalogErr := n.fetchHostedCatalog(catalogCtx)
catalogCancel()

View File

@@ -229,6 +229,162 @@ func TestNvidiaCatalogResourceRejectsMalformedDeprecation(t *testing.T) {
}
}
func withNvidiaProviderManager(t *testing.T, models []*Model) {
t.Helper()
saved := providerManager
providerManager = &ProviderManager{
Providers: []Provider{{
Name: "NVIDIA",
URL: map[string]string{"default": "https://integrate.api.nvidia.com/v1"},
Models: models,
}},
}
t.Cleanup(func() { providerManager = saved })
}
func TestNvidiaResolveEndpointUsesModelURL(t *testing.T) {
const modelURL = "https://integrate.api.nvidia.com/v1/meta/llama-3.2-11b-vision-instruct"
withNvidiaProviderManager(t, []*Model{{
Name: "meta/llama-3.2-11b-vision-instruct",
URL: modelURL,
}})
driver := NewNvidiaModel(
map[string]string{"default": "https://integrate.api.nvidia.com/v1"},
URLSuffix{Chat: "chat/completions"},
)
got, err := driver.resolveEndpoint("meta/llama-3.2-11b-vision-instruct", &APIConfig{}, "chat/completions")
if err != nil {
t.Fatalf("resolveEndpoint() error = %v", err)
}
if got != modelURL {
t.Fatalf("resolveEndpoint() = %q, want %q", got, modelURL)
}
}
func TestNvidiaResolveEndpointFallsBackToAssembly(t *testing.T) {
withNvidiaProviderManager(t, []*Model{{Name: "meta/llama-3.1-8b-instruct"}})
driver := NewNvidiaModel(
map[string]string{"default": "https://integrate.api.nvidia.com/v1/"},
URLSuffix{Chat: "chat/completions", Embedding: "embeddings"},
)
got, err := driver.resolveEndpoint("meta/llama-3.1-8b-instruct", &APIConfig{}, "chat/completions")
if err != nil {
t.Fatalf("resolveEndpoint() error = %v", err)
}
if want := "https://integrate.api.nvidia.com/v1/chat/completions"; got != want {
t.Fatalf("resolveEndpoint() = %q, want %q", got, want)
}
got, err = driver.resolveEndpoint("meta/llama-3.1-8b-instruct", &APIConfig{}, "embeddings")
if err != nil {
t.Fatalf("resolveEndpoint() error = %v", err)
}
if want := "https://integrate.api.nvidia.com/v1/embeddings"; got != want {
t.Fatalf("resolveEndpoint() = %q, want %q", got, want)
}
}
func TestNvidiaResolveEndpointIgnoresManagerWhenNil(t *testing.T) {
saved := providerManager
providerManager = nil
defer func() { providerManager = saved }()
driver := NewNvidiaModel(
map[string]string{"default": "https://integrate.api.nvidia.com/v1"},
URLSuffix{},
)
got, err := driver.resolveEndpoint("meta/llama-3.2-11b-vision-instruct", &APIConfig{}, "chat/completions")
if err != nil {
t.Fatalf("resolveEndpoint() error = %v", err)
}
if want := "https://integrate.api.nvidia.com/v1/chat/completions"; got != want {
t.Fatalf("resolveEndpoint() = %q, want %q", got, want)
}
}
func TestNvidiaChatUsesModelSpecificURL(t *testing.T) {
withSSRFBypass(t)
const apiKey = "nvapi-test"
const modelName = "meta/llama-3.2-11b-vision-instruct"
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/v1/meta/llama-3.2-11b-vision-instruct" {
t.Fatalf("path = %s, want model-specific chat endpoint", r.URL.Path)
}
_ = json.NewEncoder(w).Encode(map[string]interface{}{
"choices": []map[string]interface{}{{
"message": map[string]interface{}{"role": "assistant", "content": "pong"},
"finish_reason": "stop",
}},
"usage": map[string]interface{}{"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
})
}))
defer server.Close()
withNvidiaProviderManager(t, []*Model{{
Name: modelName,
URL: server.URL + "/v1/meta/llama-3.2-11b-vision-instruct",
}})
driver := NewNvidiaModel(
map[string]string{"default": server.URL + "/v1"},
URLSuffix{Chat: "chat/completions"},
)
resp, err := driver.ChatWithMessages(
context.Background(),
modelName,
[]Message{{Role: "user", Content: "hi"}},
&APIConfig{ApiKey: ptr(apiKey)},
nil,
nil,
)
if err != nil {
t.Fatalf("ChatWithMessages() error = %v", err)
}
if resp.Answer == nil || *resp.Answer != "pong" {
t.Fatalf("answer = %v, want pong", resp.Answer)
}
}
func TestNvidiaEmbedUsesModelSpecificURL(t *testing.T) {
withSSRFBypass(t)
const apiKey = "nvapi-test"
const modelName = "nvidia/nv-embed-v1"
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != "/v1/embeddings/nvidia/nv-embed-v1" {
t.Fatalf("path = %s, want model-specific embed endpoint", r.URL.Path)
}
_ = json.NewEncoder(w).Encode(map[string]interface{}{
"data": []map[string]interface{}{{"index": 0, "embedding": []float64{0.1, 0.2}}},
})
}))
defer server.Close()
withNvidiaProviderManager(t, []*Model{{
Name: modelName,
URL: server.URL + "/v1/embeddings/nvidia/nv-embed-v1",
}})
driver := NewNvidiaModel(
map[string]string{"default": server.URL + "/v1"},
URLSuffix{Embedding: "embeddings"},
)
namePtr := modelName
got, err := driver.Embed(
context.Background(),
&namePtr,
[]string{"hello"},
&APIConfig{ApiKey: ptr(apiKey)},
nil,
nil,
)
if err != nil {
t.Fatalf("Embed() error = %v", err)
}
if len(got) != 1 || len(got[0].Embedding) != 2 {
t.Fatalf("embedding = %#v, want 1 vector of length 2", got)
}
}
func ptr[T any](value T) *T {
return &value
}