From 1b02abd4877202ab717d05e77b0252bc0c29196d Mon Sep 17 00:00:00 2001 From: mkaaad <119158371+mkaaad@users.noreply.github.com> Date: Wed, 12 Aug 2026 14:17:13 +0800 Subject: [PATCH] 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. --- conf/models/nvidia.json | 2 + internal/entity/models/model.go | 1 + internal/entity/models/nvidia.go | 43 ++++--- internal/entity/models/nvidia_test.go | 156 ++++++++++++++++++++++++++ 4 files changed, 189 insertions(+), 13 deletions(-) diff --git a/conf/models/nvidia.json b/conf/models/nvidia.json index a0c87b9736..5b10ec9822 100644 --- a/conf/models/nvidia.json +++ b/conf/models/nvidia.json @@ -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" diff --git a/internal/entity/models/model.go b/internal/entity/models/model.go index 330e0aa16e..dba48e769b 100644 --- a/internal/entity/models/model.go +++ b/internal/entity/models/model.go @@ -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"` diff --git a/internal/entity/models/nvidia.go b/internal/entity/models/nvidia.go index c8e6a00665..02da13f467 100644 --- a/internal/entity/models/nvidia.go +++ b/internal/entity/models/nvidia.go @@ -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() diff --git a/internal/entity/models/nvidia_test.go b/internal/entity/models/nvidia_test.go index 61659e05a1..e02e9a969d 100644 --- a/internal/entity/models/nvidia_test.go +++ b/internal/entity/models/nvidia_test.go @@ -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 }